Compare commits

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

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

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

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

Implemented in both postgres and libsql backends.

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

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

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

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

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

Co-Authored-By: Claude Opus 4.6 <[email protected]>
2026-02-12 04:35:05 -08:00
107 changed files with 4277 additions and 7521 deletions
+1 -33
View File
@@ -198,9 +198,8 @@ When designing new features or systems, always prefer generic/extensible archite
### Error Handling
- Use `thiserror` for error types in `error.rs`
- Never use `.unwrap()` or `.expect()` in production code (tests are fine)
- Never use `.unwrap()` in production code (tests are fine)
- Map errors with context: `.map_err(|e| SomeError::Variant { reason: e.to_string() })?`
- Before committing, grep for `.unwrap()` and `.expect(` in changed files to catch violations mechanically
### Async
- All I/O is async with tokio
@@ -638,37 +637,6 @@ RUST_LOG=ironclaw=debug,tower_http=debug cargo run
- Keep functions focused, extract helpers when logic is reused
- Comments for non-obvious logic only
## Review & Fix Discipline
Hard-won lessons from code review -- follow these when fixing bugs or addressing review feedback.
### Fix the pattern, not just the instance
When a reviewer flags a bug (e.g., TOCTOU race in INSERT + SELECT-back), search the entire codebase for all instances of that same pattern. A fix in `SecretsStore::create()` that doesn't also fix `WasmToolStore::store()` is half a fix.
### Propagate architectural fixes to satellite types
If a core type changes its concurrency model (e.g., `LibSqlBackend` switches to connection-per-operation), every type that was handed a resource from the old model (e.g., `LibSqlSecretsStore`, `LibSqlWasmToolStore` holding a single `Connection`) must also be updated. Grep for the old type across the codebase.
### Schema translation is more than DDL
When translating a database schema between backends (PostgreSQL to libSQL, etc.), check for:
- **Indexes** -- diff `CREATE INDEX` statements between the two schemas
- **Seed data** -- check for `INSERT INTO` in migrations (e.g., `leak_detection_patterns`)
- **Semantic differences** -- document where SQL functions behave differently (e.g., `json_patch` vs `jsonb_set`)
### Feature flag testing
When adding feature-gated code, test compilation with each feature in isolation:
```bash
cargo check # default features
cargo check --no-default-features --features libsql # libsql only
cargo check --all-features # all features
```
Dead code behind the wrong `#[cfg]` gate will only show up when building with a single feature.
### Mechanical verification before committing
Run these checks on changed files before committing:
- `grep -rnE '\.unwrap\(|\.expect\(' <files>` -- no panics in production
- `grep -rn 'super::' <files>` -- use `crate::` imports
- If you fixed a pattern bug, `grep` for other instances of that pattern across `src/`
## Workspace & Memory System
Inspired by [OpenClaw](https://github.com/openclaw/openclaw), the workspace provides persistent memory for agents with a flexible filesystem-like structure.
Generated
+30 -176
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"
@@ -2311,11 +2210,11 @@ dependencies = [
"hyper 1.8.1",
"hyper-util",
"rustls",
"rustls-native-certs",
"rustls-pki-types",
"tokio",
"tokio-rustls",
"tower-service",
"webpki-roots",
]
[[package]]
@@ -2368,7 +2267,7 @@ dependencies = [
"tokio",
"tower-service",
"tracing",
"windows-registry 0.6.1",
"windows-registry",
]
[[package]]
@@ -2602,7 +2501,6 @@ dependencies = [
"blake3",
"bollard",
"bytes",
"chromiumoxide",
"chrono",
"clap",
"cron",
@@ -2650,7 +2548,6 @@ dependencies = [
"tower-http 0.6.8",
"tracing",
"tracing-subscriber",
"url",
"urlencoding",
"uuid",
"wasmparser 0.220.1",
@@ -2799,7 +2696,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 +3308,7 @@ dependencies = [
"libc",
"redox_syscall 0.5.18",
"smallvec",
"windows-link 0.2.1",
"windows-link",
]
[[package]]
@@ -4048,7 +3945,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",
@@ -4136,7 +4033,6 @@ dependencies = [
"pin-project-lite",
"quinn",
"rustls",
"rustls-native-certs",
"rustls-pki-types",
"serde",
"serde_json",
@@ -4154,6 +4050,7 @@ dependencies = [
"wasm-bindgen-futures",
"wasm-streams",
"web-sys",
"webpki-roots",
]
[[package]]
@@ -6284,7 +6181,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",
]
@@ -6340,6 +6237,15 @@ dependencies = [
"wasm-bindgen",
]
[[package]]
name = "webpki-roots"
version = "1.0.5"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "12bed680863276c63889429bfd6cab3b99943659923822de1c8a39c49e4d722c"
dependencies = [
"rustls-pki-types",
]
[[package]]
name = "which"
version = "4.4.2"
@@ -6352,17 +6258,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 +6291,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 +6367,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 +6394,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 +6417,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 +6426,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 +6471,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 +6511,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 +6669,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"
+3 -7
View File
@@ -2,7 +2,7 @@
name = "ironclaw"
version = "0.1.3"
edition = "2024"
rust-version = "1.92"
rust-version = "1.85"
description = "Secure personal AI assistant that protects your data and expands its capabilities on the fly"
authors = ["NEAR AI <[email protected]>"]
license = "MIT OR Apache-2.0"
@@ -22,7 +22,7 @@ tokio-stream = { version = "0.1", features = ["sync"] }
futures = "0.3"
# HTTP client
reqwest = { version = "0.12", default-features = false, features = ["json", "rustls-tls-native-roots", "stream"] }
reqwest = { version = "0.12", default-features = false, features = ["json", "rustls-tls", "stream"] }
# Serialization
serde = { version = "1", features = ["derive"] }
@@ -84,8 +84,7 @@ fs4 = "0.6"
# Secrecy for sensitive values
secrecy = { version = "0.10", features = ["serde"] }
# URL parsing and encoding
url = "2"
# URL encoding for OAuth flow
urlencoding = "2"
# Open URLs in browser
@@ -122,9 +121,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"
-46
View File
@@ -1,46 +0,0 @@
# Multi-stage Dockerfile for the IronClaw agent (cloud deployment).
#
# Build:
# docker build --platform linux/amd64 -t ironclaw:latest .
#
# Run:
# docker run --env-file .env -p 3000:3000 ironclaw:latest
# Stage 1: Build
FROM rust:1.92-slim-bookworm AS builder
RUN apt-get update && apt-get install -y --no-install-recommends \
pkg-config libssl-dev cmake gcc g++ \
&& rm -rf /var/lib/apt/lists/*
WORKDIR /app
# Copy manifests first for layer caching
COPY Cargo.toml Cargo.lock ./
# Copy source and build artifacts
COPY src/ src/
COPY migrations/ migrations/
COPY wit/ wit/
RUN cargo build --release --bin ironclaw
# Stage 2: Runtime
FROM debian:bookworm-slim
RUN apt-get update && apt-get install -y --no-install-recommends \
ca-certificates libssl3 \
&& rm -rf /var/lib/apt/lists/*
COPY --from=builder /app/target/release/ironclaw /usr/local/bin/ironclaw
COPY --from=builder /app/migrations /app/migrations
# Non-root user
RUN useradd -m -u 1000 -s /bin/bash ironclaw
USER ironclaw
EXPOSE 3000
ENV RUST_LOG=ironclaw=info
ENTRYPOINT ["ironclaw"]
+2 -2
View File
@@ -9,7 +9,7 @@
# The image includes common development tools so workers can build software,
# run tests, and execute shell commands.
FROM rust:1.92-bookworm AS builder
FROM rust:1.85-bookworm AS builder
WORKDIR /build
COPY . .
@@ -40,7 +40,7 @@ RUN apt-get update && apt-get install -y --no-install-recommends \
ENV RUSTUP_HOME=/usr/local/rustup \
CARGO_HOME=/usr/local/cargo \
PATH=/usr/local/cargo/bin:$PATH
RUN curl --proto '=https' --tlsv1.2 -sSf https://sh.rustup.rs | sh -s -- -y --default-toolchain 1.92.0 \
RUN curl --proto '=https' --tlsv1.2 -sSf https://sh.rustup.rs | sh -s -- -y --default-toolchain 1.85.0 \
&& chmod -R a+r /usr/local/rustup /usr/local/cargo
# Install Claude Code CLI (for claude-bridge mode)
+3 -3
View File
@@ -133,7 +133,7 @@ This document tracks feature parity between IronClaw (Rust implementation) and O
|---------|----------|----------|-------|
| Pi agent runtime | ✅ | | IronClaw uses custom runtime |
| RPC-based execution | ✅ | ✅ | Orchestrator/worker pattern |
| Multi-provider failover | ✅ | | `FailoverProvider` tries providers sequentially on retryable errors |
| Multi-provider failover | ✅ | | Provider fallback chains |
| Per-sender sessions | ✅ | ✅ | |
| Global sessions | ✅ | ❌ | Optional shared context |
| Session pruning | ✅ | ❌ | Auto cleanup old sessions |
@@ -173,7 +173,7 @@ This document tracks feature parity between IronClaw (Rust implementation) and O
| Feature | OpenClaw | IronClaw | Notes |
|---------|----------|----------|-------|
| Auto-discovery | ✅ | ❌ | |
| Failover chains | ✅ | | `FailoverProvider` with configurable `fallback_model` |
| Failover chains | ✅ | | Provider fallback |
| Cooldown management | ✅ | ❌ | Skip failed providers |
| Per-session model override | ✅ | ✅ | Model selector in TUI |
| Model selection UI | ✅ | ✅ | TUI keyboard shortcut |
@@ -419,7 +419,7 @@ This document tracks feature parity between IronClaw (Rust implementation) and O
- ❌ Slack channel (real implementation)
- ✅ Telegram channel (WASM, DM pairing, caption, /start)
- ❌ WhatsApp channel
- Multi-provider failover (`FailoverProvider` with retryable error classification)
- Multi-provider failover
- ❌ Hooks system (beforeInbound, beforeToolCall, etc.)
### P2 - Medium Priority
+36 -36
View File
@@ -181,42 +181,42 @@ External content passes through multiple security layers:
## Architecture
```
┌────────────────────────────────────────────────────────────────┐
│ Channels │
│ ┌──────┐ ┌──────┐ ┌─────────────┐ ┌─────────────┐ │
│ │ REPL │ │ HTTP │ │WASM Channels│ │ Web Gateway │ │
│ └──┬───┘ └──┬───┘ └──────┬──────┘ │ (SSE + WS) │ │
│ │ │ │ └──────┬──────┘ │
│ └─────────┴──────────────┴────────────────┘ │
│ │ │
│ ┌─────────▼─────────┐ │
│ │ Agent Loop │ Intent routing │
│ └────┬─────────┬───┘ │
│ │ │ │
│ ┌──────────▼───┐ ┌──▼──────────────┐ │
│ │ Scheduler │ │ Routines Engine │ │
│ │(parallel jobs)│ │(cron, event, wh) │ │
│ └──────┬───────┘ └────────┬─────────┘ │
│ │ │ │
│ ┌─────────────┼───────────────────┘ │
│ │ │ │
│ ┌───▼────┐ ┌────▼────────────────┐ │
│ │ Local │ │ Orchestrator │ │
│ │Workers │ │ ┌───────────────┐ │ │
│ │(in-proc)│ │ │ Docker Sandbox│ │ │
│ └───┬────┘ │ │ Containers │ │ │
│ │ │ │ ┌───────────┐ │ │ │
│ │ │ │ │Worker / CC│ │ │ │
│ │ │ │ └───────────┘ │ │ │
│ │ │ └───────────────┘ │ │
│ │ └─────────┬───────────┘ │
│ └──────────────────┤ │
│ │ │
│ ┌───────────▼──────────┐ │
│ │ Tool Registry │ │
│ │ Built-in, MCP, WASM │ │
│ └──────────────────────┘ │
└────────────────────────────────────────────────────────────────┘
┌────────────────────────────────────────────────────────────────────
│ Channels
│ ┌──────┐ ┌──────┐ ┌─────────────┐ ┌─────────────┐
│ │ REPL │ │ HTTP │ │WASM Channels│ │ Web Gateway │
│ └──┬───┘ └──┬───┘ └──────┬──────┘ │ (SSE + WS) │
│ │ │ │ └──────┬──────┘
│ └─────────┴──────────────┴────────────────┘
│ │
│ ┌─────────▼─────────┐
│ │ Agent Loop │ Intent routing
│ └────┬─────────┬───┘
│ │ │
│ ┌──────────▼───┐ ┌──▼──────────────┐
│ │ Scheduler │ │ Routines Engine │
│ │(parallel jobs)│ │(cron, event, wh) │
│ └──────┬───────┘ └────────┬─────────┘
│ │ │
│ ┌─────────────┼───────────────────┘
│ │ │
│ ┌───▼────┐ ┌────▼────────────────┐
│ │ Local │ │ Orchestrator │
│ │Workers │ │ ┌───────────────┐ │
│ │(in-proc)│ │ │ Docker Sandbox│ │
│ └───┬────┘ │ │ Containers │ │
│ │ │ │ ┌───────────┐ │ │
│ │ │ │ │Worker / CC│ │ │
│ │ │ │ └───────────┘ │ │
│ │ │ └───────────────┘ │
│ │ └─────────┬───────────┘
│ └──────────────────┤
│ │
│ ┌───────────▼──────────┐
│ │ Tool Registry │
│ │ Built-in, MCP, WASM │
│ └──────────────────────┘
└────────────────────────────────────────────────────────────────────
```
### Core Components
-13
View File
@@ -1,13 +0,0 @@
[Unit]
Description=Cloud SQL Auth Proxy
After=network.target
[Service]
Type=simple
DynamicUser=yes
ExecStart=/usr/local/bin/cloud-sql-proxy ironclaw-prod:us-central1:ironclaw-db --port=5432
Restart=always
RestartSec=5
[Install]
WantedBy=multi-user.target
-27
View File
@@ -1,27 +0,0 @@
# WARNING: Replace all CHANGE_ME values before deploying.
# Do not use placeholder passwords in production.
DATABASE_URL=postgres://ironclaw:CHANGE_ME@localhost:5432/ironclaw
# NEAR AI
NEARAI_SESSION_TOKEN=CHANGE_ME
NEARAI_MODEL=claude-3-5-sonnet-20241022
NEARAI_BASE_URL=https://cloud-api.near.ai
NEARAI_AUTH_URL=https://private.near.ai
NEARAI_API_MODE=chat_completions
# Agent
AGENT_NAME=ironclaw
CLI_ENABLED=false
# Web Gateway
GATEWAY_ENABLED=true
# 0.0.0.0 binds to all interfaces (required for Docker --network=host).
# Use 127.0.0.1 if running outside Docker or for local-only access.
GATEWAY_HOST=0.0.0.0
GATEWAY_PORT=3000
GATEWAY_AUTH_TOKEN=CHANGE_ME
# Disabled for initial deploy
SANDBOX_ENABLED=false
HEARTBEAT_ENABLED=false
EMBEDDING_ENABLED=false
-20
View File
@@ -1,20 +0,0 @@
[Unit]
Description=IronClaw AI Assistant
After=cloud-sql-proxy.service docker.service
Requires=cloud-sql-proxy.service
[Service]
Type=simple
ExecStartPre=/usr/bin/docker pull us-central1-docker.pkg.dev/ironclaw-prod/ironclaw/agent:latest
ExecStart=/usr/bin/docker run --rm \
--name ironclaw \
--env-file /opt/ironclaw/.env \
--network=host \
us-central1-docker.pkg.dev/ironclaw-prod/ironclaw/agent:latest \
--no-onboard
ExecStop=/usr/bin/docker stop ironclaw
Restart=always
RestartSec=10
[Install]
WantedBy=multi-user.target
-68
View File
@@ -1,68 +0,0 @@
#!/usr/bin/env bash
# VM bootstrap script for IronClaw on GCP Compute Engine.
#
# Run on a fresh Debian 12 VM after SSH:
# sudo bash setup.sh
#
# Prerequisites:
# - VM has the ironclaw-vm service account attached
# - Cloud SQL Auth Proxy accessible via IAM
# - Artifact Registry image pushed
set -euo pipefail
# Must run as root
if [ "$(id -u)" -ne 0 ]; then
echo "ERROR: This script must be run as root (sudo bash setup.sh)"
exit 1
fi
echo "==> Installing Docker"
apt-get update
apt-get install -y docker.io
systemctl enable docker
systemctl start docker
echo "==> Installing Cloud SQL Auth Proxy"
curl -fsSL -o /usr/local/bin/cloud-sql-proxy \
https://storage.googleapis.com/cloud-sql-connectors/cloud-sql-proxy/v2.14.3/cloud-sql-proxy.linux.amd64
chmod +x /usr/local/bin/cloud-sql-proxy
echo "==> Installing systemd services"
cp /tmp/deploy/cloud-sql-proxy.service /etc/systemd/system/
cp /tmp/deploy/ironclaw.service /etc/systemd/system/
systemctl daemon-reload
echo "==> Starting Cloud SQL Auth Proxy"
systemctl enable cloud-sql-proxy
systemctl start cloud-sql-proxy
echo "==> Configuring Docker registry auth"
# The VM service account provides Artifact Registry access
gcloud auth configure-docker us-central1-docker.pkg.dev --quiet
echo "==> Creating config directory"
# Owned by root, readable only by root. Docker reads --env-file as root
# before dropping to uid 1000 (ironclaw) inside the container.
mkdir -p /opt/ironclaw
chmod 700 /opt/ironclaw
if [ ! -f /opt/ironclaw/.env ]; then
echo "WARNING: /opt/ironclaw/.env does not exist."
echo "Create it with your configuration before starting IronClaw."
echo "See deploy/env.example for the required variables."
echo ""
echo "Then run: systemctl enable ironclaw && systemctl start ironclaw"
else
chmod 600 /opt/ironclaw/.env
echo "==> Starting IronClaw"
systemctl enable ironclaw
systemctl start ironclaw
fi
echo "==> Setup complete"
echo ""
echo "Verify with:"
echo " systemctl status cloud-sql-proxy"
echo " systemctl status ironclaw"
echo " docker logs ironclaw"
+1
View File
@@ -81,6 +81,7 @@ async fn main() -> anyhow::Result<()> {
let session = create_session_manager(SessionConfig {
auth_base_url: config.llm.nearai.auth_base_url.clone(),
session_path: config.llm.nearai.session_path.clone(),
..Default::default()
})
.await;
let llm = create_llm_provider(&config.llm, session)?;
+148 -219
View File
@@ -28,7 +28,7 @@ use crate::tools::ToolRegistry;
use crate::workspace::Workspace;
/// Collapse a tool output string into a single-line preview for display.
pub(crate) fn truncate_for_preview(output: &str, max_chars: usize) -> String {
fn truncate_for_preview(output: &str, max_chars: usize) -> String {
let collapsed: String = output
.chars()
.take(max_chars + 50)
@@ -37,14 +37,8 @@ pub(crate) fn truncate_for_preview(output: &str, max_chars: usize) -> String {
.split_whitespace()
.collect::<Vec<_>>()
.join(" ");
// char_indices gives us byte offsets at char boundaries, so the slice is always valid UTF-8.
if collapsed.chars().count() > max_chars {
let byte_offset = collapsed
.char_indices()
.nth(max_chars)
.map(|(i, _)| i)
.unwrap_or(collapsed.len());
format!("{}...", &collapsed[..byte_offset])
if collapsed.len() > max_chars {
format!("{}...", &collapsed[..max_chars])
} else {
collapsed
}
@@ -660,17 +654,19 @@ impl Agent {
}
// Restore response chain from conversation metadata
if let Some(store) = self.store()
&& let Ok(Some(metadata)) = store.get_conversation_metadata(thread_uuid).await
&& let Some(rid) = metadata
.get("last_response_id")
.and_then(|v| v.as_str())
.map(String::from)
{
thread.last_response_id = Some(rid.clone());
self.llm()
.seed_response_chain(&thread_uuid.to_string(), rid);
tracing::debug!("Restored response chain for thread {}", thread_uuid);
if let Some(store) = self.store() {
if let Ok(Some(metadata)) = store.get_conversation_metadata(thread_uuid).await {
if let Some(rid) = metadata
.get("last_response_id")
.and_then(|v| v.as_str())
.map(String::from)
{
thread.last_response_id = Some(rid.clone());
self.llm()
.seed_response_chain(&thread_uuid.to_string(), rid);
tracing::debug!("Restored response chain for thread {}", thread_uuid);
}
}
}
// Insert into session and register with session manager
@@ -958,12 +954,13 @@ impl Agent {
return;
}
if let Some(ref resp) = response
&& let Err(e) = store
if let Some(ref resp) = response {
if let Err(e) = store
.add_conversation_message(thread_id, "assistant", resp)
.await
{
tracing::warn!("Failed to persist assistant message: {}", e);
{
tracing::warn!("Failed to persist assistant message: {}", e);
}
}
});
}
@@ -1061,14 +1058,14 @@ impl Agent {
// Check if interrupted
{
let sess = session.lock().await;
if let Some(thread) = sess.threads.get(&thread_id)
&& thread.state == ThreadState::Interrupted
{
return Err(crate::error::JobError::ContextError {
id: thread_id,
reason: "Interrupted".to_string(),
if let Some(thread) = sess.threads.get(&thread_id) {
if thread.state == ThreadState::Interrupted {
return Err(crate::error::JobError::ContextError {
id: thread_id,
reason: "Interrupted".to_string(),
}
.into());
}
.into());
}
}
@@ -1143,11 +1140,11 @@ impl Agent {
// Record tool calls in the thread
{
let mut sess = session.lock().await;
if let Some(thread) = sess.threads.get_mut(&thread_id)
&& let Some(turn) = thread.last_turn_mut()
{
for tc in &tool_calls {
turn.record_tool_call(&tc.name, tc.arguments.clone());
if let Some(thread) = sess.threads.get_mut(&thread_id) {
if let Some(turn) = thread.last_turn_mut() {
for tc in &tool_calls {
turn.record_tool_call(&tc.name, tc.arguments.clone());
}
}
}
}
@@ -1155,56 +1152,54 @@ impl Agent {
// Execute each tool (with approval checking)
for tc in tool_calls {
// Check if tool requires approval
if let Some(tool) = self.tools().get(&tc.name).await
&& tool.requires_approval()
{
// Check if auto-approved for this session
let mut is_auto_approved = {
let sess = session.lock().await;
sess.is_tool_auto_approved(&tc.name)
};
// 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)
{
tracing::info!(
"Shell command '{}' requires explicit approval despite auto-approve",
cmd.chars().take(80).collect::<String>()
);
is_auto_approved = false;
}
if !is_auto_approved {
// Need approval - store pending request and return
let pending = PendingApproval {
request_id: Uuid::new_v4(),
tool_name: tc.name.clone(),
parameters: tc.arguments.clone(),
description: tool.description().to_string(),
tool_call_id: tc.id.clone(),
context_messages: context_messages.clone(),
if let Some(tool) = self.tools().get(&tc.name).await {
if tool.requires_approval() {
// Check if auto-approved for this session
let mut is_auto_approved = {
let sess = session.lock().await;
sess.is_tool_auto_approved(&tc.name)
};
return Ok(AgenticLoopResult::NeedApproval { pending });
// For shell commands, override auto-approval for
// destructive patterns that should always require
// explicit per-invocation approval.
if is_auto_approved && tc.name == "shell" {
if let Some(cmd) = tc
.arguments
.as_str()
.and_then(|s| {
serde_json::from_str::<serde_json::Value>(s).ok()
})
.and_then(|v| {
v.get("command")
.and_then(|c| c.as_str().map(String::from))
})
{
if crate::tools::builtin::shell::requires_explicit_approval(
&cmd,
) {
tracing::info!(
"Shell command '{}' requires explicit approval despite auto-approve",
cmd.chars().take(80).collect::<String>()
);
is_auto_approved = false;
}
}
}
if !is_auto_approved {
// Need approval - store pending request and return
let pending = PendingApproval {
request_id: Uuid::new_v4(),
tool_name: tc.name.clone(),
parameters: tc.arguments.clone(),
description: tool.description().to_string(),
tool_call_id: tc.id.clone(),
context_messages: context_messages.clone(),
};
return Ok(AgenticLoopResult::NeedApproval { pending });
}
}
}
@@ -1235,34 +1230,34 @@ impl Agent {
)
.await;
if let Ok(ref output) = tool_result
&& !output.is_empty()
{
let _ = self
.channels
.send_status(
&message.channel,
StatusUpdate::ToolResult {
name: tc.name.clone(),
preview: output.clone(),
},
&message.metadata,
)
.await;
if let Ok(ref output) = tool_result {
if !output.is_empty() {
let _ = self
.channels
.send_status(
&message.channel,
StatusUpdate::ToolResult {
name: tc.name.clone(),
preview: truncate_for_preview(output, 200),
},
&message.metadata,
)
.await;
}
}
// Record result in thread
{
let mut sess = session.lock().await;
if let Some(thread) = sess.threads.get_mut(&thread_id)
&& let Some(turn) = thread.last_turn_mut()
{
match &tool_result {
Ok(output) => {
turn.record_tool_result(serde_json::json!(output));
}
Err(e) => {
turn.record_tool_error(e.to_string());
if let Some(thread) = sess.threads.get_mut(&thread_id) {
if let Some(turn) = thread.last_turn_mut() {
match &tool_result {
Ok(output) => {
turn.record_tool_result(serde_json::json!(output));
}
Err(e) => {
turn.record_tool_error(e.to_string());
}
}
}
}
@@ -1645,17 +1640,17 @@ impl Agent {
};
// Verify request ID if provided
if let Some(req_id) = request_id
&& req_id != pending.request_id
{
// Put it back and return error
let mut sess = session.lock().await;
if let Some(thread) = sess.threads.get_mut(&thread_id) {
thread.await_approval(pending);
if let Some(req_id) = request_id {
if req_id != pending.request_id {
// Put it back and return error
let mut sess = session.lock().await;
if let Some(thread) = sess.threads.get_mut(&thread_id) {
thread.await_approval(pending);
}
return Ok(SubmissionResult::error(
"Request ID mismatch. Use the correct request ID.",
));
}
return Ok(SubmissionResult::error(
"Request ID mismatch. Use the correct request ID.",
));
}
if approved {
@@ -1709,20 +1704,20 @@ impl Agent {
)
.await;
if let Ok(ref output) = tool_result
&& !output.is_empty()
{
let _ = self
.channels
.send_status(
&message.channel,
StatusUpdate::ToolResult {
name: pending.tool_name.clone(),
preview: output.clone(),
},
&message.metadata,
)
.await;
if let Ok(ref output) = tool_result {
if !output.is_empty() {
let _ = self
.channels
.send_status(
&message.channel,
StatusUpdate::ToolResult {
name: pending.tool_name.clone(),
preview: truncate_for_preview(output, 200),
},
&message.metadata,
)
.await;
}
}
// Build context including the tool result
@@ -1731,15 +1726,15 @@ impl Agent {
// Record result in thread
{
let mut sess = session.lock().await;
if let Some(thread) = sess.threads.get_mut(&thread_id)
&& let Some(turn) = thread.last_turn_mut()
{
match &tool_result {
Ok(output) => {
turn.record_tool_result(serde_json::json!(output));
}
Err(e) => {
turn.record_tool_error(e.to_string());
if let Some(thread) = sess.threads.get_mut(&thread_id) {
if let Some(turn) = thread.last_turn_mut() {
match &tool_result {
Ok(output) => {
turn.record_tool_result(serde_json::json!(output));
}
Err(e) => {
turn.record_tool_error(e.to_string());
}
}
}
}
@@ -2099,15 +2094,15 @@ impl Agent {
}
// Persist new job to database (fire-and-forget)
if let Some(store) = self.store()
&& let Ok(ctx) = self.context_manager.get_context(job_id).await
{
let store = store.clone();
tokio::spawn(async move {
if let Err(e) = store.save_job(&ctx).await {
tracing::warn!("Failed to persist new job {}: {}", job_id, e);
}
});
if let Some(store) = self.store() {
if let Ok(ctx) = self.context_manager.get_context(job_id).await {
let store = store.clone();
tokio::spawn(async move {
if let Err(e) = store.save_job(&ctx).await {
tracing::warn!("Failed to persist new job {}: {}", job_id, e);
}
});
}
}
// Schedule for execution
@@ -2187,10 +2182,10 @@ impl Agent {
let mut output = String::from("Jobs:\n");
for job_id in jobs {
if let Ok(ctx) = self.context_manager.get_context(job_id).await
&& ctx.user_id == user_id
{
output.push_str(&format!(" {} - {} ({:?})\n", job_id, ctx.title, ctx.state));
if let Ok(ctx) = self.context_manager.get_context(job_id).await {
if ctx.user_id == user_id {
output.push_str(&format!(" {} - {} ({:?})\n", job_id, ctx.title, ctx.state));
}
}
}
@@ -2641,70 +2636,4 @@ mod tests {
assert!(detect_auth_awaiting("tool_activate", &result).is_none());
}
// --- truncate_for_preview tests ---
use super::truncate_for_preview;
#[test]
fn test_truncate_short_input() {
assert_eq!(truncate_for_preview("hello", 10), "hello");
}
#[test]
fn test_truncate_empty_input() {
assert_eq!(truncate_for_preview("", 10), "");
}
#[test]
fn test_truncate_exact_length() {
assert_eq!(truncate_for_preview("hello", 5), "hello");
}
#[test]
fn test_truncate_over_limit() {
let result = truncate_for_preview("hello world, this is long", 10);
assert!(result.ends_with("..."));
// "hello worl" = 10 chars + "..."
assert_eq!(result, "hello worl...");
}
#[test]
fn test_truncate_collapses_newlines() {
let result = truncate_for_preview("line1\nline2\nline3", 100);
assert!(!result.contains('\n'));
assert_eq!(result, "line1 line2 line3");
}
#[test]
fn test_truncate_collapses_whitespace() {
let result = truncate_for_preview("hello world", 100);
assert_eq!(result, "hello world");
}
#[test]
fn test_truncate_multibyte_utf8() {
// Each emoji is 4 bytes. Truncating at char boundary must not panic.
let input = "😀😁😂🤣😃😄😅😆😉😊";
let result = truncate_for_preview(input, 5);
assert!(result.ends_with("..."));
// First 5 chars = 5 emoji
assert_eq!(result, "😀😁😂🤣😃...");
}
#[test]
fn test_truncate_cjk_characters() {
// CJK chars are 3 bytes each in UTF-8.
let input = "你好世界测试数据很长的字符串";
let result = truncate_for_preview(input, 4);
assert_eq!(result, "你好世界...");
}
#[test]
fn test_truncate_mixed_multibyte_and_ascii() {
let input = "hello 世界 foo";
let result = truncate_for_preview(input, 8);
// 'h','e','l','l','o',' ','世','界' = 8 chars
assert_eq!(result, "hello 世界...");
}
}
-1
View File
@@ -26,7 +26,6 @@ pub mod task;
pub mod undo;
pub mod worker;
pub(crate) use agent_loop::truncate_for_preview;
pub use agent_loop::{Agent, AgentDeps};
pub use compaction::{CompactionResult, ContextCompactor};
pub use context_monitor::{CompactionStrategy, ContextBreakdown, ContextMonitor};
+3 -2
View File
@@ -103,9 +103,10 @@ impl RoutineEngine {
if let Trigger::Event {
channel: Some(ch), ..
} = &routine.trigger
&& ch != &message.channel
{
continue;
if ch != &message.channel {
continue;
}
}
// Regex match
+18 -18
View File
@@ -119,25 +119,25 @@ impl SelfRepair for DefaultSelfRepair {
let mut stuck_jobs = Vec::new();
for job_id in stuck_ids {
if let Ok(ctx) = self.context_manager.get_context(job_id).await
&& ctx.state == JobState::Stuck
{
let stuck_duration = ctx
.started_at
.map(|start| {
let now = Utc::now();
let duration = now.signed_duration_since(start);
Duration::from_secs(duration.num_seconds().max(0) as u64)
})
.unwrap_or_default();
if let Ok(ctx) = self.context_manager.get_context(job_id).await {
if ctx.state == JobState::Stuck {
let stuck_duration = ctx
.started_at
.map(|start| {
let now = Utc::now();
let duration = now.signed_duration_since(start);
Duration::from_secs(duration.num_seconds().max(0) as u64)
})
.unwrap_or_default();
stuck_jobs.push(StuckJob {
job_id,
last_activity: ctx.started_at.unwrap_or(ctx.created_at),
stuck_duration,
last_error: None,
repair_attempts: ctx.repair_attempts,
});
stuck_jobs.push(StuckJob {
job_id,
last_activity: ctx.started_at.unwrap_or(ctx.created_at),
stuck_duration,
last_error: None,
repair_attempts: ctx.repair_attempts,
});
}
}
}
+5 -5
View File
@@ -346,11 +346,11 @@ impl Thread {
let mut turn = Turn::new(turn_number, &msg.content);
// Check if next is assistant response
if let Some(next) = iter.peek()
&& next.role == crate::llm::Role::Assistant
{
let response = iter.next().expect("peeked");
turn.complete(&response.content);
if let Some(next) = iter.peek() {
if next.role == crate::llm::Role::Assistant {
let response = iter.next().expect("peeked");
turn.complete(&response.content);
}
}
self.turns.push(turn);
+4 -4
View File
@@ -199,10 +199,10 @@ impl SessionManager {
{
let sessions = self.sessions.read().await;
for user_id in &stale_users {
if let Some(session) = sessions.get(user_id)
&& let Ok(sess) = session.try_lock()
{
stale_thread_ids.extend(sess.threads.keys());
if let Some(session) = sessions.get(user_id) {
if let Ok(sess) = session.try_lock() {
stale_thread_ids.extend(sess.threads.keys());
}
}
}
}
+14 -13
View File
@@ -93,26 +93,27 @@ impl SubmissionParser {
// /thread <uuid> - switch thread
if let Some(rest) = lower.strip_prefix("/thread ") {
let rest = rest.trim();
if rest != "new"
&& let Ok(id) = Uuid::parse_str(rest)
{
return Submission::SwitchThread { thread_id: id };
if rest != "new" {
if let Ok(id) = Uuid::parse_str(rest) {
return Submission::SwitchThread { thread_id: id };
}
}
}
// /resume <uuid> - resume from checkpoint
if let Some(rest) = lower.strip_prefix("/resume ")
&& let Ok(id) = Uuid::parse_str(rest.trim())
{
return Submission::Resume { checkpoint_id: id };
if let Some(rest) = lower.strip_prefix("/resume ") {
if let Ok(id) = Uuid::parse_str(rest.trim()) {
return Submission::Resume { checkpoint_id: id };
}
}
// Try structured JSON approval (from web gateway's /api/chat/approval endpoint)
if trimmed.starts_with('{')
&& let Ok(submission) = serde_json::from_str::<Submission>(trimmed)
&& matches!(submission, Submission::ExecApproval { .. })
{
return submission;
if trimmed.starts_with('{') {
if let Ok(submission) = serde_json::from_str::<Submission>(trimmed) {
if matches!(submission, Submission::ExecApproval { .. }) {
return submission;
}
}
}
// Approval responses (simple yes/no/always for pending approvals)
+8 -30
View File
@@ -227,11 +227,11 @@ Report when the job is complete or if you encounter issues you cannot resolve."#
}
// Check for cancellation
if let Ok(ctx) = self.context_manager().get_context(self.job_id).await
&& ctx.state == JobState::Cancelled
{
tracing::info!("Worker for job {} detected cancellation", self.job_id);
return Ok(());
if let Ok(ctx) = self.context_manager().get_context(self.job_id).await {
if ctx.state == JobState::Cancelled {
tracing::info!("Worker for job {} detected cancellation", self.job_id);
return Ok(());
}
}
iteration += 1;
@@ -299,7 +299,6 @@ Report when the job is complete or if you encounter issues you cannot resolve."#
parameters: tc.arguments.clone(),
reasoning: String::new(),
alternatives: vec![],
tool_call_id: tc.id.clone(),
};
self.process_tool_result(reason_ctx, &selection, result)
@@ -566,7 +565,7 @@ Report when the job is complete or if you encounter issues you cannot resolve."#
);
reason_ctx.messages.push(ChatMessage::tool_result(
&selection.tool_call_id,
"tool_call_id",
&selection.tool_name,
wrapped,
));
@@ -598,7 +597,7 @@ Report when the job is complete or if you encounter issues you cannot resolve."#
}
reason_ctx.messages.push(ChatMessage::tool_result(
&selection.tool_call_id,
"tool_call_id",
&selection.tool_name,
format!("Error: {}", e),
));
@@ -648,15 +647,12 @@ Report when the job is complete or if you encounter issues you cannot resolve."#
.execute_tool(&action.tool_name, &action.parameters)
.await;
// Create a synthetic ToolSelection for process_tool_result.
// Plan actions don't originate from an LLM tool_call response so
// there is no real tool_call_id; generate a unique one.
// Create a synthetic ToolSelection for process_tool_result
let selection = ToolSelection {
tool_name: action.tool_name.clone(),
parameters: action.parameters.clone(),
reasoning: action.reasoning.clone(),
alternatives: vec![],
tool_call_id: format!("plan_{}_{}", self.job_id, i),
};
// Process the result
@@ -778,26 +774,8 @@ impl From<TaskOutput> for Result<String, Error> {
#[cfg(test)]
mod tests {
use crate::llm::ToolSelection;
use crate::util::llm_signals_completion;
#[test]
fn test_tool_selection_preserves_call_id() {
let selection = ToolSelection {
tool_name: "memory_search".to_string(),
parameters: serde_json::json!({"query": "test"}),
reasoning: "Need to search memory".to_string(),
alternatives: vec![],
tool_call_id: "call_abc123".to_string(),
};
assert_eq!(selection.tool_call_id, "call_abc123");
assert_ne!(
selection.tool_call_id, "tool_call_id",
"tool_call_id must not be the hardcoded placeholder string"
);
}
#[test]
fn test_completion_positive_signals() {
assert!(llm_signals_completion("The job is complete."));
+171 -234
View File
@@ -1,128 +1,147 @@
//! Bootstrap helpers for IronClaw.
//! Bootstrap configuration for IronClaw.
//!
//! The only setting that truly needs disk persistence before the database is
//! available is `DATABASE_URL` (chicken-and-egg: can't connect to DB without
//! it). Everything else is auto-detected or read from env vars.
//! These are the only settings that MUST live on disk because they're needed
//! before the database connection is established. Everything else lives in the
//! `settings` table in PostgreSQL.
//!
//! File: `~/.ironclaw/.env` (standard dotenvy format)
//! File: `~/.ironclaw/bootstrap.json`
use std::path::PathBuf;
/// Path to the IronClaw-specific `.env` file: `~/.ironclaw/.env`.
pub fn ironclaw_env_path() -> PathBuf {
dirs::home_dir()
.unwrap_or_else(|| PathBuf::from("."))
.join(".ironclaw")
.join(".env")
use serde::{Deserialize, Serialize};
use crate::settings::KeySource;
/// Minimal config needed to connect to the database and decrypt secrets.
///
/// This is the only JSON file IronClaw reads from disk at startup.
/// All other configuration lives in the `settings` table in PostgreSQL.
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct BootstrapConfig {
/// Database connection URL (postgres://...).
#[serde(default)]
pub database_url: Option<String>,
/// Database connection pool size.
#[serde(default)]
pub database_pool_size: Option<usize>,
/// Source for the secrets master key.
#[serde(default)]
pub secrets_master_key_source: KeySource,
/// Whether onboarding wizard has been completed.
#[serde(default)]
pub onboard_completed: bool,
}
/// Load env vars from `~/.ironclaw/.env` (in addition to the standard `.env`).
///
/// Call this **after** `dotenvy::dotenv()` so that the standard `./.env`
/// takes priority over `~/.ironclaw/.env`. dotenvy never overwrites
/// existing env vars, so the effective priority is:
///
/// explicit env vars > `./.env` > `~/.ironclaw/.env`
///
/// If `~/.ironclaw/.env` doesn't exist but the legacy `bootstrap.json` does,
/// extracts `DATABASE_URL` from it and writes the `.env` file (one-time
/// upgrade from the old config format).
pub fn load_ironclaw_env() {
let path = ironclaw_env_path();
if !path.exists() {
// One-time upgrade: extract DATABASE_URL from legacy bootstrap.json
migrate_bootstrap_json_to_env(&path);
}
if path.exists() {
let _ = dotenvy::from_path(&path);
}
}
/// If `bootstrap.json` exists, pull `database_url` out of it and write `.env`.
fn migrate_bootstrap_json_to_env(env_path: &std::path::Path) {
let ironclaw_dir = env_path
.parent()
.unwrap_or_else(|| std::path::Path::new("."));
let bootstrap_path = ironclaw_dir.join("bootstrap.json");
if !bootstrap_path.exists() {
return;
}
let content = match std::fs::read_to_string(&bootstrap_path) {
Ok(c) => c,
Err(_) => return,
};
// Minimal parse: just grab database_url from the JSON
let parsed: serde_json::Value = match serde_json::from_str(&content) {
Ok(v) => v,
Err(_) => return,
};
if let Some(url) = parsed.get("database_url").and_then(|v| v.as_str()) {
if let Some(parent) = env_path.parent()
&& let Err(e) = std::fs::create_dir_all(parent)
{
eprintln!("Warning: failed to create {}: {}", parent.display(), e);
return;
impl Default for BootstrapConfig {
fn default() -> Self {
Self {
database_url: None,
database_pool_size: None,
secrets_master_key_source: KeySource::None,
onboard_completed: false,
}
if let Err(e) = std::fs::write(env_path, format!("DATABASE_URL=\"{}\"\n", url)) {
eprintln!("Warning: failed to migrate bootstrap.json to .env: {}", e);
return;
}
}
impl BootstrapConfig {
/// Default bootstrap file path: `~/.ironclaw/bootstrap.json`.
pub fn default_path() -> PathBuf {
dirs::home_dir()
.unwrap_or_else(|| PathBuf::from("."))
.join(".ironclaw")
.join("bootstrap.json")
}
/// Legacy settings.json path (for migration detection).
pub fn legacy_settings_path() -> PathBuf {
dirs::home_dir()
.unwrap_or_else(|| PathBuf::from("."))
.join(".ironclaw")
.join("settings.json")
}
/// Load from the default path, falling back to legacy settings.json,
/// then to defaults if neither exists.
pub fn load() -> Self {
let bootstrap_path = Self::default_path();
if bootstrap_path.exists() {
return Self::load_from(&bootstrap_path);
}
rename_to_migrated(&bootstrap_path);
eprintln!(
"Migrated DATABASE_URL from bootstrap.json to {}",
env_path.display()
);
// Fall back to legacy settings.json (extract just the 4 bootstrap fields)
let legacy_path = Self::legacy_settings_path();
if legacy_path.exists() {
return Self::load_from_legacy(&legacy_path);
}
Self::default()
}
/// Load from a specific path.
pub fn load_from(path: &PathBuf) -> Self {
match std::fs::read_to_string(path) {
Ok(data) => serde_json::from_str(&data).unwrap_or_default(),
Err(_) => Self::default(),
}
}
/// Extract bootstrap fields from a legacy settings.json.
fn load_from_legacy(path: &PathBuf) -> Self {
match std::fs::read_to_string(path) {
Ok(data) => {
// The legacy Settings struct is a superset; serde will ignore extra fields.
serde_json::from_str(&data).unwrap_or_default()
}
Err(_) => Self::default(),
}
}
/// Save to the default path.
pub fn save(&self) -> std::io::Result<()> {
self.save_to(&Self::default_path())
}
/// Save to a specific path.
pub fn save_to(&self, path: &PathBuf) -> std::io::Result<()> {
if let Some(parent) = path.parent() {
std::fs::create_dir_all(parent)?;
}
let json = serde_json::to_string_pretty(self)
.map_err(|e| std::io::Error::new(std::io::ErrorKind::InvalidData, e.to_string()))?;
std::fs::write(path, json)
}
}
/// Write `DATABASE_URL` to `~/.ironclaw/.env`.
/// One-time migration from disk config files to the database settings table.
///
/// Creates the parent directory if it doesn't exist.
/// The value is 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<()> {
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))
}
/// One-time migration of legacy `~/.ironclaw/settings.json` into the database.
/// On first boot after upgrade, checks if:
/// 1. `~/.ironclaw/settings.json` exists
/// 2. The DB settings table is empty for this user
///
/// Only runs when a `settings.json` exists on disk AND the DB has no settings
/// yet. After the wizard writes directly to the DB, this path is only hit by
/// users upgrading from the old disk-only configuration.
///
/// After syncing, renames `settings.json` to `.migrated` so it won't trigger again.
/// If both conditions hold, migrates settings, MCP servers, and session data
/// to the database, writes `bootstrap.json`, and renames old files to `.migrated`.
pub async fn migrate_disk_to_db(
store: &dyn crate::db::Database,
user_id: &str,
) -> Result<(), MigrationError> {
let ironclaw_dir = dirs::home_dir()
.unwrap_or_else(|| PathBuf::from("."))
.join(".ironclaw");
let legacy_settings_path = ironclaw_dir.join("settings.json");
let legacy_settings_path = BootstrapConfig::legacy_settings_path();
if !legacy_settings_path.exists() {
tracing::debug!("No legacy settings.json found, skipping disk-to-DB migration");
return Ok(());
}
// If DB already has settings, this is not a first boot, the wizard already
// wrote directly to the DB. Just clean up the stale file.
// Only migrate if DB is empty for this user
let has_settings = store.has_settings(user_id).await.map_err(|e| {
MigrationError::Database(format!("Failed to check existing settings: {}", e))
})?;
if has_settings {
tracing::info!("DB already has settings, renaming stale settings.json");
rename_to_migrated(&legacy_settings_path);
tracing::debug!(
"DB already has settings for user '{}', skipping migration",
user_id
);
return Ok(());
}
@@ -141,14 +160,22 @@ pub async fn migrate_disk_to_db(
tracing::info!("Migrated {} settings to database", db_map.len());
}
// 2. Write DATABASE_URL to ~/.ironclaw/.env
if let Some(ref url) = settings.database_url {
save_database_url(url)
.map_err(|e| MigrationError::Io(format!("Failed to write .env: {}", e)))?;
tracing::info!("Wrote DATABASE_URL to {}", ironclaw_env_path().display());
}
// 2. Write bootstrap.json with the 4 essential fields
let bootstrap = BootstrapConfig {
database_url: settings.database_url.clone(),
database_pool_size: settings.database_pool_size,
secrets_master_key_source: settings.secrets_master_key_source,
onboard_completed: settings.onboard_completed,
};
bootstrap
.save()
.map_err(|e| MigrationError::Io(format!("Failed to write bootstrap.json: {}", e)))?;
tracing::info!("Wrote bootstrap.json");
// 3. Migrate mcp-servers.json if it exists
let ironclaw_dir = dirs::home_dir()
.unwrap_or_else(|| PathBuf::from("."))
.join(".ironclaw");
let mcp_path = ironclaw_dir.join("mcp-servers.json");
if mcp_path.exists() {
match std::fs::read_to_string(&mcp_path) {
@@ -209,19 +236,12 @@ pub async fn migrate_disk_to_db(
// 5. Rename settings.json to .migrated (don't delete, safety net)
rename_to_migrated(&legacy_settings_path);
// 6. Clean up old bootstrap.json if it exists (superseded by .env)
let old_bootstrap = ironclaw_dir.join("bootstrap.json");
if old_bootstrap.exists() {
rename_to_migrated(&old_bootstrap);
tracing::info!("Renamed old bootstrap.json to .migrated");
}
tracing::info!("Disk-to-DB migration complete");
Ok(())
}
/// Rename a file to `<name>.migrated` as a safety net.
fn rename_to_migrated(path: &std::path::Path) {
fn rename_to_migrated(path: &PathBuf) {
let mut migrated = path.as_os_str().to_owned();
migrated.push(".migrated");
if let Err(e) = std::fs::rename(path, &migrated) {
@@ -244,145 +264,62 @@ mod tests {
use tempfile::tempdir;
#[test]
fn test_save_and_load_database_url() {
fn test_bootstrap_save_load() {
let dir = tempdir().unwrap();
let env_path = dir.path().join(".env");
let path = dir.path().join("bootstrap.json");
// Write in the quoted format that save_database_url uses
let url = "postgres://localhost:5432/ironclaw_test";
std::fs::write(&env_path, format!("DATABASE_URL=\"{}\"\n", url)).unwrap();
let config = BootstrapConfig {
database_url: Some("postgres://localhost/test".to_string()),
database_pool_size: Some(5),
secrets_master_key_source: KeySource::Keychain,
onboard_completed: true,
};
// Verify the content is a valid dotenv line (quoted)
let content = std::fs::read_to_string(&env_path).unwrap();
config.save_to(&path).unwrap();
let loaded = BootstrapConfig::load_from(&path);
assert_eq!(
content,
"DATABASE_URL=\"postgres://localhost:5432/ironclaw_test\"\n"
loaded.database_url,
Some("postgres://localhost/test".to_string())
);
// Verify dotenvy can parse it (strips quotes automatically)
let parsed: Vec<(String, String)> = dotenvy::from_path_iter(&env_path)
.unwrap()
.filter_map(|r| r.ok())
.collect();
assert_eq!(parsed.len(), 1);
assert_eq!(parsed[0].0, "DATABASE_URL");
assert_eq!(parsed[0].1, url);
assert_eq!(loaded.database_pool_size, Some(5));
assert_eq!(loaded.secrets_master_key_source, KeySource::Keychain);
assert!(loaded.onboard_completed);
}
#[test]
fn test_save_database_url_with_hash_in_password() {
fn test_bootstrap_from_legacy_settings() {
let dir = tempdir().unwrap();
let env_path = dir.path().join(".env");
let path = dir.path().join("settings.json");
// URLs with # in the password are common (URL-encoded special chars).
// Without quoting, dotenvy treats # as a comment delimiter.
let url = "postgres://user:p%23ss@localhost:5432/ironclaw";
std::fs::write(&env_path, format!("DATABASE_URL=\"{}\"\n", url)).unwrap();
let parsed: Vec<(String, String)> = dotenvy::from_path_iter(&env_path)
.unwrap()
.filter_map(|r| r.ok())
.collect();
assert_eq!(parsed.len(), 1);
assert_eq!(parsed[0].0, "DATABASE_URL");
assert_eq!(parsed[0].1, url);
}
#[test]
fn test_save_database_url_creates_parent_dirs() {
let dir = tempdir().unwrap();
let nested = dir.path().join("deep").join("nested");
let env_path = nested.join(".env");
// Parent doesn't exist yet
assert!(!nested.exists());
// The global function uses a fixed path, so we test the logic directly
std::fs::create_dir_all(&nested).unwrap();
std::fs::write(&env_path, "DATABASE_URL=postgres://test\n").unwrap();
assert!(env_path.exists());
let content = std::fs::read_to_string(&env_path).unwrap();
assert!(content.contains("DATABASE_URL=postgres://test"));
}
#[test]
fn test_ironclaw_env_path() {
let path = ironclaw_env_path();
assert!(path.ends_with(".ironclaw/.env"));
}
#[test]
fn test_migrate_bootstrap_json_to_env() {
let dir = tempdir().unwrap();
let env_path = dir.path().join(".env");
let bootstrap_path = dir.path().join("bootstrap.json");
// Write a legacy bootstrap.json
let bootstrap_json = serde_json::json!({
"database_url": "postgres://localhost/ironclaw_upgrade",
"database_pool_size": 5,
// Write a legacy settings.json with many extra fields
let legacy = serde_json::json!({
"database_url": "postgres://localhost/ironclaw",
"database_pool_size": 10,
"secrets_master_key_source": "keychain",
"onboard_completed": true
"onboard_completed": true,
"selected_model": "claude-3-5-sonnet",
"agent": { "name": "testbot", "max_parallel_jobs": 3 },
"heartbeat": { "enabled": true }
});
std::fs::write(
&bootstrap_path,
serde_json::to_string_pretty(&bootstrap_json).unwrap(),
)
.unwrap();
std::fs::write(&path, serde_json::to_string_pretty(&legacy).unwrap()).unwrap();
assert!(!env_path.exists());
assert!(bootstrap_path.exists());
// Run the migration
migrate_bootstrap_json_to_env(&env_path);
// .env should now exist with DATABASE_URL
assert!(env_path.exists());
let content = std::fs::read_to_string(&env_path).unwrap();
let config = BootstrapConfig::load_from_legacy(&path);
assert_eq!(
content,
"DATABASE_URL=\"postgres://localhost/ironclaw_upgrade\"\n"
config.database_url,
Some("postgres://localhost/ironclaw".to_string())
);
// bootstrap.json should be renamed to .migrated
assert!(!bootstrap_path.exists());
assert!(dir.path().join("bootstrap.json.migrated").exists());
assert_eq!(config.database_pool_size, Some(10));
assert_eq!(config.secrets_master_key_source, KeySource::Keychain);
assert!(config.onboard_completed);
}
#[test]
fn test_migrate_bootstrap_json_no_database_url() {
let dir = tempdir().unwrap();
let env_path = dir.path().join(".env");
let bootstrap_path = dir.path().join("bootstrap.json");
// bootstrap.json with no database_url
let bootstrap_json = serde_json::json!({
"onboard_completed": false
});
std::fs::write(
&bootstrap_path,
serde_json::to_string_pretty(&bootstrap_json).unwrap(),
)
.unwrap();
migrate_bootstrap_json_to_env(&env_path);
// .env should NOT be created
assert!(!env_path.exists());
// bootstrap.json should remain (no migration happened)
assert!(bootstrap_path.exists());
}
#[test]
fn test_migrate_bootstrap_json_missing() {
let dir = tempdir().unwrap();
let env_path = dir.path().join(".env");
// No bootstrap.json at all
migrate_bootstrap_json_to_env(&env_path);
// Nothing should happen
assert!(!env_path.exists());
fn test_bootstrap_defaults() {
let config = BootstrapConfig::default();
assert!(config.database_url.is_none());
assert!(config.database_pool_size.is_none());
assert_eq!(config.secrets_master_key_source, KeySource::None);
assert!(!config.onboard_completed);
}
}
+7 -17
View File
@@ -33,16 +33,9 @@ use termimad::MadSkin;
use tokio::sync::mpsc;
use tokio_stream::wrappers::ReceiverStream;
use crate::agent::truncate_for_preview;
use crate::channels::{Channel, IncomingMessage, MessageStream, OutgoingResponse, StatusUpdate};
use crate::error::ChannelError;
/// Max characters for tool result previews in the terminal.
const CLI_TOOL_RESULT_MAX: usize = 200;
/// Max characters for thinking/status messages in the terminal.
const CLI_STATUS_MAX: usize = 200;
/// Slash commands available in the REPL.
const SLASH_COMMANDS: &[&str] = &[
"/help",
@@ -268,7 +261,7 @@ impl Channel for ReplChannel {
std::thread::spawn(move || {
// Single message mode: send it and return
if let Some(msg) = single_message {
let incoming = IncomingMessage::new("repl", "default", &msg);
let incoming = IncomingMessage::new("repl", "user", &msg);
let _ = tx.blocking_send(incoming);
return;
}
@@ -336,21 +329,21 @@ impl Channel for ReplChannel {
_ => {}
}
let msg = IncomingMessage::new("repl", "default", line);
let msg = IncomingMessage::new("repl", "user", line);
if tx.blocking_send(msg).is_err() {
break;
}
}
Err(ReadlineError::Interrupted) => {
// Ctrl+C: send /interrupt
let msg = IncomingMessage::new("repl", "default", "/interrupt");
let msg = IncomingMessage::new("repl", "user", "/interrupt");
if tx.blocking_send(msg).is_err() {
break;
}
}
Err(ReadlineError::Eof) => {
// Ctrl+D: send /quit so the agent loop runs graceful shutdown
let msg = IncomingMessage::new("repl", "default", "/quit");
let msg = IncomingMessage::new("repl", "user", "/quit");
let _ = tx.blocking_send(msg);
break;
}
@@ -407,8 +400,7 @@ impl Channel for ReplChannel {
match status {
StatusUpdate::Thinking(msg) => {
let display = truncate_for_preview(&msg, CLI_STATUS_MAX);
eprintln!(" \x1b[90m\u{25CB} {display}\x1b[0m");
eprintln!(" \x1b[90m\u{25CB} {msg}\x1b[0m");
}
StatusUpdate::ToolStarted { name } => {
eprintln!(" \x1b[33m\u{25CB} {name}\x1b[0m");
@@ -421,8 +413,7 @@ impl Channel for ReplChannel {
}
}
StatusUpdate::ToolResult { name: _, preview } => {
let display = truncate_for_preview(&preview, CLI_TOOL_RESULT_MAX);
eprintln!(" \x1b[90m{display}\x1b[0m");
eprintln!(" \x1b[90m{preview}\x1b[0m");
}
StatusUpdate::StreamChunk(chunk) => {
// Print separator on the false-to-true transition
@@ -447,8 +438,7 @@ impl Channel for ReplChannel {
}
StatusUpdate::Status(msg) => {
if debug || msg.contains("approval") || msg.contains("Approval") {
let display = truncate_for_preview(&msg, CLI_STATUS_MAX);
eprintln!(" \x1b[90m{display}\x1b[0m");
eprintln!(" \x1b[90m{msg}\x1b[0m");
}
}
StatusUpdate::ApprovalNeeded {
+43 -112
View File
@@ -76,9 +76,6 @@ struct ChannelStoreData {
credentials: HashMap<String, String>,
/// Pairing store for DM pairing (guest access control).
pairing_store: Arc<PairingStore>,
/// Dedicated tokio runtime for HTTP requests, lazily initialized.
/// Reused across multiple `http_request` calls within one execution.
http_runtime: Option<tokio::runtime::Runtime>,
}
impl ChannelStoreData {
@@ -99,7 +96,6 @@ impl ChannelStoreData {
table: ResourceTable::new(),
credentials,
pairing_store,
http_runtime: None,
}
}
@@ -138,13 +134,13 @@ impl ChannelStoreData {
if result.contains('{') && result.contains('}') {
// Only warn if it looks like an unresolved placeholder (not JSON braces)
let brace_pattern = regex::Regex::new(r"\{[A-Z_]+\}").ok();
if let Some(re) = brace_pattern
&& re.is_match(&result)
{
tracing::warn!(
context = %context,
"String may contain unresolved credential placeholders"
);
if let Some(re) = brace_pattern {
if re.is_match(&result) {
tracing::warn!(
context = %context,
"String may contain unresolved credential placeholders"
);
}
}
}
@@ -287,25 +283,10 @@ impl near::agent::channel_host::Host for ChannelStoreData {
.map(|h| h.max_response_bytes)
.unwrap_or(10 * 1024 * 1024);
// Make the HTTP request using a dedicated single-threaded runtime.
// We're inside spawn_blocking, so we can't rely on the main runtime's
// I/O driver (it may be busy with WASM compilation or other startup work).
// A dedicated runtime gives us our own I/O driver and avoids contention.
// The runtime is lazily created and reused across calls within one execution.
if self.http_runtime.is_none() {
self.http_runtime = Some(
tokio::runtime::Builder::new_current_thread()
.enable_all()
.build()
.map_err(|e| format!("Failed to create HTTP runtime: {e}"))?,
);
}
let rt = self.http_runtime.as_ref().expect("just initialized");
let result = rt.block_on(async {
let client = reqwest::Client::builder()
.connect_timeout(std::time::Duration::from_secs(10))
.build()
.map_err(|e| format!("Failed to build HTTP client: {e}"))?;
// Make the HTTP request using blocking I/O
// We're already in a spawn_blocking context, so we can use block_on
let result = tokio::runtime::Handle::current().block_on(async {
let client = reqwest::Client::new();
let mut request = match method.to_uppercase().as_str() {
"GET" => client.get(&url),
@@ -327,9 +308,9 @@ impl near::agent::channel_host::Host for ChannelStoreData {
request = request.body(body_bytes);
}
// Send request with caller-specified timeout (default 30s, max 5min).
let timeout_ms = timeout_ms.unwrap_or(30_000).min(300_000) as u64;
let timeout = std::time::Duration::from_millis(timeout_ms);
// Send request with caller-specified timeout (default 30s).
// Cap at callback_timeout to prevent outliving the host wrapper.
let timeout = std::time::Duration::from_millis(timeout_ms.unwrap_or(30_000) as u64);
let response = request.timeout(timeout).send().await.map_err(|e| {
// Walk the full error chain so we get the actual root cause
// (DNS, TLS, connection refused, etc.) instead of just
@@ -357,13 +338,13 @@ impl near::agent::channel_host::Host for ChannelStoreData {
// Enforce max response body size to prevent memory exhaustion.
let max_response = max_response_bytes;
if let Some(cl) = response.content_length()
&& cl as usize > max_response
{
return Err(format!(
"Response body too large: {} bytes exceeds limit of {} bytes",
cl, max_response
));
if let Some(cl) = response.content_length() {
if cl as usize > max_response {
return Err(format!(
"Response body too large: {} bytes exceeds limit of {} bytes",
cl, max_response
));
}
}
let body = response
.bytes()
@@ -814,21 +795,7 @@ impl WasmChannel {
.await;
match result {
Ok(Ok((config, mut host_state))) => {
// Surface WASM guest logs (errors/warnings from webhook setup, etc.)
for entry in host_state.take_logs() {
match entry.level {
crate::tools::wasm::LogLevel::Error => {
tracing::error!(channel = %self.name, "{}", entry.message);
}
crate::tools::wasm::LogLevel::Warn => {
tracing::warn!(channel = %self.name, "{}", entry.message);
}
_ => {
tracing::debug!(channel = %self.name, "{}", entry.message);
}
}
}
Ok(Ok((config, _host_state))) => {
tracing::info!(
channel = %self.name,
display_name = %config.display_name,
@@ -1528,8 +1495,8 @@ impl WasmChannel {
match result {
Ok(emitted_messages) => {
// Process any emitted messages
if !emitted_messages.is_empty()
&& let Err(e) = Self::dispatch_emitted_messages(
if !emitted_messages.is_empty() {
if let Err(e) = Self::dispatch_emitted_messages(
&channel_name,
emitted_messages,
&message_tx,
@@ -1541,6 +1508,7 @@ impl WasmChannel {
"Failed to dispatch emitted messages from poll"
);
}
}
}
Err(e) => {
tracing::warn!(
@@ -1770,22 +1738,22 @@ impl Channel for WasmChannel {
*self.endpoints.write().await = endpoints;
// Start polling if configured
if let Some(poll_config) = &config.poll
&& poll_config.enabled
{
let interval = self
.capabilities
.validate_poll_interval(poll_config.interval_ms)
.map_err(|e| ChannelError::StartupFailed {
name: self.name.clone(),
reason: e,
})?;
if let Some(poll_config) = &config.poll {
if poll_config.enabled {
let interval = self
.capabilities
.validate_poll_interval(poll_config.interval_ms)
.map_err(|e| ChannelError::StartupFailed {
name: self.name.clone(),
reason: e,
})?;
// Create shutdown channel for polling and store the sender to keep it alive
let (poll_shutdown_tx, poll_shutdown_rx) = oneshot::channel();
*self.poll_shutdown_tx.write().await = Some(poll_shutdown_tx);
// Create shutdown channel for polling and store the sender to keep it alive
let (poll_shutdown_tx, poll_shutdown_rx) = oneshot::channel();
*self.poll_shutdown_tx.write().await = Some(poll_shutdown_tx);
self.start_polling(Duration::from_millis(interval as u64), poll_shutdown_rx);
self.start_polling(Duration::from_millis(interval as u64), poll_shutdown_rx);
}
}
tracing::info!(
@@ -2648,52 +2616,15 @@ mod tests {
assert_eq!(store.redact_credentials(input), input);
}
/// Verify that WASM HTTP host functions work using a dedicated
/// current-thread runtime inside spawn_blocking.
/// Verify that the block_on-inside-spawn_blocking pattern used by the WASM
/// channel HTTP host function doesn't deadlock or panic.
#[tokio::test]
async fn test_dedicated_runtime_inside_spawn_blocking() {
async fn test_block_on_inside_spawn_blocking_does_not_deadlock() {
let result = tokio::task::spawn_blocking(|| {
let rt = tokio::runtime::Builder::new_current_thread()
.enable_all()
.build()
.expect("failed to build runtime");
rt.block_on(async { 42 })
tokio::runtime::Handle::current().block_on(async { 42 })
})
.await
.expect("spawn_blocking panicked");
assert_eq!(result, 42);
}
/// Verify a real HTTP request works using the dedicated-runtime pattern.
/// This catches DNS, TLS, and I/O driver issues that trivial tests miss.
#[tokio::test]
#[ignore] // requires network
async fn test_dedicated_runtime_real_http() {
let result = tokio::task::spawn_blocking(|| {
let rt = tokio::runtime::Builder::new_current_thread()
.enable_all()
.build()
.expect("failed to build runtime");
rt.block_on(async {
let client = reqwest::Client::builder()
.connect_timeout(std::time::Duration::from_secs(10))
.build()
.expect("failed to build client");
let resp = client
.get("https://api.telegram.org/bot000/getMe")
.timeout(std::time::Duration::from_secs(10))
.send()
.await;
match resp {
Ok(r) => r.status().as_u16(),
Err(e) if e.is_timeout() => panic!("request timed out: {e}"),
Err(e) => panic!("unexpected error: {e}"),
}
})
})
.await
.expect("spawn_blocking panicked");
// 404 because "000" is not a valid bot token
assert_eq!(result, 404);
}
}
+12 -10
View File
@@ -25,21 +25,23 @@ pub async fn auth_middleware(
next: Next,
) -> Response {
// Try Authorization header first (constant-time comparison)
if let Some(auth_header) = headers.get("authorization")
&& let Ok(value) = auth_header.to_str()
&& let Some(token) = value.strip_prefix("Bearer ")
&& bool::from(token.as_bytes().ct_eq(auth.token.as_bytes()))
{
return next.run(request).await;
if let Some(auth_header) = headers.get("authorization") {
if let Ok(value) = auth_header.to_str() {
if let Some(token) = value.strip_prefix("Bearer ") {
if bool::from(token.as_bytes().ct_eq(auth.token.as_bytes())) {
return next.run(request).await;
}
}
}
}
// Fall back to query parameter for SSE EventSource (constant-time comparison)
if let Some(query) = request.uri().query() {
for pair in query.split('&') {
if let Some(token) = pair.strip_prefix("token=")
&& bool::from(token.as_bytes().ct_eq(auth.token.as_bytes()))
{
return next.run(request).await;
if let Some(token) = pair.strip_prefix("token=") {
if bool::from(token.as_bytes().ct_eq(auth.token.as_bytes())) {
return next.run(request).await;
}
}
}
}
+8 -8
View File
@@ -473,10 +473,10 @@ pub async fn chat_completions_handler(
if let Some(mt) = req.max_tokens {
tool_req = tool_req.with_max_tokens(mt);
}
if let Some(ref tc) = req.tool_choice
&& let Some(choice) = normalize_tool_choice(tc)
{
tool_req = tool_req.with_tool_choice(choice);
if let Some(ref tc) = req.tool_choice {
if let Some(choice) = normalize_tool_choice(tc) {
tool_req = tool_req.with_tool_choice(choice);
}
}
let resp = llm
@@ -591,10 +591,10 @@ async fn handle_streaming(
if let Some(mt) = req.max_tokens {
tool_req = tool_req.with_max_tokens(mt);
}
if let Some(ref tc) = req.tool_choice
&& let Some(choice) = normalize_tool_choice(tc)
{
tool_req = tool_req.with_tool_choice(choice);
if let Some(ref tc) = req.tool_choice {
if let Some(choice) = normalize_tool_choice(tc) {
tool_req = tool_req.with_tool_choice(choice);
}
}
LlmResult::WithTools(
llm.complete_with_tools(tool_req)
+151 -150
View File
@@ -525,10 +525,10 @@ pub async fn clear_auth_mode(state: &GatewayState) {
if let Some(ref sm) = state.session_manager {
let session = sm.get_or_create_session(&state.user_id).await;
let mut sess = session.lock().await;
if let Some(thread_id) = sess.active_thread
&& let Some(thread) = sess.threads.get_mut(&thread_id)
{
thread.pending_auth = None;
if let Some(thread_id) = sess.active_thread {
if let Some(thread) = sess.threads.get_mut(&thread_id) {
thread.pending_auth = None;
}
}
}
}
@@ -626,69 +626,69 @@ async fn chat_history_handler(
// Verify the thread belongs to the authenticated user before returning any data.
// In-memory threads are already scoped by user via session_manager, but DB
// lookups could expose another user's conversation if the UUID is guessed.
if query.thread_id.is_some()
&& let Some(ref store) = state.store
{
let owned = store
.conversation_belongs_to_user(thread_id, &state.user_id)
.await
.unwrap_or(false);
if !owned && !sess.threads.contains_key(&thread_id) {
return Err((StatusCode::NOT_FOUND, "Thread not found".to_string()));
if query.thread_id.is_some() {
if let Some(ref store) = state.store {
let owned = store
.conversation_belongs_to_user(thread_id, &state.user_id)
.await
.unwrap_or(false);
if !owned && !sess.threads.contains_key(&thread_id) {
return Err((StatusCode::NOT_FOUND, "Thread not found".to_string()));
}
}
}
// For paginated requests (before cursor set), always go to DB
if before_cursor.is_some()
&& let Some(ref store) = state.store
{
let (messages, has_more) = store
.list_conversation_messages_paginated(thread_id, before_cursor, limit as i64)
.await
.map_err(|e| (StatusCode::INTERNAL_SERVER_ERROR, e.to_string()))?;
if before_cursor.is_some() {
if let Some(ref store) = state.store {
let (messages, has_more) = store
.list_conversation_messages_paginated(thread_id, before_cursor, limit as i64)
.await
.map_err(|e| (StatusCode::INTERNAL_SERVER_ERROR, e.to_string()))?;
let oldest_timestamp = messages.first().map(|m| m.created_at.to_rfc3339());
let turns = build_turns_from_db_messages(&messages);
return Ok(Json(HistoryResponse {
thread_id,
turns,
has_more,
oldest_timestamp,
}));
let oldest_timestamp = messages.first().map(|m| m.created_at.to_rfc3339());
let turns = build_turns_from_db_messages(&messages);
return Ok(Json(HistoryResponse {
thread_id,
turns,
has_more,
oldest_timestamp,
}));
}
}
// Try in-memory first (freshest data for active threads)
if let Some(thread) = sess.threads.get(&thread_id)
&& !thread.turns.is_empty()
{
let turns: Vec<TurnInfo> = thread
.turns
.iter()
.map(|t| TurnInfo {
turn_number: t.turn_number,
user_input: t.user_input.clone(),
response: t.response.clone(),
state: format!("{:?}", t.state),
started_at: t.started_at.to_rfc3339(),
completed_at: t.completed_at.map(|dt| dt.to_rfc3339()),
tool_calls: t
.tool_calls
.iter()
.map(|tc| ToolCallInfo {
name: tc.name.clone(),
has_result: tc.result.is_some(),
has_error: tc.error.is_some(),
})
.collect(),
})
.collect();
if let Some(thread) = sess.threads.get(&thread_id) {
if !thread.turns.is_empty() {
let turns: Vec<TurnInfo> = thread
.turns
.iter()
.map(|t| TurnInfo {
turn_number: t.turn_number,
user_input: t.user_input.clone(),
response: t.response.clone(),
state: format!("{:?}", t.state),
started_at: t.started_at.to_rfc3339(),
completed_at: t.completed_at.map(|dt| dt.to_rfc3339()),
tool_calls: t
.tool_calls
.iter()
.map(|tc| ToolCallInfo {
name: tc.name.clone(),
has_result: tc.result.is_some(),
has_error: tc.error.is_some(),
})
.collect(),
})
.collect();
return Ok(Json(HistoryResponse {
thread_id,
turns,
has_more: false,
oldest_timestamp: None,
}));
return Ok(Json(HistoryResponse {
thread_id,
turns,
has_more: false,
oldest_timestamp: None,
}));
}
}
// Fall back to DB for historical threads not in memory (paginated)
@@ -738,12 +738,12 @@ fn build_turns_from_db_messages(messages: &[crate::history::ConversationMessage]
};
// Check if next message is an assistant response
if let Some(next) = iter.peek()
&& next.role == "assistant"
{
let assistant_msg = iter.next().expect("peeked");
turn.response = Some(assistant_msg.content.clone());
turn.completed_at = Some(assistant_msg.created_at.to_rfc3339());
if let Some(next) = iter.peek() {
if next.role == "assistant" {
let assistant_msg = iter.next().expect("peeked");
turn.response = Some(assistant_msg.content.clone());
turn.completed_at = Some(assistant_msg.created_at.to_rfc3339());
}
}
// Incomplete turn (user message without response)
@@ -1126,65 +1126,65 @@ async fn jobs_detail_handler(
.map_err(|_| (StatusCode::BAD_REQUEST, "Invalid job ID".to_string()))?;
// Try sandbox job from DB first, scoped to the authenticated user.
if let Some(ref store) = state.store
&& let Ok(Some(job)) = store.get_sandbox_job(job_id).await
{
if job.user_id != state.user_id {
return Err((StatusCode::NOT_FOUND, "Job not found".to_string()));
}
let browse_id = std::path::Path::new(&job.project_dir)
.file_name()
.map(|n| n.to_string_lossy().to_string())
.unwrap_or_else(|| job.id.to_string());
if let Some(ref store) = state.store {
if let Ok(Some(job)) = store.get_sandbox_job(job_id).await {
if job.user_id != state.user_id {
return Err((StatusCode::NOT_FOUND, "Job not found".to_string()));
}
let browse_id = std::path::Path::new(&job.project_dir)
.file_name()
.map(|n| n.to_string_lossy().to_string())
.unwrap_or_else(|| job.id.to_string());
let ui_state = match job.status.as_str() {
"creating" => "pending",
"running" => "in_progress",
s => s,
};
let ui_state = match job.status.as_str() {
"creating" => "pending",
"running" => "in_progress",
s => s,
};
let elapsed_secs = job.started_at.map(|start| {
let end = job.completed_at.unwrap_or_else(chrono::Utc::now);
(end - start).num_seconds().max(0) as u64
});
// Synthesize transitions from timestamps.
let mut transitions = Vec::new();
if let Some(started) = job.started_at {
transitions.push(TransitionInfo {
from: "creating".to_string(),
to: "running".to_string(),
timestamp: started.to_rfc3339(),
reason: None,
let elapsed_secs = job.started_at.map(|start| {
let end = job.completed_at.unwrap_or_else(chrono::Utc::now);
(end - start).num_seconds().max(0) as u64
});
}
if let Some(completed) = job.completed_at {
transitions.push(TransitionInfo {
from: "running".to_string(),
to: job.status.clone(),
timestamp: completed.to_rfc3339(),
reason: job.failure_reason.clone(),
});
}
return Ok(Json(JobDetailResponse {
id: job.id,
title: job.task.clone(),
description: String::new(),
state: ui_state.to_string(),
user_id: job.user_id.clone(),
created_at: job.created_at.to_rfc3339(),
started_at: job.started_at.map(|dt| dt.to_rfc3339()),
completed_at: job.completed_at.map(|dt| dt.to_rfc3339()),
elapsed_secs,
project_dir: Some(job.project_dir.clone()),
browse_url: Some(format!("/projects/{}/", browse_id)),
job_mode: {
let mode = store.get_sandbox_job_mode(job.id).await.ok().flatten();
mode.filter(|m| m != "worker")
},
transitions,
}));
// Synthesize transitions from timestamps.
let mut transitions = Vec::new();
if let Some(started) = job.started_at {
transitions.push(TransitionInfo {
from: "creating".to_string(),
to: "running".to_string(),
timestamp: started.to_rfc3339(),
reason: None,
});
}
if let Some(completed) = job.completed_at {
transitions.push(TransitionInfo {
from: "running".to_string(),
to: job.status.clone(),
timestamp: completed.to_rfc3339(),
reason: job.failure_reason.clone(),
});
}
return Ok(Json(JobDetailResponse {
id: job.id,
title: job.task.clone(),
description: String::new(),
state: ui_state.to_string(),
user_id: job.user_id.clone(),
created_at: job.created_at.to_rfc3339(),
started_at: job.started_at.map(|dt| dt.to_rfc3339()),
completed_at: job.completed_at.map(|dt| dt.to_rfc3339()),
elapsed_secs,
project_dir: Some(job.project_dir.clone()),
browse_url: Some(format!("/projects/{}/", browse_id)),
job_mode: {
let mode = store.get_sandbox_job_mode(job.id).await.ok().flatten();
mode.filter(|m| m != "worker")
},
transitions,
}));
}
}
Err((StatusCode::NOT_FOUND, "Job not found".to_string()))
@@ -1198,35 +1198,35 @@ async fn jobs_cancel_handler(
.map_err(|_| (StatusCode::BAD_REQUEST, "Invalid job ID".to_string()))?;
// Try sandbox job cancellation, scoped to the authenticated user.
if let Some(ref store) = state.store
&& let Ok(Some(job)) = store.get_sandbox_job(job_id).await
{
if job.user_id != state.user_id {
return Err((StatusCode::NOT_FOUND, "Job not found".to_string()));
}
if job.status == "running" || job.status == "creating" {
// Stop the container if we have a job manager.
if let Some(ref jm) = state.job_manager
&& let Err(e) = jm.stop_job(job_id).await
{
tracing::warn!(job_id = %job_id, error = %e, "Failed to stop container during cancellation");
if let Some(ref store) = state.store {
if let Ok(Some(job)) = store.get_sandbox_job(job_id).await {
if job.user_id != state.user_id {
return Err((StatusCode::NOT_FOUND, "Job not found".to_string()));
}
store
.update_sandbox_job_status(
job_id,
"failed",
Some(false),
Some("Cancelled by user"),
None,
Some(chrono::Utc::now()),
)
.await
.map_err(|e| (StatusCode::INTERNAL_SERVER_ERROR, e.to_string()))?;
if job.status == "running" || job.status == "creating" {
// Stop the container if we have a job manager.
if let Some(ref jm) = state.job_manager {
if let Err(e) = jm.stop_job(job_id).await {
tracing::warn!(job_id = %job_id, error = %e, "Failed to stop container during cancellation");
}
}
store
.update_sandbox_job_status(
job_id,
"failed",
Some(false),
Some("Cancelled by user"),
None,
Some(chrono::Utc::now()),
)
.await
.map_err(|e| (StatusCode::INTERNAL_SERVER_ERROR, e.to_string()))?;
}
return Ok(Json(serde_json::json!({
"status": "cancelled",
"job_id": job_id,
})));
}
return Ok(Json(serde_json::json!({
"status": "cancelled",
"job_id": job_id,
})));
}
Err((StatusCode::NOT_FOUND, "Job not found".to_string()))
@@ -1334,13 +1334,14 @@ async fn jobs_prompt_handler(
.map_err(|_| (StatusCode::BAD_REQUEST, "Invalid job ID".to_string()))?;
// Verify user owns this job.
if let Some(ref store) = state.store
&& !store
if let Some(ref store) = state.store {
if !store
.sandbox_job_belongs_to_user(job_id, &state.user_id)
.await
.unwrap_or(false)
{
return Err((StatusCode::NOT_FOUND, "Job not found".to_string()));
{
return Err((StatusCode::NOT_FOUND, "Job not found".to_string()));
}
}
let content = body
+60 -28
View File
@@ -48,6 +48,8 @@ pub enum ConfigCommand {
/// Connects to the database to read/write settings. Falls back to disk
/// if the database is not available.
pub async fn run_config_command(cmd: ConfigCommand) -> anyhow::Result<()> {
let _ = dotenvy::dotenv();
// Try to connect to the DB for settings access
let db: Option<Arc<dyn crate::db::Database>> = match connect_db().await {
Ok(d) => Some(d),
@@ -90,7 +92,7 @@ async fn load_settings(store: Option<&dyn crate::db::Database>) -> Settings {
_ => {}
}
}
Settings::default()
Settings::load()
}
/// List all settings.
@@ -108,10 +110,10 @@ async fn list_settings(
println!();
for (key, value) in all {
if let Some(ref f) = filter
&& !key.starts_with(f)
{
continue;
if let Some(ref f) = filter {
if !key.starts_with(f) {
continue;
}
}
let display_value = if value.len() > 60 {
@@ -153,17 +155,19 @@ async fn set_setting(
.set(path, value)
.map_err(|e| anyhow::anyhow!("{}", e))?;
let store = store.ok_or_else(|| {
anyhow::anyhow!("Database connection required to save settings. Check DATABASE_URL.")
})?;
let json_value = match serde_json::from_str::<serde_json::Value>(value) {
Ok(v) => v,
Err(_) => serde_json::Value::String(value.to_string()),
};
store
.set_setting(DEFAULT_USER_ID, path, &json_value)
.await
.map_err(|e| anyhow::anyhow!("Failed to save to database: {}", e))?;
// Save to DB if available, otherwise disk
if let Some(store) = store {
let json_value = match serde_json::from_str::<serde_json::Value>(value) {
Ok(v) => v,
Err(_) => serde_json::Value::String(value.to_string()),
};
store
.set_setting(DEFAULT_USER_ID, path, &json_value)
.await
.map_err(|e| anyhow::anyhow!("Failed to save to database: {}", e))?;
} else {
settings.save()?;
}
println!("Set {} = {}", path, value);
Ok(())
@@ -176,13 +180,17 @@ async fn reset_setting(store: Option<&dyn crate::db::Database>, path: &str) -> a
.get(path)
.ok_or_else(|| anyhow::anyhow!("Unknown setting: {}", path))?;
let store = store.ok_or_else(|| {
anyhow::anyhow!("Database connection required to reset settings. Check DATABASE_URL.")
})?;
store
.delete_setting(DEFAULT_USER_ID, path)
.await
.map_err(|e| anyhow::anyhow!("Failed to delete setting from database: {}", e))?;
// Delete from DB (falling back to default) or reset on disk
if let Some(store) = store {
store
.delete_setting(DEFAULT_USER_ID, path)
.await
.map_err(|e| anyhow::anyhow!("Failed to delete setting from database: {}", e))?;
} else {
let mut settings = Settings::load();
settings.reset(path).map_err(|e| anyhow::anyhow!("{}", e))?;
settings.save()?;
}
println!("Reset {} to default: {}", path, default_value);
Ok(())
@@ -192,13 +200,37 @@ async fn reset_setting(store: Option<&dyn crate::db::Database>, path: &str) -> a
fn show_path(has_db: bool) -> anyhow::Result<()> {
if has_db {
println!("Settings stored in: database (settings table)");
println!(
"Bootstrap config: {}",
crate::bootstrap::BootstrapConfig::default_path().display()
);
} else {
println!("Settings stored in: PostgreSQL (not connected, using defaults)");
let path = Settings::default_path();
println!("Settings stored in: {} (disk fallback)", path.display());
if path.exists() {
let metadata = std::fs::metadata(&path)?;
println!(" Size: {} bytes", metadata.len());
if let Ok(modified) = metadata.modified() {
use std::time::SystemTime;
let duration = SystemTime::now()
.duration_since(modified)
.unwrap_or_default();
let secs = duration.as_secs();
if secs < 60 {
println!(" Modified: {} seconds ago", secs);
} else if secs < 3600 {
println!(" Modified: {} minutes ago", secs / 60);
} else if secs < 86400 {
println!(" Modified: {} hours ago", secs / 3600);
} else {
println!(" Modified: {} days ago", secs / 86400);
}
}
} else {
println!(" (does not exist, using defaults)");
}
}
println!(
"Env config: {}",
crate::bootstrap::ironclaw_env_path().display()
);
Ok(())
}
-1
View File
@@ -12,7 +12,6 @@
mod config;
mod mcp;
pub mod memory;
pub mod oauth_defaults;
mod pairing;
pub mod status;
mod tool;
-343
View File
@@ -1,343 +0,0 @@
//! Shared OAuth infrastructure: built-in credentials, callback server, landing pages.
//!
//! Every OAuth flow in the codebase (WASM tool auth, MCP server auth, NEAR AI login)
//! uses the same callback port, landing page, and listener logic from this module.
//!
//! # Built-in Credentials
//!
//! Many CLI tools (gcloud, rclone, gdrive) ship with default OAuth credentials
//! so users don't need to register their own OAuth app. Google explicitly
//! documents that client_secret for "Desktop App" / "Installed App" types
//! is NOT actually secret.
//!
//! Default credentials are hardcoded below. They can be overridden at:
//!
//! - **Compile time**: Set IRONCLAW_GOOGLE_CLIENT_ID / IRONCLAW_GOOGLE_CLIENT_SECRET
//! env vars before building to replace the hardcoded defaults.
//! - **Runtime**: Users can set GOOGLE_OAUTH_CLIENT_ID / GOOGLE_OAUTH_CLIENT_SECRET
//! env vars, which take priority over built-in defaults.
use std::time::Duration;
use tokio::io::{AsyncBufReadExt, AsyncWriteExt, BufReader};
use tokio::net::TcpListener;
// ── Built-in credentials ────────────────────────────────────────────────
pub struct OAuthCredentials {
pub client_id: &'static str,
pub client_secret: &'static str,
}
/// Google OAuth "Desktop App" credentials, shared across all Google tools.
/// Compile-time env vars override the hardcoded defaults below.
const GOOGLE_CLIENT_ID: &str = match option_env!("IRONCLAW_GOOGLE_CLIENT_ID") {
Some(v) => v,
None => "564604149681-efo25d43rs85v0tibdepsmdv5dsrhhr0.apps.googleusercontent.com",
};
const GOOGLE_CLIENT_SECRET: &str = match option_env!("IRONCLAW_GOOGLE_CLIENT_SECRET") {
Some(v) => v,
None => "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",
};
/// Returns built-in OAuth credentials for a provider, keyed by secret_name.
///
/// The secret_name comes from the tool's capabilities.json `auth.secret_name` field.
/// Returns `None` if no built-in credentials are configured for that provider.
pub fn builtin_credentials(secret_name: &str) -> Option<OAuthCredentials> {
match secret_name {
"google_oauth_token" => Some(OAuthCredentials {
client_id: GOOGLE_CLIENT_ID,
client_secret: GOOGLE_CLIENT_SECRET,
}),
_ => None,
}
}
// ── Shared callback server ──────────────────────────────────────────────
/// Fixed port for all OAuth callbacks.
///
/// Every redirect URI registered with providers must use this port:
/// `http://localhost:9876/callback` (or `/auth/callback` for NEAR AI).
pub const OAUTH_CALLBACK_PORT: u16 = 9876;
/// Error from the OAuth callback listener.
#[derive(Debug, thiserror::Error)]
pub enum OAuthCallbackError {
#[error("Port {0} is in use (another auth flow running?): {1}")]
PortInUse(u16, String),
#[error("Authorization denied by user")]
Denied,
#[error("Timed out waiting for authorization")]
Timeout,
#[error("IO error: {0}")]
Io(String),
}
/// Bind the OAuth callback listener on the fixed port.
///
/// 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.
pub async fn bind_callback_listener() -> Result<TcpListener, OAuthCallbackError> {
let ipv6_addr = format!("[::1]:{}", OAUTH_CALLBACK_PORT);
match TcpListener::bind(&ipv6_addr).await {
Ok(listener) => return Ok(listener),
Err(e) if e.kind() == std::io::ErrorKind::AddrInUse => {
return Err(OAuthCallbackError::PortInUse(
OAUTH_CALLBACK_PORT,
e.to_string(),
));
}
Err(_) => {
// IPv6 not available on this host, fall back to IPv4
}
}
TcpListener::bind(format!("127.0.0.1:{}", OAUTH_CALLBACK_PORT))
.await
.map_err(|e| {
if e.kind() == std::io::ErrorKind::AddrInUse {
OAuthCallbackError::PortInUse(OAUTH_CALLBACK_PORT, e.to_string())
} else {
OAuthCallbackError::Io(e.to_string())
}
})
}
/// Wait for an OAuth callback and extract a query parameter value.
///
/// Listens for a GET request matching `path_prefix` (e.g., "/callback" or "/auth/callback"),
/// extracts the value of `param_name` (e.g., "code" or "token"), and shows a branded
/// landing page using `display_name` (e.g., "Google", "Notion", "NEAR AI").
///
/// Times out after 5 minutes.
pub async fn wait_for_callback(
listener: TcpListener,
path_prefix: &str,
param_name: &str,
display_name: &str,
) -> Result<String, OAuthCallbackError> {
let path_prefix = path_prefix.to_string();
let param_name = param_name.to_string();
let display_name = display_name.to_string();
tokio::time::timeout(Duration::from_secs(300), async move {
loop {
let (mut socket, _) = listener
.accept()
.await
.map_err(|e| OAuthCallbackError::Io(e.to_string()))?;
let mut reader = BufReader::new(&mut socket);
let mut request_line = String::new();
reader
.read_line(&mut request_line)
.await
.map_err(|e| OAuthCallbackError::Io(e.to_string()))?;
if let Some(path) = request_line.split_whitespace().nth(1)
&& path.starts_with(&path_prefix)
&& let Some(query) = path.split('?').nth(1)
{
// Check for error first
if query.contains("error=") {
let html = landing_html(&display_name, false);
let response = format!(
"HTTP/1.1 400 Bad Request\r\n\
Content-Type: text/html; charset=utf-8\r\n\
Connection: close\r\n\
\r\n\
{}",
html
);
let _ = socket.write_all(response.as_bytes()).await;
return Err(OAuthCallbackError::Denied);
}
// Look for the target parameter
for param in query.split('&') {
let parts: Vec<&str> = param.splitn(2, '=').collect();
if parts.len() == 2 && parts[0] == param_name {
let value = urlencoding::decode(parts[1])
.unwrap_or_else(|_| parts[1].into())
.into_owned();
let html = landing_html(&display_name, true);
let response = format!(
"HTTP/1.1 200 OK\r\n\
Content-Type: text/html; charset=utf-8\r\n\
Connection: close\r\n\
\r\n\
{}",
html
);
let _ = socket.write_all(response.as_bytes()).await;
let _ = socket.shutdown().await;
return Ok(value);
}
}
}
// Not the callback we're looking for
let response = "HTTP/1.1 404 Not Found\r\nConnection: close\r\n\r\n";
let _ = socket.write_all(response.as_bytes()).await;
}
})
.await
.map_err(|_| OAuthCallbackError::Timeout)?
}
/// Escape a string for safe interpolation into HTML content.
fn html_escape(s: &str) -> String {
let mut out = String::with_capacity(s.len());
for c in s.chars() {
match c {
'&' => out.push_str("&amp;"),
'<' => out.push_str("&lt;"),
'>' => out.push_str("&gt;"),
'"' => out.push_str("&quot;"),
'\'' => out.push_str("&#x27;"),
_ => out.push(c),
}
}
out
}
/// HTML landing page shown in the browser after an OAuth redirect.
pub fn landing_html(provider_name: &str, success: bool) -> String {
let safe_name = html_escape(provider_name);
let (icon, heading, subtitle, accent) = if success {
(
r##"<div style="width:64px;height:64px;border-radius:50%;background:#22c55e;display:flex;align-items:center;justify-content:center;margin:0 auto 24px">
<svg width="32" height="32" viewBox="0 0 24 24" fill="none" stroke="#fff" stroke-width="3" stroke-linecap="round" stroke-linejoin="round"><polyline points="20 6 9 17 4 12"/></svg>
</div>"##,
format!("{} Connected", safe_name),
"You can close this window and return to your terminal.",
"#22c55e",
)
} else {
(
r##"<div style="width:64px;height:64px;border-radius:50%;background:#ef4444;display:flex;align-items:center;justify-content:center;margin:0 auto 24px">
<svg width="32" height="32" viewBox="0 0 24 24" fill="none" stroke="#fff" stroke-width="3" stroke-linecap="round" stroke-linejoin="round"><line x1="18" y1="6" x2="6" y2="18"/><line x1="6" y1="6" x2="18" y2="18"/></svg>
</div>"##,
"Authorization Failed".to_string(),
"The request was denied. You can close this window and try again.",
"#ef4444",
)
};
format!(
r#"<!DOCTYPE html>
<html lang="en">
<head>
<meta charset="utf-8">
<meta name="viewport" content="width=device-width,initial-scale=1">
<title>IronClaw - {heading}</title>
<style>
* {{ margin:0; padding:0; box-sizing:border-box }}
body {{
font-family: -apple-system, BlinkMacSystemFont, "Segoe UI", Roboto, Helvetica, Arial, sans-serif;
background: #0a0a0a;
color: #e5e5e5;
display: flex;
justify-content: center;
align-items: center;
min-height: 100vh;
}}
.card {{
text-align: center;
padding: 48px 40px;
max-width: 420px;
border: 1px solid #262626;
border-radius: 16px;
background: #141414;
}}
h1 {{
font-size: 22px;
font-weight: 600;
margin-bottom: 8px;
color: #fafafa;
}}
p {{
font-size: 14px;
color: #a3a3a3;
line-height: 1.5;
}}
.accent {{ color: {accent}; }}
.brand {{
margin-top: 32px;
font-size: 12px;
color: #525252;
letter-spacing: 0.5px;
text-transform: uppercase;
}}
</style>
</head>
<body>
<div class="card">
{icon}
<h1>{heading}</h1>
<p>{subtitle}</p>
<div class="brand">IronClaw</div>
</div>
</body>
</html>"#,
heading = heading,
icon = icon,
subtitle = subtitle,
accent = accent,
)
}
#[cfg(test)]
mod tests {
use crate::cli::oauth_defaults::{builtin_credentials, landing_html};
#[test]
fn test_unknown_provider_returns_none() {
assert!(builtin_credentials("unknown_token").is_none());
}
#[test]
fn test_google_returns_based_on_compile_env() {
let creds = builtin_credentials("google_oauth_token");
assert!(creds.is_some());
let creds = creds.unwrap();
assert!(!creds.client_id.is_empty());
assert!(!creds.client_secret.is_empty());
}
#[test]
fn test_landing_html_success_contains_key_elements() {
let html = landing_html("Google", true);
assert!(html.contains("Google Connected"));
assert!(html.contains("charset"));
assert!(html.contains("IronClaw"));
assert!(html.contains("#22c55e")); // green accent
assert!(!html.contains("Failed"));
}
#[test]
fn test_landing_html_escapes_provider_name() {
let html = landing_html("<script>alert(1)</script>", true);
assert!(!html.contains("<script>"));
assert!(html.contains("&lt;script&gt;"));
}
#[test]
fn test_landing_html_error_contains_key_elements() {
let html = landing_html("Notion", false);
assert!(html.contains("Authorization Failed"));
assert!(html.contains("charset"));
assert!(html.contains("IronClaw"));
assert!(html.contains("#ef4444")); // red accent
assert!(!html.contains("Connected"));
}
}
+17 -15
View File
@@ -9,7 +9,7 @@ use crate::settings::Settings;
/// Run the status command, printing system health info.
pub async fn run_status_command() -> anyhow::Result<()> {
let settings = Settings::default();
let settings = Settings::load();
println!("IronClaw Status");
println!("===============\n");
@@ -22,9 +22,10 @@ pub async fn run_status_command() -> anyhow::Result<()> {
);
// Database
let db_url_set = std::env::var("DATABASE_URL").is_ok();
let db_url_set = settings.database_url.is_some() || std::env::var("DATABASE_URL").is_ok();
print!(" Database: ");
if db_url_set {
// Try to connect
match check_database().await {
Ok(()) => println!("connected"),
Err(e) => println!("error ({})", e),
@@ -42,14 +43,13 @@ pub async fn run_status_command() -> anyhow::Result<()> {
println!("not found (run `ironclaw onboard`)");
}
// Secrets (auto-detect: env var or keychain)
// Secrets
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 {
println!("configured (env)");
} else if has_keychain {
println!("configured (keychain)");
let secrets_configured = settings.secrets_master_key_source != crate::settings::KeySource::None
|| std::env::var("SECRETS_MASTER_KEY").is_ok()
|| crate::secrets::keychain::has_master_key().await;
if secrets_configured {
println!("configured ({:?})", settings.secrets_master_key_source);
} else {
println!("not configured");
}
@@ -129,18 +129,20 @@ pub async fn run_status_command() -> anyhow::Result<()> {
Err(_) => println!("none configured"),
}
// Config path
println!(
"\n Config: {}",
crate::bootstrap::ironclaw_env_path().display()
);
// Settings path
println!("\n Settings: {}", Settings::default_path().display());
Ok(())
}
#[cfg(feature = "postgres")]
async fn check_database() -> anyhow::Result<()> {
let url = std::env::var("DATABASE_URL").map_err(|_| anyhow::anyhow!("DATABASE_URL not set"))?;
let _ = dotenvy::dotenv();
let settings = Settings::load();
let url = std::env::var("DATABASE_URL")
.ok()
.or(settings.database_url)
.ok_or_else(|| anyhow::anyhow!("no URL"))?;
let config: deadpool_postgres::Config = deadpool_postgres::Config {
url: Some(url),
+153 -175
View File
@@ -423,11 +423,11 @@ async fn extract_crate_name(cargo_toml: &Path) -> anyhow::Result<String> {
// Simple TOML parsing for [package] name
for line in content.lines() {
let line = line.trim();
if line.starts_with("name")
&& let Some((_, value)) = line.split_once('=')
{
let name = value.trim().trim_matches('"').trim_matches('\'');
return Ok(name.to_string());
if line.starts_with("name") {
if let Some((_, value)) = line.split_once('=') {
let name = value.trim().trim_matches('"').trim_matches('\'');
return Ok(name.to_string());
}
}
}
@@ -491,10 +491,10 @@ async fn list_tools(dir: Option<PathBuf>, verbose: bool) -> anyhow::Result<()> {
if has_caps {
let caps_path = path.with_extension("capabilities.json");
if let Ok(content) = fs::read_to_string(&caps_path).await
&& let Ok(caps) = CapabilitiesFile::from_json(&content)
{
print_capabilities_summary(&caps);
if let Ok(content) = fs::read_to_string(&caps_path).await {
if let Ok(caps) = CapabilitiesFile::from_json(&content) {
print_capabilities_summary(&caps);
}
}
}
println!();
@@ -607,16 +607,16 @@ fn print_capabilities_summary(caps: &CapabilitiesFile) {
}
}
if let Some(ref secrets) = caps.secrets
&& !secrets.allowed_names.is_empty()
{
parts.push(format!("secrets: {}", secrets.allowed_names.len()));
if let Some(ref secrets) = caps.secrets {
if !secrets.allowed_names.is_empty() {
parts.push(format!("secrets: {}", secrets.allowed_names.len()));
}
}
if let Some(ref ws) = caps.workspace
&& !ws.allowed_prefixes.is_empty()
{
parts.push("workspace: read".to_string());
if let Some(ref ws) = caps.workspace {
if !ws.allowed_prefixes.is_empty() {
parts.push("workspace: read".to_string());
}
}
if !parts.is_empty() {
@@ -653,30 +653,30 @@ fn print_capabilities_detail(caps: &CapabilitiesFile) {
}
}
if let Some(ref secrets) = caps.secrets
&& !secrets.allowed_names.is_empty()
{
println!(" Secrets (existence check only):");
for name in &secrets.allowed_names {
println!(" {}", name);
if let Some(ref secrets) = caps.secrets {
if !secrets.allowed_names.is_empty() {
println!(" Secrets (existence check only):");
for name in &secrets.allowed_names {
println!(" {}", name);
}
}
}
if let Some(ref tool_invoke) = caps.tool_invoke
&& !tool_invoke.aliases.is_empty()
{
println!(" Tool aliases:");
for (alias, real_name) in &tool_invoke.aliases {
println!(" {} -> {}", alias, real_name);
if let Some(ref tool_invoke) = caps.tool_invoke {
if !tool_invoke.aliases.is_empty() {
println!(" Tool aliases:");
for (alias, real_name) in &tool_invoke.aliases {
println!(" {} -> {}", alias, real_name);
}
}
}
if let Some(ref ws) = caps.workspace
&& !ws.allowed_prefixes.is_empty()
{
println!(" Workspace read prefixes:");
for prefix in &ws.allowed_prefixes {
println!(" {}", prefix);
if let Some(ref ws) = caps.workspace {
if !ws.allowed_prefixes.is_empty() {
println!(" Workspace read prefixes:");
for prefix in &ws.allowed_prefixes {
println!(" {}", prefix);
}
}
}
}
@@ -802,100 +802,48 @@ async fn auth_tool(name: String, dir: Option<PathBuf>, user_id: String) -> anyho
}
// Check for environment variable
if let Some(ref env_var) = auth.env_var
&& let Ok(token) = std::env::var(env_var)
&& !token.is_empty()
{
println!(" Found {} in environment.", env_var);
println!();
if let Some(ref env_var) = auth.env_var {
if let Ok(token) = std::env::var(env_var) {
if !token.is_empty() {
println!(" Found {} in environment.", env_var);
println!();
// Validate if endpoint is provided
if let Some(ref validation) = auth.validation_endpoint {
print!(" Validating token...");
std::io::stdout().flush()?;
// Validate if endpoint is provided
if let Some(ref validation) = auth.validation_endpoint {
print!(" Validating token...");
std::io::stdout().flush()?;
match validate_token(&token, validation, &auth.secret_name).await {
Ok(()) => {
println!("");
}
Err(e) => {
println!("");
println!(" Validation failed: {}", e);
println!();
println!(" Falling back to manual entry...");
return auth_tool_manual(secrets_store.as_ref(), &user_id, &auth).await;
match validate_token(&token, validation, &auth.secret_name).await {
Ok(()) => {
println!("");
}
Err(e) => {
println!("");
println!(" Validation failed: {}", e);
println!();
println!(" Falling back to manual entry...");
return auth_tool_manual(secrets_store.as_ref(), &user_id, &auth).await;
}
}
}
// Save the token
save_token(secrets_store.as_ref(), &user_id, &auth, &token).await?;
print_success(display_name);
return Ok(());
}
}
// Save the token
save_token(secrets_store.as_ref(), &user_id, &auth, &token, None, None).await?;
print_success(display_name);
return Ok(());
}
// Check for OAuth configuration
if let Some(ref oauth) = auth.oauth {
// For providers with shared tokens (e.g., all Google tools share google_oauth_token),
// combine scopes from all installed tools so one auth covers everything.
let combined = combine_provider_scopes(&tools_dir, &auth.secret_name, oauth).await;
if combined.scopes.len() > oauth.scopes.len() {
let extra = combined.scopes.len() - oauth.scopes.len();
println!(
" Including scopes from {} other installed tool(s) sharing this credential.",
extra
);
println!();
}
return auth_tool_oauth(secrets_store.as_ref(), &user_id, &auth, &combined).await;
return auth_tool_oauth(secrets_store.as_ref(), &user_id, &auth, oauth).await;
}
// Fall back to manual entry
auth_tool_manual(secrets_store.as_ref(), &user_id, &auth).await
}
/// Scan the tools directory for all capabilities files sharing the same secret_name
/// and combine their OAuth scopes. This way, authing any Google tool requests scopes
/// for ALL installed Google tools, so one login covers everything.
async fn combine_provider_scopes(
tools_dir: &Path,
secret_name: &str,
base_oauth: &crate::tools::wasm::OAuthConfigSchema,
) -> crate::tools::wasm::OAuthConfigSchema {
let mut all_scopes: std::collections::HashSet<String> =
base_oauth.scopes.iter().cloned().collect();
if let Ok(mut entries) = tokio::fs::read_dir(tools_dir).await {
while let Ok(Some(entry)) = entries.next_entry().await {
let path = entry.path();
if path.extension().and_then(|e| e.to_str()) != Some("json") {
continue;
}
let name = path
.file_name()
.and_then(|n| n.to_str())
.unwrap_or_default();
if !name.ends_with(".capabilities.json") {
continue;
}
if let Ok(content) = tokio::fs::read_to_string(&path).await
&& let Ok(caps) = CapabilitiesFile::from_json(&content)
&& let Some(auth) = &caps.auth
&& auth.secret_name == secret_name
&& let Some(oauth) = &auth.oauth
{
all_scopes.extend(oauth.scopes.iter().cloned());
}
}
}
let mut combined = base_oauth.clone();
combined.scopes = all_scopes.into_iter().collect();
combined.scopes.sort(); // deterministic ordering
combined
}
/// OAuth browser-based login flow.
async fn auth_tool_oauth(
store: &(dyn SecretsStore + Send + Sync),
@@ -906,14 +854,12 @@ async fn auth_tool_oauth(
use base64::{Engine, engine::general_purpose::URL_SAFE_NO_PAD};
use rand::RngCore;
use sha2::{Digest, Sha256};
use crate::cli::oauth_defaults::{self, OAUTH_CALLBACK_PORT};
use tokio::io::{AsyncBufReadExt, AsyncWriteExt, BufReader};
use tokio::net::TcpListener;
let display_name = auth.display_name.as_deref().unwrap_or(&auth.secret_name);
// Get client_id: capabilities file > runtime env var > built-in defaults
let builtin = oauth_defaults::builtin_credentials(&auth.secret_name);
// Get client_id from config or env
let client_id = oauth
.client_id
.clone()
@@ -923,32 +869,41 @@ async fn auth_tool_oauth(
.as_ref()
.and_then(|env| std::env::var(env).ok())
})
.or_else(|| builtin.as_ref().map(|c| c.client_id.to_string()))
.ok_or_else(|| {
anyhow::anyhow!(
"OAuth client_id not configured.\n\
Set {} env var, or build with IRONCLAW_GOOGLE_CLIENT_ID.",
oauth.client_id_env.as_deref().unwrap_or("the client_id")
Set it in the capabilities file or via environment variable."
)
})?;
// Get client_secret: capabilities file > runtime env var > built-in defaults
let client_secret = oauth
.client_secret
.clone()
.or_else(|| {
oauth
.client_secret_env
.as_ref()
.and_then(|env| std::env::var(env).ok())
})
.or_else(|| builtin.as_ref().map(|c| c.client_secret.to_string()));
// Get client_secret if provided
let client_secret = oauth.client_secret.clone().or_else(|| {
oauth
.client_secret_env
.as_ref()
.and_then(|env| std::env::var(env).ok())
});
println!(" Starting OAuth authentication...");
println!();
let listener = oauth_defaults::bind_callback_listener().await?;
let redirect_uri = format!("http://localhost:{}/callback", OAUTH_CALLBACK_PORT);
// Find an available port for the callback
let mut listener = None;
let mut port = 0;
for p in 9876..=9886 {
match TcpListener::bind(format!("127.0.0.1:{}", p)).await {
Ok(l) => {
listener = Some(l);
port = p;
break;
}
Err(_) => continue,
}
}
let listener = listener.ok_or_else(|| anyhow::anyhow!("Could not find available port"))?;
let redirect_uri = format!("http://localhost:{}/callback", port);
// Generate PKCE verifier and challenge
let (code_verifier, code_challenge) = if oauth.use_pkce {
@@ -1007,8 +962,65 @@ async fn auth_tool_oauth(
println!(" Waiting for authorization...");
let code =
oauth_defaults::wait_for_callback(listener, "/callback", "code", display_name).await?;
// Wait for callback with timeout
let timeout = std::time::Duration::from_secs(300);
let code = tokio::time::timeout(timeout, async {
loop {
let (mut socket, _) = listener.accept().await?;
let mut reader = BufReader::new(&mut socket);
let mut request_line = String::new();
reader.read_line(&mut request_line).await?;
// Parse GET /callback?code=xxx HTTP/1.1
if let Some(path) = request_line.split_whitespace().nth(1) {
if path.starts_with("/callback") {
if let Some(query) = path.split('?').nth(1) {
for param in query.split('&') {
let parts: Vec<&str> = param.splitn(2, '=').collect();
if parts.len() == 2 && parts[0] == "code" {
let code = urlencoding::decode(parts[1])
.unwrap_or_else(|_| parts[1].into())
.into_owned();
// Send success response
let response = format!(
"HTTP/1.1 200 OK\r\n\
Content-Type: text/html\r\n\
\r\n\
<!DOCTYPE html><html><body style=\"font-family: sans-serif; \
display: flex; justify-content: center; align-items: center; \
height: 100vh; margin: 0; background: #191919; color: white;\">\
<div style=\"text-align: center;\">\
<h1> {} Connected!</h1>\
<p>You can close this window.</p>\
</div></body></html>",
display_name
);
let _ = socket.write_all(response.as_bytes()).await;
let _ = socket.shutdown().await;
return Ok::<_, anyhow::Error>(code);
}
}
// Check for error
if query.contains("error=") {
let response =
"HTTP/1.1 400 Bad Request\r\n\r\nAuthorization denied";
let _ = socket.write_all(response.as_bytes()).await;
return Err(anyhow::anyhow!("Authorization denied by user"));
}
}
}
}
let response = "HTTP/1.1 404 Not Found\r\n\r\n";
let _ = socket.write_all(response.as_bytes()).await;
}
})
.await
.map_err(|_| anyhow::anyhow!("Timed out waiting for authorization"))??;
println!();
println!(" Exchanging code for token...");
@@ -1059,19 +1071,8 @@ async fn auth_tool_oauth(
)
})?;
let refresh_token = token_data.get("refresh_token").and_then(|v| v.as_str());
let expires_in = token_data.get("expires_in").and_then(|v| v.as_u64());
// Save the token (with refresh token and expiry if provided)
save_token(
store,
user_id,
auth,
access_token,
refresh_token,
expires_in,
)
.await?;
// Save the token
save_token(store, user_id, auth, access_token).await?;
// Extract any additional info for display
let workspace_name = token_data
@@ -1173,8 +1174,8 @@ async fn auth_tool_manual(
}
}
// Save the token (manual path: no refresh token or expiry)
save_token(store, user_id, auth, &token, None, None).await?;
// Save the token
save_token(store, user_id, auth, &token).await?;
print_success(display_name);
Ok(())
}
@@ -1265,16 +1266,11 @@ async fn validate_token(
}
/// Save token to secrets store.
///
/// Optionally stores a refresh token (as `{secret_name}_refresh_token`) and
/// sets `expires_at` on the access token so the runtime can auto-refresh.
async fn save_token(
store: &(dyn SecretsStore + Send + Sync),
user_id: &str,
auth: &crate::tools::wasm::AuthCapabilitySchema,
token: &str,
refresh_token: Option<&str>,
expires_in: Option<u64>,
) -> anyhow::Result<()> {
let mut params = CreateSecretParams::new(&auth.secret_name, token);
@@ -1282,29 +1278,11 @@ async fn save_token(
params = params.with_provider(provider);
}
if let Some(secs) = expires_in {
let expires_at = chrono::Utc::now() + chrono::Duration::seconds(secs as i64);
params = params.with_expiry(expires_at);
}
store
.create(user_id, params)
.await
.map_err(|e| anyhow::anyhow!("Failed to save token: {}", e))?;
// Store refresh token separately (no expiry, it's long-lived)
if let Some(rt) = refresh_token {
let refresh_name = format!("{}_refresh_token", auth.secret_name);
let mut refresh_params = CreateSecretParams::new(&refresh_name, rt);
if let Some(ref provider) = auth.provider {
refresh_params = refresh_params.with_provider(provider);
}
store
.create(user_id, refresh_params)
.await
.map_err(|e| anyhow::anyhow!("Failed to save refresh token: {}", e))?;
}
Ok(())
}
+66 -53
View File
@@ -1,9 +1,9 @@
//! Configuration for IronClaw.
//!
//! Settings are loaded with priority: env var > database > default.
//! `DATABASE_URL` lives in `~/.ironclaw/.env` (loaded via dotenvy early
//! in startup). Everything else comes from env vars, the DB settings
//! table, or auto-detection.
//! The database replaces the old `settings.json` file for all settings
//! except the 4 bootstrap fields (database_url, pool_size, secrets key
//! source, onboard_completed) which live in `~/.ironclaw/bootstrap.json`.
use std::path::PathBuf;
use std::time::Duration;
@@ -40,9 +40,9 @@ impl Config {
pub async fn from_db(
store: &dyn crate::db::Database,
user_id: &str,
bootstrap: &crate::bootstrap::BootstrapConfig,
) -> Result<Self, ConfigError> {
let _ = dotenvy::dotenv();
crate::bootstrap::load_ironclaw_env();
// Load all settings from DB into a Settings struct
let db_settings = match store.get_all_settings(user_id).await {
@@ -53,7 +53,7 @@ impl Config {
}
};
Self::build(&db_settings).await
Self::build(bootstrap, &db_settings).await
}
/// Load configuration from environment variables only (no database).
@@ -61,20 +61,20 @@ impl Config {
/// Used during early startup before the database is connected,
/// and by CLI commands that don't have DB access.
/// Falls back to legacy `settings.json` on disk if present.
///
/// Loads both `./.env` (standard, higher priority) and `~/.ironclaw/.env`
/// (lower priority) via dotenvy, which never overwrites existing vars.
pub async fn from_env() -> Result<Self, ConfigError> {
let _ = dotenvy::dotenv();
crate::bootstrap::load_ironclaw_env();
let bootstrap = crate::bootstrap::BootstrapConfig::load();
let settings = Settings::load();
Self::build(&settings).await
Self::build(&bootstrap, &settings).await
}
/// Build config from settings (shared by from_env and from_db).
async fn build(settings: &Settings) -> Result<Self, ConfigError> {
/// Build config from bootstrap + settings (shared by from_env and from_db).
async fn build(
bootstrap: &crate::bootstrap::BootstrapConfig,
settings: &Settings,
) -> Result<Self, ConfigError> {
Ok(Self {
database: DatabaseConfig::resolve()?,
database: DatabaseConfig::resolve(bootstrap)?,
llm: LlmConfig::resolve(settings)?,
embeddings: EmbeddingsConfig::resolve(settings)?,
tunnel: TunnelConfig::resolve(settings)?,
@@ -82,7 +82,7 @@ impl Config {
agent: AgentConfig::resolve(settings)?,
safety: SafetyConfig::resolve()?,
wasm: WasmConfig::resolve()?,
secrets: SecretsConfig::resolve().await?,
secrets: SecretsConfig::resolve(bootstrap).await?,
builder: BuilderModeConfig::resolve()?,
heartbeat: HeartbeatConfig::resolve(settings)?,
routines: RoutineConfig::resolve()?,
@@ -107,13 +107,13 @@ impl TunnelConfig {
let public_url = optional_env("TUNNEL_URL")?
.or_else(|| settings.tunnel.public_url.clone().filter(|s| !s.is_empty()));
if let Some(ref url) = public_url
&& !url.starts_with("https://")
{
return Err(ConfigError::InvalidValue {
key: "TUNNEL_URL".to_string(),
message: "must start with https:// (webhooks require HTTPS)".to_string(),
});
if let Some(ref url) = public_url {
if !url.starts_with("https://") {
return Err(ConfigError::InvalidValue {
key: "TUNNEL_URL".to_string(),
message: "must start with https:// (webhooks require HTTPS)".to_string(),
});
}
}
Ok(Self { public_url })
@@ -179,7 +179,7 @@ pub struct DatabaseConfig {
}
impl DatabaseConfig {
fn resolve() -> Result<Self, ConfigError> {
fn resolve(bootstrap: &crate::bootstrap::BootstrapConfig) -> Result<Self, ConfigError> {
let backend: DatabaseBackend = if let Some(b) = optional_env("DATABASE_BACKEND")? {
b.parse().map_err(|e| ConfigError::InvalidValue {
key: "DATABASE_BACKEND".to_string(),
@@ -191,8 +191,8 @@ impl DatabaseConfig {
// PostgreSQL URL is required only when using the postgres backend.
// For libsql backend, default to an empty placeholder.
// DATABASE_URL is loaded from ~/.ironclaw/.env via dotenvy early in startup.
let url = optional_env("DATABASE_URL")?
.or_else(|| bootstrap.database_url.clone())
.or_else(|| {
if backend == DatabaseBackend::LibSql {
Some("unused://libsql".to_string())
@@ -205,7 +205,15 @@ impl DatabaseConfig {
hint: "Run 'ironclaw onboard' or set DATABASE_URL environment variable".to_string(),
})?;
let pool_size = parse_optional_env("DATABASE_POOL_SIZE", 10)?;
let pool_size = optional_env("DATABASE_POOL_SIZE")?
.map(|s| s.parse())
.transpose()
.map_err(|e| ConfigError::InvalidValue {
key: "DATABASE_POOL_SIZE".to_string(),
message: format!("must be a positive integer: {e}"),
})?
.or(bootstrap.database_pool_size)
.unwrap_or(10);
let libsql_path = optional_env("LIBSQL_PATH")?.map(PathBuf::from).or_else(|| {
if backend == DatabaseBackend::LibSql {
@@ -389,15 +397,6 @@ pub struct NearAiConfig {
pub api_mode: NearAiApiMode,
/// API key for cloud-api (required for chat_completions mode)
pub api_key: Option<SecretString>,
/// Optional fallback model for failover (default: None).
/// When set, a secondary provider is created with this model and wrapped
/// in a `FailoverProvider` so transient errors on the primary model
/// automatically fall through to the fallback.
pub fallback_model: Option<String>,
/// Maximum number of retries for transient errors (default: 3).
/// With the default of 3, the provider makes up to 4 total attempts
/// (1 initial + 3 retries) before giving up.
pub max_retries: u32,
}
impl LlmConfig {
@@ -442,8 +441,6 @@ impl LlmConfig {
.unwrap_or_else(default_session_path),
api_mode,
api_key: nearai_api_key,
fallback_model: optional_env("NEARAI_FALLBACK_MODEL")?,
max_retries: parse_optional_env("NEARAI_MAX_RETRIES", 3)?,
};
// Resolve provider-specific configs based on backend
@@ -856,35 +853,51 @@ impl std::fmt::Debug for SecretsConfig {
}
impl SecretsConfig {
/// Auto-detect secrets master key from env var, then OS keychain.
///
/// Sequential probe: SECRETS_MASTER_KEY env var first, then OS keychain.
/// No saved "source" needed; just try each source in order.
async fn resolve() -> Result<Self, ConfigError> {
async fn resolve(bootstrap: &crate::bootstrap::BootstrapConfig) -> Result<Self, ConfigError> {
use crate::settings::KeySource;
let (master_key, source) = if let Some(env_key) = optional_env("SECRETS_MASTER_KEY")? {
(Some(SecretString::from(env_key)), KeySource::Env)
} else {
// Probe the OS keychain; if a key is stored, use it
match crate::secrets::keychain::get_master_key().await {
Ok(key_bytes) => {
let key_hex: String = key_bytes.iter().map(|b| format!("{:02x}", b)).collect();
(Some(SecretString::from(key_hex)), KeySource::Keychain)
match bootstrap.secrets_master_key_source {
KeySource::Keychain => {
// Try to load from OS keychain (async on Linux)
match crate::secrets::keychain::get_master_key().await {
Ok(key_bytes) => {
let key_hex: String =
key_bytes.iter().map(|b| format!("{:02x}", b)).collect();
(Some(SecretString::from(key_hex)), KeySource::Keychain)
}
Err(_) => {
// Keychain configured but key not found
// This might happen if keychain was cleared
tracing::warn!(
"Secrets configured for keychain but key not found. \
Run 'ironclaw onboard' to reconfigure."
);
(None, KeySource::None)
}
}
}
Err(_) => (None, KeySource::None),
KeySource::Env => {
tracing::warn!(
"Secrets configured for env var but SECRETS_MASTER_KEY not set."
);
(None, KeySource::None)
}
KeySource::None => (None, KeySource::None),
}
};
let enabled = master_key.is_some();
if let Some(ref key) = master_key
&& key.expose_secret().len() < 32
{
return Err(ConfigError::InvalidValue {
key: "SECRETS_MASTER_KEY".to_string(),
message: "must be at least 32 bytes for AES-256-GCM".to_string(),
});
if let Some(ref key) = master_key {
if key.expose_secret().len() < 32 {
return Err(ConfigError::InvalidValue {
key: "SECRETS_MASTER_KEY".to_string(),
message: "must be at least 32 bytes for AES-256-GCM".to_string(),
});
}
}
Ok(Self {
+8 -8
View File
@@ -772,10 +772,10 @@ impl Database for LibSqlBackend {
.await
.map_err(|e| DatabaseError::Query(e.to_string()))?
{
if let Ok(id_str) = row.get::<String>(0)
&& let Ok(id) = id_str.parse()
{
ids.push(id);
if let Ok(id_str) = row.get::<String>(0) {
if let Ok(id) = id_str.parse() {
ids.push(id);
}
}
}
Ok(ids)
@@ -2199,10 +2199,10 @@ impl Database for LibSqlBackend {
e.content_preview = None;
}
// Update to latest timestamp
if let (Some(existing), Some(new)) = (&e.updated_at, &updated_at)
&& new > existing
{
e.updated_at = Some(*new);
if let (Some(existing), Some(new)) = (&e.updated_at, &updated_at) {
if new > existing {
e.updated_at = Some(*new);
}
}
})
.or_insert(WorkspaceEntry {
+5 -1
View File
@@ -20,6 +20,10 @@ impl CostEstimator {
// Default tool costs (in USD or equivalent)
tool_costs.insert("http".to_string(), dec!(0.0001)); // API call
tool_costs.insert("marketplace".to_string(), dec!(0.01)); // Gas costs
tool_costs.insert("ecommerce".to_string(), dec!(0.001)); // API call
tool_costs.insert("taskrabbit".to_string(), dec!(0.0)); // Cost comes from task itself
tool_costs.insert("restaurant".to_string(), dec!(0.001)); // API call
tool_costs.insert("echo".to_string(), dec!(0.0)); // Free
tool_costs.insert("time".to_string(), dec!(0.0)); // Free
tool_costs.insert("json".to_string(), dec!(0.0)); // Free
@@ -70,7 +74,7 @@ mod tests {
let estimator = CostEstimator::new();
assert_eq!(estimator.estimate_tool("echo"), dec!(0.0));
assert_eq!(estimator.estimate_tool("http"), dec!(0.0001));
assert_eq!(estimator.estimate_tool("marketplace"), dec!(0.01));
assert!(estimator.estimate_tool("unknown") > dec!(0.0));
}
+4
View File
@@ -16,6 +16,10 @@ impl TimeEstimator {
// Default tool durations
tool_durations.insert("http".to_string(), Duration::from_secs(5));
tool_durations.insert("marketplace".to_string(), Duration::from_secs(10));
tool_durations.insert("ecommerce".to_string(), Duration::from_secs(8));
tool_durations.insert("taskrabbit".to_string(), Duration::from_secs(30)); // Just API, not task itself
tool_durations.insert("restaurant".to_string(), Duration::from_secs(5));
tool_durations.insert("echo".to_string(), Duration::from_millis(10));
tool_durations.insert("time".to_string(), Duration::from_millis(1));
tool_durations.insert("json".to_string(), Duration::from_millis(5));
+6 -5
View File
@@ -144,11 +144,12 @@ impl SuccessEvaluator for RuleBasedEvaluator {
// Check for critical errors
for action in actions.iter().filter(|a| !a.success) {
if let Some(ref error) = action.error
&& (error.to_lowercase().contains("critical")
|| error.to_lowercase().contains("fatal"))
{
issues.push(format!("Critical error in {}: {}", action.tool_name, error));
if let Some(ref error) = action.error {
if error.to_lowercase().contains("critical")
|| error.to_lowercase().contains("fatal")
{
issues.push(format!("Critical error in {}: {}", action.tool_name, error));
}
}
}
+27 -27
View File
@@ -492,13 +492,13 @@ impl ExtensionManager {
}
// Check Content-Length header before downloading the full body
if let Some(len) = response.content_length()
&& len as usize > MAX_WASM_SIZE
{
return Err(ExtensionError::InstallFailed(format!(
"WASM binary too large ({} bytes, max {} bytes)",
len, MAX_WASM_SIZE
)));
if let Some(len) = response.content_length() {
if len as usize > MAX_WASM_SIZE {
return Err(ExtensionError::InstallFailed(format!(
"WASM binary too large ({} bytes, max {} bytes)",
len, MAX_WASM_SIZE
)));
}
}
let bytes = response
@@ -768,27 +768,27 @@ impl ExtensionManager {
};
// Check env var first
if let Some(ref env_var) = auth.env_var
&& let Ok(value) = std::env::var(env_var)
{
// Store the env var value as a secret
let params =
CreateSecretParams::new(&auth.secret_name, &value).with_provider(name.to_string());
self.secrets
.create(&self.user_id, params)
.await
.map_err(|e| ExtensionError::AuthFailed(e.to_string()))?;
if let Some(ref env_var) = auth.env_var {
if let Ok(value) = std::env::var(env_var) {
// Store the env var value as a secret
let params = CreateSecretParams::new(&auth.secret_name, &value)
.with_provider(name.to_string());
self.secrets
.create(&self.user_id, params)
.await
.map_err(|e| ExtensionError::AuthFailed(e.to_string()))?;
return Ok(AuthResult {
name: name.to_string(),
kind: ExtensionKind::WasmTool,
auth_url: None,
callback_type: None,
instructions: None,
setup_url: None,
awaiting_token: false,
status: "authenticated".to_string(),
});
return Ok(AuthResult {
name: name.to_string(),
kind: ExtensionKind::WasmTool,
auth_url: None,
callback_type: None,
instructions: None,
setup_url: None,
awaiting_token: false,
status: "authenticated".to_string(),
});
}
}
// Check if already authenticated
-5
View File
@@ -36,11 +36,6 @@ pub struct Store {
#[cfg(feature = "postgres")]
impl Store {
/// Wrap an existing pool (useful when the caller already has a connection).
pub fn from_pool(pool: Pool) -> Self {
Self { pool }
}
/// Create a new store and connect to the database.
pub async fn new(config: &DatabaseConfig) -> Result<Self, DatabaseError> {
let mut cfg = Config::new();
-1
View File
@@ -59,7 +59,6 @@ pub mod secrets;
pub mod settings;
pub mod setup;
pub mod tools;
pub mod tracing_fmt;
pub mod util;
pub mod worker;
pub mod workspace;
-483
View File
@@ -1,483 +0,0 @@
//! Multi-provider LLM failover.
//!
//! Wraps multiple LlmProvider instances and tries each in sequence
//! until one succeeds. Transparent to callers --- same LlmProvider trait.
use std::future::Future;
use std::sync::Arc;
use std::sync::atomic::{AtomicUsize, Ordering};
use async_trait::async_trait;
use rust_decimal::Decimal;
use crate::error::LlmError;
use crate::llm::provider::{
CompletionRequest, CompletionResponse, LlmProvider, ToolCompletionRequest,
ToolCompletionResponse,
};
/// Returns `true` if the error is transient and the request should be retried
/// on the next provider in the failover chain.
///
/// Retryable: `RequestFailed`, `RateLimited`, `InvalidResponse`,
/// `SessionRenewalFailed`, `ModelNotAvailable`, `Http`, `Io`.
///
/// `ModelNotAvailable` is retryable because the next provider in the chain may
/// offer a different model, so it's worth trying.
///
/// Non-retryable errors (`AuthFailed`, `SessionExpired`, `ContextLengthExceeded`)
/// propagate immediately because a different provider won't fix them.
fn is_retryable(err: &LlmError) -> bool {
matches!(
err,
LlmError::RequestFailed { .. }
| LlmError::RateLimited { .. }
| LlmError::InvalidResponse { .. }
| LlmError::SessionRenewalFailed { .. }
// ModelNotAvailable is retryable: the next provider may offer a different model.
| LlmError::ModelNotAvailable { .. }
| LlmError::Http(_)
| LlmError::Io(_)
)
}
/// 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.
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,
}
impl FailoverProvider {
/// Create a new failover provider.
///
/// Returns an error if `providers` is empty.
pub fn new(providers: Vec<Arc<dyn LlmProvider>>) -> Result<Self, LlmError> {
if providers.is_empty() {
return Err(LlmError::RequestFailed {
provider: "failover".to_string(),
reason: "FailoverProvider requires at least one provider".to_string(),
});
}
Ok(Self {
providers,
last_used: AtomicUsize::new(0),
})
}
/// Try each provider in sequence until one succeeds or all fail.
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 mut last_error: Option<LlmError> = None;
for (i, provider) in self.providers.iter().enumerate() {
let result = call(Arc::clone(provider)).await;
match result {
Ok(response) => {
self.last_used.store(i, Ordering::Relaxed);
return Ok(response);
}
Err(err) => {
if !is_retryable(&err) {
return Err(err);
}
if i + 1 < self.providers.len() {
tracing::warn!(
provider = %provider.model_name(),
error = %err,
next_provider = %self.providers[i + 1].model_name(),
"Provider failed with retryable error, trying next provider"
);
}
last_error = Some(err);
}
}
}
// SAFETY: providers is non-empty (checked in `new`), so at least one
// iteration ran and `last_error` is `Some`.
Err(last_error.expect("providers list is non-empty"))
}
}
#[async_trait]
impl LlmProvider for FailoverProvider {
fn model_name(&self) -> &str {
self.providers[self.last_used.load(Ordering::Relaxed)].model_name()
}
fn cost_per_token(&self) -> (Decimal, Decimal) {
self.providers[self.last_used.load(Ordering::Relaxed)].cost_per_token()
}
async fn complete(&self, request: CompletionRequest) -> Result<CompletionResponse, LlmError> {
self.try_providers(|provider| {
let req = request.clone();
async move { provider.complete(req).await }
})
.await
}
async fn complete_with_tools(
&self,
request: ToolCompletionRequest,
) -> Result<ToolCompletionResponse, LlmError> {
self.try_providers(|provider| {
let req = request.clone();
async move { provider.complete_with_tools(req).await }
})
.await
}
async fn list_models(&self) -> Result<Vec<String>, LlmError> {
let mut all_models = Vec::new();
for provider in &self.providers {
match provider.list_models().await {
Ok(models) => all_models.extend(models),
Err(err) => {
tracing::warn!(
provider = %provider.model_name(),
error = %err,
"Failed to list models from provider, skipping"
);
}
}
}
all_models.sort();
all_models.dedup();
Ok(all_models)
}
}
#[cfg(test)]
mod tests {
use super::*;
use std::sync::Mutex;
use std::time::Duration;
use crate::llm::provider::{CompletionResponse, FinishReason, ToolCompletionResponse};
/// A mock LLM provider that returns a predetermined result.
struct MockProvider {
name: String,
input_cost: Decimal,
output_cost: Decimal,
complete_result: Mutex<Option<Result<CompletionResponse, LlmError>>>,
tool_complete_result: Mutex<Option<Result<ToolCompletionResponse, LlmError>>>,
}
impl MockProvider {
fn succeeding(name: &str, content: &str) -> Self {
Self {
name: name.to_string(),
input_cost: Decimal::ZERO,
output_cost: Decimal::ZERO,
complete_result: Mutex::new(Some(Ok(CompletionResponse {
content: content.to_string(),
input_tokens: 10,
output_tokens: 5,
finish_reason: FinishReason::Stop,
response_id: None,
}))),
tool_complete_result: Mutex::new(Some(Ok(ToolCompletionResponse {
content: Some(content.to_string()),
tool_calls: vec![],
input_tokens: 10,
output_tokens: 5,
finish_reason: FinishReason::Stop,
response_id: None,
}))),
}
}
fn succeeding_with_cost(
name: &str,
content: &str,
input_cost: Decimal,
output_cost: Decimal,
) -> Self {
Self {
input_cost,
output_cost,
..Self::succeeding(name, content)
}
}
fn failing_retryable(name: &str) -> Self {
Self {
name: name.to_string(),
input_cost: Decimal::ZERO,
output_cost: Decimal::ZERO,
complete_result: Mutex::new(Some(Err(LlmError::RequestFailed {
provider: name.to_string(),
reason: "server error".to_string(),
}))),
tool_complete_result: Mutex::new(Some(Err(LlmError::RequestFailed {
provider: name.to_string(),
reason: "server error".to_string(),
}))),
}
}
fn failing_non_retryable(name: &str) -> Self {
Self {
name: name.to_string(),
input_cost: Decimal::ZERO,
output_cost: Decimal::ZERO,
complete_result: Mutex::new(Some(Err(LlmError::AuthFailed {
provider: name.to_string(),
}))),
tool_complete_result: Mutex::new(Some(Err(LlmError::AuthFailed {
provider: name.to_string(),
}))),
}
}
fn failing_rate_limited(name: &str) -> Self {
Self {
name: name.to_string(),
input_cost: Decimal::ZERO,
output_cost: Decimal::ZERO,
complete_result: Mutex::new(Some(Err(LlmError::RateLimited {
provider: name.to_string(),
retry_after: Some(Duration::from_secs(30)),
}))),
tool_complete_result: Mutex::new(Some(Err(LlmError::RateLimited {
provider: name.to_string(),
retry_after: Some(Duration::from_secs(30)),
}))),
}
}
}
#[async_trait]
impl LlmProvider for MockProvider {
fn model_name(&self) -> &str {
&self.name
}
fn cost_per_token(&self) -> (Decimal, Decimal) {
(self.input_cost, self.output_cost)
}
async fn complete(
&self,
_request: CompletionRequest,
) -> Result<CompletionResponse, LlmError> {
self.complete_result
.lock()
.unwrap()
.take()
.expect("MockProvider::complete called more than once")
}
async fn complete_with_tools(
&self,
_request: ToolCompletionRequest,
) -> Result<ToolCompletionResponse, LlmError> {
self.tool_complete_result
.lock()
.unwrap()
.take()
.expect("MockProvider::complete_with_tools called more than once")
}
async fn list_models(&self) -> Result<Vec<String>, LlmError> {
Ok(vec![self.name.clone()])
}
}
fn make_request() -> CompletionRequest {
CompletionRequest::new(vec![crate::llm::ChatMessage::user("hello")])
}
fn make_tool_request() -> ToolCompletionRequest {
ToolCompletionRequest::new(vec![crate::llm::ChatMessage::user("hello")], vec![])
}
// Test 1: Primary succeeds, no failover occurs.
#[tokio::test]
async fn primary_succeeds_no_failover() {
let primary = Arc::new(MockProvider::succeeding("primary", "primary response"));
let fallback = Arc::new(MockProvider::succeeding("fallback", "fallback response"));
let failover = FailoverProvider::new(vec![primary, fallback]).unwrap();
let response = failover.complete(make_request()).await.unwrap();
assert_eq!(response.content, "primary response");
}
// Test 2: Primary fails with retryable error, fallback succeeds.
#[tokio::test]
async fn primary_fails_retryable_fallback_succeeds() {
let primary = Arc::new(MockProvider::failing_retryable("primary"));
let fallback = Arc::new(MockProvider::succeeding("fallback", "fallback response"));
let failover = FailoverProvider::new(vec![primary, fallback]).unwrap();
let response = failover.complete(make_request()).await.unwrap();
assert_eq!(response.content, "fallback response");
}
// Test 3: All providers fail, returns last error.
#[tokio::test]
async fn all_providers_fail_returns_last_error() {
let primary = Arc::new(MockProvider::failing_retryable("primary"));
let fallback = Arc::new(MockProvider::failing_retryable("fallback"));
let failover = FailoverProvider::new(vec![primary, fallback]).unwrap();
let err = failover.complete(make_request()).await.unwrap_err();
match err {
LlmError::RequestFailed { provider, .. } => {
assert_eq!(provider, "fallback");
}
other => panic!("expected RequestFailed, got: {other:?}"),
}
}
// Test 4: Non-retryable error fails immediately, no failover.
#[tokio::test]
async fn non_retryable_error_fails_immediately() {
let primary = Arc::new(MockProvider::failing_non_retryable("primary"));
let fallback = Arc::new(MockProvider::succeeding("fallback", "fallback response"));
let failover = FailoverProvider::new(vec![primary, fallback]).unwrap();
let err = failover.complete(make_request()).await.unwrap_err();
match err {
LlmError::AuthFailed { provider } => {
assert_eq!(provider, "primary");
}
other => panic!("expected AuthFailed, got: {other:?}"),
}
}
// Test 5: Three providers, first two fail (retryable), third succeeds.
#[tokio::test]
async fn three_providers_first_two_fail_third_succeeds() {
let p1 = Arc::new(MockProvider::failing_retryable("provider-1"));
let p2 = Arc::new(MockProvider::failing_rate_limited("provider-2"));
let p3 = Arc::new(MockProvider::succeeding("provider-3", "third time lucky"));
let failover = FailoverProvider::new(vec![p1, p2, p3]).unwrap();
let response = failover.complete(make_request()).await.unwrap();
assert_eq!(response.content, "third time lucky");
}
// Test: complete_with_tools follows same failover logic.
#[tokio::test]
async fn complete_with_tools_failover() {
let primary = Arc::new(MockProvider::failing_retryable("primary"));
let fallback = Arc::new(MockProvider::succeeding("fallback", "tools fallback"));
let failover = FailoverProvider::new(vec![primary, fallback]).unwrap();
let response = failover
.complete_with_tools(make_tool_request())
.await
.unwrap();
assert_eq!(response.content.as_deref(), Some("tools fallback"));
}
// Test: model_name and cost_per_token reflect the last-used provider.
#[tokio::test]
async fn model_name_and_cost_track_last_used_provider() {
let fallback_cost = Decimal::new(15, 6); // 0.000015
let primary = Arc::new(MockProvider::failing_retryable("primary-model"));
let fallback = Arc::new(MockProvider::succeeding_with_cost(
"fallback-model",
"ok",
fallback_cost,
fallback_cost,
));
let failover = FailoverProvider::new(vec![primary, fallback]).unwrap();
// Before any call, defaults to primary (index 0).
assert_eq!(failover.model_name(), "primary-model");
assert_eq!(failover.cost_per_token(), (Decimal::ZERO, Decimal::ZERO));
// After failover, should reflect the fallback provider.
let _ = failover.complete(make_request()).await.unwrap();
assert_eq!(failover.model_name(), "fallback-model");
assert_eq!(failover.cost_per_token(), (fallback_cost, fallback_cost));
}
// Test: list_models aggregates from all providers.
#[tokio::test]
async fn list_models_aggregates_all() {
let p1 = Arc::new(MockProvider::succeeding("model-a", "ok"));
let p2 = Arc::new(MockProvider::succeeding("model-b", "ok"));
let failover = FailoverProvider::new(vec![p1, p2]).unwrap();
let models = failover.list_models().await.unwrap();
assert!(models.contains(&"model-a".to_string()));
assert!(models.contains(&"model-b".to_string()));
}
// Test: is_retryable correctly classifies errors.
#[test]
fn retryable_classification() {
// Retryable
assert!(is_retryable(&LlmError::RequestFailed {
provider: "p".into(),
reason: "err".into(),
}));
assert!(is_retryable(&LlmError::RateLimited {
provider: "p".into(),
retry_after: None,
}));
assert!(is_retryable(&LlmError::InvalidResponse {
provider: "p".into(),
reason: "bad json".into(),
}));
assert!(is_retryable(&LlmError::SessionRenewalFailed {
provider: "p".into(),
reason: "timeout".into(),
}));
assert!(is_retryable(&LlmError::Io(std::io::Error::new(
std::io::ErrorKind::ConnectionReset,
"reset"
))));
assert!(is_retryable(&LlmError::ModelNotAvailable {
provider: "p".into(),
model: "m".into(),
}));
// Non-retryable
assert!(!is_retryable(&LlmError::AuthFailed {
provider: "p".into(),
}));
assert!(!is_retryable(&LlmError::SessionExpired {
provider: "p".into(),
}));
assert!(!is_retryable(&LlmError::ContextLengthExceeded {
used: 100_000,
limit: 50_000,
}));
}
// Test: empty providers list returns error (not panic).
#[test]
fn empty_providers_returns_error() {
let result = FailoverProvider::new(vec![]);
assert!(result.is_err());
}
}
+12 -22
View File
@@ -8,16 +8,13 @@
//! - **OpenAI-compatible**: Any endpoint that speaks the OpenAI API
mod costs;
pub mod failover;
mod nearai;
mod nearai_chat;
mod provider;
mod reasoning;
mod retry;
mod rig_adapter;
pub mod session;
pub use failover::FailoverProvider;
pub use nearai::{ModelInfo, NearAiProvider};
pub use nearai_chat::NearAiChatProvider;
pub use provider::{
@@ -36,7 +33,7 @@ use std::sync::Arc;
use rig::client::CompletionClient;
use secrecy::ExposeSecret;
use crate::config::{LlmBackend, LlmConfig, NearAiApiMode, NearAiConfig};
use crate::config::{LlmBackend, LlmConfig, NearAiApiMode};
use crate::error::LlmError;
/// Create an LLM provider based on configuration.
@@ -49,7 +46,7 @@ pub fn create_llm_provider(
session: Arc<SessionManager>,
) -> Result<Arc<dyn LlmProvider>, LlmError> {
match config.backend {
LlmBackend::NearAi => create_llm_provider_with_config(&config.nearai, session),
LlmBackend::NearAi => create_nearai_provider(config, session),
LlmBackend::OpenAi => create_openai_provider(config),
LlmBackend::Anthropic => create_anthropic_provider(config),
LlmBackend::Ollama => create_ollama_provider(config),
@@ -57,28 +54,21 @@ pub fn create_llm_provider(
}
}
/// Create an LLM provider from a `NearAiConfig` directly.
///
/// This is useful when constructing additional providers for failover,
/// where only the model name differs from the primary config.
pub fn create_llm_provider_with_config(
config: &NearAiConfig,
fn create_nearai_provider(
config: &LlmConfig,
session: Arc<SessionManager>,
) -> Result<Arc<dyn LlmProvider>, LlmError> {
match config.api_mode {
match config.nearai.api_mode {
NearAiApiMode::Responses => {
tracing::info!(
model = %config.model,
"Using Responses API (chat-api) with session auth"
);
Ok(Arc::new(NearAiProvider::new(config.clone(), session)))
tracing::info!("Using NEAR AI Responses API (chat-api) with session auth");
Ok(Arc::new(NearAiProvider::new(
config.nearai.clone(),
session,
)))
}
NearAiApiMode::ChatCompletions => {
tracing::info!(
model = %config.model,
"Using Chat Completions API (cloud-api) with API key auth"
);
Ok(Arc::new(NearAiChatProvider::new(config.clone())?))
tracing::info!("Using NEAR AI Chat Completions API (cloud-api) with API key auth");
Ok(Arc::new(NearAiChatProvider::new(config.nearai.clone())?))
}
}
}
+96 -138
View File
@@ -19,7 +19,6 @@ use crate::llm::provider::{
ChatMessage, CompletionRequest, CompletionResponse, FinishReason, LlmProvider, Role, ToolCall,
ToolCompletionRequest, ToolCompletionResponse,
};
use crate::llm::retry::{is_retryable_status, retry_backoff_delay};
use crate::llm::session::SessionManager;
/// Information about an available model from NEAR AI API.
@@ -210,20 +209,20 @@ impl NearAiProvider {
data: Option<Vec<ModelEntry>>,
}
if let Ok(resp) = serde_json::from_str::<ModelsResponse>(&response_text)
&& let Some(entries) = resp.models.or(resp.data)
{
let models: Vec<ModelInfo> = entries
.into_iter()
.filter_map(|e| {
e.get_name().map(|name| ModelInfo {
name,
provider: None,
if let Ok(resp) = serde_json::from_str::<ModelsResponse>(&response_text) {
if let Some(entries) = resp.models.or(resp.data) {
let models: Vec<ModelInfo> = entries
.into_iter()
.filter_map(|e| {
e.get_name().map(|name| ModelInfo {
name,
provider: None,
})
})
})
.collect();
if !models.is_empty() {
return Ok(models);
.collect();
if !models.is_empty() {
return Ok(models);
}
}
}
@@ -271,139 +270,88 @@ impl NearAiProvider {
}
}
/// Inner request implementation with retry logic for transient errors.
///
/// Retries on HTTP 429, 500, 502, 503, 504 with exponential backoff.
/// Does not retry on client errors (400, 401, 403, 404) or parse errors.
/// Inner request implementation without retry logic.
async fn send_request_inner<T: Serialize + std::fmt::Debug, R: for<'de> Deserialize<'de>>(
&self,
path: &str,
body: &T,
) -> Result<R, LlmError> {
let url = self.api_url(path);
let max_retries = self.config.max_retries;
let token = self.session.get_token().await?;
for attempt in 0..=max_retries {
let token = self.session.get_token().await?;
tracing::debug!("Sending request to NEAR AI: {}", url);
tracing::debug!("Request body: {:?}", body);
tracing::debug!(
"Sending request to NEAR AI: {} (attempt {})",
url,
attempt + 1
);
tracing::debug!("Request body: {:?}", body);
let response = self
.client
.post(&url)
.header("Authorization", format!("Bearer {}", token.expose_secret()))
.header("Content-Type", "application/json")
.json(body)
.send()
.await
.map_err(|e| {
tracing::error!("NEAR AI request failed: {}", e);
e
})?;
let response = self
.client
.post(&url)
.header("Authorization", format!("Bearer {}", token.expose_secret()))
.header("Content-Type", "application/json")
.json(body)
.send()
.await;
let status = response.status();
let response_text = response.text().await.unwrap_or_default();
let response = match response {
Ok(r) => r,
Err(e) => {
tracing::error!("NEAR AI request failed: {}", e);
// Network errors (timeout, connection refused) are transient
if attempt < max_retries {
let delay = retry_backoff_delay(attempt);
tracing::warn!(
"NEAR AI request error (attempt {}/{}), retrying in {:?}: {}",
attempt + 1,
max_retries + 1,
delay,
e,
);
tokio::time::sleep(delay).await;
continue;
}
return Err(e.into());
}
};
tracing::debug!("NEAR AI response status: {}", status);
tracing::debug!("NEAR AI response body: {}", response_text);
let status = response.status();
let response_text = response.text().await.unwrap_or_default();
if !status.is_success() {
// Check for session expiration (401 with specific message patterns)
if status.as_u16() == 401 {
let is_session_expired = response_text.to_lowercase().contains("session")
&& (response_text.to_lowercase().contains("expired")
|| response_text.to_lowercase().contains("invalid"));
tracing::debug!("NEAR AI response status: {}", status);
tracing::debug!("NEAR AI response body: {}", response_text);
if !status.is_success() {
let status_code = status.as_u16();
// Check for session expiration (401 with specific message patterns)
if status_code == 401 {
let lower = response_text.to_lowercase();
let is_session_expired = lower.contains("session")
&& (lower.contains("expired") || lower.contains("invalid"));
if is_session_expired {
return Err(LlmError::SessionExpired {
provider: "nearai".to_string(),
});
}
// Generic 401 -- not retryable
return Err(LlmError::AuthFailed {
if is_session_expired {
return Err(LlmError::SessionExpired {
provider: "nearai".to_string(),
});
}
// Check if this is a transient error worth retrying
if is_retryable_status(status_code) && attempt < max_retries {
let delay = retry_backoff_delay(attempt);
tracing::warn!(
"NEAR AI returned HTTP {} (attempt {}/{}), retrying in {:?}",
status_code,
attempt + 1,
max_retries + 1,
delay,
);
tokio::time::sleep(delay).await;
continue;
}
// Non-retryable error or exhausted retries
if let Ok(error) = serde_json::from_str::<NearAiErrorResponse>(&response_text) {
if status_code == 429 {
return Err(LlmError::RateLimited {
provider: "nearai".to_string(),
retry_after: None,
});
}
return Err(LlmError::RequestFailed {
provider: "nearai".to_string(),
reason: error.error,
});
}
return Err(LlmError::RequestFailed {
// Generic 401 without session expiration indication
return Err(LlmError::AuthFailed {
provider: "nearai".to_string(),
reason: format!("HTTP {}: {}", status, response_text),
});
}
// Success -- parse the response
return match serde_json::from_str::<R>(&response_text) {
Ok(parsed) => Ok(parsed),
Err(e) => {
tracing::debug!("Response is not expected JSON format: {}", e);
tracing::debug!("Will try alternative parsing in caller");
Err(LlmError::InvalidResponse {
// Try to parse as JSON error
if let Ok(error) = serde_json::from_str::<NearAiErrorResponse>(&response_text) {
if status.as_u16() == 429 {
return Err(LlmError::RateLimited {
provider: "nearai".to_string(),
reason: format!("Parse error: {}. Raw: {}", e, response_text),
})
retry_after: None,
});
}
};
return Err(LlmError::RequestFailed {
provider: "nearai".to_string(),
reason: error.error,
});
}
return Err(LlmError::RequestFailed {
provider: "nearai".to_string(),
reason: format!("HTTP {}: {}", status, response_text),
});
}
// This is unreachable because the loop always returns, but the compiler
// cannot prove that. Return a generic error as a safety net.
Err(LlmError::RequestFailed {
provider: "nearai".to_string(),
reason: "retry loop exited unexpectedly".to_string(),
})
// Try to parse as our expected type
match serde_json::from_str::<R>(&response_text) {
Ok(parsed) => Ok(parsed),
Err(e) => {
tracing::debug!("Response is not expected JSON format: {}", e);
tracing::debug!("Will try alternative parsing in caller");
Err(LlmError::InvalidResponse {
provider: "nearai".to_string(),
reason: format!("Parse error: {}. Raw: {}", e, response_text),
})
}
}
}
}
@@ -508,7 +456,7 @@ impl LlmProvider for NearAiProvider {
Err(e) => return Err(e),
};
tracing::debug!("NEAR AI response: output_items={}", response.output.len());
tracing::debug!("NEAR AI response: {:?}", response);
// Extract text from response output
// Try multiple formats since API response shape may vary
@@ -516,6 +464,11 @@ impl LlmProvider for NearAiProvider {
.output
.iter()
.filter_map(|item| {
tracing::debug!(
"Processing output item: type={}, text={:?}",
item.item_type,
item.text
);
if item.item_type == "message" {
// First check for direct text field on item
if let Some(ref text) = item.text {
@@ -526,6 +479,11 @@ impl LlmProvider for NearAiProvider {
contents
.iter()
.filter_map(|c| {
tracing::debug!(
"Content item: type={}, text={:?}",
c.content_type,
c.text
);
// Accept various content types that might contain text
match c.content_type.as_str() {
"output_text" | "text" => c.text.clone(),
@@ -736,21 +694,21 @@ impl LlmProvider for NearAiProvider {
}
}
}
} else if item.item_type == "function_call"
&& let (Some(name), Some(call_id)) = (&item.name, &item.call_id)
{
// Parse arguments JSON string into Value
let arguments = item
.arguments
.as_ref()
.and_then(|s| serde_json::from_str(s).ok())
.unwrap_or(serde_json::Value::Object(Default::default()));
} else if item.item_type == "function_call" {
if let (Some(name), Some(call_id)) = (&item.name, &item.call_id) {
// Parse arguments JSON string into Value
let arguments = item
.arguments
.as_ref()
.and_then(|s| serde_json::from_str(s).ok())
.unwrap_or(serde_json::Value::Object(Default::default()));
tool_calls.push(ToolCall {
id: call_id.clone(),
name: name.clone(),
arguments,
});
tool_calls.push(ToolCall {
id: call_id.clone(),
name: name.clone(),
arguments,
});
}
}
}
+41 -95
View File
@@ -16,7 +16,6 @@ use crate::llm::provider::{
ChatMessage, CompletionRequest, CompletionResponse, FinishReason, LlmProvider, ModelMetadata,
Role, ToolCall, ToolCompletionRequest, ToolCompletionResponse,
};
use crate::llm::retry::{is_retryable_status, retry_backoff_delay};
/// NEAR AI Chat Completions API provider.
pub struct NearAiChatProvider {
@@ -63,116 +62,63 @@ impl NearAiChatProvider {
.unwrap_or_default()
}
/// Send a request to the chat completions API with retry on transient errors.
///
/// Retries on HTTP 429, 500, 502, 503, 504 with exponential backoff.
/// Does not retry on client errors (400, 401, 403, 404) or parse errors.
/// Send a request to the chat completions API.
async fn send_request<T: Serialize, R: for<'de> Deserialize<'de>>(
&self,
body: &T,
) -> Result<R, LlmError> {
let url = self.api_url("chat/completions");
let max_retries = self.config.max_retries;
for attempt in 0..=max_retries {
tracing::debug!(
"Sending request to NEAR AI Chat: {} (attempt {})",
url,
attempt + 1,
);
tracing::debug!("Sending request to NEAR AI Chat: {}", url);
if tracing::enabled!(tracing::Level::DEBUG)
&& let Ok(json) = serde_json::to_string(body)
{
tracing::debug!("NEAR AI Chat request body: {}", json);
}
// Log the request body for debugging tool call issues
if let Ok(json) = serde_json::to_string(body) {
tracing::debug!("NEAR AI Chat request body: {}", json);
}
let response = self
.client
.post(&url)
.header("Authorization", format!("Bearer {}", self.api_key()))
.header("Content-Type", "application/json")
.json(body)
.send()
.await;
let response = match response {
Ok(r) => r,
Err(e) => {
tracing::error!("NEAR AI Chat request failed: {}", e);
if attempt < max_retries {
let delay = retry_backoff_delay(attempt);
tracing::warn!(
"NEAR AI Chat request error (attempt {}/{}), retrying in {:?}: {}",
attempt + 1,
max_retries + 1,
delay,
e,
);
tokio::time::sleep(delay).await;
continue;
}
return Err(LlmError::RequestFailed {
provider: "nearai_chat".to_string(),
reason: e.to_string(),
});
let response = self
.client
.post(&url)
.header("Authorization", format!("Bearer {}", self.api_key()))
.header("Content-Type", "application/json")
.json(body)
.send()
.await
.map_err(|e| {
tracing::error!("NEAR AI Chat request failed: {}", e);
LlmError::RequestFailed {
provider: "nearai_chat".to_string(),
reason: e.to_string(),
}
};
let status = response.status();
let response_text = response.text().await.unwrap_or_default();
})?;
tracing::debug!("NEAR AI Chat response status: {}", status);
tracing::debug!("NEAR AI Chat response body: {}", response_text);
let status = response.status();
let response_text = response.text().await.unwrap_or_default();
if !status.is_success() {
let status_code = status.as_u16();
// Auth errors are not retryable
if status_code == 401 {
return Err(LlmError::AuthFailed {
provider: "nearai_chat".to_string(),
});
}
tracing::debug!("NEAR AI Chat response status: {}", status);
tracing::debug!("NEAR AI Chat response body: {}", response_text);
// Transient errors: retry with backoff
if is_retryable_status(status_code) && attempt < max_retries {
let delay = retry_backoff_delay(attempt);
tracing::warn!(
"NEAR AI Chat returned HTTP {} (attempt {}/{}), retrying in {:?}",
status_code,
attempt + 1,
max_retries + 1,
delay,
);
tokio::time::sleep(delay).await;
continue;
}
// Non-retryable or exhausted retries
if status_code == 429 {
return Err(LlmError::RateLimited {
provider: "nearai_chat".to_string(),
retry_after: None,
});
}
return Err(LlmError::RequestFailed {
if !status.is_success() {
if status.as_u16() == 401 {
return Err(LlmError::AuthFailed {
provider: "nearai_chat".to_string(),
reason: format!("HTTP {}: {}", status, response_text),
});
}
// Success — parse the response
return serde_json::from_str(&response_text).map_err(|e| LlmError::InvalidResponse {
if status.as_u16() == 429 {
return Err(LlmError::RateLimited {
provider: "nearai_chat".to_string(),
retry_after: None,
});
}
return Err(LlmError::RequestFailed {
provider: "nearai_chat".to_string(),
reason: format!("JSON parse error: {}. Raw: {}", e, response_text),
reason: format!("HTTP {}: {}", status, response_text),
});
}
// Safety net: unreachable because the loop always returns
Err(LlmError::RequestFailed {
serde_json::from_str(&response_text).map_err(|e| LlmError::InvalidResponse {
provider: "nearai_chat".to_string(),
reason: "retry loop exited unexpectedly".to_string(),
reason: format!("JSON parse error: {}. Raw: {}", e, response_text),
})
}
@@ -449,10 +395,10 @@ fn flatten_tool_messages(messages: Vec<ChatCompletionMessage>) -> Vec<ChatComple
if let (true, Some(calls)) = (msg.role == "assistant", &msg.tool_calls) {
// Convert assistant tool_calls into descriptive text
let mut parts: Vec<String> = Vec::new();
if let Some(ref text) = msg.content
&& !text.is_empty()
{
parts.push(text.clone());
if let Some(ref text) = msg.content {
if !text.is_empty() {
parts.push(text.clone());
}
}
for tc in calls {
parts.push(format!(
+15 -21
View File
@@ -113,12 +113,6 @@ pub struct ToolSelection {
pub reasoning: String,
/// Alternative tools considered.
pub alternatives: Vec<String>,
/// The tool call ID from the LLM response.
///
/// OpenAI-compatible providers assign each tool call a unique ID that must
/// be echoed back in the corresponding tool result message. Without this,
/// the provider cannot match results to their originating calls.
pub tool_call_id: String,
}
/// Token usage from a single LLM call.
@@ -250,7 +244,6 @@ impl Reasoning {
parameters: tool_call.arguments,
reasoning: reasoning.clone(),
alternatives: vec![],
tool_call_id: tool_call.id,
})
.collect();
@@ -588,20 +581,21 @@ fn recover_tool_calls_from_content(
}
// Try JSON first: {"name":"x","arguments":{}}
if let Ok(parsed) = serde_json::from_str::<serde_json::Value>(inner)
&& let Some(name) = parsed.get("name").and_then(|v| v.as_str())
&& tool_names.contains(name)
{
let arguments = parsed
.get("arguments")
.cloned()
.unwrap_or(serde_json::Value::Object(Default::default()));
calls.push(ToolCall {
id: format!("recovered_{}", calls.len()),
name: name.to_string(),
arguments,
});
continue;
if let Ok(parsed) = serde_json::from_str::<serde_json::Value>(inner) {
if let Some(name) = parsed.get("name").and_then(|v| v.as_str()) {
if tool_names.contains(name) {
let arguments = parsed
.get("arguments")
.cloned()
.unwrap_or(serde_json::Value::Object(Default::default()));
calls.push(ToolCall {
id: format!("recovered_{}", calls.len()),
name: name.to_string(),
arguments,
});
continue;
}
}
}
// Bare tool name (e.g. "<tool_call>tool_list</tool_call>")
-96
View File
@@ -1,96 +0,0 @@
//! Shared retry helpers for LLM providers.
//!
//! Provides exponential backoff with jitter and retryable status classification
//! used by both `NearAiProvider` and `NearAiChatProvider`.
use std::time::Duration;
use rand::Rng;
/// Returns `true` if the HTTP status code is transient and worth retrying.
pub(crate) fn is_retryable_status(status: u16) -> bool {
matches!(status, 429 | 500 | 502 | 503 | 504)
}
/// Calculate exponential backoff delay with random jitter.
///
/// Base delay is 1 second, doubled each attempt, with +/-25% jitter.
/// - attempt 0: ~1s (0.75s - 1.25s)
/// - attempt 1: ~2s (1.5s - 2.5s)
/// - attempt 2: ~4s (3.0s - 5.0s)
pub(crate) fn retry_backoff_delay(attempt: u32) -> Duration {
let base_ms: u64 = 1000u64.saturating_mul(2u64.saturating_pow(attempt));
let jitter_range = base_ms / 4; // 25%
let jitter = if jitter_range > 0 {
let offset = rand::thread_rng().gen_range(0..=jitter_range * 2);
offset as i64 - jitter_range as i64
} else {
0
};
let delay_ms = (base_ms as i64 + jitter).max(100) as u64;
Duration::from_millis(delay_ms)
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_is_retryable_status() {
// Transient errors should be retryable
assert!(is_retryable_status(429));
assert!(is_retryable_status(500));
assert!(is_retryable_status(502));
assert!(is_retryable_status(503));
assert!(is_retryable_status(504));
// Client errors should not be retryable
assert!(!is_retryable_status(400));
assert!(!is_retryable_status(401));
assert!(!is_retryable_status(403));
assert!(!is_retryable_status(404));
assert!(!is_retryable_status(422));
// Success codes should not be retryable
assert!(!is_retryable_status(200));
assert!(!is_retryable_status(201));
}
#[test]
fn test_retry_backoff_delay_exponential_growth() {
// Run multiple samples to verify the range, accounting for jitter
for _ in 0..20 {
let d0 = retry_backoff_delay(0);
let d1 = retry_backoff_delay(1);
let d2 = retry_backoff_delay(2);
// Attempt 0: base 1000ms, jitter +/-250ms -> [750, 1250]
assert!(d0.as_millis() >= 750, "attempt 0 too low: {:?}", d0);
assert!(d0.as_millis() <= 1250, "attempt 0 too high: {:?}", d0);
// Attempt 1: base 2000ms, jitter +/-500ms -> [1500, 2500]
assert!(d1.as_millis() >= 1500, "attempt 1 too low: {:?}", d1);
assert!(d1.as_millis() <= 2500, "attempt 1 too high: {:?}", d1);
// Attempt 2: base 4000ms, jitter +/-1000ms -> [3000, 5000]
assert!(d2.as_millis() >= 3000, "attempt 2 too low: {:?}", d2);
assert!(d2.as_millis() <= 5000, "attempt 2 too high: {:?}", d2);
}
}
#[test]
fn test_retry_backoff_delay_minimum() {
// Even at attempt 0, delay should be at least 100ms (the minimum floor)
for _ in 0..20 {
let delay = retry_backoff_delay(0);
assert!(delay.as_millis() >= 100);
}
}
#[test]
fn test_retry_backoff_delay_no_overflow() {
// Very high attempt numbers should not panic from overflow
let delay = retry_backoff_delay(30);
assert!(delay.as_millis() >= 100);
}
}
+180 -35
View File
@@ -31,6 +31,8 @@ pub struct SessionConfig {
pub auth_base_url: String,
/// Path to session file (e.g., ~/.ironclaw/session.json).
pub session_path: PathBuf,
/// Port range for OAuth callback server.
pub callback_port_range: (u16, u16),
}
impl Default for SessionConfig {
@@ -38,6 +40,7 @@ impl Default for SessionConfig {
Self {
auth_base_url: "https://private.near.ai".to_string(),
session_path: default_session_path(),
callback_port_range: (9876, 9886),
}
}
}
@@ -80,16 +83,16 @@ impl SessionManager {
};
// Try to load existing session synchronously during construction
if let Ok(data) = std::fs::read_to_string(&manager.config.session_path)
&& let Ok(session) = serde_json::from_str::<SessionData>(&data)
{
// We can't await here, so we use try_write
if let Ok(mut guard) = manager.token.try_write() {
*guard = Some(SecretString::from(session.session_token));
tracing::info!(
"Loaded session token from {}",
manager.config.session_path.display()
);
if let Ok(data) = std::fs::read_to_string(&manager.config.session_path) {
if let Ok(session) = serde_json::from_str::<SessionData>(&data) {
// We can't await here, so we use try_write
if let Ok(mut guard) = manager.token.try_write() {
*guard = Some(SecretString::from(session.session_token));
tracing::info!(
"Loaded session token from {}",
manager.config.session_path.display()
);
}
}
}
@@ -219,21 +222,38 @@ impl SessionManager {
/// Start the OAuth login flow.
///
/// 1. Bind the fixed callback port
/// 1. Find an available port for the callback server
/// 2. Print the auth URL and attempt to open browser
/// 3. Wait for OAuth callback with session token
/// 4. Save and return the token
async fn initiate_login(&self) -> Result<(), LlmError> {
use crate::cli::oauth_defaults::{self, OAUTH_CALLBACK_PORT};
use tokio::io::{AsyncBufReadExt, AsyncWriteExt, BufReader};
use tokio::net::TcpListener;
let listener = oauth_defaults::bind_callback_listener()
.await
.map_err(|e| LlmError::SessionRenewalFailed {
provider: "nearai".to_string(),
reason: e.to_string(),
})?;
// Find an available port
let mut listener = None;
let mut port = 0;
let callback_url = format!("http://127.0.0.1:{}", OAUTH_CALLBACK_PORT);
for p in self.config.callback_port_range.0..=self.config.callback_port_range.1 {
match TcpListener::bind(format!("127.0.0.1:{}", p)).await {
Ok(l) => {
listener = Some(l);
port = p;
break;
}
Err(_) => continue,
}
}
let listener = listener.ok_or_else(|| LlmError::SessionRenewalFailed {
provider: "nearai".to_string(),
reason: format!(
"Could not find available port in range {}-{}",
self.config.callback_port_range.0, self.config.callback_port_range.1
),
})?;
let callback_url = format!("http://127.0.0.1:{}", port);
// Show auth provider menu
println!();
@@ -313,16 +333,138 @@ impl SessionManager {
println!();
println!("Waiting for authentication...");
// The NEAR AI API redirects to: {frontend_callback}/auth/callback?token=X&...
let session_token =
oauth_defaults::wait_for_callback(listener, "/auth/callback", "token", "NEAR AI")
.await
.map_err(|e| LlmError::SessionRenewalFailed {
provider: "nearai".to_string(),
reason: e.to_string(),
// Wait for callback with timeout
// The API redirects to: {frontend_callback}/auth/callback?token=X&session_id=X&expires_at=X&is_new_user=X
let timeout = std::time::Duration::from_secs(300); // 5 minutes
let selected_provider = auth_provider.to_string();
let (session_token, auth_provider) = tokio::time::timeout(timeout, async move {
loop {
let (mut socket, _) = listener.accept().await.map_err(|e| {
LlmError::SessionRenewalFailed {
provider: "nearai".to_string(),
reason: format!("Failed to accept connection: {}", e),
}
})?;
let auth_provider = Some(auth_provider.to_string());
let mut reader = BufReader::new(&mut socket);
let mut request_line = String::new();
reader.read_line(&mut request_line).await.map_err(|e| {
LlmError::SessionRenewalFailed {
provider: "nearai".to_string(),
reason: format!("Failed to read request: {}", e),
}
})?;
// Parse GET /auth/callback?token=xxx&session_id=xxx&expires_at=xxx&is_new_user=xxx HTTP/1.1
if let Some(path) = request_line.split_whitespace().nth(1) {
if path.starts_with("/auth/callback") {
// Parse query parameters
if let Some(query) = path.split('?').nth(1) {
let mut token = None;
for param in query.split('&') {
let parts: Vec<&str> = param.splitn(2, '=').collect();
if parts.len() == 2 && parts[0] == "token" {
token = Some(
urlencoding::decode(parts[1])
.unwrap_or_else(|_| parts[1].into())
.into_owned(),
);
}
}
if let Some(token) = token {
// Send success response with nice styling
let response = concat!(
"HTTP/1.1 200 OK\r\n",
"Content-Type: text/html; charset=utf-8\r\n",
"Connection: close\r\n",
"\r\n",
"<!DOCTYPE html>\n",
"<html>\n",
"<head>\n",
" <meta charset=\"utf-8\">\n",
" <title>NEAR AI - Authentication Successful</title>\n",
" <style>\n",
" * { margin: 0; padding: 0; box-sizing: border-box; }\n",
" body {\n",
" font-family: -apple-system, BlinkMacSystemFont, 'Segoe UI', Roboto, sans-serif;\n",
" background: linear-gradient(135deg, #1a1a2e 0%, #16213e 100%);\n",
" min-height: 100vh;\n",
" display: flex;\n",
" align-items: center;\n",
" justify-content: center;\n",
" color: #fff;\n",
" }\n",
" .container {\n",
" text-align: center;\n",
" padding: 3rem;\n",
" background: rgba(255,255,255,0.05);\n",
" border-radius: 16px;\n",
" backdrop-filter: blur(10px);\n",
" border: 1px solid rgba(255,255,255,0.1);\n",
" max-width: 400px;\n",
" }\n",
" .checkmark {\n",
" width: 80px;\n",
" height: 80px;\n",
" background: linear-gradient(135deg, #00d9a5 0%, #00b386 100%);\n",
" border-radius: 50%;\n",
" display: flex;\n",
" align-items: center;\n",
" justify-content: center;\n",
" margin: 0 auto 1.5rem;\n",
" font-size: 40px;\n",
" }\n",
" h1 {\n",
" font-size: 1.5rem;\n",
" font-weight: 600;\n",
" margin-bottom: 0.75rem;\n",
" }\n",
" p {\n",
" color: rgba(255,255,255,0.7);\n",
" font-size: 0.95rem;\n",
" line-height: 1.5;\n",
" }\n",
" .brand {\n",
" margin-top: 2rem;\n",
" padding-top: 1.5rem;\n",
" border-top: 1px solid rgba(255,255,255,0.1);\n",
" font-size: 0.8rem;\n",
" color: rgba(255,255,255,0.4);\n",
" }\n",
" </style>\n",
"</head>\n",
"<body>\n",
" <div class=\"container\">\n",
" <div class=\"checkmark\">&#10003;</div>\n",
" <h1>Authentication Successful</h1>\n",
" <p>You can close this window and return to the terminal.</p>\n",
" <div class=\"brand\">NEAR AI Agent</div>\n",
" </div>\n",
"</body>\n",
"</html>"
);
let _ = socket.write_all(response.as_bytes()).await;
let _ = socket.shutdown().await;
return Ok::<_, LlmError>((token, Some(selected_provider.clone())));
}
}
}
}
// Not the callback we're looking for, send 404
let response = "HTTP/1.1 404 Not Found\r\nConnection: close\r\n\r\n";
let _ = socket.write_all(response.as_bytes()).await;
}
})
.await
.map_err(|_| LlmError::SessionRenewalFailed {
provider: "nearai".to_string(),
reason: "Authentication timed out after 5 minutes".to_string(),
})??;
// Save the token
self.save_session(&session_token, auth_provider.as_deref())
@@ -500,14 +642,15 @@ pub async fn create_session_manager(config: SessionConfig) -> Arc<SessionManager
let manager = SessionManager::new_async(config).await;
// Check for legacy env var and migrate if present and no file token
if !manager.has_token().await
&& let Ok(token) = std::env::var("NEARAI_SESSION_TOKEN")
&& !token.is_empty()
{
tracing::info!("Migrating session token from NEARAI_SESSION_TOKEN env var to file");
manager.set_token(SecretString::from(token.clone())).await;
if let Err(e) = manager.save_session(&token, None).await {
tracing::warn!("Failed to save migrated session: {}", e);
if !manager.has_token().await {
if let Ok(token) = std::env::var("NEARAI_SESSION_TOKEN") {
if !token.is_empty() {
tracing::info!("Migrating session token from NEARAI_SESSION_TOKEN env var to file");
manager.set_token(SecretString::from(token.clone())).await;
if let Err(e) = manager.save_session(&token, None).await {
tracing::warn!("Failed to save migrated session: {}", e);
}
}
}
}
@@ -528,6 +671,7 @@ mod tests {
let config = SessionConfig {
auth_base_url: "https://example.com".to_string(),
session_path: session_path.clone(),
callback_port_range: (9900, 9910),
};
let manager = SessionManager::new_async(config.clone()).await;
@@ -568,6 +712,7 @@ mod tests {
let config = SessionConfig {
auth_base_url: "https://example.com".to_string(),
session_path: dir.path().join("nonexistent.json"),
callback_port_range: (9900, 9910),
};
let manager = SessionManager::new_async(config).await;
+66 -78
View File
@@ -22,10 +22,7 @@ use ironclaw::{
config::Config,
context::ContextManager,
extensions::ExtensionManager,
llm::{
FailoverProvider, LlmProvider, SessionConfig, create_llm_provider,
create_llm_provider_with_config, create_session_manager,
},
llm::{SessionConfig, create_llm_provider, create_session_manager},
orchestrator::{
ContainerJobConfig, ContainerJobManager, OrchestratorApi, TokenStore,
api::OrchestratorState,
@@ -93,6 +90,7 @@ async fn main() -> anyhow::Result<()> {
.init();
// Memory commands need database (and optionally embeddings)
let _ = dotenvy::dotenv();
let config = Config::from_env()
.await
.map_err(|e| anyhow::anyhow!("{}", e))?;
@@ -101,6 +99,7 @@ async fn main() -> anyhow::Result<()> {
let session = ironclaw::llm::create_session_manager(ironclaw::llm::SessionConfig {
auth_base_url: config.llm.nearai.auth_base_url.clone(),
session_path: config.llm.nearai.session_path.clone(),
..Default::default()
})
.await;
@@ -153,6 +152,7 @@ async fn main() -> anyhow::Result<()> {
return run_pairing_command(pairing_cmd.clone()).map_err(|e| anyhow::anyhow!("{}", e));
}
Some(Command::Status) => {
let _ = dotenvy::dotenv();
tracing_subscriber::fmt()
.with_env_filter(
EnvFilter::try_from_default_env().unwrap_or_else(|_| EnvFilter::new("warn")),
@@ -243,10 +243,8 @@ async fn main() -> anyhow::Result<()> {
skip_auth,
channels_only,
}) => {
// Load .env files before running onboarding wizard.
// Standard ./.env first (higher priority), then ~/.ironclaw/.env.
// Load .env before running onboarding wizard
let _ = dotenvy::dotenv();
ironclaw::bootstrap::load_ironclaw_env();
#[cfg(any(feature = "postgres", feature = "libsql"))]
{
@@ -269,23 +267,23 @@ async fn main() -> anyhow::Result<()> {
}
}
// Load .env files early so DATABASE_URL (and any other vars) are
// available to all subsequent env-based config resolution.
// Standard ./.env first (higher priority), then ~/.ironclaw/.env.
// Load .env if present
let _ = dotenvy::dotenv();
ironclaw::bootstrap::load_ironclaw_env();
// Enhanced first-run detection
#[cfg(any(feature = "postgres", feature = "libsql"))]
if !cli.no_onboard
&& let Some(reason) = check_onboard_needed()
{
println!("Onboarding needed: {}", reason);
println!();
let mut wizard = SetupWizard::new();
wizard.run().await?;
if !cli.no_onboard {
if let Some(reason) = check_onboard_needed().await {
println!("Onboarding needed: {}", reason);
println!();
let mut wizard = SetupWizard::new();
wizard.run().await?;
}
}
// Load bootstrap config (4 fields that must live on disk)
let bootstrap = ironclaw::bootstrap::BootstrapConfig::load();
// Load initial config from env + disk (before DB is available)
let mut config = match Config::from_env().await {
Ok(c) => c,
@@ -305,6 +303,7 @@ async fn main() -> anyhow::Result<()> {
let session_config = SessionConfig {
auth_base_url: config.llm.nearai.auth_base_url.clone(),
session_path: config.llm.nearai.session_path.clone(),
..Default::default()
};
let session = create_session_manager(session_config).await;
@@ -315,7 +314,7 @@ async fn main() -> anyhow::Result<()> {
// Initialize tracing
let env_filter = EnvFilter::try_from_default_env()
.unwrap_or_else(|_| EnvFilter::new("ironclaw=info,tower_http=warn"));
.unwrap_or_else(|_| EnvFilter::new("ironclaw=info,tower_http=debug"));
// Create log broadcaster before tracing init so the WebLogLayer can capture all events.
// This gets wired to the gateway's /api/logs/events SSE endpoint later.
@@ -323,11 +322,7 @@ async fn main() -> anyhow::Result<()> {
tracing_subscriber::registry()
.with(env_filter)
.with(
tracing_subscriber::fmt::layer()
.with_target(false)
.with_writer(ironclaw::tracing_fmt::TruncatingStderr::default()),
)
.with(tracing_subscriber::fmt::layer().with_target(false))
.with(WebLogLayer::new(Arc::clone(&log_broadcaster)))
.init();
@@ -423,7 +418,7 @@ async fn main() -> anyhow::Result<()> {
}
// Reload config from DB now that we have a connection.
match Config::from_db(db.as_ref(), "default").await {
match Config::from_db(db.as_ref(), "default", &bootstrap).await {
Ok(db_config) => {
config = db_config;
tracing::info!("Configuration reloaded from database");
@@ -448,27 +443,6 @@ async fn main() -> anyhow::Result<()> {
let llm = create_llm_provider(&config.llm, session.clone())?;
tracing::info!("LLM provider initialized: {}", llm.model_name());
// Wrap in failover if a fallback model is configured
let llm: Arc<dyn LlmProvider> =
if let Some(fallback_model) = config.llm.nearai.fallback_model.as_ref() {
if fallback_model == &config.llm.nearai.model {
tracing::warn!(
"fallback_model is the same as primary model, failover may not be effective"
);
}
let mut fallback_config = config.llm.nearai.clone();
fallback_config.model = fallback_model.clone();
let fallback = create_llm_provider_with_config(&fallback_config, session.clone())?;
tracing::info!(
primary = %llm.model_name(),
fallback = %fallback.model_name(),
"LLM failover enabled"
);
Arc::new(FailoverProvider::new(vec![llm, fallback])?)
} else {
llm
};
// Initialize safety layer
let safety = Arc::new(SafetyLayer::new(&config.safety));
tracing::info!("Safety layer initialized");
@@ -605,10 +579,7 @@ async fn main() -> anyhow::Result<()> {
// Both register into the shared ToolRegistry (RwLock-based) so concurrent writes are safe.
let wasm_tools_future = async {
if let Some(ref runtime) = wasm_tool_runtime {
let mut loader = WasmToolLoader::new(Arc::clone(runtime), Arc::clone(&tools));
if let Some(ref secrets) = secrets_store {
loader = loader.with_secrets_store(Arc::clone(secrets));
}
let loader = WasmToolLoader::new(Arc::clone(runtime), Arc::clone(&tools));
// Load installed tools from ~/.ironclaw/tools/
match loader.load_from_dir(&config.wasm.tools_dir).await {
@@ -931,13 +902,13 @@ async fn main() -> anyhow::Result<()> {
// Inject owner_id for Telegram so the bot only responds
// to the bound user account.
if channel_name == "telegram"
&& let Some(owner_id) = config.channels.telegram_owner_id
{
config_updates.insert(
"owner_id".to_string(),
serde_json::json!(owner_id),
);
if channel_name == "telegram" {
if let Some(owner_id) = config.channels.telegram_owner_id {
config_updates.insert(
"owner_id".to_string(),
serde_json::json!(owner_id),
);
}
}
if !config_updates.is_empty() {
@@ -1028,23 +999,23 @@ async fn main() -> anyhow::Result<()> {
// Extract its routes for the unified server; the channel itself just
// provides the mpsc stream.
let mut webhook_server_addr: Option<std::net::SocketAddr> = None;
if !cli.cli_only
&& let Some(ref http_config) = config.channels.http
{
let http_channel = HttpChannel::new(http_config.clone());
webhook_routes.push(http_channel.routes());
let (host, port) = http_channel.addr();
webhook_server_addr = Some(
format!("{}:{}", host, port)
.parse()
.expect("HttpConfig host:port must be a valid SocketAddr"),
);
channels.add(Box::new(http_channel));
tracing::info!(
"HTTP channel enabled on {}:{}",
http_config.host,
http_config.port
);
if !cli.cli_only {
if let Some(ref http_config) = config.channels.http {
let http_channel = HttpChannel::new(http_config.clone());
webhook_routes.push(http_channel.routes());
let (host, port) = http_channel.addr();
webhook_server_addr = Some(
format!("{}:{}", host, port)
.parse()
.expect("HttpConfig host:port must be a valid SocketAddr"),
);
channels.add(Box::new(http_channel));
tracing::info!(
"HTTP channel enabled on {}:{}",
http_config.host,
http_config.port
);
}
}
// Start the unified webhook server if any routes were registered.
@@ -1195,11 +1166,13 @@ async fn main() -> anyhow::Result<()> {
/// Check if onboarding is needed and return the reason.
///
/// Returns `Some(reason)` if onboarding should be triggered, `None` otherwise.
/// Called after `load_ironclaw_env()`, so DATABASE_URL from `~/.ironclaw/.env`
/// is already in the environment.
#[cfg(any(feature = "postgres", feature = "libsql"))]
fn check_onboard_needed() -> Option<&'static str> {
let has_db = std::env::var("DATABASE_URL").is_ok()
async fn check_onboard_needed() -> Option<&'static str> {
let bootstrap = ironclaw::bootstrap::BootstrapConfig::load();
// Database not configured (and not in env)
let has_db = bootstrap.database_url.is_some()
|| std::env::var("DATABASE_URL").is_ok()
|| std::env::var("LIBSQL_PATH").is_ok()
|| ironclaw::config::default_libsql_path().exists();
@@ -1207,6 +1180,21 @@ fn check_onboard_needed() -> Option<&'static str> {
return Some("Database not configured");
}
// Secrets not configured (and not in env)
if bootstrap.secrets_master_key_source == ironclaw::settings::KeySource::None
&& std::env::var("SECRETS_MASTER_KEY").is_err()
&& !ironclaw::secrets::keychain::has_master_key().await
{
// Only require secrets setup if user hasn't explicitly disabled it
// For now, we don't require it for first run
}
// First run (onboarding never completed and no session)
let session_path = ironclaw::llm::session::default_session_path();
if !bootstrap.onboard_completed && !session_path.exists() {
return Some("First run");
}
None
}
+12 -12
View File
@@ -202,7 +202,7 @@ async fn report_complete(
State(state): State<OrchestratorState>,
Path(job_id): Path<Uuid>,
Json(report): Json<CompletionReport>,
) -> Result<Json<serde_json::Value>, StatusCode> {
) -> Result<StatusCode, StatusCode> {
if report.success {
tracing::info!(
job_id = %job_id,
@@ -223,7 +223,7 @@ async fn report_complete(
};
let _ = state.job_manager.complete_job(job_id, result).await;
Ok(Json(serde_json::json!({"status": "ok"})))
Ok(StatusCode::OK)
}
// -- Sandbox job event handlers --
@@ -339,16 +339,16 @@ async fn get_prompt_handler(
Path(job_id): Path<Uuid>,
) -> Result<(StatusCode, Json<serde_json::Value>), StatusCode> {
let mut queue = state.prompt_queue.lock().await;
if let Some(prompts) = queue.get_mut(&job_id)
&& let Some(prompt) = prompts.pop_front()
{
return Ok((
StatusCode::OK,
Json(serde_json::json!({
"content": prompt.content,
"done": prompt.done,
})),
));
if let Some(prompts) = queue.get_mut(&job_id) {
if let Some(prompt) = prompts.pop_front() {
return Ok((
StatusCode::OK,
Json(serde_json::json!({
"content": prompt.content,
"done": prompt.done,
})),
));
}
}
// Return 204 with an empty body. The Json wrapper requires some value
+38 -38
View File
@@ -229,17 +229,17 @@ impl ContainerJobManager {
.unwrap_or_else(|| PathBuf::from("."))
.join(".ironclaw")
.join("projects");
if let Ok(canonical_base) = projects_base.canonicalize()
&& !canonical.starts_with(&canonical_base)
{
return Err(OrchestratorError::ContainerCreationFailed {
job_id,
reason: format!(
"project directory {} is outside allowed base {}",
canonical.display(),
canonical_base.display()
),
});
if let Ok(canonical_base) = projects_base.canonicalize() {
if !canonical.starts_with(&canonical_base) {
return Err(OrchestratorError::ContainerCreationFailed {
job_id,
reason: format!(
"project directory {} is outside allowed base {}",
canonical.display(),
canonical_base.display()
),
});
}
}
binds.push(format!("{}:/workspace:rw", canonical.display()));
env_vec.push("IRONCLAW_WORKSPACE=/workspace".to_string());
@@ -442,36 +442,36 @@ impl ContainerJobManager {
let containers = self.containers.read().await;
containers.get(&job_id).map(|h| h.container_id.clone())
};
if let Some(cid) = container_id
&& !cid.is_empty()
{
match connect_docker().await {
Ok(docker) => {
if let Err(e) = docker
.stop_container(
&cid,
Some(bollard::container::StopContainerOptions { t: 5 }),
)
.await
{
tracing::warn!(job_id = %job_id, error = %e, "Failed to stop completed container");
if let Some(cid) = container_id {
if !cid.is_empty() {
match connect_docker().await {
Ok(docker) => {
if let Err(e) = docker
.stop_container(
&cid,
Some(bollard::container::StopContainerOptions { t: 5 }),
)
.await
{
tracing::warn!(job_id = %job_id, error = %e, "Failed to stop completed container");
}
if let Err(e) = docker
.remove_container(
&cid,
Some(bollard::container::RemoveContainerOptions {
force: true,
..Default::default()
}),
)
.await
{
tracing::warn!(job_id = %job_id, error = %e, "Failed to remove completed container");
}
}
if let Err(e) = docker
.remove_container(
&cid,
Some(bollard::container::RemoveContainerOptions {
force: true,
..Default::default()
}),
)
.await
{
tracing::warn!(job_id = %job_id, error = %e, "Failed to remove completed container");
Err(e) => {
tracing::warn!(job_id = %job_id, error = %e, "Failed to connect to Docker for container cleanup");
}
}
Err(e) => {
tracing::warn!(job_id = %job_id, error = %e, "Failed to connect to Docker for container cleanup");
}
}
}
self.token_store.revoke(job_id).await;
+4 -4
View File
@@ -147,10 +147,10 @@ impl LeakDetector {
// Build prefix matcher for patterns that start with a known prefix
let mut prefixes = Vec::new();
for (idx, pattern) in patterns.iter().enumerate() {
if let Some(prefix) = extract_literal_prefix(pattern.regex.as_str())
&& prefix.len() >= 3
{
prefixes.push((prefix, idx));
if let Some(prefix) = extract_literal_prefix(pattern.regex.as_str()) {
if prefix.len() >= 3 {
prefixes.push((prefix, idx));
}
}
}
+7 -6
View File
@@ -494,10 +494,10 @@ impl ContainerRunner {
/// 3. `~/.docker/run/docker.sock` (Docker Desktop on macOS)
pub async fn connect_docker() -> Result<Docker> {
// First try bollard defaults (checks DOCKER_HOST, then /var/run/docker.sock)
if let Ok(docker) = Docker::connect_with_local_defaults()
&& docker.ping().await.is_ok()
{
return Ok(docker);
if let Ok(docker) = Docker::connect_with_local_defaults() {
if docker.ping().await.is_ok() {
return Ok(docker);
}
}
// Try Docker Desktop socket (macOS)
@@ -507,9 +507,10 @@ pub async fn connect_docker() -> Result<Docker> {
let sock_str = desktop_sock.to_string_lossy();
if let Ok(docker) =
Docker::connect_with_socket(&sock_str, 120, bollard::API_DEFAULT_VERSION)
&& docker.ping().await.is_ok()
{
return Ok(docker);
if docker.ping().await.is_ok() {
return Ok(docker);
}
}
}
}
+9 -9
View File
@@ -259,11 +259,11 @@ async fn handle_connect(
let decision = state.decider.decide(&network_req).await;
if !decision.is_allowed()
&& let NetworkDecision::Deny { reason } = decision
{
tracing::info!("Proxy: blocked CONNECT {} - {}", host, reason);
return error_response(StatusCode::FORBIDDEN, reason);
if !decision.is_allowed() {
if let NetworkDecision::Deny { reason } = decision {
tracing::info!("Proxy: blocked CONNECT {} - {}", host, reason);
return error_response(StatusCode::FORBIDDEN, reason);
}
}
tracing::debug!("Proxy: allowing CONNECT to {}", host);
@@ -294,10 +294,10 @@ async fn forward_request(
// Copy headers (except hop-by-hop headers)
for (name, value) in req.headers() {
if !is_hop_by_hop_header(name.as_str())
&& let Ok(v) = value.to_str()
{
builder = builder.header(name.as_str(), v);
if !is_hop_by_hop_header(name.as_str()) {
if let Ok(v) = value.to_str() {
builder = builder.header(name.as_str(), v);
}
}
}
+5 -4
View File
@@ -109,11 +109,12 @@ impl NetworkPolicyDecider for DefaultPolicyDecider {
async fn decide(&self, request: &NetworkRequest) -> NetworkDecision {
// First check if the domain is allowed
let validation = self.allowlist.is_allowed(&request.host);
if !validation.is_allowed()
&& let crate::sandbox::proxy::allowlist::DomainValidationResult::Denied(reason) =
if !validation.is_allowed() {
if let crate::sandbox::proxy::allowlist::DomainValidationResult::Denied(reason) =
validation
{
return NetworkDecision::Deny { reason };
{
return NetworkDecision::Deny { reason };
}
}
// Check if we need to inject credentials
+1 -1
View File
@@ -261,7 +261,7 @@ pub use platform::{delete_master_key, get_master_key, has_master_key, store_mast
/// Parse a hex string to bytes.
fn hex_to_bytes(hex: &str) -> Result<Vec<u8>, SecretError> {
if !hex.len().is_multiple_of(2) {
if hex.len() % 2 != 0 {
return Err(SecretError::KeychainError(
"Invalid hex string length".to_string(),
));
+22 -59
View File
@@ -153,10 +153,10 @@ impl SecretsStore for PostgresSecretsStore {
let secret = row_to_secret(&r);
// Check expiration
if let Some(expires_at) = secret.expires_at
&& expires_at < Utc::now()
{
return Err(SecretError::Expired);
if let Some(expires_at) = secret.expires_at {
if expires_at < Utc::now() {
return Err(SecretError::Expired);
}
}
Ok(secret)
@@ -276,10 +276,10 @@ impl SecretsStore for PostgresSecretsStore {
}
// Simple glob: * matches any suffix
if let Some(prefix) = pattern.strip_suffix('*')
&& secret_name.starts_with(prefix)
{
return Ok(true);
if let Some(prefix) = pattern.strip_suffix('*') {
if secret_name.starts_with(prefix) {
return Ok(true);
}
}
}
@@ -432,10 +432,10 @@ impl SecretsStore for LibSqlSecretsStore {
Some(row) => {
let secret = libsql_row_to_secret(&row)?;
if let Some(expires_at) = secret.expires_at
&& expires_at < Utc::now()
{
return Err(SecretError::Expired);
if let Some(expires_at) = secret.expires_at {
if expires_at < Utc::now() {
return Err(SecretError::Expired);
}
}
Ok(secret)
@@ -541,10 +541,10 @@ impl SecretsStore for LibSqlSecretsStore {
return Ok(true);
}
if let Some(prefix) = pattern.strip_suffix('*')
&& secret_name.starts_with(prefix)
{
return Ok(true);
if let Some(prefix) = pattern.strip_suffix('*') {
if secret_name.starts_with(prefix) {
return Ok(true);
}
}
}
@@ -695,21 +695,12 @@ pub mod testing {
}
async fn get(&self, user_id: &str, name: &str) -> Result<Secret, SecretError> {
let secret = self
.secrets
self.secrets
.read()
.await
.get(&(user_id.to_string(), name.to_string()))
.cloned()
.ok_or_else(|| SecretError::NotFound(name.to_string()))?;
if let Some(expires_at) = secret.expires_at
&& expires_at < Utc::now()
{
return Err(SecretError::Expired);
}
Ok(secret)
.ok_or_else(|| SecretError::NotFound(name.to_string()))
}
async fn get_decrypted(
@@ -770,10 +761,10 @@ pub mod testing {
if pattern == secret_name {
return Ok(true);
}
if let Some(prefix) = pattern.strip_suffix('*')
&& secret_name.starts_with(prefix)
{
return Ok(true);
if let Some(prefix) = pattern.strip_suffix('*') {
if secret_name.starts_with(prefix) {
return Ok(true);
}
}
}
Ok(false)
@@ -898,34 +889,6 @@ mod tests {
);
}
#[tokio::test]
async fn test_expired_secret_returns_error() {
let store = test_store();
let expires_at = chrono::Utc::now() - chrono::Duration::hours(1);
let params = CreateSecretParams::new("expired_key", "value").with_expiry(expires_at);
store.create("user1", params).await.unwrap();
let result = store.get("user1", "expired_key").await;
assert!(result.is_err());
assert!(matches!(
result.unwrap_err(),
crate::secrets::SecretError::Expired
));
}
#[tokio::test]
async fn test_non_expired_secret_succeeds() {
let store = test_store();
let expires_at = chrono::Utc::now() + chrono::Duration::hours(1);
let params = CreateSecretParams::new("fresh_key", "value").with_expiry(expires_at);
store.create("user1", params).await.unwrap();
let result = store.get("user1", "fresh_key").await;
assert!(result.is_ok());
}
#[tokio::test]
async fn test_user_isolation() {
let store = test_store();
+72 -18
View File
@@ -499,6 +499,14 @@ impl Default for BuilderSettings {
}
impl Settings {
/// Get the default settings file path (~/.ironclaw/settings.json).
pub fn default_path() -> PathBuf {
dirs::home_dir()
.unwrap_or_else(|| PathBuf::from("."))
.join(".ironclaw")
.join("settings.json")
}
/// Reconstruct Settings from a flat key-value map (as stored in the DB).
///
/// Each key is a dotted path (e.g., "agent.name"), value is a JSONB value.
@@ -544,27 +552,50 @@ impl Settings {
map
}
/// Get the default settings file path (~/.ironclaw/settings.json).
pub fn default_path() -> std::path::PathBuf {
dirs::home_dir()
.unwrap_or_else(|| std::path::PathBuf::from("."))
.join(".ironclaw")
.join("settings.json")
}
/// Load settings from disk, returning default if not found.
pub fn load() -> Self {
Self::load_from(&Self::default_path())
}
/// Load settings from a specific path (used by bootstrap legacy migration).
pub fn load_from(path: &std::path::Path) -> Self {
/// Load settings from a specific path.
pub fn load_from(path: &PathBuf) -> Self {
match std::fs::read_to_string(path) {
Ok(data) => serde_json::from_str(&data).unwrap_or_default(),
Err(_) => Self::default(),
}
}
/// Save settings to disk.
pub fn save(&self) -> std::io::Result<()> {
self.save_to(&Self::default_path())
}
/// Save settings to a specific path.
pub fn save_to(&self, path: &PathBuf) -> std::io::Result<()> {
// Ensure parent directory exists
if let Some(parent) = path.parent() {
std::fs::create_dir_all(parent)?;
}
let json = serde_json::to_string_pretty(self)
.map_err(|e| std::io::Error::new(std::io::ErrorKind::InvalidData, e.to_string()))?;
std::fs::write(path, json)
}
/// Get the selected model, falling back to the provided default.
pub fn model_or(&self, default: &str) -> String {
self.selected_model
.clone()
.unwrap_or_else(|| default.to_string())
}
/// Set the selected model and save.
pub fn set_model(&mut self, model: &str) -> std::io::Result<()> {
self.selected_model = Some(model.to_string());
self.save()
}
/// Get a setting value by dotted path (e.g., "agent.max_parallel_jobs").
pub fn get(&self, path: &str) -> Option<String> {
let json = serde_json::to_value(self).ok()?;
@@ -749,22 +780,42 @@ fn collect_settings(
#[cfg(test)]
mod tests {
use super::*;
use tempfile::tempdir;
#[test]
fn test_db_map_round_trip() {
fn test_settings_save_load() {
let dir = tempdir().unwrap();
let path = dir.path().join("settings.json");
let settings = Settings {
selected_model: Some("claude-3-5-sonnet-20241022".to_string()),
..Default::default()
};
let map = settings.to_db_map();
let restored = Settings::from_db_map(&map);
settings.save_to(&path).unwrap();
let loaded = Settings::load_from(&path);
assert_eq!(
restored.selected_model,
loaded.selected_model,
Some("claude-3-5-sonnet-20241022".to_string())
);
}
#[test]
fn test_model_or_default() {
let settings = Settings::default();
assert_eq!(
settings.model_or("default-model"),
"default-model".to_string()
);
let settings = Settings {
selected_model: Some("my-model".to_string()),
..Default::default()
};
assert_eq!(settings.model_or("default-model"), "my-model".to_string());
}
#[test]
fn test_get_setting() {
let settings = Settings::default();
@@ -835,13 +886,16 @@ mod tests {
}
#[test]
fn test_telegram_owner_id_db_round_trip() {
fn test_telegram_owner_id_round_trip() {
let dir = tempdir().unwrap();
let path = dir.path().join("settings.json");
let mut settings = Settings::default();
settings.channels.telegram_owner_id = Some(123456789);
settings.save_to(&path).unwrap();
let map = settings.to_db_map();
let restored = Settings::from_db_map(&map);
assert_eq!(restored.channels.telegram_owner_id, Some(123456789));
let loaded = Settings::load_from(&path);
assert_eq!(loaded.channels.telegram_owner_id, Some(123456789));
}
#[test]
+47 -43
View File
@@ -15,7 +15,7 @@ use serde::Deserialize;
#[cfg(feature = "postgres")]
use crate::secrets::SecretsCrypto;
use crate::secrets::{CreateSecretParams, SecretsStore};
use crate::settings::{Settings, TunnelSettings};
use crate::settings::Settings;
use crate::setup::prompts::{
confirm, input, optional_input, print_error, print_info, print_success, secret_input,
};
@@ -131,10 +131,7 @@ struct TelegramUpdateUser {
/// 2. Entering the bot token
/// 3. Validating the token
/// 4. Saving the token to the database
pub async fn setup_telegram(
secrets: &SecretsContext,
settings: &Settings,
) -> Result<TelegramSetupResult, String> {
pub async fn setup_telegram(secrets: &SecretsContext) -> Result<TelegramSetupResult, String> {
println!("Telegram Setup:");
println!();
print_info("To create a Telegram bot:");
@@ -148,8 +145,8 @@ pub async fn setup_telegram(
print_info("Existing Telegram token found in database.");
if !confirm("Replace existing token?", false).map_err(|e| e.to_string())? {
// Still offer to configure webhook secret and owner binding
let webhook_secret = setup_telegram_webhook_secret(secrets, &settings.tunnel).await?;
let owner_id = bind_telegram_owner_flow(secrets, settings).await?;
let webhook_secret = setup_telegram_webhook_secret(secrets).await?;
let owner_id = bind_telegram_owner_flow(secrets).await?;
return Ok(TelegramSetupResult {
enabled: true,
bot_username: None,
@@ -179,7 +176,7 @@ pub async fn setup_telegram(
let owner_id = bind_telegram_owner(&token).await?;
// Offer webhook secret configuration
let webhook_secret = setup_telegram_webhook_secret(secrets, &settings.tunnel).await?;
let webhook_secret = setup_telegram_webhook_secret(secrets).await?;
Ok(TelegramSetupResult {
enabled: true,
@@ -192,7 +189,7 @@ pub async fn setup_telegram(
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
Box::pin(setup_telegram(secrets)).await
} else {
Ok(TelegramSetupResult {
enabled: false,
@@ -266,32 +263,32 @@ async fn bind_telegram_owner(token: &SecretString) -> Result<Option<i64>, String
// Find the first message with a sender
for update in &body.result {
if let Some(ref msg) = update.message
&& let Some(ref from) = msg.from
{
let display_name = from
.username
.as_ref()
.map(|u| format!("@{}", u))
.unwrap_or_else(|| from.first_name.clone());
if let Some(ref msg) = update.message {
if let Some(ref from) = msg.from {
let display_name = from
.username
.as_ref()
.map(|u| format!("@{}", u))
.unwrap_or_else(|| from.first_name.clone());
print_success(&format!(
"Received message from {} (ID: {})",
display_name, from.id
));
print_success(&format!(
"Received message from {} (ID: {})",
display_name, from.id
));
// Acknowledge the update so it doesn't pile up
let ack_url = format!(
"https://api.telegram.org/bot{}/getUpdates",
token.expose_secret()
);
let _ = client
.get(&ack_url)
.query(&[("offset", &(update.update_id + 1).to_string())])
.send()
.await;
// Acknowledge the update so it doesn't pile up
let ack_url = format!(
"https://api.telegram.org/bot{}/getUpdates",
token.expose_secret()
);
let _ = client
.get(&ack_url)
.query(&[("offset", &(update.update_id + 1).to_string())])
.send()
.await;
return Ok(Some(from.id));
return Ok(Some(from.id));
}
}
}
}
@@ -304,10 +301,9 @@ async fn bind_telegram_owner(token: &SecretString) -> Result<Option<i64>, String
/// Bind flow when the token already exists (reads from secrets store).
///
/// Retrieves the saved bot token and delegates to `bind_telegram_owner`.
async fn bind_telegram_owner_flow(
secrets: &SecretsContext,
settings: &Settings,
) -> Result<Option<i64>, String> {
async fn bind_telegram_owner_flow(secrets: &SecretsContext) -> Result<Option<i64>, String> {
// Check current settings first
let settings = Settings::load();
if settings.channels.telegram_owner_id.is_some() {
print_info("Bot is already bound to a Telegram account.");
if !confirm("Re-bind to a different account?", false).map_err(|e| e.to_string())? {
@@ -325,7 +321,9 @@ 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() -> Result<Option<String>, String> {
// Check if already configured
let settings = Settings::load();
if let Some(ref url) = settings.tunnel.public_url {
print_info(&format!("Existing tunnel configured: {}", url));
if !confirm("Change tunnel configuration?", false).map_err(|e| e.to_string())? {
@@ -364,7 +362,14 @@ pub fn setup_tunnel(settings: &Settings) -> Result<Option<String>, String> {
// Remove trailing slash if present
let tunnel_url = tunnel_url.trim_end_matches('/').to_string();
print_success(&format!("Tunnel URL configured: {}", tunnel_url));
// Save to settings
let mut settings = Settings::load();
settings.tunnel.public_url = Some(tunnel_url.clone());
settings
.save()
.map_err(|e| format!("Failed to save settings: {}", e))?;
print_success(&format!("Tunnel URL saved: {}", tunnel_url));
print_info("");
print_info("Make sure your tunnel is running before starting the agent.");
print_info("You can also set TUNNEL_URL environment variable to override.");
@@ -375,11 +380,10 @@ pub fn setup_tunnel(settings: &Settings) -> Result<Option<String>, String> {
/// Set up Telegram webhook secret for signature validation.
///
/// Returns the webhook secret if configured.
async fn setup_telegram_webhook_secret(
secrets: &SecretsContext,
tunnel: &TunnelSettings,
) -> Result<Option<String>, String> {
if tunnel.public_url.is_none() {
async fn setup_telegram_webhook_secret(secrets: &SecretsContext) -> Result<Option<String>, String> {
// Check if tunnel is configured
let settings = Settings::load();
if settings.tunnel.public_url.is_none() {
print_info("");
print_info("No tunnel configured. Telegram will use polling mode (30s+ delay).");
print_info("Run setup again to configure a tunnel for instant delivery.");
+4 -5
View File
@@ -54,11 +54,10 @@ pub fn select_one(prompt: &str, options: &[&str]) -> io::Result<usize> {
}
// Parse number
if let Ok(num) = input.parse::<usize>()
&& num >= 1
&& num <= options.len()
{
return Ok(num - 1);
if let Ok(num) = input.parse::<usize>() {
if num >= 1 && num <= options.len() {
return Ok(num - 1);
}
}
writeln!(
+27 -90
View File
@@ -83,7 +83,7 @@ impl SetupWizard {
pub fn new() -> Self {
Self {
config: SetupConfig::default(),
settings: Settings::default(),
settings: Settings::load(),
session_manager: None,
#[cfg(feature = "postgres")]
db_pool: None,
@@ -97,7 +97,7 @@ impl SetupWizard {
pub fn with_config(config: SetupConfig) -> Self {
Self {
config,
settings: Settings::default(),
settings: Settings::load(),
session_manager: None,
#[cfg(feature = "postgres")]
db_pool: None,
@@ -158,7 +158,7 @@ impl SetupWizard {
}
// Save settings and print summary
self.save_and_summarize().await?;
self.save_and_summarize()?;
Ok(())
}
@@ -534,17 +534,17 @@ impl SetupWizard {
/// Step 3: NEAR AI authentication.
async fn step_authentication(&mut self) -> Result<(), SetupError> {
// Check if we already have a session
if let Some(ref session) = self.session_manager
&& session.has_token().await
{
print_info("Existing session found. Validating...");
match session.ensure_authenticated().await {
Ok(()) => {
print_success("Session valid");
return Ok(());
}
Err(e) => {
print_info(&format!("Session invalid: {}. Re-authenticating...", e));
if let Some(ref session) = self.session_manager {
if session.has_token().await {
print_info("Existing session found. Validating...");
match session.ensure_authenticated().await {
Ok(()) => {
print_success("Session valid");
return Ok(());
}
Err(e) => {
print_info(&format!("Session invalid: {}. Re-authenticating...", e));
}
}
}
}
@@ -653,8 +653,6 @@ impl SetupWizard {
session_path: crate::llm::session::default_session_path(),
api_mode: crate::config::NearAiApiMode::Responses,
api_key: None,
fallback_model: None,
max_retries: 3,
},
openai: None,
anthropic: None,
@@ -818,7 +816,7 @@ impl SetupWizard {
/// Step 6: Channel configuration.
async fn step_channels(&mut self) -> Result<(), SetupError> {
// First, configure tunnel (shared across all channels that need webhooks)
match setup_tunnel(&self.settings) {
match setup_tunnel() {
Ok(Some(url)) => {
self.settings.tunnel.public_url = Some(url);
}
@@ -883,10 +881,11 @@ impl SetupWizard {
&installed_names,
)
.await?
&& !installed.is_empty()
{
print_success(&format!("Installed channels: {}", installed.join(", ")));
discovered_channels = discover_wasm_channels(&channels_dir).await;
if !installed.is_empty() {
print_success(&format!("Installed channels: {}", installed.join(", ")));
discovered_channels = discover_wasm_channels(&channels_dir).await;
}
}
// Determine if we need secrets context
@@ -934,9 +933,8 @@ impl SetupWizard {
.await
.map_err(SetupError::Channel)?
} else if channel_name == "telegram" {
let telegram_result = setup_telegram(ctx, &self.settings)
.await
.map_err(SetupError::Channel)?;
let telegram_result =
setup_telegram(ctx).await.map_err(SetupError::Channel)?;
if let Some(owner_id) = telegram_result.owner_id {
self.settings.channels.telegram_owner_id = Some(owner_id);
}
@@ -1019,77 +1017,16 @@ impl SetupWizard {
Ok(())
}
/// Save settings to the database and `~/.ironclaw/.env`, then print summary.
async fn save_and_summarize(&mut self) -> Result<(), SetupError> {
/// Save settings and print summary.
fn save_and_summarize(&mut self) -> Result<(), SetupError> {
self.settings.onboard_completed = true;
// Write all settings to the database (whichever backend is active).
{
let db_map = self.settings.to_db_map();
let saved = false;
#[cfg(feature = "postgres")]
let saved = if !saved {
if let Some(ref pool) = self.db_pool {
let store = crate::history::Store::from_pool(pool.clone());
store
.set_all_settings("default", &db_map)
.await
.map_err(|e| {
SetupError::Database(format!(
"Failed to save settings to database: {}",
e
))
})?;
true
} else {
false
}
} else {
saved
};
#[cfg(feature = "libsql")]
let saved = if !saved {
if let Some(ref backend) = self.db_backend {
use crate::db::Database as _;
backend
.set_all_settings("default", &db_map)
.await
.map_err(|e| {
SetupError::Database(format!(
"Failed to save settings to database: {}",
e
))
})?;
true
} else {
false
}
} else {
saved
};
if !saved {
return Err(SetupError::Database(
"No database connection, cannot save settings".to_string(),
));
}
}
// Save DATABASE_URL to ~/.ironclaw/.env (the only field that needs
// disk persistence before the DB is available).
if let Some(ref url) = self.settings.database_url {
crate::bootstrap::save_database_url(url).map_err(|e| {
SetupError::Io(std::io::Error::other(format!(
"Failed to save DATABASE_URL to .env: {}",
e
)))
})?;
}
self.settings
.save()
.map_err(|e| std::io::Error::other(format!("Failed to save settings: {}", e)))?;
println!();
print_success("Configuration saved to database");
print_success("Configuration saved to ~/.ironclaw/");
println!();
// Print summary
+27 -27
View File
@@ -326,20 +326,20 @@ impl TestHarness {
}
// Verify expected output
if let Some(ref expected) = test.expected_output
&& &actual != expected
{
return TestResult {
name: test.name.clone(),
passed: false,
duration,
error: Some(format!(
"Output mismatch:\nExpected: {}\nActual: {}",
serde_json::to_string_pretty(expected).unwrap_or_default(),
serde_json::to_string_pretty(&actual).unwrap_or_default()
)),
actual_output: Some(actual),
};
if let Some(ref expected) = test.expected_output {
if &actual != expected {
return TestResult {
name: test.name.clone(),
passed: false,
duration,
error: Some(format!(
"Output mismatch:\nExpected: {}\nActual: {}",
serde_json::to_string_pretty(expected).unwrap_or_default(),
serde_json::to_string_pretty(&actual).unwrap_or_default()
)),
actual_output: Some(actual),
};
}
}
// Verify expected fields
@@ -357,19 +357,19 @@ impl TestHarness {
};
}
if let Some(ref expected_value) = field.value
&& field_value != Some(expected_value)
{
return TestResult {
name: test.name.clone(),
passed: false,
duration,
error: Some(format!(
"Field '{}' mismatch: expected {:?}, got {:?}",
field.path, expected_value, field_value
)),
actual_output: Some(actual),
};
if let Some(ref expected_value) = field.value {
if field_value != Some(expected_value) {
return TestResult {
name: test.name.clone(),
passed: false,
duration,
error: Some(format!(
"Field '{}' mismatch: expected {:?}, got {:?}",
field.path, expected_value, field_value
)),
actual_output: Some(actual),
};
}
}
}
}
-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);
}
}
}
+7 -2
View File
@@ -3,7 +3,7 @@
use async_trait::async_trait;
use crate::context::JobContext;
use crate::tools::tool::{Tool, ToolError, ToolOutput, require_str};
use crate::tools::tool::{Tool, ToolError, ToolOutput};
/// Simple echo tool for testing.
pub struct EchoTool;
@@ -38,7 +38,12 @@ impl Tool for EchoTool {
) -> Result<ToolOutput, ToolError> {
let start = std::time::Instant::now();
let message = require_str(&params, "message")?;
let message = params
.get("message")
.and_then(|v| v.as_str())
.ok_or_else(|| {
ToolError::InvalidParameters("missing 'message' parameter".to_string())
})?;
Ok(ToolOutput::text(message, start.elapsed()))
}
+136
View File
@@ -0,0 +1,136 @@
//! E-commerce tool for shopping and price comparison.
use async_trait::async_trait;
use crate::context::JobContext;
use crate::tools::tool::{Tool, ToolError, ToolOutput};
/// Tool for e-commerce operations (Amazon, price comparison, etc.).
pub struct EcommerceTool {
// TODO: Add API clients
}
impl EcommerceTool {
/// Create a new e-commerce tool.
pub fn new() -> Self {
Self {}
}
}
impl Default for EcommerceTool {
fn default() -> Self {
Self::new()
}
}
#[async_trait]
impl Tool for EcommerceTool {
fn name(&self) -> &str {
"ecommerce"
}
fn description(&self) -> &str {
"Search products, compare prices, and find deals across e-commerce platforms."
}
fn parameters_schema(&self) -> serde_json::Value {
serde_json::json!({
"type": "object",
"properties": {
"action": {
"type": "string",
"enum": ["search", "get_product", "compare_prices", "track_price"],
"description": "The e-commerce action to perform"
},
"query": {
"type": "string",
"description": "Search query (for search action)"
},
"product_id": {
"type": "string",
"description": "Product ID or ASIN (for get_product, compare_prices)"
},
"platform": {
"type": "string",
"enum": ["amazon", "ebay", "walmart", "all"],
"description": "E-commerce platform to search"
},
"max_price": {
"type": "number",
"description": "Maximum price filter"
},
"category": {
"type": "string",
"description": "Product category filter"
}
},
"required": ["action"]
})
}
async fn execute(
&self,
params: serde_json::Value,
_ctx: &JobContext,
) -> Result<ToolOutput, ToolError> {
let start = std::time::Instant::now();
let action = params
.get("action")
.and_then(|v| v.as_str())
.ok_or_else(|| {
ToolError::InvalidParameters("missing 'action' parameter".to_string())
})?;
// TODO: Implement actual e-commerce API integrations
let result = match action {
"search" => {
let query = params.get("query").and_then(|v| v.as_str()).unwrap_or("");
serde_json::json!({
"query": query,
"results": [],
"message": "E-commerce integration not yet implemented"
})
}
"get_product" => {
let product_id = params
.get("product_id")
.and_then(|v| v.as_str())
.ok_or_else(|| {
ToolError::InvalidParameters("missing 'product_id' parameter".to_string())
})?;
serde_json::json!({
"product_id": product_id,
"found": false,
"message": "E-commerce integration not yet implemented"
})
}
"compare_prices" => {
serde_json::json!({
"prices": [],
"message": "E-commerce integration not yet implemented"
})
}
"track_price" => {
serde_json::json!({
"tracking": false,
"message": "E-commerce integration not yet implemented"
})
}
_ => {
return Err(ToolError::InvalidParameters(format!(
"unknown action: {}",
action
)));
}
};
Ok(ToolOutput::success(result, start.elapsed()))
}
fn requires_sanitization(&self) -> bool {
true // External e-commerce data
}
}
+17 -5
View File
@@ -9,7 +9,7 @@ use async_trait::async_trait;
use crate::context::JobContext;
use crate::extensions::{ExtensionKind, ExtensionManager};
use crate::tools::tool::{Tool, ToolError, ToolOutput, require_str};
use crate::tools::tool::{Tool, ToolError, ToolOutput};
// ── tool_search ──────────────────────────────────────────────────────────
@@ -133,7 +133,10 @@ impl Tool for ToolInstallTool {
) -> Result<ToolOutput, ToolError> {
let start = std::time::Instant::now();
let name = require_str(&params, "name")?;
let name = params
.get("name")
.and_then(|v| v.as_str())
.ok_or_else(|| ToolError::InvalidParameters("name is required".to_string()))?;
let url = params.get("url").and_then(|v| v.as_str());
@@ -207,7 +210,10 @@ impl Tool for ToolAuthTool {
) -> Result<ToolOutput, ToolError> {
let start = std::time::Instant::now();
let name = require_str(&params, "name")?;
let name = params
.get("name")
.and_then(|v| v.as_str())
.ok_or_else(|| ToolError::InvalidParameters("name is required".to_string()))?;
let result = self
.manager
@@ -300,7 +306,10 @@ impl Tool for ToolActivateTool {
) -> Result<ToolOutput, ToolError> {
let start = std::time::Instant::now();
let name = require_str(&params, "name")?;
let name = params
.get("name")
.and_then(|v| v.as_str())
.ok_or_else(|| ToolError::InvalidParameters("name is required".to_string()))?;
match self.manager.activate(name).await {
Ok(result) => {
@@ -462,7 +471,10 @@ impl Tool for ToolRemoveTool {
) -> Result<ToolOutput, ToolError> {
let start = std::time::Instant::now();
let name = require_str(&params, "name")?;
let name = params
.get("name")
.and_then(|v| v.as_str())
.ok_or_else(|| ToolError::InvalidParameters("name is required".to_string()))?;
let message = self
.manager
+25 -7
View File
@@ -11,7 +11,7 @@ use async_trait::async_trait;
use tokio::fs;
use crate::context::JobContext;
use crate::tools::tool::{Tool, ToolDomain, ToolError, ToolOutput, require_str};
use crate::tools::tool::{Tool, ToolDomain, ToolError, ToolOutput};
use crate::workspace::paths as ws_paths;
/// Well-known workspace filenames that must go through memory_write, not write_file.
@@ -203,7 +203,10 @@ impl Tool for ReadFileTool {
params: serde_json::Value,
_ctx: &JobContext,
) -> Result<ToolOutput, ToolError> {
let path_str = require_str(&params, "path")?;
let path_str = params
.get("path")
.and_then(|v| v.as_str())
.ok_or_else(|| ToolError::InvalidParameters("missing 'path' parameter".into()))?;
let offset = params.get("offset").and_then(|v| v.as_u64()).unwrap_or(0) as usize;
let limit = params.get("limit").and_then(|v| v.as_u64());
@@ -325,7 +328,10 @@ impl Tool for WriteFileTool {
params: serde_json::Value,
_ctx: &JobContext,
) -> Result<ToolOutput, ToolError> {
let path_str = require_str(&params, "path")?;
let path_str = params
.get("path")
.and_then(|v| v.as_str())
.ok_or_else(|| ToolError::InvalidParameters("missing 'path' parameter".into()))?;
// Reject workspace paths: these live in the database, not on disk.
if is_workspace_path(path_str) {
@@ -336,7 +342,10 @@ impl Tool for WriteFileTool {
)));
}
let content = require_str(&params, "content")?;
let content = params
.get("content")
.and_then(|v| v.as_str())
.ok_or_else(|| ToolError::InvalidParameters("missing 'content' parameter".into()))?;
let start = std::time::Instant::now();
@@ -641,11 +650,20 @@ impl Tool for ApplyPatchTool {
params: serde_json::Value,
_ctx: &JobContext,
) -> Result<ToolOutput, ToolError> {
let path_str = require_str(&params, "path")?;
let path_str = params
.get("path")
.and_then(|v| v.as_str())
.ok_or_else(|| ToolError::InvalidParameters("missing 'path' parameter".into()))?;
let old_string = require_str(&params, "old_string")?;
let old_string = params
.get("old_string")
.and_then(|v| v.as_str())
.ok_or_else(|| ToolError::InvalidParameters("missing 'old_string' parameter".into()))?;
let new_string = require_str(&params, "new_string")?;
let new_string = params
.get("new_string")
.and_then(|v| v.as_str())
.ok_or_else(|| ToolError::InvalidParameters("missing 'new_string' parameter".into()))?;
let replace_all = params
.get("replace_all")
+17 -9
View File
@@ -9,7 +9,7 @@ use reqwest::Client;
use crate::context::JobContext;
use crate::safety::LeakDetector;
use crate::tools::tool::{Tool, ToolError, ToolOutput, require_str};
use crate::tools::tool::{Tool, ToolError, ToolOutput};
/// Maximum response body size (5 MB). Prevents OOM from unbounded responses.
const MAX_RESPONSE_SIZE: usize = 5 * 1024 * 1024;
@@ -54,12 +54,12 @@ fn validate_url(url: &str) -> Result<reqwest::Url, ToolError> {
}
// Check literal IP addresses
if let Ok(ip) = host.parse::<IpAddr>()
&& is_disallowed_ip(&ip)
{
return Err(ToolError::NotAuthorized(
"private or local IPs are not allowed".to_string(),
));
if let Ok(ip) = host.parse::<IpAddr>() {
if is_disallowed_ip(&ip) {
return Err(ToolError::NotAuthorized(
"private or local IPs are not allowed".to_string(),
));
}
}
// Resolve hostname and check all resolved IPs against the blocklist.
@@ -154,9 +154,17 @@ impl Tool for HttpTool {
) -> Result<ToolOutput, ToolError> {
let start = std::time::Instant::now();
let method = require_str(&params, "method")?;
let method = params
.get("method")
.and_then(|v| v.as_str())
.ok_or_else(|| {
ToolError::InvalidParameters("missing 'method' parameter".to_string())
})?;
let url = require_str(&params, "url")?;
let url = params
.get("url")
.and_then(|v| v.as_str())
.ok_or_else(|| ToolError::InvalidParameters("missing 'url' parameter".to_string()))?;
let parsed_url = validate_url(url)?;
// Parse headers
+31 -17
View File
@@ -18,7 +18,7 @@ use crate::context::{ContextManager, JobContext, JobState};
use crate::db::Database;
use crate::history::SandboxJobRecord;
use crate::orchestrator::job_manager::{ContainerJobManager, JobMode};
use crate::tools::tool::{Tool, ToolError, ToolOutput, require_str};
use crate::tools::tool::{Tool, ToolError, ToolOutput};
/// Tool for creating a new job.
///
@@ -158,18 +158,18 @@ impl CreateJobTool {
});
// Persist the job mode to DB
if mode == JobMode::ClaudeCode
&& let Some(store) = self.store.clone()
{
let job_id_copy = job_id;
tokio::spawn(async move {
if let Err(e) = store
.update_sandbox_job_mode(job_id_copy, "claude_code")
.await
{
tracing::warn!(job_id = %job_id_copy, "Failed to set job mode: {}", e);
}
});
if mode == JobMode::ClaudeCode {
if let Some(store) = self.store.clone() {
let job_id_copy = job_id;
tokio::spawn(async move {
if let Err(e) = store
.update_sandbox_job_mode(job_id_copy, "claude_code")
.await
{
tracing::warn!(job_id = %job_id_copy, "Failed to set job mode: {}", e);
}
});
}
}
// Create the container job with the pre-determined job_id.
@@ -467,9 +467,17 @@ impl Tool for CreateJobTool {
params: serde_json::Value,
ctx: &JobContext,
) -> Result<ToolOutput, ToolError> {
let title = require_str(&params, "title")?;
let title = params
.get("title")
.and_then(|v| v.as_str())
.ok_or_else(|| ToolError::InvalidParameters("missing 'title' parameter".into()))?;
let description = require_str(&params, "description")?;
let description = params
.get("description")
.and_then(|v| v.as_str())
.ok_or_else(|| {
ToolError::InvalidParameters("missing 'description' parameter".into())
})?;
if self.sandbox_enabled() {
let wait = params.get("wait").and_then(|v| v.as_bool()).unwrap_or(true);
@@ -627,7 +635,10 @@ impl Tool for JobStatusTool {
let start = std::time::Instant::now();
let requester_id = ctx.user_id.clone();
let job_id_str = require_str(&params, "job_id")?;
let job_id_str = params
.get("job_id")
.and_then(|v| v.as_str())
.ok_or_else(|| ToolError::InvalidParameters("missing 'job_id' parameter".into()))?;
let job_id = Uuid::parse_str(job_id_str).map_err(|_| {
ToolError::InvalidParameters(format!("invalid job ID format: {}", job_id_str))
@@ -709,7 +720,10 @@ impl Tool for CancelJobTool {
let start = std::time::Instant::now();
let requester_id = ctx.user_id.clone();
let job_id_str = require_str(&params, "job_id")?;
let job_id_str = params
.get("job_id")
.and_then(|v| v.as_str())
.ok_or_else(|| ToolError::InvalidParameters("missing 'job_id' parameter".into()))?;
let job_id = Uuid::parse_str(job_id_str).map_err(|_| {
ToolError::InvalidParameters(format!("invalid job ID format: {}", job_id_str))
+10 -3
View File
@@ -3,7 +3,7 @@
use async_trait::async_trait;
use crate::context::JobContext;
use crate::tools::tool::{Tool, ToolError, ToolOutput, require_param, require_str};
use crate::tools::tool::{Tool, ToolError, ToolOutput};
/// Tool for JSON manipulation (parse, query, transform).
pub struct JsonTool;
@@ -46,9 +46,16 @@ impl Tool for JsonTool {
) -> Result<ToolOutput, ToolError> {
let start = std::time::Instant::now();
let operation = require_str(&params, "operation")?;
let operation = params
.get("operation")
.and_then(|v| v.as_str())
.ok_or_else(|| {
ToolError::InvalidParameters("missing 'operation' parameter".to_string())
})?;
let data = require_param(&params, "data")?;
let data = params
.get("data")
.ok_or_else(|| ToolError::InvalidParameters("missing 'data' parameter".to_string()))?;
let result = match operation {
"parse" => {
+160
View File
@@ -0,0 +1,160 @@
//! NEAR AI Marketplace tool.
use async_trait::async_trait;
use rust_decimal::Decimal;
use crate::context::JobContext;
use crate::tools::tool::{Tool, ToolError, ToolOutput};
/// Tool for interacting with the NEAR AI marketplace.
pub struct MarketplaceTool {
// TODO: Add marketplace client
}
impl MarketplaceTool {
/// Create a new marketplace tool.
pub fn new() -> Self {
Self {}
}
}
impl Default for MarketplaceTool {
fn default() -> Self {
Self::new()
}
}
#[async_trait]
impl Tool for MarketplaceTool {
fn name(&self) -> &str {
"marketplace"
}
fn description(&self) -> &str {
"Interact with the NEAR AI marketplace: search jobs, submit bids, deliver work."
}
fn parameters_schema(&self) -> serde_json::Value {
serde_json::json!({
"type": "object",
"properties": {
"action": {
"type": "string",
"enum": ["search_jobs", "get_job", "submit_bid", "accept_job", "submit_work", "get_status"],
"description": "The marketplace action to perform"
},
"job_id": {
"type": "string",
"description": "Job ID (for get_job, submit_bid, accept_job, submit_work)"
},
"query": {
"type": "string",
"description": "Search query (for search_jobs)"
},
"category": {
"type": "string",
"description": "Job category filter (for search_jobs)"
},
"bid_amount": {
"type": "number",
"description": "Bid amount in NEAR (for submit_bid)"
},
"work_url": {
"type": "string",
"description": "URL to submitted work (for submit_work)"
},
"work_description": {
"type": "string",
"description": "Description of completed work (for submit_work)"
}
},
"required": ["action"]
})
}
async fn execute(
&self,
params: serde_json::Value,
_ctx: &JobContext,
) -> Result<ToolOutput, ToolError> {
let start = std::time::Instant::now();
let action = params
.get("action")
.and_then(|v| v.as_str())
.ok_or_else(|| {
ToolError::InvalidParameters("missing 'action' parameter".to_string())
})?;
// TODO: Implement actual marketplace integration
let result = match action {
"search_jobs" => {
// Placeholder response
serde_json::json!({
"jobs": [],
"total": 0,
"message": "Marketplace integration not yet implemented"
})
}
"get_job" => {
let job_id = params
.get("job_id")
.and_then(|v| v.as_str())
.ok_or_else(|| {
ToolError::InvalidParameters("missing 'job_id' parameter".to_string())
})?;
serde_json::json!({
"job_id": job_id,
"status": "not_found",
"message": "Marketplace integration not yet implemented"
})
}
"submit_bid" => {
serde_json::json!({
"success": false,
"message": "Marketplace integration not yet implemented"
})
}
"accept_job" => {
serde_json::json!({
"success": false,
"message": "Marketplace integration not yet implemented"
})
}
"submit_work" => {
serde_json::json!({
"success": false,
"message": "Marketplace integration not yet implemented"
})
}
"get_status" => {
serde_json::json!({
"connected": false,
"message": "Marketplace integration not yet implemented"
})
}
_ => {
return Err(ToolError::InvalidParameters(format!(
"unknown action: {}",
action
)));
}
};
Ok(ToolOutput::success(result, start.elapsed()))
}
fn estimated_cost(&self, params: &serde_json::Value) -> Option<Decimal> {
// Bidding has a cost
if params.get("action").and_then(|v| v.as_str()) == Some("submit_bid") {
Some(Decimal::new(1, 2)) // 0.01 NEAR gas cost
} else {
None
}
}
fn requires_sanitization(&self) -> bool {
true // External marketplace data
}
}
+15 -4
View File
@@ -17,7 +17,7 @@ use std::sync::Arc;
use async_trait::async_trait;
use crate::context::JobContext;
use crate::tools::tool::{Tool, ToolError, ToolOutput, require_str};
use crate::tools::tool::{Tool, ToolError, ToolOutput};
use crate::workspace::{Workspace, paths};
/// Identity files that the LLM must not overwrite via tool calls.
@@ -81,7 +81,10 @@ impl Tool for MemorySearchTool {
) -> Result<ToolOutput, ToolError> {
let start = std::time::Instant::now();
let query = require_str(&params, "query")?;
let query = params
.get("query")
.and_then(|v| v.as_str())
.ok_or_else(|| ToolError::InvalidParameters("missing 'query' parameter".to_string()))?;
let limit = params
.get("limit")
@@ -173,7 +176,12 @@ impl Tool for MemoryWriteTool {
) -> Result<ToolOutput, ToolError> {
let start = std::time::Instant::now();
let content = require_str(&params, "content")?;
let content = params
.get("content")
.and_then(|v| v.as_str())
.ok_or_else(|| {
ToolError::InvalidParameters("missing 'content' parameter".to_string())
})?;
if content.trim().is_empty() {
return Err(ToolError::InvalidParameters(
@@ -329,7 +337,10 @@ impl Tool for MemoryReadTool {
) -> Result<ToolOutput, ToolError> {
let start = std::time::Instant::now();
let path = require_str(&params, "path")?;
let path = params
.get("path")
.and_then(|v| v.as_str())
.ok_or_else(|| ToolError::InvalidParameters("missing 'path' parameter".to_string()))?;
let doc = self
.workspace
+8 -3
View File
@@ -1,20 +1,22 @@
//! Built-in tools that come with the agent.
mod browser;
mod echo;
mod ecommerce;
pub mod extension_tools;
mod file;
mod http;
mod job;
mod json;
mod marketplace;
mod memory;
mod restaurant;
pub mod routine;
pub(crate) mod shell;
mod taskrabbit;
mod time;
pub use browser::BrowserTool;
pub use browser::session::find_chrome;
pub use echo::EchoTool;
pub use ecommerce::EcommerceTool;
pub use extension_tools::{
ToolActivateTool, ToolAuthTool, ToolInstallTool, ToolListTool, ToolRemoveTool, ToolSearchTool,
};
@@ -22,9 +24,12 @@ pub use file::{ApplyPatchTool, ListDirTool, ReadFileTool, WriteFileTool};
pub use http::HttpTool;
pub use job::{CancelJobTool, CreateJobTool, JobStatusTool, ListJobsTool};
pub use json::JsonTool;
pub use marketplace::MarketplaceTool;
pub use memory::{MemoryReadTool, MemorySearchTool, MemoryTreeTool, MemoryWriteTool};
pub use restaurant::RestaurantTool;
pub use routine::{
RoutineCreateTool, RoutineDeleteTool, RoutineHistoryTool, RoutineListTool, RoutineUpdateTool,
};
pub use shell::ShellTool;
pub use taskrabbit::TaskRabbitTool;
pub use time::TimeTool;
+172
View File
@@ -0,0 +1,172 @@
//! Restaurant reservation tool.
use async_trait::async_trait;
use crate::context::JobContext;
use crate::tools::tool::{Tool, ToolError, ToolOutput};
/// Tool for restaurant reservations (OpenTable, Resy, etc.).
pub struct RestaurantTool {
// TODO: Add reservation API clients
}
impl RestaurantTool {
/// Create a new restaurant tool.
pub fn new() -> Self {
Self {}
}
}
impl Default for RestaurantTool {
fn default() -> Self {
Self::new()
}
}
#[async_trait]
impl Tool for RestaurantTool {
fn name(&self) -> &str {
"restaurant"
}
fn description(&self) -> &str {
"Search restaurants, check availability, and make reservations via OpenTable, Resy, etc."
}
fn parameters_schema(&self) -> serde_json::Value {
serde_json::json!({
"type": "object",
"properties": {
"action": {
"type": "string",
"enum": ["search", "check_availability", "make_reservation", "cancel_reservation", "get_reservation"],
"description": "The restaurant action to perform"
},
"query": {
"type": "string",
"description": "Search query (cuisine type, restaurant name, etc.)"
},
"location": {
"type": "object",
"properties": {
"city": { "type": "string" },
"neighborhood": { "type": "string" },
"latitude": { "type": "number" },
"longitude": { "type": "number" }
},
"description": "Location to search near"
},
"date": {
"type": "string",
"description": "Reservation date (YYYY-MM-DD)"
},
"time": {
"type": "string",
"description": "Preferred time (HH:MM)"
},
"party_size": {
"type": "integer",
"description": "Number of guests"
},
"restaurant_id": {
"type": "string",
"description": "Restaurant ID (for check_availability, make_reservation)"
},
"reservation_id": {
"type": "string",
"description": "Reservation ID (for cancel_reservation, get_reservation)"
},
"guest_name": {
"type": "string",
"description": "Name for the reservation"
},
"guest_phone": {
"type": "string",
"description": "Phone number for the reservation"
},
"guest_email": {
"type": "string",
"description": "Email for the reservation"
}
},
"required": ["action"]
})
}
async fn execute(
&self,
params: serde_json::Value,
_ctx: &JobContext,
) -> Result<ToolOutput, ToolError> {
let start = std::time::Instant::now();
let action = params
.get("action")
.and_then(|v| v.as_str())
.ok_or_else(|| {
ToolError::InvalidParameters("missing 'action' parameter".to_string())
})?;
// TODO: Implement actual restaurant reservation API integrations
let result = match action {
"search" => {
let query = params.get("query").and_then(|v| v.as_str()).unwrap_or("");
serde_json::json!({
"query": query,
"restaurants": [],
"message": "Restaurant integration not yet implemented"
})
}
"check_availability" => {
let restaurant_id = params
.get("restaurant_id")
.and_then(|v| v.as_str())
.ok_or_else(|| {
ToolError::InvalidParameters(
"missing 'restaurant_id' parameter".to_string(),
)
})?;
serde_json::json!({
"restaurant_id": restaurant_id,
"available_times": [],
"message": "Restaurant integration not yet implemented"
})
}
"make_reservation" => {
serde_json::json!({
"success": false,
"message": "Restaurant integration not yet implemented"
})
}
"cancel_reservation" => {
serde_json::json!({
"cancelled": false,
"message": "Restaurant integration not yet implemented"
})
}
"get_reservation" => {
let reservation_id = params.get("reservation_id").and_then(|v| v.as_str());
serde_json::json!({
"reservation_id": reservation_id,
"found": false,
"message": "Restaurant integration not yet implemented"
})
}
_ => {
return Err(ToolError::InvalidParameters(format!(
"unknown action: {}",
action
)));
}
};
Ok(ToolOutput::success(result, start.elapsed()))
}
fn requires_sanitization(&self) -> bool {
true // External restaurant data
}
}
+25 -7
View File
@@ -20,7 +20,7 @@ use crate::agent::routine::{
use crate::agent::routine_engine::RoutineEngine;
use crate::context::JobContext;
use crate::db::Database;
use crate::tools::tool::{Tool, ToolError, ToolOutput, require_str};
use crate::tools::tool::{Tool, ToolError, ToolOutput};
// ==================== routine_create ====================
@@ -106,16 +106,25 @@ impl Tool for RoutineCreateTool {
) -> Result<ToolOutput, ToolError> {
let start = std::time::Instant::now();
let name = require_str(&params, "name")?;
let name = params
.get("name")
.and_then(|v| v.as_str())
.ok_or_else(|| ToolError::InvalidParameters("missing 'name'".to_string()))?;
let description = params
.get("description")
.and_then(|v| v.as_str())
.unwrap_or("");
let trigger_type = require_str(&params, "trigger_type")?;
let trigger_type = params
.get("trigger_type")
.and_then(|v| v.as_str())
.ok_or_else(|| ToolError::InvalidParameters("missing 'trigger_type'".to_string()))?;
let prompt = require_str(&params, "prompt")?;
let prompt = params
.get("prompt")
.and_then(|v| v.as_str())
.ok_or_else(|| ToolError::InvalidParameters("missing 'prompt'".to_string()))?;
// Build trigger
let trigger = match trigger_type {
@@ -399,7 +408,10 @@ impl Tool for RoutineUpdateTool {
) -> Result<ToolOutput, ToolError> {
let start = std::time::Instant::now();
let name = require_str(&params, "name")?;
let name = params
.get("name")
.and_then(|v| v.as_str())
.ok_or_else(|| ToolError::InvalidParameters("missing 'name'".to_string()))?;
let mut routine = self
.store
@@ -502,7 +514,10 @@ impl Tool for RoutineDeleteTool {
) -> Result<ToolOutput, ToolError> {
let start = std::time::Instant::now();
let name = require_str(&params, "name")?;
let name = params
.get("name")
.and_then(|v| v.as_str())
.ok_or_else(|| ToolError::InvalidParameters("missing 'name'".to_string()))?;
let routine = self
.store
@@ -580,7 +595,10 @@ impl Tool for RoutineHistoryTool {
) -> Result<ToolOutput, ToolError> {
let start = std::time::Instant::now();
let name = require_str(&params, "name")?;
let name = params
.get("name")
.and_then(|v| v.as_str())
.ok_or_else(|| ToolError::InvalidParameters("missing 'name'".to_string()))?;
let limit = params
.get("limit")
+11 -50
View File
@@ -30,7 +30,7 @@ use tokio::process::Command;
use crate::context::JobContext;
use crate::sandbox::{SandboxManager, SandboxPolicy};
use crate::tools::tool::{Tool, ToolDomain, ToolError, ToolOutput, require_str};
use crate::tools::tool::{Tool, ToolDomain, ToolError, ToolOutput};
/// Maximum output size before truncation (64KB).
const MAX_OUTPUT_SIZE: usize = 64 * 1024;
@@ -343,12 +343,12 @@ impl ShellTool {
// Use sandbox if configured; fail-closed (never silently fall through
// to unsandboxed execution when sandbox was intended).
if let Some(ref sandbox) = self.sandbox
&& (sandbox.is_initialized() || sandbox.config().enabled)
{
return self
.execute_sandboxed(sandbox, cmd, &cwd, timeout_duration)
.await;
if let Some(ref sandbox) = self.sandbox {
if sandbox.is_initialized() || sandbox.config().enabled {
return self
.execute_sandboxed(sandbox, cmd, &cwd, timeout_duration)
.await;
}
}
// Only execute directly when no sandbox was configured at all.
@@ -401,7 +401,10 @@ impl Tool for ShellTool {
params: serde_json::Value,
_ctx: &JobContext,
) -> Result<ToolOutput, ToolError> {
let command = require_str(&params, "command")?;
let command = params
.get("command")
.and_then(|v| v.as_str())
.ok_or_else(|| ToolError::InvalidParameters("missing 'command' parameter".into()))?;
let workdir = params.get("workdir").and_then(|v| v.as_str());
let timeout = params.get("timeout").and_then(|v| v.as_u64());
@@ -524,48 +527,6 @@ mod tests {
));
}
/// Replicate the extraction logic from agent_loop.rs to prove it works
/// when `arguments` is a `serde_json::Value::Object` (the common case
/// that was previously broken because `Value::Object.as_str()` returns None).
#[test]
fn test_destructive_command_extraction_from_object_args() {
let arguments = serde_json::json!({"command": "rm -rf /tmp/stuff"});
let cmd = arguments
.get("command")
.and_then(|c| c.as_str().map(String::from))
.or_else(|| {
arguments
.as_str()
.and_then(|s| serde_json::from_str::<serde_json::Value>(s).ok())
.and_then(|v| v.get("command").and_then(|c| c.as_str().map(String::from)))
});
assert_eq!(cmd.as_deref(), Some("rm -rf /tmp/stuff"));
assert!(requires_explicit_approval(cmd.as_deref().unwrap()));
}
/// Verify extraction still works when `arguments` is a JSON string
/// (rare, but possible if the LLM provider returns string-encoded JSON).
#[test]
fn test_destructive_command_extraction_from_string_args() {
let arguments =
serde_json::Value::String(r#"{"command": "git push --force origin main"}"#.to_string());
let cmd = arguments
.get("command")
.and_then(|c| c.as_str().map(String::from))
.or_else(|| {
arguments
.as_str()
.and_then(|s| serde_json::from_str::<serde_json::Value>(s).ok())
.and_then(|v| v.get("command").and_then(|c| c.as_str().map(String::from)))
});
assert_eq!(cmd.as_deref(), Some("git push --force origin main"));
assert!(requires_explicit_approval(cmd.as_deref().unwrap()));
}
#[test]
fn test_sandbox_policy_builder() {
let tool = ShellTool::new()
+157
View File
@@ -0,0 +1,157 @@
//! TaskRabbit tool for real-world task delegation.
use async_trait::async_trait;
use rust_decimal::Decimal;
use crate::context::JobContext;
use crate::tools::tool::{Tool, ToolError, ToolOutput};
/// Tool for delegating real-world tasks via TaskRabbit.
pub struct TaskRabbitTool {
// TODO: Add TaskRabbit API client
}
impl TaskRabbitTool {
/// Create a new TaskRabbit tool.
pub fn new() -> Self {
Self {}
}
}
impl Default for TaskRabbitTool {
fn default() -> Self {
Self::new()
}
}
#[async_trait]
impl Tool for TaskRabbitTool {
fn name(&self) -> &str {
"taskrabbit"
}
fn description(&self) -> &str {
"Delegate real-world tasks to TaskRabbit taskers (delivery, assembly, cleaning, etc.)."
}
fn parameters_schema(&self) -> serde_json::Value {
serde_json::json!({
"type": "object",
"properties": {
"action": {
"type": "string",
"enum": ["search_taskers", "get_quote", "book_task", "get_status", "cancel_task"],
"description": "The TaskRabbit action to perform"
},
"task_type": {
"type": "string",
"enum": ["delivery", "assembly", "moving", "cleaning", "handyman", "other"],
"description": "Type of task"
},
"description": {
"type": "string",
"description": "Detailed description of the task"
},
"location": {
"type": "object",
"properties": {
"address": { "type": "string" },
"city": { "type": "string" },
"state": { "type": "string" },
"zip": { "type": "string" }
},
"description": "Location for the task"
},
"scheduled_time": {
"type": "string",
"description": "ISO 8601 datetime for when the task should be performed"
},
"budget": {
"type": "number",
"description": "Maximum budget for the task in USD"
},
"task_id": {
"type": "string",
"description": "Task ID (for get_status, cancel_task)"
}
},
"required": ["action"]
})
}
async fn execute(
&self,
params: serde_json::Value,
_ctx: &JobContext,
) -> Result<ToolOutput, ToolError> {
let start = std::time::Instant::now();
let action = params
.get("action")
.and_then(|v| v.as_str())
.ok_or_else(|| {
ToolError::InvalidParameters("missing 'action' parameter".to_string())
})?;
// TODO: Implement actual TaskRabbit API integration
let result = match action {
"search_taskers" => {
serde_json::json!({
"taskers": [],
"message": "TaskRabbit integration not yet implemented"
})
}
"get_quote" => {
serde_json::json!({
"quotes": [],
"message": "TaskRabbit integration not yet implemented"
})
}
"book_task" => {
serde_json::json!({
"booked": false,
"message": "TaskRabbit integration not yet implemented"
})
}
"get_status" => {
let task_id = params.get("task_id").and_then(|v| v.as_str());
serde_json::json!({
"task_id": task_id,
"status": "unknown",
"message": "TaskRabbit integration not yet implemented"
})
}
"cancel_task" => {
serde_json::json!({
"cancelled": false,
"message": "TaskRabbit integration not yet implemented"
})
}
_ => {
return Err(ToolError::InvalidParameters(format!(
"unknown action: {}",
action
)));
}
};
Ok(ToolOutput::success(result, start.elapsed()))
}
fn estimated_cost(&self, params: &serde_json::Value) -> Option<Decimal> {
// Booking a task has associated costs
if params.get("action").and_then(|v| v.as_str()) == Some("book_task") {
params
.get("budget")
.and_then(|v| v.as_f64())
.map(|b| Decimal::try_from(b).unwrap_or_default())
} else {
None
}
}
fn requires_sanitization(&self) -> bool {
true // External TaskRabbit data
}
}
+25 -5
View File
@@ -4,7 +4,7 @@ use async_trait::async_trait;
use chrono::{DateTime, Utc};
use crate::context::JobContext;
use crate::tools::tool::{Tool, ToolError, ToolOutput, require_str};
use crate::tools::tool::{Tool, ToolError, ToolOutput};
/// Tool for getting current time and date operations.
pub struct TimeTool;
@@ -52,7 +52,12 @@ impl Tool for TimeTool {
) -> Result<ToolOutput, ToolError> {
let start = std::time::Instant::now();
let operation = require_str(&params, "operation")?;
let operation = params
.get("operation")
.and_then(|v| v.as_str())
.ok_or_else(|| {
ToolError::InvalidParameters("missing 'operation' parameter".to_string())
})?;
let result = match operation {
"now" => {
@@ -64,7 +69,12 @@ impl Tool for TimeTool {
})
}
"parse" => {
let timestamp = require_str(&params, "timestamp")?;
let timestamp = params
.get("timestamp")
.and_then(|v| v.as_str())
.ok_or_else(|| {
ToolError::InvalidParameters("missing 'timestamp' parameter".to_string())
})?;
let dt: DateTime<Utc> = timestamp.parse().map_err(|e| {
ToolError::InvalidParameters(format!("invalid timestamp: {}", e))
@@ -77,9 +87,19 @@ impl Tool for TimeTool {
})
}
"diff" => {
let ts1 = require_str(&params, "timestamp")?;
let ts1 = params
.get("timestamp")
.and_then(|v| v.as_str())
.ok_or_else(|| {
ToolError::InvalidParameters("missing 'timestamp' parameter".to_string())
})?;
let ts2 = require_str(&params, "timestamp2")?;
let ts2 = params
.get("timestamp2")
.and_then(|v| v.as_str())
.ok_or_else(|| {
ToolError::InvalidParameters("missing 'timestamp2' parameter".to_string())
})?;
let dt1: DateTime<Utc> = ts1.parse().map_err(|e| {
ToolError::InvalidParameters(format!("invalid timestamp: {}", e))
+72 -15
View File
@@ -11,9 +11,9 @@ use base64::{Engine, engine::general_purpose::URL_SAFE_NO_PAD};
use rand::RngCore;
use serde::{Deserialize, Serialize};
use sha2::{Digest, Sha256};
use tokio::io::{AsyncBufReadExt, AsyncWriteExt, BufReader};
use tokio::net::TcpListener;
use crate::cli::oauth_defaults::{self, OAUTH_CALLBACK_PORT};
use crate::secrets::{CreateSecretParams, SecretsStore};
use crate::tools::mcp::config::McpServerConfig;
@@ -466,12 +466,14 @@ pub async fn authorize_mcp_server(
Ok(token)
}
/// Bind the OAuth callback listener on the shared fixed port.
/// Find an available port for the OAuth callback.
pub async fn find_available_port() -> Result<(TcpListener, u16), AuthError> {
let listener = oauth_defaults::bind_callback_listener()
.await
.map_err(|_| AuthError::PortUnavailable)?;
Ok((listener, OAUTH_CALLBACK_PORT))
for port in 9876..=9886 {
if let Ok(listener) = TcpListener::bind(format!("127.0.0.1:{}", port)).await {
return Ok((listener, port));
}
}
Err(AuthError::PortUnavailable)
}
/// Build the authorization URL with all required parameters.
@@ -520,16 +522,71 @@ pub async fn wait_for_authorization_callback(
listener: TcpListener,
server_name: &str,
) -> Result<String, AuthError> {
oauth_defaults::wait_for_callback(listener, "/callback", "code", server_name)
.await
.map_err(|e| match e {
oauth_defaults::OAuthCallbackError::Denied => AuthError::AuthorizationDenied,
oauth_defaults::OAuthCallbackError::Timeout => AuthError::Timeout,
oauth_defaults::OAuthCallbackError::PortInUse(_, msg) => {
AuthError::Http(format!("Port error: {}", msg))
let timeout = Duration::from_secs(300);
tokio::time::timeout(timeout, async {
loop {
let (mut socket, _) = listener
.accept()
.await
.map_err(|e| AuthError::Http(e.to_string()))?;
let mut reader = BufReader::new(&mut socket);
let mut request_line = String::new();
reader
.read_line(&mut request_line)
.await
.map_err(|e| AuthError::Http(e.to_string()))?;
// Parse GET /callback?code=xxx HTTP/1.1
if let Some(path) = request_line.split_whitespace().nth(1) {
if path.starts_with("/callback") {
if let Some(query) = path.split('?').nth(1) {
// Check for error first
if query.contains("error=") {
let response = "HTTP/1.1 400 Bad Request\r\n\r\nAuthorization denied";
let _ = socket.write_all(response.as_bytes()).await;
return Err(AuthError::AuthorizationDenied);
}
// Look for code
for param in query.split('&') {
let parts: Vec<&str> = param.splitn(2, '=').collect();
if parts.len() == 2 && parts[0] == "code" {
let code = urlencoding::decode(parts[1])
.unwrap_or_else(|_| parts[1].into())
.into_owned();
// Send success response
let response = format!(
"HTTP/1.1 200 OK\r\n\
Content-Type: text/html\r\n\
\r\n\
<!DOCTYPE html><html><body style=\"font-family: sans-serif; \
display: flex; justify-content: center; align-items: center; \
height: 100vh; margin: 0; background: #191919; color: white;\">\
<div style=\"text-align: center;\">\
<h1> {} Connected!</h1>\
<p>You can close this window.</p>\
</div></body></html>",
server_name
);
let _ = socket.write_all(response.as_bytes()).await;
let _ = socket.shutdown().await;
return Ok(code);
}
}
}
}
}
oauth_defaults::OAuthCallbackError::Io(msg) => AuthError::Http(msg),
})
let response = "HTTP/1.1 404 Not Found\r\n\r\n";
let _ = socket.write_all(response.as_bytes()).await;
}
})
.await
.map_err(|_| AuthError::Timeout)?
}
/// Exchange the authorization code for an access token.
+43 -44
View File
@@ -184,46 +184,44 @@ impl McpClient {
}
// Add Mcp-Session-Id header if we have a session
if let Some(ref session_manager) = self.session_manager
&& let Some(session_id) = session_manager.get_session_id(&self.server_name).await
{
req_builder = req_builder.header("Mcp-Session-Id", session_id);
if let Some(ref session_manager) = self.session_manager {
if let Some(session_id) = session_manager.get_session_id(&self.server_name).await {
req_builder = req_builder.header("Mcp-Session-Id", session_id);
}
}
let response = req_builder.send().await.map_err(|e| {
let mut chain = format!("MCP request failed: {}", e);
let mut source = std::error::Error::source(&e);
while let Some(cause) = source {
chain.push_str(&format!(" -> {}", cause));
source = cause.source();
}
ToolError::ExternalService(chain)
})?;
let response = req_builder
.send()
.await
.map_err(|e| ToolError::ExternalService(format!("MCP request failed: {}", e)))?;
// Check for 401 Unauthorized - try to refresh token on first attempt
if response.status() == reqwest::StatusCode::UNAUTHORIZED {
if attempt == 0 {
// Try to refresh the token
if let Some(ref secrets) = self.secrets
&& let Some(ref config) = self.server_config
{
tracing::debug!(
"MCP token expired, attempting refresh for '{}'",
self.server_name
);
match refresh_access_token(config, secrets, &self.user_id).await {
Ok(_) => {
tracing::info!("MCP token refreshed for '{}'", self.server_name);
// Continue to next iteration to retry with new token
continue;
}
Err(e) => {
tracing::debug!(
"Token refresh failed for '{}': {}",
self.server_name,
e
);
// Fall through to return auth error
if let Some(ref secrets) = self.secrets {
if let Some(ref config) = self.server_config {
tracing::debug!(
"MCP token expired, attempting refresh for '{}'",
self.server_name
);
match refresh_access_token(config, secrets, &self.user_id).await {
Ok(_) => {
tracing::info!(
"MCP token refreshed for '{}'",
self.server_name
);
// Continue to next iteration to retry with new token
continue;
}
Err(e) => {
tracing::debug!(
"Token refresh failed for '{}': {}",
self.server_name,
e
);
// Fall through to return auth error
}
}
}
}
@@ -247,15 +245,16 @@ impl McpClient {
/// Parse the HTTP response into an MCP response.
async fn parse_response(&self, response: reqwest::Response) -> Result<McpResponse, ToolError> {
// Extract session ID from response header
if let Some(ref session_manager) = self.session_manager
&& let Some(session_id) = response
if let Some(ref session_manager) = self.session_manager {
if let Some(session_id) = response
.headers()
.get("Mcp-Session-Id")
.and_then(|v| v.to_str().ok())
{
session_manager
.update_session_id(&self.server_name, Some(session_id.to_string()))
.await;
{
session_manager
.update_session_id(&self.server_name, Some(session_id.to_string()))
.await;
}
}
if !response.status().is_success() {
@@ -317,11 +316,11 @@ impl McpClient {
/// This should be called once per session to establish capabilities.
pub async fn initialize(&self) -> Result<InitializeResult, ToolError> {
// Check if already initialized
if let Some(ref session_manager) = self.session_manager
&& session_manager.is_initialized(&self.server_name).await
{
// Return cached/default capabilities
return Ok(InitializeResult::default());
if let Some(ref session_manager) = self.session_manager {
if session_manager.is_initialized(&self.server_name).await {
// Return cached/default capabilities
return Ok(InitializeResult::default());
}
}
// Ensure we have a session
+1 -82
View File
@@ -88,18 +88,8 @@ impl McpServerConfig {
}
/// Check if this server requires authentication.
///
/// Returns true if OAuth is pre-configured OR if this is a remote HTTPS server
/// (which likely supports Dynamic Client Registration even without pre-configured OAuth).
pub fn requires_auth(&self) -> bool {
if self.oauth.is_some() {
return true;
}
// Remote HTTPS servers need auth handling (DCR, token refresh, 401 detection).
// Localhost/127.0.0.1 servers are assumed to be dev servers without auth.
let url_lower = self.url.to_lowercase();
let is_localhost = is_localhost_url(&url_lower);
url_lower.starts_with("https://") && !is_localhost
self.oauth.is_some()
}
/// Get the secret name used to store the access token.
@@ -412,43 +402,11 @@ pub async fn remove_mcp_server_db(
Ok(())
}
/// Check if a URL points to a loopback address (localhost, 127.0.0.1, [::1]).
///
/// Uses `url::Url` for proper parsing so edge cases (IPv6, userinfo, ports)
/// are handled correctly without manual string splitting.
fn is_localhost_url(url: &str) -> bool {
let Ok(parsed) = url::Url::parse(url) else {
return false;
};
match parsed.host() {
Some(url::Host::Domain(d)) => d.eq_ignore_ascii_case("localhost"),
Some(url::Host::Ipv4(ip)) => ip.is_loopback(),
Some(url::Host::Ipv6(ip)) => ip.is_loopback(),
None => false,
}
}
#[cfg(test)]
mod tests {
use super::*;
use tempfile::tempdir;
#[test]
fn test_is_localhost_url() {
assert!(is_localhost_url("http://localhost:3000/path"));
assert!(is_localhost_url("https://localhost/path"));
assert!(is_localhost_url("http://127.0.0.1:8080"));
assert!(is_localhost_url("http://127.0.0.1"));
assert!(!is_localhost_url("https://notlocalhost.com/path"));
assert!(!is_localhost_url("https://example-localhost.io"));
assert!(!is_localhost_url("https://mcp.notion.com"));
assert!(is_localhost_url("http://user:pass@localhost:3000/path"));
// IPv6 loopback
assert!(is_localhost_url("http://[::1]:8080/path"));
assert!(is_localhost_url("http://[::1]/path"));
assert!(!is_localhost_url("http://[::2]:8080/path"));
}
#[test]
fn test_server_config_validation() {
// Valid HTTPS server
@@ -556,43 +514,4 @@ mod tests {
"mcp_notion_refresh_token"
);
}
#[test]
fn test_requires_auth_with_oauth() {
let config = McpServerConfig::new("notion", "https://mcp.notion.com")
.with_oauth(OAuthConfig::new("client-123"));
assert!(config.requires_auth());
}
#[test]
fn test_requires_auth_remote_https_without_oauth() {
// Remote HTTPS servers need auth even without pre-configured OAuth (DCR)
let config = McpServerConfig::new("github-copilot", "https://api.githubcopilot.com/mcp/");
assert!(config.requires_auth());
let config = McpServerConfig::new("notion", "https://mcp.notion.com");
assert!(config.requires_auth());
}
#[test]
fn test_requires_auth_localhost_no_auth() {
// Localhost servers are dev servers, no auth needed
let config = McpServerConfig::new("local", "http://localhost:8080");
assert!(!config.requires_auth());
let config = McpServerConfig::new("local", "http://127.0.0.1:3000/mcp");
assert!(!config.requires_auth());
// Even HTTPS localhost doesn't require auth
let config = McpServerConfig::new("local", "https://localhost:8443");
assert!(!config.requires_auth());
}
#[test]
fn test_requires_auth_http_remote_no_auth() {
// HTTP remote servers won't pass validation, but if they existed
// they wouldn't trigger HTTPS auth detection
let config = McpServerConfig::new("bad", "http://mcp.example.com");
assert!(!config.requires_auth());
}
}
+11 -25
View File
@@ -11,18 +11,17 @@ use crate::extensions::ExtensionManager;
use crate::llm::{LlmProvider, ToolDefinition};
use crate::orchestrator::job_manager::ContainerJobManager;
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::{
Capabilities, OAuthRefreshConfig, ResourceLimits, WasmError, WasmStorageError, WasmToolRuntime,
WasmToolStore, WasmToolWrapper,
Capabilities, ResourceLimits, WasmError, WasmStorageError, WasmToolRuntime, WasmToolStore,
WasmToolWrapper,
};
use crate::workspace::Workspace;
@@ -97,10 +96,10 @@ impl ToolRegistry {
if let Ok(mut tools) = self.tools.try_write() {
tools.insert(name.clone(), tool);
// Mark as built-in so it can't be shadowed later
if PROTECTED_TOOL_NAMES.contains(&name.as_str())
&& let Ok(mut builtins) = self.builtin_names.try_write()
{
builtins.insert(name.clone());
if PROTECTED_TOOL_NAMES.contains(&name.as_str()) {
if let Ok(mut builtins) = self.builtin_names.try_write() {
builtins.insert(name.clone());
}
}
tracing::debug!("Registered tool: {}", name);
}
@@ -218,9 +217,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.
@@ -368,12 +366,6 @@ impl ToolRegistry {
if let Some(s) = reg.schema {
wrapper = wrapper.with_schema(s);
}
if let Some(store) = reg.secrets_store {
wrapper = wrapper.with_secrets_store(store);
}
if let Some(oauth) = reg.oauth_refresh {
wrapper = wrapper.with_oauth_refresh(oauth);
}
// Register the tool
self.register(Arc::new(wrapper)).await;
@@ -429,8 +421,6 @@ impl ToolRegistry {
limits: None,
description: Some(&tool_with_binary.tool.description),
schema: Some(tool_with_binary.tool.parameters_schema.clone()),
secrets_store: None,
oauth_refresh: None,
})
.await
.map_err(WasmRegistrationError::Wasm)?;
@@ -472,10 +462,6 @@ pub struct WasmToolRegistration<'a> {
pub description: Option<&'a str>,
/// Optional parameter schema override.
pub schema: Option<serde_json::Value>,
/// Secrets store for credential injection at request time.
pub secrets_store: Option<Arc<dyn SecretsStore + Send + Sync>>,
/// OAuth refresh configuration for auto-refreshing expired tokens.
pub oauth_refresh: Option<OAuthRefreshConfig>,
}
impl Default for ToolRegistry {
+6 -59
View File
@@ -199,28 +199,6 @@ pub trait Tool: Send + Sync {
}
}
/// Extract a required string parameter from a JSON object.
///
/// Returns `ToolError::InvalidParameters` if the key is missing or not a string.
pub fn require_str<'a>(params: &'a serde_json::Value, name: &str) -> Result<&'a str, ToolError> {
params
.get(name)
.and_then(|v| v.as_str())
.ok_or_else(|| ToolError::InvalidParameters(format!("missing '{}' parameter", name)))
}
/// Extract a required parameter of any type from a JSON object.
///
/// Returns `ToolError::InvalidParameters` if the key is missing.
pub fn require_param<'a>(
params: &'a serde_json::Value,
name: &str,
) -> Result<&'a serde_json::Value, ToolError> {
params
.get(name)
.ok_or_else(|| ToolError::InvalidParameters(format!("missing '{}' parameter", name)))
}
#[cfg(test)]
mod tests {
use super::*;
@@ -257,7 +235,12 @@ mod tests {
params: serde_json::Value,
_ctx: &JobContext,
) -> Result<ToolOutput, ToolError> {
let message = require_str(&params, "message")?;
let message = params
.get("message")
.and_then(|v| v.as_str())
.ok_or_else(|| {
ToolError::InvalidParameters("missing 'message' parameter".to_string())
})?;
Ok(ToolOutput::text(message, Duration::from_millis(1)))
}
@@ -294,40 +277,4 @@ mod tests {
let tool = EchoTool;
assert_eq!(tool.execution_timeout(), Duration::from_secs(60));
}
#[test]
fn test_require_str_present() {
let params = serde_json::json!({"name": "alice"});
assert_eq!(require_str(&params, "name").unwrap(), "alice");
}
#[test]
fn test_require_str_missing() {
let params = serde_json::json!({});
let err = require_str(&params, "name").unwrap_err();
assert!(err.to_string().contains("missing 'name'"));
}
#[test]
fn test_require_str_wrong_type() {
let params = serde_json::json!({"name": 42});
let err = require_str(&params, "name").unwrap_err();
assert!(err.to_string().contains("missing 'name'"));
}
#[test]
fn test_require_param_present() {
let params = serde_json::json!({"data": [1, 2, 3]});
assert_eq!(
require_param(&params, "data").unwrap(),
&serde_json::json!([1, 2, 3])
);
}
#[test]
fn test_require_param_missing() {
let params = serde_json::json!({});
let err = require_param(&params, "data").unwrap_err();
assert!(err.to_string().contains("missing 'data'"));
}
}
+15 -16
View File
@@ -209,10 +209,10 @@ impl EndpointPattern {
}
// Check path prefix
if let Some(ref prefix) = self.path_prefix
&& !url_path.starts_with(prefix)
{
return false;
if let Some(ref prefix) = self.path_prefix {
if !url_path.starts_with(prefix) {
return false;
}
}
// Check method
@@ -237,14 +237,13 @@ impl EndpointPattern {
}
// Support wildcard: *.example.com matches sub.example.com
if let Some(suffix) = self.host.strip_prefix("*.")
&& url_host.ends_with(suffix)
&& url_host.len() > suffix.len()
{
// Ensure there's a dot before the suffix (or it's the whole thing)
let prefix = &url_host[..url_host.len() - suffix.len()];
if prefix.ends_with('.') || prefix.is_empty() {
return true;
if let Some(suffix) = self.host.strip_prefix("*.") {
if url_host.ends_with(suffix) && url_host.len() > suffix.len() {
// Ensure there's a dot before the suffix (or it's the whole thing)
let prefix = &url_host[..url_host.len() - suffix.len()];
if prefix.ends_with('.') || prefix.is_empty() {
return true;
}
}
}
@@ -292,10 +291,10 @@ impl SecretsCapability {
if pattern == name {
return true;
}
if let Some(prefix) = pattern.strip_suffix('*')
&& name.starts_with(prefix)
{
return true;
if let Some(prefix) = pattern.strip_suffix('*') {
if name.starts_with(prefix) {
return true;
}
}
}
false
+12 -13
View File
@@ -158,10 +158,10 @@ impl CredentialInjector {
if pattern == name {
return true;
}
if let Some(prefix) = pattern.strip_suffix('*')
&& name.starts_with(prefix)
{
return true;
if let Some(prefix) = pattern.strip_suffix('*') {
if name.starts_with(prefix) {
return true;
}
}
}
false
@@ -169,7 +169,7 @@ impl CredentialInjector {
}
/// Inject a single credential into the result.
pub(crate) fn inject_credential(
fn inject_credential(
result: &mut InjectedCredentials,
location: &CredentialLocation,
secret: &DecryptedSecret,
@@ -208,19 +208,18 @@ pub(crate) fn inject_credential(
}
/// Check if a host matches a pattern (supports wildcards).
pub(crate) fn host_matches_pattern(host: &str, pattern: &str) -> bool {
fn host_matches_pattern(host: &str, pattern: &str) -> bool {
if pattern == host {
return true;
}
// Support wildcard: *.example.com matches sub.example.com
if let Some(suffix) = pattern.strip_prefix("*.")
&& host.ends_with(suffix)
&& host.len() > suffix.len()
{
let prefix = &host[..host.len() - suffix.len()];
if prefix.ends_with('.') || prefix.is_empty() {
return true;
if let Some(suffix) = pattern.strip_prefix("*.") {
if host.ends_with(suffix) && host.len() > suffix.len() {
let prefix = &host[..host.len() - suffix.len()];
if prefix.ends_with('.') || prefix.is_empty() {
return true;
}
}
}
+7 -179
View File
@@ -39,11 +39,10 @@ use std::sync::Arc;
use tokio::fs;
use crate::secrets::SecretsStore;
use crate::tools::registry::{ToolRegistry, WasmRegistrationError, WasmToolRegistration};
use crate::tools::wasm::capabilities_schema::CapabilitiesFile;
use crate::tools::wasm::{
Capabilities, OAuthRefreshConfig, WasmError, WasmStorageError, WasmToolRuntime, WasmToolStore,
Capabilities, WasmError, WasmStorageError, WasmToolRuntime, WasmToolStore,
};
/// Error during WASM tool loading.
@@ -78,23 +77,12 @@ pub enum WasmLoadError {
pub struct WasmToolLoader {
runtime: Arc<WasmToolRuntime>,
registry: Arc<ToolRegistry>,
secrets_store: Option<Arc<dyn SecretsStore + Send + Sync>>,
}
impl WasmToolLoader {
/// Create a new loader with the given runtime and registry.
pub fn new(runtime: Arc<WasmToolRuntime>, registry: Arc<ToolRegistry>) -> Self {
Self {
runtime,
registry,
secrets_store: None,
}
}
/// Set the secrets store for credential injection in WASM tools.
pub fn with_secrets_store(mut self, store: Arc<dyn SecretsStore + Send + Sync>) -> Self {
self.secrets_store = Some(store);
self
Self { runtime, registry }
}
/// Load a single WASM tool from a file pair.
@@ -120,24 +108,22 @@ impl WasmToolLoader {
}
let wasm_bytes = fs::read(wasm_path).await?;
// Read capabilities (optional) and extract OAuth refresh config
let (capabilities, oauth_refresh) = if let Some(cap_path) = capabilities_path {
// Read capabilities (optional)
let capabilities = if let Some(cap_path) = capabilities_path {
if cap_path.exists() {
let cap_bytes = fs::read(cap_path).await?;
let cap_file = CapabilitiesFile::from_bytes(&cap_bytes)
.map_err(|e| WasmLoadError::InvalidCapabilities(e.to_string()))?;
let caps = cap_file.to_capabilities();
let oauth = resolve_oauth_refresh_config(&cap_file);
(caps, oauth)
cap_file.to_capabilities()
} else {
tracing::warn!(
path = %cap_path.display(),
"Capabilities file not found, using default (no permissions)"
);
(Capabilities::default(), None)
Capabilities::default()
}
} else {
(Capabilities::default(), None)
Capabilities::default()
};
// Register the tool
@@ -150,8 +136,6 @@ impl WasmToolLoader {
limits: None,
description: None,
schema: None,
secrets_store: self.secrets_store.clone(),
oauth_refresh,
})
.await?;
@@ -309,50 +293,6 @@ impl WasmToolLoader {
}
}
/// Extract OAuth refresh configuration from a parsed capabilities file.
///
/// Returns `None` if there's no `auth.oauth` section or if the client_id
/// can't be resolved from any source (inline, env var, or built-in defaults).
///
/// Fallback chain for client_id:
/// `oauth.client_id` > env var (`oauth.client_id_env`) > `builtin_credentials()`
fn resolve_oauth_refresh_config(cap_file: &CapabilitiesFile) -> Option<OAuthRefreshConfig> {
let auth = cap_file.auth.as_ref()?;
let oauth = auth.oauth.as_ref()?;
let builtin = crate::cli::oauth_defaults::builtin_credentials(&auth.secret_name);
let client_id = oauth
.client_id
.clone()
.or_else(|| {
oauth
.client_id_env
.as_ref()
.and_then(|env| std::env::var(env).ok())
})
.or_else(|| builtin.as_ref().map(|c| c.client_id.to_string()))?;
let client_secret = oauth
.client_secret
.clone()
.or_else(|| {
oauth
.client_secret_env
.as_ref()
.and_then(|env| std::env::var(env).ok())
})
.or_else(|| builtin.as_ref().map(|c| c.client_secret.to_string()));
Some(OAuthRefreshConfig {
token_url: oauth.token_url.clone(),
client_id,
client_secret,
secret_name: auth.secret_name.clone(),
provider: auth.provider.clone(),
})
}
/// Results from loading multiple tools.
#[derive(Debug, Default)]
pub struct LoadResults {
@@ -678,116 +618,4 @@ mod tests {
);
}
}
#[test]
fn test_resolve_oauth_refresh_config_with_oauth() {
use crate::tools::wasm::capabilities_schema::{
AuthCapabilitySchema, CapabilitiesFile, OAuthConfigSchema,
};
let caps = CapabilitiesFile {
auth: Some(AuthCapabilitySchema {
secret_name: "google_oauth_token".to_string(),
provider: Some("google".to_string()),
oauth: Some(OAuthConfigSchema {
authorization_url: "https://accounts.google.com/o/oauth2/v2/auth".to_string(),
token_url: "https://oauth2.googleapis.com/token".to_string(),
client_id: Some("test-client-id".to_string()),
client_secret: Some("test-client-secret".to_string()),
..Default::default()
}),
..Default::default()
}),
..Default::default()
};
let config = super::resolve_oauth_refresh_config(&caps);
assert!(config.is_some());
let config = config.unwrap();
assert_eq!(config.token_url, "https://oauth2.googleapis.com/token");
assert_eq!(config.client_id, "test-client-id");
assert_eq!(config.client_secret, Some("test-client-secret".to_string()));
assert_eq!(config.secret_name, "google_oauth_token");
assert_eq!(config.provider, Some("google".to_string()));
}
#[test]
fn test_resolve_oauth_refresh_config_no_auth() {
use crate::tools::wasm::capabilities_schema::CapabilitiesFile;
let caps = CapabilitiesFile::default();
let config = super::resolve_oauth_refresh_config(&caps);
assert!(config.is_none());
}
#[test]
fn test_resolve_oauth_refresh_config_no_oauth() {
use crate::tools::wasm::capabilities_schema::{AuthCapabilitySchema, CapabilitiesFile};
let caps = CapabilitiesFile {
auth: Some(AuthCapabilitySchema {
secret_name: "manual_token".to_string(),
..Default::default()
}),
..Default::default()
};
let config = super::resolve_oauth_refresh_config(&caps);
assert!(config.is_none());
}
#[test]
fn test_resolve_oauth_refresh_config_no_client_id() {
use crate::tools::wasm::capabilities_schema::{
AuthCapabilitySchema, CapabilitiesFile, OAuthConfigSchema,
};
// A non-Google provider with no client_id anywhere should return None
let caps = CapabilitiesFile {
auth: Some(AuthCapabilitySchema {
secret_name: "unknown_provider_token".to_string(),
oauth: Some(OAuthConfigSchema {
authorization_url: "https://example.com/auth".to_string(),
token_url: "https://example.com/token".to_string(),
// No client_id, no client_id_env, no builtin
..Default::default()
}),
..Default::default()
}),
..Default::default()
};
let config = super::resolve_oauth_refresh_config(&caps);
assert!(config.is_none());
}
#[test]
fn test_resolve_oauth_refresh_config_builtin_google() {
use crate::tools::wasm::capabilities_schema::{
AuthCapabilitySchema, CapabilitiesFile, OAuthConfigSchema,
};
// google_oauth_token should fall back to built-in credentials
let caps = CapabilitiesFile {
auth: Some(AuthCapabilitySchema {
secret_name: "google_oauth_token".to_string(),
provider: Some("google".to_string()),
oauth: Some(OAuthConfigSchema {
authorization_url: "https://accounts.google.com/o/oauth2/v2/auth".to_string(),
token_url: "https://oauth2.googleapis.com/token".to_string(),
// No inline client_id, should fall back to builtin
..Default::default()
}),
..Default::default()
}),
..Default::default()
};
let config = super::resolve_oauth_refresh_config(&caps);
assert!(config.is_some());
let config = config.unwrap();
assert!(!config.client_id.is_empty());
assert!(config.client_secret.is_some());
}
}
+1 -1
View File
@@ -94,7 +94,7 @@ pub use limits::{
WasmResourceLimiter,
};
pub use runtime::{PreparedModule, WasmRuntimeConfig, WasmToolRuntime};
pub use wrapper::{OAuthRefreshConfig, WasmToolWrapper};
pub use wrapper::WasmToolWrapper;
// Capabilities (V2)
pub use capabilities::{
-1
View File
@@ -953,7 +953,6 @@ fn libsql_row_to_tool_with_offset(row: &libsql::Row) -> Result<StoredWasmTool, W
}
#[cfg(feature = "libsql")]
#[allow(clippy::too_many_arguments)]
fn libsql_row_to_tool_at(
row: &libsql::Row,
id_idx: i32,
+21 -884
View File
File diff suppressed because it is too large Load Diff
-319
View File
@@ -1,319 +0,0 @@
//! Truncating terminal writer for tracing.
//!
//! Tracing events from LLM providers can dump 10KB+ JSON bodies to stderr.
//! Rather than truncating at every call site (fragile, easy to miss), we
//! handle it at the writer level: the fmt layer gets a `TruncatingStderr`
//! that caps each event before flushing, while the web gateway `WebLogLayer`
//! still sees the full, untruncated content.
//!
//! ```text
//! tracing::debug!("body: {huge_json}")
//! |
//! v
//! tracing_subscriber::registry()
//! |
//! +-- fmt::layer().with_writer(TruncatingStderr) <-- caps at 500B
//! | \-- stderr (truncated)
//! |
//! \-- WebLogLayer (unchanged)
//! \-- SSE broadcast (full)
//! ```
use std::io::{self, Write};
use tracing_subscriber::fmt::MakeWriter;
/// Maximum bytes per tracing event written to the terminal.
const TERMINAL_MAX_EVENT_BYTES: usize = 500;
/// A `MakeWriter` that creates per-event buffers which truncate on flush.
///
/// Each call to `make_writer()` returns an `EventBuffer`. All `write()`
/// calls accumulate into the buffer. When the buffer drops (after the fmt
/// layer finishes writing one event), it flushes to stderr, truncating if
/// the total exceeds `TERMINAL_MAX_EVENT_BYTES`.
#[derive(Clone)]
pub struct TruncatingStderr {
max_bytes: usize,
}
impl Default for TruncatingStderr {
fn default() -> Self {
Self {
max_bytes: TERMINAL_MAX_EVENT_BYTES,
}
}
}
impl TruncatingStderr {
#[cfg(test)]
fn with_max_bytes(max_bytes: usize) -> Self {
Self { max_bytes }
}
}
impl<'a> MakeWriter<'a> for TruncatingStderr {
type Writer = EventBuffer;
fn make_writer(&'a self) -> Self::Writer {
EventBuffer {
buf: Vec::with_capacity(256),
max_bytes: self.max_bytes,
#[cfg(test)]
sink: None,
}
}
}
/// Per-event buffer that truncates on drop.
pub struct EventBuffer {
buf: Vec<u8>,
max_bytes: usize,
/// Test-only: capture output instead of writing to stderr.
#[cfg(test)]
sink: Option<std::sync::Arc<std::sync::Mutex<Vec<u8>>>>,
}
impl Write for EventBuffer {
fn write(&mut self, data: &[u8]) -> io::Result<usize> {
self.buf.extend_from_slice(data);
Ok(data.len())
}
fn flush(&mut self) -> io::Result<()> {
Ok(())
}
}
/// Find the last valid UTF-8 char boundary at or before `pos` in `bytes`.
///
/// Walks backwards from `pos` until we find a byte that isn't a UTF-8
/// continuation byte (0x80..0xBF). Returns 0 if the entire prefix is
/// somehow invalid (shouldn't happen with valid UTF-8 input from tracing).
fn utf8_floor(bytes: &[u8], pos: usize) -> usize {
let mut i = pos;
// UTF-8 continuation bytes have the form 10xxxxxx (0x80..0xBF).
// Walk backwards past them to find the start of the last character.
while i > 0 && bytes[i] & 0xC0 == 0x80 {
i -= 1;
}
i
}
impl Drop for EventBuffer {
fn drop(&mut self) {
if self.buf.is_empty() {
return;
}
let output = if self.buf.len() <= self.max_bytes {
&self.buf[..]
} else {
// Truncate at a UTF-8 safe boundary
let cut = utf8_floor(&self.buf, self.max_bytes);
let suffix = format!("...[{}B total]\n", self.buf.len());
let mut truncated = Vec::with_capacity(cut + suffix.len());
// Strip trailing newline from the cut portion (we add our own via suffix)
let cut_slice = &self.buf[..cut];
let trimmed = if cut_slice.last() == Some(&b'\n') {
&cut_slice[..cut_slice.len() - 1]
} else {
cut_slice
};
truncated.extend_from_slice(trimmed);
truncated.extend_from_slice(suffix.as_bytes());
#[cfg(test)]
if let Some(ref sink) = self.sink {
let mut s = sink.lock().expect("test sink lock poisoned");
s.extend_from_slice(&truncated);
return;
}
let _ = io::stderr().write_all(&truncated);
return;
};
#[cfg(test)]
if let Some(ref sink) = self.sink {
let mut s = sink.lock().expect("test sink lock poisoned");
s.extend_from_slice(output);
return;
}
let _ = io::stderr().write_all(output);
}
}
#[cfg(test)]
mod tests {
use std::sync::{Arc, Mutex};
use crate::tracing_fmt::{EventBuffer, TruncatingStderr, utf8_floor};
use std::io::Write;
/// Helper: create an EventBuffer that captures output to a shared Vec
/// instead of writing to stderr.
fn test_buffer(max_bytes: usize) -> (EventBuffer, Arc<Mutex<Vec<u8>>>) {
let sink = Arc::new(Mutex::new(Vec::new()));
let buf = EventBuffer {
buf: Vec::new(),
max_bytes,
sink: Some(Arc::clone(&sink)),
};
(buf, sink)
}
#[test]
fn test_short_event_not_truncated() {
let (mut buf, sink) = test_buffer(500);
buf.write_all(b"hello world\n").unwrap();
drop(buf);
let output = sink.lock().unwrap();
assert_eq!(&*output, b"hello world\n");
}
#[test]
fn test_long_event_truncated() {
let (mut buf, sink) = test_buffer(20);
let data = "abcdefghijklmnopqrstuvwxyz0123456789\n";
buf.write_all(data.as_bytes()).unwrap();
let total = data.len();
drop(buf);
let output = sink.lock().unwrap();
let output_str = String::from_utf8_lossy(&output);
// Should contain the suffix with total byte count
assert!(
output_str.contains(&format!("...[{}B total]", total)),
"expected truncation suffix, got: {}",
output_str
);
// Should be shorter than the original
assert!(output.len() < total);
}
#[test]
fn test_utf8_boundary_safe() {
// "Helloé" = [72, 101, 108, 108, 111, 195, 169]
// ^-- 2-byte UTF-8 char
// If we truncate at 6 bytes, we'd land in the middle of 'é'.
// utf8_floor should back up to byte 5 (start of 'é' = 195).
let (mut buf, sink) = test_buffer(6);
let data = "Helloé world";
buf.write_all(data.as_bytes()).unwrap();
drop(buf);
let output = sink.lock().unwrap();
let output_str = String::from_utf8(output.clone());
assert!(
output_str.is_ok(),
"output should be valid UTF-8, got bytes: {:?}",
&*output
);
let s = output_str.unwrap();
assert!(
s.contains("...["),
"should be truncated with suffix, got: {}",
s
);
// The truncated prefix must be valid UTF-8 up to the cut point.
// "Hello" (5 bytes) is the last valid cut before the 2-byte é.
assert!(
s.starts_with("Hello"),
"should start with 'Hello', got: {}",
s
);
}
#[test]
fn test_utf8_floor_basic() {
// ASCII: every byte is a valid boundary
assert_eq!(utf8_floor(b"hello", 3), 3);
// 2-byte UTF-8 char é = [0xC3, 0xA9]
// Landing on the continuation byte (0xA9) should back up to 0xC3
let bytes = "".as_bytes(); // [72, 0xC3, 0xA9]
assert_eq!(utf8_floor(bytes, 2), 1); // backs up to start of é
// 3-byte UTF-8 char (e.g. あ = [0xE3, 0x81, 0x82])
let bytes = "aあ".as_bytes(); // [97, 0xE3, 0x81, 0x82]
assert_eq!(utf8_floor(bytes, 2), 1); // backs up past continuation to 0xE3
assert_eq!(utf8_floor(bytes, 3), 1); // same: 0x82 is continuation, 0x81 is too
}
#[test]
fn test_multiple_writes_accumulated() {
let (mut buf, sink) = test_buffer(500);
buf.write_all(b"hello ").unwrap();
buf.write_all(b"world\n").unwrap();
drop(buf);
let output = sink.lock().unwrap();
assert_eq!(&*output, b"hello world\n");
}
#[test]
fn test_empty_buffer_no_output() {
let (_buf, sink) = test_buffer(500);
// drop without writing
drop(_buf);
let output = sink.lock().unwrap();
assert!(output.is_empty());
}
#[test]
fn test_default_max_bytes() {
let writer = TruncatingStderr::default();
assert_eq!(writer.max_bytes, 500);
}
#[test]
fn test_custom_max_bytes() {
let writer = TruncatingStderr::with_max_bytes(100);
assert_eq!(writer.max_bytes, 100);
}
#[test]
fn test_exactly_at_limit_not_truncated() {
let (mut buf, sink) = test_buffer(5);
buf.write_all(b"hello").unwrap();
drop(buf);
let output = sink.lock().unwrap();
assert_eq!(&*output, b"hello");
}
#[test]
fn test_one_over_limit_truncated() {
let (mut buf, sink) = test_buffer(5);
buf.write_all(b"hello!").unwrap();
drop(buf);
let output = sink.lock().unwrap();
let s = String::from_utf8_lossy(&output);
assert!(s.contains("...[6B total]"), "got: {}", s);
}
#[test]
fn test_4byte_utf8_boundary() {
// 4-byte UTF-8 char: 𝄞 (musical symbol) = [0xF0, 0x9D, 0x84, 0x9E]
let data = "AB𝄞CD";
// bytes: [65, 66, 0xF0, 0x9D, 0x84, 0x9E, 67, 68]
// Truncating at byte 4 lands in the middle of the 4-byte char
let (mut buf, sink) = test_buffer(4);
buf.write_all(data.as_bytes()).unwrap();
drop(buf);
let output = sink.lock().unwrap();
let s = String::from_utf8(output.clone());
assert!(s.is_ok(), "output must be valid UTF-8, got: {:?}", &*output);
let s = s.unwrap();
// Should back up to byte 2 (just "AB"), since bytes 2..5 are all part of 𝄞
assert!(s.starts_with("AB"), "expected 'AB', got: {}", s);
assert!(s.contains("...["), "should be truncated, got: {}", s);
}
}
+70 -54
View File
@@ -129,15 +129,11 @@ impl WorkerHttpClient {
format!("{}/worker/{}/{}", self.orchestrator_url, self.job_id, path)
}
/// Send a GET request, check the status, and deserialize the JSON body.
async fn get_json<T: serde::de::DeserializeOwned>(
&self,
path: &str,
context: &str,
) -> Result<T, WorkerError> {
/// Fetch the job description from the orchestrator.
pub async fn get_job(&self) -> Result<JobDescription, WorkerError> {
let resp = self
.client
.get(self.url(path))
.get(self.url("job"))
.bearer_auth(&self.token)
.send()
.await
@@ -149,51 +145,15 @@ impl WorkerHttpClient {
if !resp.status().is_success() {
return Err(WorkerError::OrchestratorRejected {
job_id: self.job_id,
reason: format!("{} returned {}", context, resp.status()),
reason: format!("GET /job returned {}", resp.status()),
});
}
resp.json().await.map_err(|e| WorkerError::LlmProxyFailed {
reason: format!("{}: failed to parse response: {}", context, e),
reason: format!("failed to parse job description: {}", e),
})
}
/// Send a POST request with a JSON body, check the status, and deserialize the response.
async fn post_json<B: Serialize, T: serde::de::DeserializeOwned>(
&self,
path: &str,
body: &B,
context: &str,
) -> Result<T, WorkerError> {
let resp = self
.client
.post(self.url(path))
.bearer_auth(&self.token)
.json(body)
.send()
.await
.map_err(|e| WorkerError::LlmProxyFailed {
reason: format!("{}: {}", context, e),
})?;
if !resp.status().is_success() {
let status = resp.status();
let body = resp.text().await.unwrap_or_default();
return Err(WorkerError::LlmProxyFailed {
reason: format!("{}: orchestrator returned {}: {}", context, status, body),
});
}
resp.json().await.map_err(|e| WorkerError::LlmProxyFailed {
reason: format!("{}: failed to parse response: {}", context, e),
})
}
/// Fetch the job description from the orchestrator.
pub async fn get_job(&self) -> Result<JobDescription, WorkerError> {
self.get_json("job", "GET /job").await
}
/// Proxy an LLM completion request through the orchestrator.
pub async fn llm_complete(
&self,
@@ -206,9 +166,29 @@ impl WorkerHttpClient {
stop_sequences: request.stop_sequences.clone(),
};
let proxy_resp: ProxyCompletionResponse = self
.post_json("llm/complete", &proxy_req, "LLM complete")
.await?;
let resp = self
.client
.post(self.url("llm/complete"))
.bearer_auth(&self.token)
.json(&proxy_req)
.send()
.await
.map_err(|e| WorkerError::LlmProxyFailed {
reason: e.to_string(),
})?;
if !resp.status().is_success() {
let status = resp.status();
let body = resp.text().await.unwrap_or_default();
return Err(WorkerError::LlmProxyFailed {
reason: format!("orchestrator returned {}: {}", status, body),
});
}
let proxy_resp: ProxyCompletionResponse =
resp.json().await.map_err(|e| WorkerError::LlmProxyFailed {
reason: format!("failed to parse LLM response: {}", e),
})?;
Ok(CompletionResponse {
content: proxy_resp.content,
@@ -232,9 +212,29 @@ impl WorkerHttpClient {
tool_choice: request.tool_choice.clone(),
};
let proxy_resp: ProxyToolCompletionResponse = self
.post_json("llm/complete_with_tools", &proxy_req, "LLM tool complete")
.await?;
let resp = self
.client
.post(self.url("llm/complete_with_tools"))
.bearer_auth(&self.token)
.json(&proxy_req)
.send()
.await
.map_err(|e| WorkerError::LlmProxyFailed {
reason: e.to_string(),
})?;
if !resp.status().is_success() {
let status = resp.status();
let body = resp.text().await.unwrap_or_default();
return Err(WorkerError::LlmProxyFailed {
reason: format!("orchestrator returned {}: {}", status, body),
});
}
let proxy_resp: ProxyToolCompletionResponse =
resp.json().await.map_err(|e| WorkerError::LlmProxyFailed {
reason: format!("failed to parse tool completion response: {}", e),
})?;
Ok(ToolCompletionResponse {
content: proxy_resp.content,
@@ -337,9 +337,25 @@ impl WorkerHttpClient {
/// Signal job completion to the orchestrator.
pub async fn report_complete(&self, report: &CompletionReport) -> Result<(), WorkerError> {
let _: serde_json::Value = self
.post_json("complete", report, "report complete")
.await?;
let resp = self
.client
.post(self.url("complete"))
.bearer_auth(&self.token)
.json(report)
.send()
.await
.map_err(|e| WorkerError::ConnectionFailed {
url: self.orchestrator_url.clone(),
reason: e.to_string(),
})?;
if !resp.status().is_success() {
return Err(WorkerError::OrchestratorRejected {
job_id: self.job_id,
reason: format!("completion report rejected: {}", resp.status()),
});
}
Ok(())
}
}
+9 -9
View File
@@ -326,15 +326,15 @@ impl ClaudeBridgeRuntime {
match serde_json::from_str::<ClaudeStreamEvent>(&line) {
Ok(event) => {
// Capture session_id from system init
if event.event_type == "system"
&& let Some(ref sid) = event.session_id
{
session_id = Some(sid.clone());
tracing::info!(
job_id = %self.config.job_id,
session_id = %sid,
"Captured Claude session ID"
);
if event.event_type == "system" {
if let Some(ref sid) = event.session_id {
session_id = Some(sid.clone());
tracing::info!(
job_id = %self.config.job_id,
session_id = %sid,
"Captured Claude session ID"
);
}
}
// Convert to our event payload and forward
+2 -3
View File
@@ -313,7 +313,6 @@ Work independently to complete this job. Report when done."#,
parameters: tc.arguments.clone(),
reasoning: String::new(),
alternatives: vec![],
tool_call_id: tc.id.clone(),
};
self.process_result(reason_ctx, &selection, result);
}
@@ -423,7 +422,7 @@ Work independently to complete this job. Report when done."#,
);
reason_ctx.messages.push(ChatMessage::tool_result(
&selection.tool_call_id,
"tool_call_id",
&selection.tool_name,
wrapped,
));
@@ -437,7 +436,7 @@ Work independently to complete this job. Report when done."#,
Err(e) => {
tracing::warn!("Tool {} failed: {}", selection.tool_name, e);
reason_ctx.messages.push(ChatMessage::tool_result(
&selection.tool_call_id,
"tool_call_id",
&selection.tool_name,
format!("Error: {}", e),
));
+13 -13
View File
@@ -532,10 +532,10 @@ impl Workspace {
];
for (path, header) in identity_files {
if let Ok(doc) = self.read(path).await
&& !doc.content.is_empty()
{
parts.push(format!("{}\n\n{}", header, doc.content));
if let Ok(doc) = self.read(path).await {
if !doc.content.is_empty() {
parts.push(format!("{}\n\n{}", header, doc.content));
}
}
}
@@ -544,15 +544,15 @@ impl Workspace {
let yesterday = today.pred_opt().unwrap_or(today);
for date in [today, yesterday] {
if let Ok(doc) = self.daily_log(date).await
&& !doc.content.is_empty()
{
let header = if date == today {
"## Today's Notes"
} else {
"## Yesterday's Notes"
};
parts.push(format!("{}\n\n{}", header, doc.content));
if let Ok(doc) = self.daily_log(date).await {
if !doc.content.is_empty() {
let header = if date == today {
"## Today's Notes"
} else {
"## Yesterday's Notes"
};
parts.push(format!("{}\n\n{}", header, doc.content));
}
}
}
+5 -5
View File
@@ -201,11 +201,11 @@ pub fn reciprocal_rank_fusion(
.collect();
// Normalize scores to 0-1 range
if let Some(max_score) = results.iter().map(|r| r.score).reduce(f32::max)
&& max_score > 0.0
{
for result in &mut results {
result.score /= max_score;
if let Some(max_score) = results.iter().map(|r| r.score).reduce(f32::max) {
if max_score > 0.0 {
for result in &mut results {
result.score /= max_score;
}
}
}
-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!");
}
+4 -4
View File
@@ -302,10 +302,10 @@ async fn test_chat_completions_streaming() {
if data == "[DONE]" {
continue;
}
if let Ok(chunk) = serde_json::from_str::<serde_json::Value>(data)
&& let Some(content) = chunk["choices"][0]["delta"]["content"].as_str()
{
full_content.push_str(content);
if let Ok(chunk) = serde_json::from_str::<serde_json::Value>(data) {
if let Some(content) = chunk["choices"][0]["delta"]["content"].as_str() {
full_content.push_str(content);
}
}
}
}
+107 -44
View File
@@ -53,56 +53,119 @@ impl exports::near::agent::tool::Guest for GmailTool {
r#"{
"type": "object",
"required": ["action"],
"properties": {
"action": {
"type": "string",
"enum": ["list_messages", "get_message", "send_message", "create_draft", "reply_to_message", "trash_message"],
"description": "The Gmail operation to perform"
"oneOf": [
{
"properties": {
"action": { "const": "list_messages" },
"query": {
"type": "string",
"description": "Gmail search query (same syntax as Gmail search box). Examples: 'is:unread', 'from:[email protected]', 'subject:meeting after:2025/01/01'"
},
"max_results": {
"type": "integer",
"description": "Maximum number of messages to return (default: 20)",
"default": 20
},
"label_ids": {
"type": "array",
"items": { "type": "string" },
"description": "Label IDs to filter by (e.g., 'INBOX', 'SENT', 'DRAFT')"
}
},
"required": ["action"]
},
"query": {
"type": "string",
"description": "Gmail search query (same syntax as Gmail search box, e.g., 'is:unread', 'from:[email protected]'). Used by: list_messages"
{
"properties": {
"action": { "const": "get_message" },
"message_id": {
"type": "string",
"description": "The message ID to retrieve"
}
},
"required": ["action", "message_id"]
},
"max_results": {
"type": "integer",
"description": "Maximum number of messages to return (default: 20). Used by: list_messages",
"default": 20
{
"properties": {
"action": { "const": "send_message" },
"to": {
"type": "string",
"description": "Recipient email address(es), comma-separated"
},
"subject": {
"type": "string",
"description": "Email subject"
},
"body": {
"type": "string",
"description": "Email body (plain text)"
},
"cc": {
"type": "string",
"description": "CC recipients, comma-separated"
},
"bcc": {
"type": "string",
"description": "BCC recipients, comma-separated"
}
},
"required": ["action", "to", "subject", "body"]
},
"label_ids": {
"type": "array",
"items": { "type": "string" },
"description": "Label IDs to filter by (e.g., 'INBOX', 'SENT', 'DRAFT'). Used by: list_messages"
{
"properties": {
"action": { "const": "create_draft" },
"to": {
"type": "string",
"description": "Recipient email address(es), comma-separated"
},
"subject": {
"type": "string",
"description": "Email subject"
},
"body": {
"type": "string",
"description": "Email body (plain text)"
},
"cc": {
"type": "string",
"description": "CC recipients, comma-separated"
},
"bcc": {
"type": "string",
"description": "BCC recipients, comma-separated"
}
},
"required": ["action", "to", "subject", "body"]
},
"message_id": {
"type": "string",
"description": "Message ID. Required for: get_message, reply_to_message, trash_message"
{
"properties": {
"action": { "const": "reply_to_message" },
"message_id": {
"type": "string",
"description": "The message ID to reply to"
},
"body": {
"type": "string",
"description": "Reply body (plain text)"
},
"reply_all": {
"type": "boolean",
"description": "If true, reply to all recipients (default: false)",
"default": false
}
},
"required": ["action", "message_id", "body"]
},
"to": {
"type": "string",
"description": "Recipient email address(es), comma-separated. Required for: send_message, create_draft"
},
"subject": {
"type": "string",
"description": "Email subject. Required for: send_message, create_draft"
},
"body": {
"type": "string",
"description": "Email body (plain text). Required for: send_message, create_draft, reply_to_message"
},
"cc": {
"type": "string",
"description": "CC recipients, comma-separated. Used by: send_message, create_draft"
},
"bcc": {
"type": "string",
"description": "BCC recipients, comma-separated. Used by: send_message, create_draft"
},
"reply_all": {
"type": "boolean",
"description": "If true, reply to all recipients (default: false). Used by: reply_to_message",
"default": false
{
"properties": {
"action": { "const": "trash_message" },
"message_id": {
"type": "string",
"description": "The message ID to move to trash"
}
},
"required": ["action", "message_id"]
}
}
]
}"#
.to_string()
}
+155 -65
View File
@@ -52,76 +52,166 @@ impl exports::near::agent::tool::Guest for GoogleCalendarTool {
r#"{
"type": "object",
"required": ["action"],
"properties": {
"action": {
"type": "string",
"enum": ["list_events", "get_event", "create_event", "update_event", "delete_event"],
"description": "The calendar operation to perform"
"oneOf": [
{
"properties": {
"action": { "const": "list_events" },
"calendar_id": {
"type": "string",
"description": "Calendar ID (default: 'primary')",
"default": "primary"
},
"time_min": {
"type": "string",
"description": "Lower bound for event start time (RFC3339, e.g., '2025-01-15T00:00:00Z')"
},
"time_max": {
"type": "string",
"description": "Upper bound for event end time (RFC3339)"
},
"max_results": {
"type": "integer",
"description": "Maximum number of events to return (default: 25)",
"default": 25
},
"query": {
"type": "string",
"description": "Free text search terms to filter events"
}
},
"required": ["action"]
},
"calendar_id": {
"type": "string",
"description": "Calendar ID (default: 'primary')",
"default": "primary"
{
"properties": {
"action": { "const": "get_event" },
"calendar_id": {
"type": "string",
"description": "Calendar ID (default: 'primary')",
"default": "primary"
},
"event_id": {
"type": "string",
"description": "The event ID to retrieve"
}
},
"required": ["action", "event_id"]
},
"event_id": {
"type": "string",
"description": "Event ID. Required for: get_event, update_event, delete_event"
{
"properties": {
"action": { "const": "create_event" },
"calendar_id": {
"type": "string",
"description": "Calendar ID (default: 'primary')",
"default": "primary"
},
"summary": {
"type": "string",
"description": "Event title"
},
"description": {
"type": "string",
"description": "Event description"
},
"location": {
"type": "string",
"description": "Event location"
},
"start_datetime": {
"type": "string",
"description": "Start time as RFC3339 (e.g., '2025-01-15T09:00:00-05:00'). Use start_date for all-day events."
},
"end_datetime": {
"type": "string",
"description": "End time as RFC3339. Use end_date for all-day events."
},
"start_date": {
"type": "string",
"description": "Start date for all-day events (e.g., '2025-01-15')"
},
"end_date": {
"type": "string",
"description": "End date for all-day events (exclusive, e.g., '2025-01-16' for a single day)"
},
"timezone": {
"type": "string",
"description": "Timezone (e.g., 'America/New_York')"
},
"attendees": {
"type": "array",
"items": { "type": "string" },
"description": "Attendee email addresses"
}
},
"required": ["action", "summary"]
},
"time_min": {
"type": "string",
"description": "Lower bound for event start time (RFC3339, e.g., '2025-01-15T00:00:00Z'). Used by: list_events"
{
"properties": {
"action": { "const": "update_event" },
"calendar_id": {
"type": "string",
"description": "Calendar ID (default: 'primary')",
"default": "primary"
},
"event_id": {
"type": "string",
"description": "The event ID to update"
},
"summary": {
"type": "string",
"description": "New event title"
},
"description": {
"type": "string",
"description": "New event description"
},
"location": {
"type": "string",
"description": "New event location"
},
"start_datetime": {
"type": "string",
"description": "New start time (RFC3339)"
},
"end_datetime": {
"type": "string",
"description": "New end time (RFC3339)"
},
"start_date": {
"type": "string",
"description": "New start date for all-day events"
},
"end_date": {
"type": "string",
"description": "New end date for all-day events"
},
"timezone": {
"type": "string",
"description": "Timezone for datetime fields"
},
"attendees": {
"type": "array",
"items": { "type": "string" },
"description": "Replace attendees with these email addresses"
}
},
"required": ["action", "event_id"]
},
"time_max": {
"type": "string",
"description": "Upper bound for event end time (RFC3339). Used by: list_events"
},
"max_results": {
"type": "integer",
"description": "Maximum number of events to return (default: 25). Used by: list_events",
"default": 25
},
"query": {
"type": "string",
"description": "Free text search terms to filter events. Used by: list_events"
},
"summary": {
"type": "string",
"description": "Event title. Required for: create_event. Optional for: update_event"
},
"description": {
"type": "string",
"description": "Event description. Used by: create_event, update_event"
},
"location": {
"type": "string",
"description": "Event location. Used by: create_event, update_event"
},
"start_datetime": {
"type": "string",
"description": "Start time (RFC3339, e.g., '2025-01-15T09:00:00-05:00'). For all-day events use start_date. Used by: create_event, update_event"
},
"end_datetime": {
"type": "string",
"description": "End time (RFC3339). For all-day events use end_date. Used by: create_event, update_event"
},
"start_date": {
"type": "string",
"description": "Start date for all-day events (e.g., '2025-01-15'). Used by: create_event, update_event"
},
"end_date": {
"type": "string",
"description": "End date for all-day events (exclusive, e.g., '2025-01-16'). Used by: create_event, update_event"
},
"timezone": {
"type": "string",
"description": "Timezone (e.g., 'America/New_York'). Used by: create_event, update_event"
},
"attendees": {
"type": "array",
"items": { "type": "string" },
"description": "Attendee email addresses. Used by: create_event, update_event"
{
"properties": {
"action": { "const": "delete_event" },
"calendar_id": {
"type": "string",
"description": "Calendar ID (default: 'primary')",
"default": "primary"
},
"event_id": {
"type": "string",
"description": "The event ID to delete"
}
},
"required": ["action", "event_id"]
}
}
]
}"#
.to_string()
}

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