chore: resolve conflicts

This commit is contained in:
italic-jinxin
2026-03-21 16:22:14 +08:00
85 changed files with 7408 additions and 1316 deletions
+19 -1
View File
@@ -4,7 +4,7 @@ DATABASE_POOL_SIZE=10
# LLM Provider
# LLM_BACKEND=nearai # default
# Possible values: nearai, ollama, openai_compatible, openai, anthropic, tinfoil
# Possible values: nearai, ollama, openai_compatible, openai, anthropic, github_copilot, tinfoil, openai_codex
# LLM_REQUEST_TIMEOUT_SECS=120 # Increase for local LLMs (Ollama, vLLM, LM Studio)
# === Anthropic Direct ===
@@ -24,6 +24,17 @@ DATABASE_POOL_SIZE=10
# LLM_USE_CODEX_AUTH=true
# CODEX_AUTH_PATH=~/.codex/auth.json
# === GitHub Copilot ===
# Uses the OAuth token from your Copilot IDE sign-in (for example
# ~/.config/github-copilot/apps.json on Linux/macOS), or run `ironclaw onboard`
# and choose the GitHub device login flow.
# LLM_BACKEND=github_copilot
# GITHUB_COPILOT_TOKEN=gho_...
# GITHUB_COPILOT_MODEL=gpt-4o
# IronClaw injects standard VS Code Copilot headers automatically.
# Optional advanced headers for custom overrides:
# GITHUB_COPILOT_EXTRA_HEADERS=Copilot-Integration-Id:vscode-chat
# === NEAR AI (Chat Completions API) ===
# Two auth modes:
# 1. Session token (default): Uses browser OAuth (GitHub/Google) on first run.
@@ -92,6 +103,13 @@ NEARAI_AUTH_URL=https://private.near.ai
# long = 1-hour TTL, 2.0× (200%) write surcharge
# ANTHROPIC_CACHE_RETENTION=short
# === OpenAI Codex (ChatGPT subscription, OAuth) ===
# LLM_BACKEND=openai_codex
# OPENAI_CODEX_MODEL=gpt-5.3-codex # default
# OPENAI_CODEX_CLIENT_ID=app_EMoamEEZ73f0CkXaXp7hrann # override (rare)
# OPENAI_CODEX_AUTH_URL=https://auth.openai.com # override (rare)
# OPENAI_CODEX_API_URL=https://chatgpt.com/backend-api/codex # override (rare)
# For full provider setup guide see docs/LLM_PROVIDERS.md
# Channel Configuration
+89 -1
View File
@@ -1,6 +1,94 @@
# Agent Rules
## Feature Parity Update Policy
## Purpose and Precedence
- `AGENTS.md` is the quick-start contract for coding agents. It is not the full architecture spec.
- Read the relevant subsystem spec before changing a complex area. When a repo spec exists, treat it as authoritative.
Start with these deeper docs as needed:
- `CLAUDE.md`
- `src/agent/CLAUDE.md`
- `src/channels/web/CLAUDE.md`
- `src/db/CLAUDE.md`
- `src/llm/CLAUDE.md`
- `src/setup/README.md`
- `src/tools/README.md`
- `src/workspace/README.md`
- `src/NETWORK_SECURITY.md`
- `tests/e2e/CLAUDE.md`
## Architecture Mental Model
- Channels normalize external input into `IncomingMessage`; `ChannelManager` merges all active channel streams.
- `Agent` owns session/thread/turn handling, submission parsing, the LLM/tool loop, approvals, routines, and background runtime behavior.
- `AppBuilder` is the composition root that wires database, secrets, LLMs, tools, workspace, extensions, skills, hooks, and cost controls before the agent starts.
- The web gateway is a browser-facing API/UI layered on top of the same agent/session/tool systems, not a separate product path.
## Where to Work
- Agent/runtime behavior: `src/agent/`
- Web gateway/API/SSE/WebSocket: `src/channels/web/`
- Persistence and DB abstractions: `src/db/`
- Setup/onboarding/configuration flow: `src/setup/`
- LLM providers and routing: `src/llm/`
- Workspace, memory, embeddings, search: `src/workspace/`
- Extensions, tools, channels, MCP, WASM: `src/extensions/`, `src/tools/`, `src/channels/`
## Ownership and Composition Rules
- Keep `src/main.rs` and `src/app.rs` orchestration-focused. Do not move module-owned logic into entrypoints.
- Module-specific initialization should live in the owning module behind a public factory/helper, not be reimplemented ad hoc.
- Keep feature-flag branching inside the module that owns the abstraction whenever possible.
- Prefer extending existing traits and registries over hardcoding one-off integration paths.
## Repo-Wide Coding Rules
- Avoid `.unwrap()` and `.expect()` in production; prefer proper error handling. They are fine in tests, and in production only for truly infallible invariants (e.g., literals/regexes) with a safety comment.
- Keep clippy clean with zero warnings.
- Prefer `crate::` imports for cross-module references.
- Use strong types and enums over stringly-typed control flow when the shape is known.
## Database, Setup, and Config Rules
- New persistence behavior must support both PostgreSQL and libSQL.
- Add new DB operations to the shared DB trait first, then implement both backends.
- Treat bootstrap config, DB-backed settings, and encrypted secrets as distinct layers; do not collapse them casually.
- If onboarding or setup behavior changes, update `src/setup/README.md` in the same branch.
- Do not break config precedence, bootstrap env loading, DB-backed config reload, or post-secrets LLM re-resolution.
## Security and Runtime Invariants
- Review any change touching listeners, routes, auth, secrets, sandboxing, approvals, or outbound HTTP with a security mindset.
- Do not weaken bearer-token auth, webhook auth, CORS/origin checks, body limits, rate limits, allowlists, or secret-handling guarantees.
- Treat Docker containers and external services as untrusted.
- Session/thread/turn state matters. Submission parsing happens before normal chat handling.
- Skills are selected deterministically. Tool approval and auth flows are special paths and must not be mixed into normal chat history carelessly.
- Persistent memory is the workspace system, not just transcript storage; preserve file-like semantics, chunking/search behavior, and identity/system-prompt loading.
## Tools, Channels, and Extensions
- Use a built-in Rust tool for core internal capabilities tightly coupled to the runtime.
- Use WASM tools or WASM channels for sandboxed extensions and plugin-style integrations.
- Use MCP for external server integrations when the capability belongs outside the main binary.
- Preserve extension lifecycle expectations: install, authenticate/configure, activate, remove.
## Docs, Parity, and Testing
- If behavior changes, update the relevant docs/specs in the same branch.
- If you change implementation status for any feature tracked in `FEATURE_PARITY.md`, update that file in the same branch.
- Do not open a PR that changes feature behavior without checking `FEATURE_PARITY.md` for needed status updates (`❌`, `🚧`, `✅`, notes, and priorities).
- Add the narrowest tests that validate the change: unit tests for local logic, integration tests for runtime/DB/routing behavior, and E2E or trace coverage for gateway, approvals, extensions, or other user-visible flows.
## Risk and Change Discipline
- Keep changes scoped; avoid broad refactors unless the task truly requires them.
- Security, database schema, runtime, worker, CI, and secrets changes are high-risk. Call out rollback risks, compatibility concerns, and hidden side effects.
- Preserve existing defaults unless the task explicitly changes them.
- Avoid unrelated file churn and generated-file edits unless required.
- Respect a dirty worktree and never revert user changes you did not make.
## Before Finishing
- Confirm whether behavior changes require updates to `FEATURE_PARITY.md`, specs, API docs, or `CHANGELOG.md`.
- Run the most targeted tests/checks that cover the change.
- Re-check security-sensitive paths when touching auth, secrets, network listeners, sandboxing, or approvals.
- Keep the final diff scoped to the task.
Generated
+7 -7
View File
@@ -2339,7 +2339,7 @@ source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "39cab71617ae0d63f51a36d69f866391735b51691dbda63cf6f96d042b63efeb"
dependencies = [
"libc",
"windows-sys 0.52.0",
"windows-sys 0.59.0",
]
[[package]]
@@ -5575,7 +5575,7 @@ dependencies = [
"errno",
"libc",
"linux-raw-sys 0.12.1",
"windows-sys 0.52.0",
"windows-sys 0.59.0",
]
[[package]]
@@ -5624,7 +5624,7 @@ dependencies = [
"once_cell",
"ring",
"rustls-pki-types",
"rustls-webpki 0.103.9",
"rustls-webpki 0.103.10",
"subtle",
"zeroize",
]
@@ -5696,9 +5696,9 @@ dependencies = [
[[package]]
name = "rustls-webpki"
version = "0.103.9"
version = "0.103.10"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "d7df23109aa6c1567d1c575b9952556388da57401e4ace1d15f79eedad0d8f53"
checksum = "df33b2b81ac578cabaf06b89b0631153a3f416b0a886e8a7a1707fb51abbd1ef"
dependencies = [
"aws-lc-rs",
"ring",
@@ -6479,10 +6479,10 @@ source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "32497e9a4c7b38532efcdebeef879707aa9f794296a4f0244f6f69e9bc8574bd"
dependencies = [
"fastrand",
"getrandom 0.3.4",
"getrandom 0.4.2",
"once_cell",
"rustix 1.1.4",
"windows-sys 0.52.0",
"windows-sys 0.59.0",
]
[[package]]
+1
View File
@@ -242,6 +242,7 @@ This document tracks feature parity between IronClaw (Rust implementation) and O
| OpenRouter | ✅ | ✅ | - | Via OpenAI-compatible provider (RigAdapter) |
| Tinfoil | ❌ | ✅ | - | Private inference provider (IronClaw-only) |
| OpenAI-compatible | ❌ | ✅ | - | Generic OpenAI-compatible endpoint (RigAdapter) |
| GitHub Copilot | ✅ | ✅ | - | Dedicated provider with OAuth token exchange (`GithubCopilotProvider`) |
| Ollama (local) | ✅ | ✅ | - | via `rig::providers::ollama` (full support) |
| Perplexity | ✅ | ❌ | P3 | Freshness parameter for web_search |
| MiniMax | ✅ | ❌ | P3 | Regional endpoint selection |
+1 -1
View File
@@ -168,7 +168,7 @@ written to `~/.ironclaw/.env` so they are available before the database connects
### Alternative LLM Providers
IronClaw defaults to NEAR AI but supports many LLM providers out of the box.
Built-in providers include **Anthropic**, **OpenAI**, **Google Gemini**, **MiniMax**,
Built-in providers include **Anthropic**, **OpenAI**, **GitHub Copilot**, **Google Gemini**, **MiniMax**,
**Mistral**, and **Ollama** (local). OpenAI-compatible services like **OpenRouter**
(300+ models), **Together AI**, **Fireworks AI**, and self-hosted servers (**vLLM**,
**LiteLLM**) are also supported.
+1 -1
View File
@@ -165,7 +165,7 @@ ironclaw onboard
### 替代 LLM 提供商
IronClaw 默认使用 NEAR AI,但开箱即用地支持多种 LLM 提供商。
内置提供商包括 **Anthropic**、**OpenAI**、**Google Gemini**、**MiniMax**、**Mistral** 和 **Ollama**(本地部署)。同时也支持 OpenAI 兼容服务,如 **OpenRouter**300+ 模型)、**Together AI**、**Fireworks AI** 以及自托管服务器(**vLLM**、**LiteLLM**)。
内置提供商包括 **Anthropic**、**OpenAI**、**GitHub Copilot**、**Google Gemini**、**MiniMax**、**Mistral** 和 **Ollama**(本地部署)。同时也支持 OpenAI 兼容服务,如 **OpenRouter**300+ 模型)、**Together AI**、**Fireworks AI** 以及自托管服务器(**vLLM**、**LiteLLM**)。
在向导中选择你的提供商,或直接设置环境变量:
+12
View File
@@ -243,6 +243,18 @@ mod tests {
assert!(wrapped.contains("Hello <world>"));
}
#[test]
fn test_wrap_for_llm_escapes_attr_chars() {
let config = SafetyConfig {
max_output_length: 100_000,
injection_check_enabled: true,
};
let safety = SafetyLayer::new(&config);
let wrapped = safety.wrap_for_llm("bad&\"<>name", "ok", false);
assert!(wrapped.contains("name=\"bad&amp;&quot;&lt;&gt;name\"")); // safety: test assertion in #[cfg(test)] module
}
#[test]
fn test_sanitize_action_forces_sanitization_when_injection_check_disabled() {
let config = SafetyConfig {
+2
View File
@@ -15,6 +15,8 @@ ignore = [
"RUSTSEC-2026-0020",
# wasmtime wasi:http/types.fields panic — mitigated by fuel limits
"RUSTSEC-2026-0021",
# rustls-webpki CRL distributionPoint matching — 0.102.8 pinned by libsql transitive dep
"RUSTSEC-2026-0049",
]
[licenses]
+29
View File
@@ -17,6 +17,7 @@ configurations.
| Yandex AI Studio | `yandex` | `YANDEX_API_KEY` | YandexGPT models |
| MiniMax | `minimax` | `MINIMAX_API_KEY` | MiniMax-M2.7 models |
| Cloudflare Workers AI | `cloudflare` | `CLOUDFLARE_API_KEY` | Access to Workers AI |
| GitHub Copilot | `github_copilot` | `GITHUB_COPILOT_TOKEN` | Multi-models |
| Ollama | `ollama` | No | Local inference |
| AWS Bedrock | `bedrock` | AWS credentials | Native Converse API |
| OpenRouter | `openai_compatible` | `LLM_API_KEY` | 300+ models |
@@ -61,6 +62,34 @@ Popular models: `gpt-4o`, `gpt-4o-mini`, `o3-mini`
---
## GitHub Copilot
GitHub Copilot exposes chat endpoint at
`https://api.githubcopilot.com`. IronClaw uses that endpoint directly through the
built-in `github_copilot` provider.
```env
LLM_BACKEND=github_copilot
GITHUB_COPILOT_TOKEN=gho_...
GITHUB_COPILOT_MODEL=gpt-4o
# Optional advanced headers if your setup needs them:
# GITHUB_COPILOT_EXTRA_HEADERS=Copilot-Integration-Id:vscode-chat
```
`ironclaw onboard` can acquire this token for you using GitHub device login. If you
already signed into Copilot through VS Code or a JetBrains IDE, you can also reuse
the `oauth_token` stored in `~/.config/github-copilot/apps.json`. If you prefer,
`LLM_BACKEND=github-copilot` also works as an alias.
Popular models vary by subscription, but `gpt-4o` is a safe default. IronClaw keeps
model entry manual for this provider because GitHub Copilot model listing may require
extra integration headers on some clients. IronClaw automatically injects the standard
VS Code identity headers (`User-Agent`, `Editor-Version`, `Editor-Plugin-Version`,
`Copilot-Integration-Id`) and lets you override them with
`GITHUB_COPILOT_EXTRA_HEADERS`.
---
## Ollama (local)
Install Ollama from [ollama.com](https://ollama.com), pull a model, then:
@@ -1,87 +0,0 @@
# Staging CI Review Issues Triage
**Date:** 2026-03-18
**Branch:** staging (HEAD `b7a1edf`)
**Total open issues:** 50
---
## Batch 1 — Critical & 100-confidence issues
| # | Title | Severity | Verdict | File(s) | Action |
|---|-------|----------|---------|---------|--------|
| 1281 | Logic inversion in Telegram auto-verification | CRITICAL:100 | **FALSE POSITIVE** (closed) | `src/channels/web/server.rs` | Different handlers with intentional different SSE behavior |
| 908 | Missing consecutive_failures reset | CRITICAL:100 | **STALE** | `src/llm/circuit_breaker.rs` | Close — `record_success()` already resets to 0 |
| 1282 | Variable shadowing fallback notification | HIGH:100 | **STALE** | `src/agent/agent_loop.rs` | Close — fixed in commit `bcc38ce` |
| 1283 | Inconsistent fallback logic DRY | HIGH:75 | **STALE** | `src/agent/agent_loop.rs` | Close — fixed in commit `bcc38ce` |
| 1178 | Workflow linting bypass for test code | CRITICAL:75 | **FALSE POSITIVE** | `.github/workflows/code_style.yml` | Close — script reads full file, not hunk headers |
---
## Remaining Batches (queued)
### Batch 2 — Retry/DRY + CI workflow issues (completed)
| # | Title | Severity | Verdict | Action |
|---|-------|----------|---------|--------|
| 1288 | DRY violation: retry-after parsing | HIGH:95 | **LEGIT** | Fixed: extracted shared `parse_retry_after()` |
| 1289 | Semantic mismatch in RFC2822 test helpers | MEDIUM:85 | **DUPLICATE** (closed) | Duplicate of #1288 |
| 1290 | Unnecessary eager `chrono::Utc::now()` call | LOW:85 | **FALSE POSITIVE** (closed) | Already deferred inside successful parse branch |
| 963 | Logical equivalence bug in workflow conditions | HIGH:100 | **FALSE POSITIVE** (closed) | Refactored condition correctly handles `workflow_call` |
| 1280 | Flaky OAuth wildcard callback tests | Flaky | **LEGIT** | Fixed: added `tokio::sync::Mutex` for env var serialization |
### Batch 3 — Routine engine + notification routing
- #1365 — too_many_arguments on RoutineEngine::new()
- #1371 — Discovery schema regeneration on every tool_info call
- #1364 — Prompt injection via unescaped channel/user in lightweight routines
- #1284 — notification_target_for_channel() assumes channel owner
### Batch 4 — Telegram/Extension Manager webhook group
- #1247 — Synchronous 120-second blocking poll in HTTP handler
- #1248 — Hardcoded channel-specific logic violates architecture
- #1249 — Telegram-specific business logic bloats ExtensionManager
- #1250 — Response success/failure logic mismatch in chat auth
- #1251 — Channel-specific configuration mappings lack extensibility
### Batch 5 — HMAC/Auth/Security
- #1034 — Signature verification not constant-time
- #1035 — Incorrect order of operations in HMAC verification
- #1036 — Double opt-in lacks runtime validation consistency
- #1037 — API breaking change: auth() signature
- #1038 — CSP policy allows CDN scripts with risky fallback
### Batch 6 — Webhook handler + config
- #1039 — Per-request HTTP client creation in hot path
- #1040 — Complex nested auth logic in webhook_handler
- #1041 — Redundant JSON deserialization in webhook handler
- #1042 — Implicit state mutation in config conversion
- #1005 — Inconsistent double opt-in enforcement
### Batch 7 — Tool schema validation / WASM bounds
- #974 — Unbounded recursion in resolve_nested()
- #975 — Unbounded recursion in validate_tool_schema()
- #976 — Unbounded description string in CapabilitiesFile
- #977 — Unbounded parameters schema JSON
- #978 — Unnecessary clone of large JSON in hot path
### Batch 8 — Tool schema + config + security
- #979 — No size limits on JSON files read
- #980 — Misleading warning condition for missing parameters
- #988 — Hardcoded CLI_ENABLED env var in systemd template
- #990 — Configuration semantics unclear for daemon mode
- #1103 — SSRF risk via configurable embedding base URL
### Batch 9 — Agent loop / job worker
- #870 — Unbounded loop without cancellation token
- #871 — Stringly-typed unsupported parameter filtering
- #873 — RwLock overhead on hot path
- #892 — JobDelegate::check_signals() treats non-terminal as terminal
- #1252 — String concatenation in hot polling loop
### Batch 10 — Agent loop perf + CI scripts
- #893 — Unnecessary parameter cloning on every tool execution
- #894 — truncate_for_preview allocates for non-truncated strings
- #895 — Tool definitions fetched every iteration without caching
- #1179 — AWK state machine never resets between hunks
- #1180 — Code fence detection logic flawed in extract_suggestions()
- #1181 — Unsafe .unwrap() in production code manifest.rs
+23
View File
@@ -77,6 +77,29 @@
"can_list_models": false
}
},
{
"id": "github_copilot",
"aliases": [
"github-copilot",
"githubcopilot",
"copilot"
],
"protocol": "github_copilot",
"default_base_url": "https://api.githubcopilot.com",
"api_key_env": "GITHUB_COPILOT_TOKEN",
"api_key_required": true,
"model_env": "GITHUB_COPILOT_MODEL",
"default_model": "gpt-4o",
"extra_headers_env": "GITHUB_COPILOT_EXTRA_HEADERS",
"description": "GitHub Copilot Chat API (OAuth token from IDE sign-in)",
"setup": {
"kind": "api_key",
"secret_name": "llm_github_copilot_token",
"key_url": "https://docs.github.com/en/copilot",
"display_name": "GitHub Copilot",
"can_list_models": false
}
},
{
"id": "tinfoil",
"aliases": [],
+62 -13
View File
@@ -10,6 +10,7 @@
use std::sync::Arc;
use futures::StreamExt;
use uuid::Uuid;
use crate::agent::context_monitor::ContextMonitor;
use crate::agent::heartbeat::spawn_heartbeat;
@@ -17,7 +18,7 @@ use crate::agent::routine_engine::{RoutineEngine, spawn_cron_ticker};
use crate::agent::self_repair::{DefaultSelfRepair, RepairResult, SelfRepair};
use crate::agent::session_manager::SessionManager;
use crate::agent::submission::{Submission, SubmissionParser, SubmissionResult};
use crate::agent::{HeartbeatConfig as AgentHeartbeatConfig, Router, Scheduler};
use crate::agent::{HeartbeatConfig as AgentHeartbeatConfig, Router, Scheduler, SchedulerDeps};
use crate::channels::{ChannelManager, IncomingMessage, OutgoingResponse};
use crate::config::{AgentConfig, HeartbeatConfig, RoutineConfig, SkillsConfig};
use crate::context::ContextManager;
@@ -227,9 +228,12 @@ impl Agent {
context_manager.clone(),
deps.llm.clone(),
deps.safety.clone(),
deps.tools.clone(),
deps.store.clone(),
deps.hooks.clone(),
SchedulerDeps {
tools: deps.tools.clone(),
extension_manager: deps.extension_manager.clone(),
store: deps.store.clone(),
hooks: deps.hooks.clone(),
},
);
if let Some(ref tx) = deps.sse_tx {
scheduler.set_sse_sender(tx.clone());
@@ -600,6 +604,7 @@ impl Agent {
Arc::clone(workspace),
notify_tx,
Some(self.scheduler.clone()),
self.deps.extension_manager.clone(),
self.tools().clone(),
self.safety().clone(),
self.deps.sandbox_readiness,
@@ -1010,15 +1015,59 @@ impl Agent {
}
}
// Resolve session and thread
let (session, thread_id) = self
.session_manager
.resolve_thread(
&message.user_id,
&message.channel,
message.conversation_scope(),
)
.await;
// Resolve session and thread. Approval submissions are allowed to
// target an already-loaded owned thread by UUID across channels so the
// web approval UI can approve work that originated from HTTP/other
// owner-scoped channels.
let approval_thread_uuid = if matches!(
submission,
Submission::ExecApproval { .. } | Submission::ApprovalResponse { .. }
) {
message
.conversation_scope()
.and_then(|thread_id| Uuid::parse_str(thread_id).ok())
} else {
None
};
let (session, thread_id) = if let Some(target_thread_id) = approval_thread_uuid {
let session = self
.session_manager
.get_or_create_session(&message.user_id)
.await;
let mut sess = session.lock().await;
if sess.threads.contains_key(&target_thread_id) {
sess.active_thread = Some(target_thread_id);
sess.last_active_at = chrono::Utc::now();
drop(sess);
self.session_manager
.register_thread(
&message.user_id,
&message.channel,
target_thread_id,
Arc::clone(&session),
)
.await;
(session, target_thread_id)
} else {
drop(sess);
self.session_manager
.resolve_thread(
&message.user_id,
&message.channel,
message.conversation_scope(),
)
.await
}
} else {
self.session_manager
.resolve_thread(
&message.user_id,
&message.channel,
message.conversation_scope(),
)
.await
};
tracing::debug!(
message_id = %message.id,
thread_id = %thread_id,
+1 -1
View File
@@ -40,7 +40,7 @@ pub use heartbeat::{HeartbeatConfig, HeartbeatResult, HeartbeatRunner, spawn_hea
pub use router::{MessageIntent, Router};
pub use routine::{Routine, RoutineAction, RoutineRun, Trigger};
pub use routine_engine::{RoutineEngine, SandboxReadiness};
pub use scheduler::Scheduler;
pub use scheduler::{Scheduler, SchedulerDeps};
pub use self_repair::{BrokenTool, RepairResult, RepairTask, SelfRepair, StuckJob};
pub use session::{PendingApproval, PendingAuth, Session, Thread, ThreadState, Turn, TurnState};
pub use session_manager::SessionManager;
+51 -233
View File
@@ -17,7 +17,7 @@
//! └──────────────┘
//! ```
use std::collections::{HashSet, hash_map::DefaultHasher};
use std::collections::hash_map::DefaultHasher;
use std::hash::{Hash, Hasher};
use std::str::FromStr;
use std::time::Duration;
@@ -28,171 +28,6 @@ use uuid::Uuid;
use crate::error::RoutineError;
pub const FULL_JOB_OWNER_ALLOWED_TOOLS_SETTING_KEY: &str = "routines.full_job_owner_allowed_tools";
pub const FULL_JOB_DEFAULT_PERMISSION_MODE_SETTING_KEY: &str =
"routines.full_job_default_permission_mode";
/// Persisted per-routine permission mode for autonomous `full_job` routines.
#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize, Default)]
#[serde(rename_all = "snake_case")]
pub enum FullJobPermissionMode {
/// Only use the routine's stored `tool_permissions`.
#[default]
Explicit,
/// Union the owner-scoped allowlist with the routine's `tool_permissions`.
InheritOwner,
}
impl FullJobPermissionMode {
pub fn as_str(self) -> &'static str {
match self {
Self::Explicit => "explicit",
Self::InheritOwner => "inherit_owner",
}
}
}
impl FromStr for FullJobPermissionMode {
type Err = ();
fn from_str(s: &str) -> Result<Self, Self::Err> {
match s {
"explicit" => Ok(Self::Explicit),
"inherit_owner" => Ok(Self::InheritOwner),
_ => Err(()),
}
}
}
/// Owner-scoped default behavior for newly-created `full_job` routines.
#[derive(Debug, Clone, Copy, PartialEq, Eq, Default)]
pub enum FullJobPermissionDefaultMode {
Explicit,
#[default]
InheritOwner,
CopyOwner,
}
impl FullJobPermissionDefaultMode {
pub fn as_str(self) -> &'static str {
match self {
Self::Explicit => "explicit",
Self::InheritOwner => "inherit_owner",
Self::CopyOwner => "copy_owner",
}
}
}
impl FromStr for FullJobPermissionDefaultMode {
type Err = ();
fn from_str(s: &str) -> Result<Self, Self::Err> {
match s {
"explicit" => Ok(Self::Explicit),
"inherit_owner" => Ok(Self::InheritOwner),
"copy_owner" => Ok(Self::CopyOwner),
_ => Err(()),
}
}
}
#[derive(Debug, Clone, PartialEq, Eq, Default)]
pub struct FullJobPermissionSettings {
pub owner_allowed_tools: Vec<String>,
pub default_mode: FullJobPermissionDefaultMode,
}
pub fn normalize_tool_names<I>(tools: I) -> Vec<String>
where
I: IntoIterator<Item = String>,
{
let mut seen = HashSet::new();
let mut normalized = Vec::new();
for tool in tools {
let trimmed = tool.trim();
if trimmed.is_empty() {
continue;
}
let normalized_name = trimmed.to_string();
if seen.insert(normalized_name.clone()) {
normalized.push(normalized_name);
}
}
normalized
}
pub fn parse_full_job_permission_mode(value: &serde_json::Value) -> FullJobPermissionMode {
value
.get("permission_mode")
.and_then(|v| v.as_str())
.and_then(|mode| FullJobPermissionMode::from_str(mode).ok())
.unwrap_or_default()
}
fn parse_owner_allowed_tools_setting(value: Option<serde_json::Value>) -> Vec<String> {
match value {
Some(serde_json::Value::Array(values)) => normalize_tool_names(
values
.into_iter()
.filter_map(|value| value.as_str().map(ToOwned::to_owned)),
),
Some(serde_json::Value::String(csv)) => normalize_tool_names(
csv.split([',', '\n'])
.map(str::trim)
.filter(|value| !value.is_empty())
.map(ToOwned::to_owned),
),
_ => Vec::new(),
}
}
fn parse_default_permission_mode_setting(
value: Option<serde_json::Value>,
) -> FullJobPermissionDefaultMode {
value
.and_then(|v| v.as_str().map(ToOwned::to_owned))
.and_then(|mode| FullJobPermissionDefaultMode::from_str(&mode).ok())
.unwrap_or_default()
}
pub async fn load_full_job_permission_settings(
store: &(dyn crate::db::SettingsStore + Sync),
user_id: &str,
) -> Result<FullJobPermissionSettings, crate::error::DatabaseError> {
let owner_allowed_tools = parse_owner_allowed_tools_setting(
store
.get_setting(user_id, FULL_JOB_OWNER_ALLOWED_TOOLS_SETTING_KEY)
.await?,
);
let default_mode = parse_default_permission_mode_setting(
store
.get_setting(user_id, FULL_JOB_DEFAULT_PERMISSION_MODE_SETTING_KEY)
.await?,
);
Ok(FullJobPermissionSettings {
owner_allowed_tools,
default_mode,
})
}
pub fn effective_full_job_tool_permissions(
permission_mode: FullJobPermissionMode,
routine_tool_permissions: &[String],
owner_allowed_tools: &[String],
) -> Vec<String> {
match permission_mode {
FullJobPermissionMode::Explicit => {
normalize_tool_names(routine_tool_permissions.iter().cloned())
}
FullJobPermissionMode::InheritOwner => normalize_tool_names(
owner_allowed_tools
.iter()
.cloned()
.chain(routine_tool_permissions.iter().cloned()),
),
}
}
/// A routine is a named, persistent, user-owned task with a trigger and an action.
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct Routine {
@@ -244,6 +79,13 @@ pub enum Trigger {
#[serde(default)]
filters: std::collections::HashMap<String, String>,
},
/// Fire on incoming webhook POST to /api/webhooks/{path}.
Webhook {
/// Optional webhook path suffix (defaults to routine id).
path: Option<String>,
/// Optional shared secret for HMAC validation.
secret: Option<String>,
},
/// Only fires via tool call or CLI.
Manual,
}
@@ -255,6 +97,7 @@ impl Trigger {
Trigger::Cron { .. } => "cron",
Trigger::Event { .. } => "event",
Trigger::SystemEvent { .. } => "system_event",
Trigger::Webhook { .. } => "webhook",
Trigger::Manual => "manual",
}
}
@@ -336,6 +179,17 @@ impl Trigger {
filters,
})
}
"webhook" => {
let path = config
.get("path")
.and_then(|v| v.as_str())
.map(String::from);
let secret = config
.get("secret")
.and_then(|v| v.as_str())
.map(String::from);
Ok(Trigger::Webhook { path, secret })
}
"manual" => Ok(Trigger::Manual),
other => Err(RoutineError::UnknownTriggerType {
trigger_type: other.to_string(),
@@ -363,6 +217,10 @@ impl Trigger {
"event_type": event_type,
"filters": filters,
}),
Trigger::Webhook { path, secret } => serde_json::json!({
"path": path,
"secret": secret,
}),
Trigger::Manual => serde_json::json!({}),
}
}
@@ -400,15 +258,6 @@ pub enum RoutineAction {
/// Max reasoning iterations (default: 10).
#[serde(default = "default_max_iterations")]
max_iterations: u32,
/// Tool names pre-authorized for `Always`-approval tools (e.g. destructive
/// shell commands, cross-channel messaging). `UnlessAutoApproved` tools are
/// automatically permitted in routine jobs without listing them here.
#[serde(default)]
tool_permissions: Vec<String>,
/// Whether this routine should inherit the owner's durable full-job
/// permission allowlist or use only its explicit `tool_permissions`.
#[serde(default)]
permission_mode: FullJobPermissionMode,
},
}
@@ -433,18 +282,6 @@ fn clamp_max_tool_rounds(value: u64) -> u32 {
value.clamp(1, MAX_TOOL_ROUNDS_LIMIT as u64) as u32
}
/// Parse a `tool_permissions` JSON array into a `Vec<String>`.
pub fn parse_tool_permissions(value: &serde_json::Value) -> Vec<String> {
normalize_tool_names(
value
.get("tool_permissions")
.and_then(|v| v.as_array())
.into_iter()
.flatten()
.filter_map(|v| v.as_str().map(String::from)),
)
}
impl RoutineAction {
/// The string tag stored in the DB action_type column.
pub fn type_tag(&self) -> &'static str {
@@ -519,14 +356,10 @@ impl RoutineAction {
.and_then(|v| v.as_u64())
.unwrap_or(default_max_iterations() as u64)
as u32;
let tool_permissions = parse_tool_permissions(&config);
let permission_mode = parse_full_job_permission_mode(&config);
Ok(RoutineAction::FullJob {
title,
description,
max_iterations,
tool_permissions,
permission_mode,
})
}
other => Err(RoutineError::UnknownActionType {
@@ -555,14 +388,10 @@ impl RoutineAction {
title,
description,
max_iterations,
tool_permissions,
permission_mode,
} => serde_json::json!({
"title": title,
"description": description,
"max_iterations": max_iterations,
"tool_permissions": tool_permissions,
"permission_mode": permission_mode,
}),
}
}
@@ -896,9 +725,8 @@ pub fn describe_cron(schedule: &str, timezone: Option<&str>) -> String {
#[cfg(test)]
mod tests {
use crate::agent::routine::{
FullJobPermissionMode, MAX_TOOL_ROUNDS_LIMIT, RoutineAction, RoutineGuardrails, RunStatus,
Trigger, content_hash, describe_cron, effective_full_job_tool_permissions, next_cron_fire,
normalize_cron_expression,
MAX_TOOL_ROUNDS_LIMIT, RoutineAction, RoutineGuardrails, RunStatus, Trigger, content_hash,
describe_cron, next_cron_fire, normalize_cron_expression,
};
#[test]
@@ -965,68 +793,50 @@ mod tests {
title: "Deploy review".to_string(),
description: "Review and deploy pending changes".to_string(),
max_iterations: 5,
tool_permissions: vec!["shell".to_string()],
permission_mode: FullJobPermissionMode::InheritOwner,
};
let json = action.to_config_json();
let parsed = RoutineAction::from_db("full_job", json).expect("parse full_job");
assert!(
matches!(parsed, RoutineAction::FullJob { title, max_iterations, tool_permissions, permission_mode, .. }
matches!(parsed, RoutineAction::FullJob { title, max_iterations, .. }
if title == "Deploy review"
&& max_iterations == 5
&& tool_permissions == vec!["shell".to_string()]
&& permission_mode == FullJobPermissionMode::InheritOwner)
&& max_iterations == 5)
);
}
#[test]
fn test_action_full_job_missing_permission_mode_defaults_to_explicit() {
fn test_action_full_job_ignores_legacy_permission_fields() {
let parsed = RoutineAction::from_db(
"full_job",
serde_json::json!({
"title": "Deploy review",
"description": "Review and deploy pending changes",
"max_iterations": 5,
"tool_permissions": ["shell"]
"tool_permissions": ["shell"],
"permission_mode": "inherit_owner"
}),
)
.expect("parse full_job");
assert!(matches!(
parsed,
RoutineAction::FullJob {
permission_mode: FullJobPermissionMode::Explicit,
ref title,
ref description,
max_iterations,
..
}
} if title == "Deploy review"
&& description == "Review and deploy pending changes"
&& max_iterations == 5
));
}
#[test]
fn test_effective_full_job_tool_permissions_inherit_owner_unions_lists() {
let resolved = effective_full_job_tool_permissions(
FullJobPermissionMode::InheritOwner,
&["shell".to_string(), "message".to_string()],
&["message".to_string(), "http".to_string()],
);
assert_eq!(
resolved,
vec![
"message".to_string(),
"http".to_string(),
"shell".to_string()
]
parsed.to_config_json(),
serde_json::json!({
"title": "Deploy review",
"description": "Review and deploy pending changes",
"max_iterations": 5,
})
);
}
#[test]
fn test_effective_full_job_tool_permissions_explicit_ignores_owner_defaults() {
let resolved = effective_full_job_tool_permissions(
FullJobPermissionMode::Explicit,
&["shell".to_string()],
&["message".to_string(), "http".to_string()],
);
assert_eq!(resolved, vec!["shell".to_string()]);
}
#[test]
fn test_run_status_display_parse() {
for status in [
@@ -1175,6 +985,14 @@ mod tests {
.type_tag(),
"system_event"
);
assert_eq!(
Trigger::Webhook {
path: None,
secret: None,
}
.type_tag(),
"webhook"
);
assert_eq!(Trigger::Manual.type_tag(), "manual");
}
+114 -80
View File
@@ -22,19 +22,20 @@ use uuid::Uuid;
use crate::agent::Scheduler;
use crate::agent::routine::{
NotifyConfig, Routine, RoutineAction, RoutineRun, RunStatus, Trigger,
effective_full_job_tool_permissions, load_full_job_permission_settings, next_cron_fire,
NotifyConfig, Routine, RoutineAction, RoutineRun, RunStatus, Trigger, next_cron_fire,
};
use crate::channels::OutgoingResponse;
use crate::config::RoutineConfig;
use crate::context::{JobContext, JobState};
use crate::db::Database;
use crate::error::RoutineError;
use crate::extensions::ExtensionManager;
use crate::llm::{
ChatMessage, CompletionRequest, FinishReason, LlmProvider, ToolCall, ToolCompletionRequest,
};
use crate::tools::{
ApprovalContext, ApprovalRequirement, ToolError, ToolRegistry, prepare_tool_params,
ToolError, ToolRegistry, autonomous_allowed_tool_names, autonomous_unavailable_message,
prepare_tool_params,
};
use crate::workspace::Workspace;
use ironclaw_safety::SafetyLayer;
@@ -69,6 +70,8 @@ pub struct RoutineEngine {
event_cache: Arc<RwLock<Vec<EventMatcher>>>,
/// Scheduler for dispatching jobs (FullJob mode).
scheduler: Option<Arc<Scheduler>>,
/// Owner-scoped extension activation state for autonomous tool resolution.
extension_manager: Option<Arc<ExtensionManager>>,
/// Tool registry for lightweight routine tool execution.
tools: Arc<ToolRegistry>,
/// Safety layer for tool output sanitization.
@@ -90,6 +93,7 @@ impl RoutineEngine {
workspace: Arc<Workspace>,
notify_tx: mpsc::Sender<OutgoingResponse>,
scheduler: Option<Arc<Scheduler>>,
extension_manager: Option<Arc<ExtensionManager>>,
tools: Arc<ToolRegistry>,
safety: Arc<SafetyLayer>,
sandbox_readiness: SandboxReadiness,
@@ -103,6 +107,7 @@ impl RoutineEngine {
running_count: Arc::new(AtomicUsize::new(0)),
event_cache: Arc::new(RwLock::new(Vec::new())),
scheduler,
extension_manager,
tools,
safety,
sandbox_readiness,
@@ -702,6 +707,92 @@ impl RoutineEngine {
notify_tx: self.notify_tx.clone(),
running_count: self.running_count.clone(),
scheduler: self.scheduler.clone(),
extension_manager: self.extension_manager.clone(),
tools: self.tools.clone(),
safety: self.safety.clone(),
sandbox_readiness: self.sandbox_readiness,
};
tokio::spawn(async move {
execute_routine(engine, routine, run).await;
});
Ok(run_id)
}
/// Fire a routine from a webhook trigger.
///
/// Similar to `fire_manual` but records the trigger as `"webhook"` with the
/// webhook path as detail. Skips ownership check (auth is via webhook secret).
/// Enforces enabled check, cooldown, and concurrent run limit.
pub async fn fire_webhook(
&self,
routine_id: Uuid,
webhook_path: &str,
) -> Result<Uuid, RoutineError> {
let routine = self
.store
.get_routine(routine_id)
.await
.map_err(|e| RoutineError::Database {
reason: e.to_string(),
})?
.ok_or(RoutineError::NotFound { id: routine_id })?;
if !routine.enabled {
return Err(RoutineError::Disabled {
name: routine.name.clone(),
});
}
if !self.check_cooldown(&routine) {
return Err(RoutineError::Cooldown {
name: routine.name.clone(),
});
}
if !self.check_concurrent(&routine).await {
return Err(RoutineError::MaxConcurrent {
name: routine.name.clone(),
});
}
if self.running_count.load(Ordering::Relaxed) >= self.config.max_concurrent_routines {
return Err(RoutineError::MaxConcurrent {
name: routine.name.clone(),
});
}
let run_id = Uuid::new_v4();
let run = RoutineRun {
id: run_id,
routine_id: routine.id,
trigger_type: "webhook".to_string(),
trigger_detail: Some(webhook_path.to_string()),
started_at: Utc::now(),
completed_at: None,
status: RunStatus::Running,
result_summary: None,
tokens_used: None,
job_id: None,
created_at: Utc::now(),
};
if let Err(e) = self.store.create_routine_run(&run).await {
return Err(RoutineError::Database {
reason: format!("failed to create run record: {e}"),
});
}
let engine = EngineContext {
config: self.config.clone(),
store: self.store.clone(),
llm: self.llm.clone(),
workspace: self.workspace.clone(),
notify_tx: self.notify_tx.clone(),
running_count: self.running_count.clone(),
scheduler: self.scheduler.clone(),
extension_manager: self.extension_manager.clone(),
tools: self.tools.clone(),
safety: self.safety.clone(),
sandbox_readiness: self.sandbox_readiness,
@@ -738,6 +829,7 @@ impl RoutineEngine {
notify_tx: self.notify_tx.clone(),
running_count: self.running_count.clone(),
scheduler: self.scheduler.clone(),
extension_manager: self.extension_manager.clone(),
tools: self.tools.clone(),
safety: self.safety.clone(),
sandbox_readiness: self.sandbox_readiness,
@@ -875,6 +967,7 @@ struct EngineContext {
notify_tx: mpsc::Sender<OutgoingResponse>,
running_count: Arc<AtomicUsize>,
scheduler: Option<Arc<Scheduler>>,
extension_manager: Option<Arc<ExtensionManager>>,
tools: Arc<ToolRegistry>,
safety: Arc<SafetyLayer>,
sandbox_readiness: SandboxReadiness,
@@ -908,15 +1001,11 @@ async fn execute_routine(ctx: EngineContext, routine: Routine, run: RoutineRun)
title,
description,
max_iterations,
tool_permissions,
permission_mode,
} => {
let execution = FullJobExecutionConfig {
title,
description,
max_iterations: *max_iterations,
tool_permissions,
permission_mode: *permission_mode,
};
execute_full_job(&ctx, &routine, &run, &execution).await
}
@@ -1048,8 +1137,6 @@ struct FullJobExecutionConfig<'a> {
title: &'a str,
description: &'a str,
max_iterations: u32,
tool_permissions: &'a [String],
permission_mode: crate::agent::routine::FullJobPermissionMode,
}
async fn execute_full_job(
@@ -1094,40 +1181,12 @@ async fn execute_full_job(
}
metadata["notify_user"] = serde_json::json!(&routine.notify.user);
let effective_permissions = match execution.permission_mode {
crate::agent::routine::FullJobPermissionMode::Explicit => {
effective_full_job_tool_permissions(
execution.permission_mode,
execution.tool_permissions,
&[],
)
}
crate::agent::routine::FullJobPermissionMode::InheritOwner => {
let owner_permissions =
load_full_job_permission_settings(ctx.store.as_ref(), &routine.user_id)
.await
.map_err(|e| RoutineError::Database {
reason: format!("failed to load routine permission settings: {e}"),
})?;
effective_full_job_tool_permissions(
execution.permission_mode,
execution.tool_permissions,
&owner_permissions.owner_allowed_tools,
)
}
};
// Build approval context: UnlessAutoApproved tools are auto-approved for routines;
// Always tools require explicit listing in the resolved effective permissions.
let approval_context = ApprovalContext::autonomous_with_tools(effective_permissions);
let job_id = scheduler
.dispatch_job_with_context(
.dispatch_job(
&routine.user_id,
execution.title,
execution.description,
Some(metadata),
approval_context,
)
.await
.map_err(|e| RoutineError::JobDispatchFailed {
@@ -1416,6 +1475,9 @@ async fn execute_lightweight_with_tools(
description: routine.name.clone(),
..Default::default()
};
let allowed_tools =
autonomous_allowed_tool_names(&ctx.tools, ctx.extension_manager.as_ref(), &routine.user_id)
.await;
loop {
iteration += 1;
@@ -1450,8 +1512,11 @@ async fn execute_lightweight_with_tools(
// Tool-enabled iteration
let tool_defs = ctx
.tools
.tool_definitions_excluding(ROUTINE_TOOL_DENYLIST)
.await;
.tool_definitions()
.await
.into_iter()
.filter(|tool| allowed_tools.contains(&tool.name))
.collect();
let request_messages = snapshot_messages_for_tool_iteration(&messages);
let request = ToolCompletionRequest::new(request_messages, tool_defs)
@@ -1486,7 +1551,7 @@ async fn execute_lightweight_with_tools(
// Execute tools sequentially
for tc in response.tool_calls {
let result = execute_routine_tool(ctx, &job_ctx, &tc).await;
let result = execute_routine_tool(ctx, &job_ctx, &allowed_tools, &tc).await;
// Sanitize and wrap result (including errors)
let result_content = match result {
@@ -1555,31 +1620,16 @@ fn snapshot_messages_for_tool_iteration(messages: &[ChatMessage]) -> Vec<ChatMes
snapshot
}
/// Tools that must never be callable from lightweight routines.
///
/// These tools pose autonomy-escalation risks: a routine could self-replicate,
/// modify its own triggers/prompts, delete other routines, or restart the agent.
const ROUTINE_TOOL_DENYLIST: &[&str] = &[
"routine_create",
"routine_update",
"routine_delete",
"routine_fire",
"restart",
];
/// Execute a single tool for a lightweight routine.
async fn execute_routine_tool(
ctx: &EngineContext,
job_ctx: &JobContext,
allowed_tools: &std::collections::HashSet<String>,
tc: &ToolCall,
) -> Result<String, Box<dyn std::error::Error + Send + Sync>> {
// Block tools that pose autonomy-escalation risks
if ROUTINE_TOOL_DENYLIST.contains(&tc.name.as_str()) {
return Err(format!(
"Tool '{}' is not available in lightweight routines",
tc.name
)
.into());
if !allowed_tools.contains(&tc.name) {
let message = autonomous_unavailable_message(&tc.name, &job_ctx.user_id);
return Err(message.into());
}
// Check if tool exists
@@ -1590,22 +1640,6 @@ async fn execute_routine_tool(
.ok_or_else(|| format!("Tool '{}' not found", tc.name))?;
let normalized_params = prepare_tool_params(tool.as_ref(), &tc.arguments);
// Check approval requirement: only allow Never tools in lightweight routines.
// UnlessAutoApproved and Always tools are blocked to prevent prompt injection attacks.
// Lightweight routines can be triggered by external events and may process untrusted data,
// making them vulnerable to prompt injection that could trick the LLM into calling
// sensitive tools. Blocking these tools entirely is the safest approach.
match tool.requires_approval(&normalized_params) {
ApprovalRequirement::Never => {}
ApprovalRequirement::UnlessAutoApproved | ApprovalRequirement::Always => {
return Err(format!(
"Tool '{}' requires manual approval and cannot be used in lightweight routines",
tc.name
)
.into());
}
}
// Validate tool parameters
let validation = ctx
.safety
@@ -2021,8 +2055,8 @@ mod tests {
];
for tool in &denylisted {
assert!(
super::ROUTINE_TOOL_DENYLIST.contains(tool),
"Tool '{}' should be in ROUTINE_TOOL_DENYLIST",
crate::tools::AUTONOMOUS_TOOL_DENYLIST.contains(tool),
"Tool '{}' should be in AUTONOMOUS_TOOL_DENYLIST",
tool
);
}
@@ -2033,8 +2067,8 @@ mod tests {
let allowed = vec!["echo", "time", "json", "http", "memory_search", "shell"];
for tool in &allowed {
assert!(
!super::ROUTINE_TOOL_DENYLIST.contains(tool),
"Tool '{}' should NOT be in ROUTINE_TOOL_DENYLIST",
!crate::tools::AUTONOMOUS_TOOL_DENYLIST.contains(tool),
"Tool '{}' should NOT be in AUTONOMOUS_TOOL_DENYLIST",
tool
);
}
+58 -19
View File
@@ -14,10 +14,14 @@ use crate::config::AgentConfig;
use crate::context::{ContextManager, JobContext, JobState};
use crate::db::Database;
use crate::error::{Error, JobError};
use crate::extensions::ExtensionManager;
use crate::hooks::HookRegistry;
use crate::llm::LlmProvider;
use crate::safety::SafetyLayer;
use crate::tools::{ApprovalContext, ToolRegistry, prepare_tool_params};
use crate::tools::{
ApprovalContext, ToolRegistry, autonomous_allowed_tool_names, autonomous_unavailable_error,
prepare_tool_params,
};
use crate::worker::job::{Worker, WorkerDeps};
/// Message to send to a worker.
@@ -45,6 +49,14 @@ struct ScheduledSubtask {
handle: JoinHandle<Result<TaskOutput, Error>>,
}
/// Shared scheduler-owned dependencies that are forwarded into autonomous runs.
pub struct SchedulerDeps {
pub tools: Arc<ToolRegistry>,
pub extension_manager: Option<Arc<ExtensionManager>>,
pub store: Option<Arc<dyn Database>>,
pub hooks: Arc<HookRegistry>,
}
/// Schedules and manages parallel job execution.
pub struct Scheduler {
config: AgentConfig,
@@ -52,6 +64,7 @@ pub struct Scheduler {
llm: Arc<dyn LlmProvider>,
safety: Arc<SafetyLayer>,
tools: Arc<ToolRegistry>,
extension_manager: Option<Arc<ExtensionManager>>,
store: Option<Arc<dyn Database>>,
hooks: Arc<HookRegistry>,
/// SSE broadcast sender for live job event streaming.
@@ -71,18 +84,17 @@ impl Scheduler {
context_manager: Arc<ContextManager>,
llm: Arc<dyn LlmProvider>,
safety: Arc<SafetyLayer>,
tools: Arc<ToolRegistry>,
store: Option<Arc<dyn Database>>,
hooks: Arc<HookRegistry>,
deps: SchedulerDeps,
) -> Self {
Self {
config,
context_manager,
llm,
safety,
tools,
store,
hooks,
tools: deps.tools,
extension_manager: deps.extension_manager,
store: deps.store,
hooks: deps.hooks,
sse_tx: None,
http_interceptor: None,
jobs: Arc::new(RwLock::new(HashMap::new())),
@@ -120,14 +132,21 @@ impl Scheduler {
description: &str,
metadata: Option<serde_json::Value>,
) -> Result<Uuid, JobError> {
self.dispatch_job_inner(user_id, title, description, metadata, None)
.await
let approval_context = self.autonomous_approval_context(user_id).await;
self.dispatch_job_inner(
user_id,
title,
description,
metadata,
Some(approval_context),
)
.await
}
/// Dispatch a job with an explicit approval context for autonomous execution.
///
/// Same as `dispatch_job`, but the worker will use the given `ApprovalContext`
/// to determine which tools are pre-approved (instead of blocking all non-`Never` tools).
/// to determine the explicit autonomous allowlist for that job.
pub async fn dispatch_job_with_context(
&self,
user_id: &str,
@@ -216,6 +235,13 @@ impl Scheduler {
Ok(job_id)
}
async fn autonomous_approval_context(&self, user_id: &str) -> ApprovalContext {
ApprovalContext::autonomous_with_tools(
autonomous_allowed_tool_names(&self.tools, self.extension_manager.as_ref(), user_id)
.await,
)
}
/// Schedule a job for execution.
pub async fn schedule(&self, job_id: Uuid) -> Result<(), JobError> {
self.schedule_with_context(job_id, None).await
@@ -518,10 +544,7 @@ impl Scheduler {
let blocked =
ApprovalContext::is_blocked_or_default(&approval_context, tool_name, requirement);
if blocked {
return Err(crate::error::ToolError::AuthRequired {
name: tool_name.to_string(),
}
.into());
return Err(autonomous_unavailable_error(tool_name, &job_ctx.user_id).into());
}
// Delegate to shared tool execution pipeline
@@ -776,7 +799,18 @@ mod tests {
let tools = Arc::new(ToolRegistry::new());
let hooks = Arc::new(HookRegistry::default());
Scheduler::new(config, cm, llm, safety, tools, None, hooks)
Scheduler::new(
config,
cm,
llm,
safety,
SchedulerDeps {
tools,
extension_manager: None,
store: None,
hooks,
},
)
}
#[tokio::test]
@@ -1003,12 +1037,14 @@ mod tests {
async fn test_execute_tool_task_autonomous_unblocks_soft() {
let (tools, cm, safety, job_id) = setup_tools_and_job().await;
// Autonomous context auto-approves UnlessAutoApproved
// Autonomous execution only allows tools explicitly in scope.
let result = Scheduler::execute_tool_task(
tools.clone(),
cm.clone(),
safety.clone(),
Some(ApprovalContext::autonomous()),
Some(ApprovalContext::autonomous_with_tools([
"soft_gate".to_string()
])),
job_id,
"soft_gate",
serde_json::json!({}),
@@ -1040,8 +1076,11 @@ mod tests {
async fn test_execute_tool_task_autonomous_with_permissions() {
let (tools, cm, safety, job_id) = setup_tools_and_job().await;
// Autonomous context with explicit permission for hard_gate
let ctx = ApprovalContext::autonomous_with_tools(["hard_gate".to_string()]);
// Autonomous context with explicit permission for both tools.
let ctx = ApprovalContext::autonomous_with_tools([
"soft_gate".to_string(),
"hard_gate".to_string(),
]);
let result = Scheduler::execute_tool_task(
tools.clone(),
+27
View File
@@ -772,6 +772,33 @@ mod tests {
assert_ne!(resolved, tid);
}
#[tokio::test]
async fn test_register_then_resolve_same_uuid_on_second_channel_reuses_thread() {
use crate::agent::session::{Session, Thread};
let manager = SessionManager::new();
let tid = Uuid::new_v4();
let session = Arc::new(Mutex::new(Session::new("user-cross")));
{
let mut sess = session.lock().await;
let thread = Thread::with_id(tid, sess.id);
sess.threads.insert(tid, thread);
}
manager
.register_thread("user-cross", "http", tid, Arc::clone(&session))
.await;
manager
.register_thread("user-cross", "gateway", tid, Arc::clone(&session))
.await;
let (_, resolved) = manager
.resolve_thread("user-cross", "gateway", Some(&tid.to_string()))
.await;
assert_eq!(resolved, tid);
}
// === QA Plan P3 - 4.2: Concurrent session stress tests ===
#[tokio::test]
+44 -5
View File
@@ -312,15 +312,23 @@ impl AppBuilder {
.create_provider(&self.config.llm.nearai.base_url, self.session.clone());
// Register memory tools if database is available
let workspace_user_id = self
.config
.channels
.gateway
.as_ref()
.map(|gw| gw.user_id.as_str())
.unwrap_or("default");
let workspace = if let Some(ref db) = self.db {
let emb_cache_config = EmbeddingCacheConfig {
max_entries: self.config.embeddings.cache_size,
};
let mut ws = Workspace::new_with_db(&self.config.owner_id, db.clone())
let mut ws = Workspace::new_with_db(workspace_user_id, db.clone())
.with_search_config(&self.config.search);
if let Some(ref emb) = embeddings {
ws = ws.with_embeddings_cached(emb.clone(), emb_cache_config);
}
ws = ws.with_memory_layers(self.config.workspace.memory_layers.clone());
let ws = Arc::new(ws);
tools.register_memory_tools(Arc::clone(&ws));
Some(ws)
@@ -528,7 +536,7 @@ impl AppBuilder {
server_name,
e
);
return;
return None;
}
};
@@ -545,6 +553,10 @@ impl AppBuilder {
tool_count,
server_name
);
return Some((
server_name,
Arc::new(client),
));
}
Err(e) => {
tracing::warn!(
@@ -575,14 +587,27 @@ impl AppBuilder {
}
}
}
None
});
}
let mut startup_clients = Vec::new();
while let Some(result) = join_set.join_next().await {
if let Err(e) = result {
tracing::warn!("MCP server loading task panicked: {}", e);
match result {
Ok(Some(client_pair)) => {
startup_clients.push(client_pair);
}
Ok(None) => {}
Err(e) => {
if e.is_panic() {
tracing::error!("MCP server loading task panicked: {}", e);
} else {
tracing::warn!("MCP server loading task failed: {}", e);
}
}
}
}
return startup_clients;
}
Err(e) => {
if matches!(
@@ -600,10 +625,12 @@ impl AppBuilder {
}
}
}
Vec::new()
}
};
let (dev_loaded_tool_names, _) = tokio::join!(wasm_tools_future, mcp_servers_future);
let (dev_loaded_tool_names, startup_mcp_clients) =
tokio::join!(wasm_tools_future, mcp_servers_future);
// Load registry catalog entries for extension discovery
let mut catalog_entries = match crate::registry::RegistryCatalog::load_or_embedded() {
@@ -665,6 +692,17 @@ impl AppBuilder {
));
tools.register_extension_tools(Arc::clone(&manager));
tracing::debug!("Extension manager initialized with in-chat discovery tools");
if !startup_mcp_clients.is_empty() {
tracing::info!(
count = startup_mcp_clients.len(),
"Injecting startup MCP clients into extension manager"
);
for (name, client) in startup_mcp_clients {
manager.inject_mcp_client(name, client).await;
}
}
Some(manager)
};
@@ -696,6 +734,7 @@ impl AppBuilder {
// fail early with a clear error instead of a confusing runtime failure.
if self.config.llm.backend != "nearai"
&& self.config.llm.backend != "bedrock"
&& self.config.llm.backend != "openai_codex"
&& self.config.llm.provider.is_none()
{
let backend = &self.config.llm.backend;
+2 -19
View File
@@ -123,25 +123,8 @@ pub async fn memory_read_handler(
}))
}
pub async fn memory_write_handler(
State(state): State<Arc<GatewayState>>,
Json(req): Json<MemoryWriteRequest>,
) -> Result<Json<MemoryWriteResponse>, (StatusCode, String)> {
let workspace = state.workspace.as_ref().ok_or((
StatusCode::SERVICE_UNAVAILABLE,
"Workspace not available".to_string(),
))?;
workspace
.write(&req.path, &req.content)
.await
.map_err(|e| (StatusCode::INTERNAL_SERVER_ERROR, e.to_string()))?;
Ok(Json(MemoryWriteResponse {
path: req.path,
status: "written",
}))
}
// memory_write_handler lives in server.rs (layer-aware version with append,
// privacy redirect, and proper error status codes).
pub async fn memory_search_handler(
State(state): State<Arc<GatewayState>>,
+1
View File
@@ -26,3 +26,4 @@ pub mod routines;
pub mod settings;
#[allow(dead_code)]
pub mod static_files;
pub mod webhooks;
+4 -45
View File
@@ -10,29 +10,11 @@ use axum::{
use serde::Deserialize;
use uuid::Uuid;
use crate::agent::routine::{
FullJobPermissionDefaultMode, FullJobPermissionMode, RoutineAction, Trigger,
effective_full_job_tool_permissions, load_full_job_permission_settings, next_cron_fire,
};
use crate::agent::routine::{Trigger, next_cron_fire};
use crate::channels::web::server::GatewayState;
use crate::channels::web::types::*;
use crate::error::RoutineError;
fn permission_mode_label(mode: FullJobPermissionMode) -> String {
match mode {
FullJobPermissionMode::Explicit => "explicit".to_string(),
FullJobPermissionMode::InheritOwner => "inherit_owner".to_string(),
}
}
fn default_permission_mode_label(mode: FullJobPermissionDefaultMode) -> String {
match mode {
FullJobPermissionDefaultMode::Explicit => "explicit".to_string(),
FullJobPermissionDefaultMode::InheritOwner => "inherit_owner".to_string(),
FullJobPermissionDefaultMode::CopyOwner => "copy_owner".to_string(),
}
}
pub async fn routines_list_handler(
State(state): State<Arc<GatewayState>>,
) -> Result<Json<RoutineListResponse>, (StatusCode, String)> {
@@ -131,30 +113,6 @@ pub async fn routines_detail_handler(
})
.collect();
let routine_info = RoutineInfo::from_routine(&routine);
let full_job_permissions = match &routine.action {
RoutineAction::FullJob {
tool_permissions,
permission_mode,
..
} => {
let owner_settings =
load_full_job_permission_settings(store.as_ref(), &routine.user_id)
.await
.map_err(|e| (StatusCode::INTERNAL_SERVER_ERROR, e.to_string()))?;
Some(FullJobPermissionInfo {
permission_mode: permission_mode_label(*permission_mode),
default_permission_mode: default_permission_mode_label(owner_settings.default_mode),
stored_tool_permissions: tool_permissions.clone(),
effective_tool_permissions: effective_full_job_tool_permissions(
*permission_mode,
tool_permissions,
&owner_settings.owner_allowed_tools,
),
owner_allowed_tools: owner_settings.owner_allowed_tools,
})
}
RoutineAction::Lightweight { .. } => None,
};
Ok(Json(RoutineDetailResponse {
id: routine.id,
@@ -173,7 +131,6 @@ pub async fn routines_detail_handler(
run_count: routine.run_count,
consecutive_failures: routine.consecutive_failures,
created_at: routine.created_at.to_rfc3339(),
full_job_permissions,
recent_runs,
}))
}
@@ -346,7 +303,9 @@ fn routine_error_status(err: &RoutineError) -> StatusCode {
match err {
RoutineError::NotFound { .. } => StatusCode::NOT_FOUND,
RoutineError::NotAuthorized { .. } => StatusCode::FORBIDDEN,
RoutineError::Disabled { .. } | RoutineError::MaxConcurrent { .. } => StatusCode::CONFLICT,
RoutineError::Disabled { .. }
| RoutineError::Cooldown { .. }
| RoutineError::MaxConcurrent { .. } => StatusCode::CONFLICT,
_ => StatusCode::INTERNAL_SERVER_ERROR,
}
}
+197
View File
@@ -0,0 +1,197 @@
//! Public webhook trigger endpoint for routine webhook triggers.
//!
//! `POST /api/webhooks/{path}` — matches the path against routines with
//! `Trigger::Webhook { path, secret }`, validates the secret via constant-time
//! comparison, and fires the matching routine through the `RoutineEngine`.
use std::sync::Arc;
use axum::{
Json,
extract::{Path, State},
http::{HeaderMap, StatusCode},
};
use subtle::ConstantTimeEq;
use crate::agent::routine::Trigger;
use crate::channels::web::server::GatewayState;
/// Validate the webhook secret for a routine.
///
/// Returns `Ok(())` if the routine has a configured secret and the provided
/// secret matches via constant-time comparison. Returns an appropriate HTTP
/// error if the secret is missing (403) or invalid (401).
fn validate_webhook_secret(
trigger: &Trigger,
provided_secret: &str,
) -> Result<(), (StatusCode, String)> {
// Require webhook secret — routines without a secret cannot be triggered via webhook
let expected_secret = match trigger {
Trigger::Webhook {
secret: Some(s), ..
} => s,
_ => {
return Err((
StatusCode::FORBIDDEN,
"Webhook secret not configured for this routine. \
Set a secret with: ironclaw routine update <id> --webhook-secret <secret>"
.to_string(),
));
}
};
if !bool::from(provided_secret.as_bytes().ct_eq(expected_secret.as_bytes())) {
return Err((
StatusCode::UNAUTHORIZED,
"Invalid webhook secret".to_string(),
));
}
Ok(())
}
/// Handle incoming webhook POST to `/api/webhooks/{path}`.
///
/// This endpoint is **public** (no gateway auth token required) but protected
/// by the per-routine webhook secret sent via the `X-Webhook-Secret` header.
pub async fn webhook_trigger_handler(
State(state): State<Arc<GatewayState>>,
Path(path): Path<String>,
headers: HeaderMap,
) -> Result<Json<serde_json::Value>, (StatusCode, String)> {
// Rate limit check
if !state.webhook_rate_limiter.check() {
return Err((
StatusCode::TOO_MANY_REQUESTS,
"Rate limit exceeded. Try again shortly.".to_string(),
));
}
let store = state.store.as_ref().ok_or((
StatusCode::SERVICE_UNAVAILABLE,
"Database not available".to_string(),
))?;
// Targeted query instead of loading all routines
let routine = store
.get_webhook_routine_by_path(&path)
.await
.map_err(|e| (StatusCode::INTERNAL_SERVER_ERROR, e.to_string()))?
.ok_or((
StatusCode::NOT_FOUND,
"No routine matches this webhook path".to_string(),
))?;
let provided_secret = headers
.get("x-webhook-secret")
.and_then(|v| v.to_str().ok())
.unwrap_or("");
validate_webhook_secret(&routine.trigger, provided_secret)?;
// Fire through the RoutineEngine so guardrails, run tracking,
// notifications, and FullJob dispatch all work correctly.
let engine = {
let guard = state.routine_engine.read().await;
guard.as_ref().cloned().ok_or((
StatusCode::SERVICE_UNAVAILABLE,
"Routine engine not available".to_string(),
))?
};
let run_id = engine.fire_webhook(routine.id, &path).await.map_err(|e| {
let status = match &e {
crate::error::RoutineError::NotFound { .. } => StatusCode::NOT_FOUND,
crate::error::RoutineError::Disabled { .. }
| crate::error::RoutineError::Cooldown { .. }
| crate::error::RoutineError::MaxConcurrent { .. } => StatusCode::CONFLICT,
_ => StatusCode::INTERNAL_SERVER_ERROR,
};
(status, e.to_string())
})?;
Ok(Json(serde_json::json!({
"status": "triggered",
"routine_id": routine.id,
"routine_name": routine.name,
"run_id": run_id,
})))
}
#[cfg(test)]
mod tests {
use super::*;
/// Routines with `secret: None` must be rejected with 403.
#[test]
fn test_validate_rejects_missing_secret() {
let trigger = Trigger::Webhook {
path: Some("my-hook".to_string()),
secret: None,
};
let result = validate_webhook_secret(&trigger, "any-secret");
let (status, msg) = result.unwrap_err();
assert_eq!(status, StatusCode::FORBIDDEN);
assert!(
msg.contains("not configured"),
"Error should tell user to configure a secret, got: {msg}"
);
}
/// Non-webhook triggers must be rejected with 403.
#[test]
fn test_validate_rejects_non_webhook_trigger() {
let trigger = Trigger::Manual;
let result = validate_webhook_secret(&trigger, "any-secret");
let (status, _) = result.unwrap_err();
assert_eq!(status, StatusCode::FORBIDDEN);
}
/// Correct secret passes validation.
#[test]
fn test_validate_accepts_correct_secret() {
let trigger = Trigger::Webhook {
path: Some("my-hook".to_string()),
secret: Some("s3cret-token".to_string()),
};
assert!(validate_webhook_secret(&trigger, "s3cret-token").is_ok());
}
/// Wrong secret returns 401.
#[test]
fn test_validate_rejects_wrong_secret() {
let trigger = Trigger::Webhook {
path: Some("my-hook".to_string()),
secret: Some("correct-secret".to_string()),
};
let result = validate_webhook_secret(&trigger, "wrong-secret");
let (status, msg) = result.unwrap_err();
assert_eq!(status, StatusCode::UNAUTHORIZED);
assert!(msg.contains("Invalid"), "Expected 'Invalid' in: {msg}");
}
/// Empty provided secret returns 401 (not a false positive).
#[test]
fn test_validate_rejects_empty_provided_secret() {
let trigger = Trigger::Webhook {
path: Some("my-hook".to_string()),
secret: Some("real-secret".to_string()),
};
let result = validate_webhook_secret(&trigger, "");
let (status, _) = result.unwrap_err();
assert_eq!(status, StatusCode::UNAUTHORIZED);
}
/// Constant-time comparison: secrets of different lengths are still rejected
/// (not short-circuited in a way that leaks length info).
#[test]
fn test_validate_rejects_different_length_secret() {
let trigger = Trigger::Webhook {
path: None,
secret: Some("short".to_string()),
};
let result = validate_webhook_secret(&trigger, "a-much-longer-secret-value");
let (status, _) = result.unwrap_err();
assert_eq!(status, StatusCode::UNAUTHORIZED);
}
}
+2
View File
@@ -98,6 +98,7 @@ impl GatewayChannel {
skill_catalog: None,
chat_rate_limiter: server::RateLimiter::new(30, 60),
oauth_rate_limiter: server::RateLimiter::new(10, 60),
webhook_rate_limiter: server::RateLimiter::new(10, 60),
registry_entries: Vec::new(),
cost_guard: None,
routine_engine: Arc::new(tokio::sync::RwLock::new(None)),
@@ -136,6 +137,7 @@ impl GatewayChannel {
skill_catalog: self.state.skill_catalog.clone(),
chat_rate_limiter: server::RateLimiter::new(30, 60),
oauth_rate_limiter: server::RateLimiter::new(10, 60),
webhook_rate_limiter: server::RateLimiter::new(10, 60),
registry_entries: self.state.registry_entries.clone(),
cost_guard: self.state.cost_guard.clone(),
routine_engine: Arc::clone(&self.state.routine_engine),
+51 -5
View File
@@ -194,6 +194,8 @@ pub struct GatewayState {
pub chat_rate_limiter: RateLimiter,
/// Rate limiter for OAuth callback endpoints (10 requests per 60 seconds).
pub oauth_rate_limiter: RateLimiter,
/// Rate limiter for webhook trigger endpoints (10 requests per 60 seconds).
pub webhook_rate_limiter: RateLimiter,
/// Registry catalog entries for the available extensions API.
/// Populated at startup from `registry/` manifests, independent of extension manager.
pub registry_entries: Vec<crate::extensions::RegistryEntry>,
@@ -237,7 +239,11 @@ pub async fn start_server(
"/oauth/slack/callback",
get(slack_relay_oauth_callback_handler),
)
.route("/relay/events", post(relay_events_handler));
.route("/relay/events", post(relay_events_handler))
.route(
"/api/webhooks/{path}",
post(crate::channels::web::handlers::webhooks::webhook_trigger_handler),
);
// Protected routes (require auth)
let auth_state = AuthState { token: auth_token };
@@ -1826,14 +1832,53 @@ async fn memory_write_handler(
"Workspace not available".to_string(),
))?;
workspace
.write(&req.path, &req.content)
.await
.map_err(|e| (StatusCode::INTERNAL_SERVER_ERROR, e.to_string()))?;
// Route through layer-aware methods when a layer is specified
if let Some(ref layer_name) = req.layer {
let result = if req.append {
workspace
.append_to_layer(layer_name, &req.path, &req.content, req.force)
.await
} else {
workspace
.write_to_layer(layer_name, &req.path, &req.content, req.force)
.await
}
.map_err(|e| {
use crate::error::WorkspaceError;
let status = match &e {
WorkspaceError::LayerNotFound { .. } => StatusCode::BAD_REQUEST,
WorkspaceError::LayerReadOnly { .. } => StatusCode::FORBIDDEN,
WorkspaceError::PrivacyRedirectFailed => StatusCode::UNPROCESSABLE_ENTITY,
_ => StatusCode::INTERNAL_SERVER_ERROR,
};
(status, e.to_string())
})?;
return Ok(Json(MemoryWriteResponse {
path: req.path,
status: "written",
redirected: Some(result.redirected),
actual_layer: Some(result.actual_layer),
}));
}
// Non-layer path: honor the append field
if req.append {
workspace
.append(&req.path, &req.content)
.await
.map_err(|e| (StatusCode::INTERNAL_SERVER_ERROR, e.to_string()))?;
} else {
workspace
.write(&req.path, &req.content)
.await
.map_err(|e| (StatusCode::INTERNAL_SERVER_ERROR, e.to_string()))?;
}
Ok(Json(MemoryWriteResponse {
path: req.path,
status: "written",
redirected: None,
actual_layer: None,
}))
}
@@ -3000,6 +3045,7 @@ mod tests {
scheduler: None,
chat_rate_limiter: RateLimiter::new(30, 60),
oauth_rate_limiter: RateLimiter::new(10, 60),
webhook_rate_limiter: RateLimiter::new(10, 60),
registry_entries: vec![],
cost_guard: None,
routine_engine: Arc::new(tokio::sync::RwLock::new(None)),
-16
View File
@@ -3942,18 +3942,6 @@ function renderRoutineDetail(routine) {
+ '<pre class="action-json">' + escapeHtml(JSON.stringify(routine.trigger, null, 2)) + '</pre></div>';
}
// Action config
if (routine.full_job_permissions) {
html += '<div class="job-description"><h3>Full Job Permissions</h3>'
+ '<div class="job-meta-grid">'
+ metaItem('Mode', routine.full_job_permissions.permission_mode)
+ metaItem('Owner Default', routine.full_job_permissions.default_permission_mode)
+ metaItem('Inherited Tools', (routine.full_job_permissions.owner_allowed_tools || []).join(', ') || '-')
+ metaItem('Stored Tools', (routine.full_job_permissions.stored_tool_permissions || []).join(', ') || '-')
+ metaItem('Effective Tools', (routine.full_job_permissions.effective_tool_permissions || []).join(', ') || '-')
+ '</div></div>';
}
html += '<div class="job-description"><h3>Action</h3>'
+ '<pre class="action-json">' + escapeHtml(JSON.stringify(routine.action, null, 2)) + '</pre></div>';
@@ -4788,10 +4776,6 @@ var AGENT_SETTINGS = [
settings: [
{ key: 'routines.max_concurrent', label: 'cfg.routines_max_concurrent.label', description: 'cfg.routines_max_concurrent.desc', type: 'number', min: 0 },
{ key: 'routines.default_cooldown_secs', label: 'cfg.routines_cooldown.label', description: 'cfg.routines_cooldown.desc', type: 'number', min: 0 },
{ key: 'routines.full_job_default_permission_mode', label: 'cfg.routines_full_job_default_mode.label', description: 'cfg.routines_full_job_default_mode.desc',
type: 'select', options: ['inherit_owner', 'explicit', 'copy_owner'] },
{ key: 'routines.full_job_owner_allowed_tools', label: 'cfg.routines_full_job_owner_tools.label', description: 'cfg.routines_full_job_owner_tools.desc',
type: 'list', placeholder: 'shell, http' },
]
},
{
-4
View File
@@ -522,10 +522,6 @@ I18n.register('en', {
'cfg.routines_max_concurrent.desc': 'Maximum routines running simultaneously',
'cfg.routines_cooldown.label': 'Default Cooldown',
'cfg.routines_cooldown.desc': 'Minimum seconds between routine fires',
'cfg.routines_full_job_default_mode.label': 'Full Job Default Mode',
'cfg.routines_full_job_default_mode.desc': 'Default permission behavior for new full_job routines. When unset, inherit_owner is used.',
'cfg.routines_full_job_owner_tools.label': 'Full Job Owner Allowlist',
'cfg.routines_full_job_owner_tools.desc': 'Comma-separated tool names that full_job routines may inherit at run time.',
// Safety settings
'cfg.safety_max_output.label': 'Max Output Length',
-4
View File
@@ -521,10 +521,6 @@ I18n.register('zh-CN', {
'cfg.routines_max_concurrent.desc': '同时运行的最大定时任务数',
'cfg.routines_cooldown.label': '默认冷却时间',
'cfg.routines_cooldown.desc': '定时任务触发间的最小秒数',
'cfg.routines_full_job_default_mode.label': '完整任务默认权限模式',
'cfg.routines_full_job_default_mode.desc': '新建 full_job 定时任务的默认权限行为。未设置时使用 inherit_owner。',
'cfg.routines_full_job_owner_tools.label': '完整任务所有者允许工具',
'cfg.routines_full_job_owner_tools.desc': '逗号分隔的工具名列表,full_job 定时任务可在运行时继承这些工具权限。',
// 安全设置
'cfg.safety_max_output.label': '最大输出长度',
+1
View File
@@ -83,6 +83,7 @@ impl TestGatewayBuilder {
scheduler: None,
chat_rate_limiter: RateLimiter::new(30, 60),
oauth_rate_limiter: RateLimiter::new(10, 60),
webhook_rate_limiter: RateLimiter::new(10, 60),
registry_entries: Vec::new(),
cost_guard: None,
routine_engine: Arc::new(tokio::sync::RwLock::new(None)),
+26 -11
View File
@@ -302,12 +302,30 @@ pub struct MemoryReadResponse {
pub struct MemoryWriteRequest {
pub path: String,
pub content: String,
/// Optional layer to write to. When present, uses `write_to_layer()`
/// which enables privacy classification and redirect.
pub layer: Option<String>,
/// When true and a layer is specified, appends to existing content
/// instead of replacing it.
#[serde(default)]
pub append: bool,
/// Skip privacy classification and write directly to the specified layer.
#[serde(default)]
pub force: bool,
}
#[derive(Debug, Serialize)]
pub struct MemoryWriteResponse {
pub path: String,
pub status: &'static str,
/// Whether the write was redirected to a different layer (e.g., sensitive
/// content redirected from shared to private).
#[serde(skip_serializing_if = "Option::is_none")]
pub redirected: Option<bool>,
/// The layer the content was actually written to (may differ from requested
/// layer if privacy redirect occurred).
#[serde(skip_serializing_if = "Option::is_none")]
pub actual_layer: Option<String>,
}
#[derive(Debug, Deserialize)]
@@ -814,6 +832,14 @@ impl RoutineInfo {
String::new(),
format!("event: {}.{}", source, event_type),
),
crate::agent::routine::Trigger::Webhook { path, .. } => {
let p = path.as_deref().unwrap_or("default");
(
"webhook".to_string(),
String::new(),
format!("webhook: /api/webhooks/{}", p),
)
}
crate::agent::routine::Trigger::Manual => (
"manual".to_string(),
String::new(),
@@ -884,20 +910,9 @@ pub struct RoutineDetailResponse {
pub run_count: u64,
pub consecutive_failures: u32,
pub created_at: String,
#[serde(skip_serializing_if = "Option::is_none")]
pub full_job_permissions: Option<FullJobPermissionInfo>,
pub recent_runs: Vec<RoutineRunInfo>,
}
#[derive(Debug, Serialize)]
pub struct FullJobPermissionInfo {
pub permission_mode: String,
pub default_permission_mode: String,
pub stored_tool_permissions: Vec<String>,
pub owner_allowed_tools: Vec<String>,
pub effective_tool_permissions: Vec<String>,
}
#[derive(Debug, Serialize)]
pub struct RoutineRunInfo {
pub id: Uuid,
+1
View File
@@ -517,6 +517,7 @@ mod tests {
skill_catalog: None,
chat_rate_limiter: crate::channels::web::server::RateLimiter::new(30, 60),
oauth_rate_limiter: crate::channels::web::server::RateLimiter::new(10, 60),
webhook_rate_limiter: crate::channels::web::server::RateLimiter::new(10, 60),
registry_entries: Vec::new(),
cost_guard: None,
routine_engine: Arc::new(tokio::sync::RwLock::new(None)),
+11
View File
@@ -239,6 +239,17 @@ pub enum Command {
)]
Import(ImportCommand),
/// Authenticate with a provider (re-login)
#[command(
about = "Authenticate with a provider",
long_about = "Re-authenticate with an LLM provider.\nExample: ironclaw login --openai-codex"
)]
Login {
/// Authenticate with OpenAI Codex (ChatGPT subscription)
#[arg(long)]
openai_codex: bool,
},
/// Run as a sandboxed worker inside a Docker container (internal use).
/// This is invoked automatically by the orchestrator, not by users directly.
#[command(hide = true)]
@@ -24,6 +24,7 @@ Commands:
status Show system status
completion Generate completions
import Import from other AI systems
login Authenticate with a provider
help Print this message or the help of the given subcommand(s)
Options:
@@ -23,6 +23,7 @@ Commands:
logs View and manage gateway logs
status Show system status
completion Generate completions
login Authenticate with a provider
help Print this message or the help of the given subcommand(s)
Options:
@@ -27,6 +27,7 @@ Commands:
status Show system status
completion Generate completions
import Import from other AI systems
login Authenticate with a provider
help Print this message or the help of the given subcommand(s)
Options:
@@ -26,6 +26,7 @@ Commands:
logs View and manage gateway logs
status Show system status
completion Generate completions
login Authenticate with a provider
help Print this message or the help of the given subcommand(s)
Options:
+5 -1
View File
@@ -111,6 +111,10 @@ impl ChannelsConfig {
let gateway_enabled = parse_bool_env("GATEWAY_ENABLED", cs.gateway_enabled)?;
let gateway = if gateway_enabled {
let user_id = optional_env("GATEWAY_USER_ID")?
.or_else(|| cs.gateway_user_id.clone())
.unwrap_or_else(|| "default".to_string());
Some(GatewayConfig {
host: optional_env("GATEWAY_HOST")?
.or_else(|| cs.gateway_host.clone())
@@ -121,7 +125,7 @@ impl ChannelsConfig {
)?,
auth_token: optional_env("GATEWAY_AUTH_TOKEN")?
.or_else(|| cs.gateway_auth_token.clone()),
user_id: owner_id.to_string(),
user_id,
})
} else {
None
+2 -5
View File
@@ -299,15 +299,12 @@ mod tests {
// SAFETY: Under ENV_MUTEX, no concurrent env access.
unsafe {
std::env::set_var("EMBEDDING_BASE_URL", "https://custom.example.com");
std::env::set_var("EMBEDDING_BASE_URL", "https://8.8.8.8");
}
let settings = Settings::default();
let config = EmbeddingsConfig::resolve(&settings).expect("resolve should succeed");
assert_eq!(
config.openai_base_url.as_deref(),
Some("https://custom.example.com")
);
assert_eq!(config.openai_base_url.as_deref(), Some("https://8.8.8.8"));
// SAFETY: Under ENV_MUTEX.
unsafe {
std::env::remove_var("EMBEDDING_BASE_URL");
+300 -7
View File
@@ -37,6 +37,7 @@ impl LlmConfig {
},
provider: None,
bedrock: None,
openai_codex: None,
request_timeout_secs: 120,
cheap_model: None,
smart_routing_cascade: false,
@@ -79,6 +80,9 @@ impl LlmConfig {
backend_lower == "nearai" || backend_lower == "near_ai" || backend_lower == "near";
let is_bedrock =
backend_lower == "bedrock" || backend_lower == "aws_bedrock" || backend_lower == "aws";
let is_openai_codex = backend_lower == "openai_codex"
|| backend_lower == "openai-codex"
|| backend_lower == "codex";
// Check custom providers defined
let custom_provider = settings
@@ -88,6 +92,7 @@ impl LlmConfig {
if !is_nearai
&& !is_bedrock
&& !is_openai_codex
&& custom_provider.is_none()
&& registry.find(&backend_lower).is_none()
{
@@ -143,8 +148,8 @@ impl LlmConfig {
smart_routing_cascade: parse_optional_env("SMART_ROUTING_CASCADE", true)?,
};
// Resolve registry provider config (for non-NearAI, non-Bedrock backends)
let provider = if is_nearai || is_bedrock {
// Resolve registry provider config (for non-NearAI, non-Bedrock, non-Codex backends)
let provider = if is_nearai || is_bedrock || is_openai_codex {
None
} else if let Some(custom) = custom_provider {
Some(Self::resolve_custom_provider(custom, settings)?)
@@ -193,6 +198,38 @@ impl LlmConfig {
None
};
// Resolve OpenAI Codex config
let openai_codex = if is_openai_codex {
// Model: OPENAI_CODEX_MODEL > OPENAI_MODEL > settings.selected_model > default
let model = optional_env("OPENAI_CODEX_MODEL")?
.or(optional_env("OPENAI_MODEL")?)
.or_else(|| settings.selected_model.clone())
.unwrap_or_else(|| "gpt-5.3-codex".to_string());
let auth_endpoint = optional_env("OPENAI_CODEX_AUTH_URL")?
.unwrap_or_else(|| "https://auth.openai.com".to_string());
validate_base_url(&auth_endpoint, "OPENAI_CODEX_AUTH_URL")?;
let api_base_url = optional_env("OPENAI_CODEX_API_URL")?
.unwrap_or_else(|| "https://chatgpt.com/backend-api/codex".to_string());
validate_base_url(&api_base_url, "OPENAI_CODEX_API_URL")?;
let client_id = optional_env("OPENAI_CODEX_CLIENT_ID")?
.unwrap_or_else(|| "app_EMoamEEZ73f0CkXaXp7hrann".to_string());
let session_path = optional_env("OPENAI_CODEX_SESSION_PATH")?
.map(PathBuf::from)
.unwrap_or_else(|| ironclaw_base_dir().join("openai_codex_session.json"));
let token_refresh_margin_secs =
parse_optional_env("OPENAI_CODEX_REFRESH_MARGIN_SECS", 300)?;
Some(OpenAiCodexConfig {
model,
auth_endpoint,
api_base_url,
client_id,
session_path,
token_refresh_margin_secs,
})
} else {
None
};
let request_timeout_secs = parse_optional_env("LLM_REQUEST_TIMEOUT_SECS", 120)?;
// Generic cheap model (works with any backend).
@@ -208,6 +245,8 @@ impl LlmConfig {
"nearai".to_string()
} else if is_bedrock {
"bedrock".to_string()
} else if is_openai_codex {
"openai_codex".to_string()
} else if let Some(ref p) = provider {
p.provider_id.clone()
} else {
@@ -217,6 +256,7 @@ impl LlmConfig {
nearai,
provider,
bedrock,
openai_codex,
request_timeout_secs,
cheap_model,
smart_routing_cascade,
@@ -434,6 +474,14 @@ impl LlmConfig {
} else {
Vec::new()
};
let extra_headers = if canonical_id == "github_copilot" {
merge_extra_headers(
crate::llm::github_copilot_auth::default_headers(),
extra_headers,
)
} else {
extra_headers
};
// Resolve OAuth token (Anthropic-specific: `claude login` flow).
// Only check for OAuth token when the provider is actually Anthropic.
@@ -518,6 +566,26 @@ fn parse_extra_headers(val: &str) -> Result<Vec<(String, String)>, ConfigError>
Ok(headers)
}
fn merge_extra_headers(
defaults: Vec<(String, String)>,
overrides: Vec<(String, String)>,
) -> Vec<(String, String)> {
let mut merged = Vec::new();
let mut positions = std::collections::HashMap::<String, usize>::new();
for (key, value) in defaults.into_iter().chain(overrides) {
let normalized = key.to_ascii_lowercase();
if let Some(existing_index) = positions.get(&normalized).copied() {
merged[existing_index] = (key, value);
} else {
positions.insert(normalized, merged.len());
merged.push((key, value));
}
}
merged
}
/// Get the default session file path (~/.ironclaw/session.json).
pub fn default_session_path() -> PathBuf {
ironclaw_base_dir().join("session.json")
@@ -649,6 +717,29 @@ mod tests {
);
}
#[test]
fn merge_extra_headers_prefers_overrides_case_insensitively() {
let merged = merge_extra_headers(
vec![
("User-Agent".to_string(), "default-agent".to_string()),
("X-Test".to_string(), "default".to_string()),
],
vec![
("user-agent".to_string(), "override-agent".to_string()),
("X-Extra".to_string(), "present".to_string()),
],
);
assert_eq!(
merged,
vec![
("user-agent".to_string(), "override-agent".to_string()),
("X-Test".to_string(), "default".to_string()),
("X-Extra".to_string(), "present".to_string()),
]
);
}
/// Clear all ollama-related env vars.
fn clear_ollama_env() {
// SAFETY: Only called under ENV_MUTEX in tests.
@@ -801,6 +892,54 @@ mod tests {
assert_eq!(provider.protocol, ProviderProtocol::OpenAiCompletions);
}
#[test]
fn registry_provider_resolves_github_copilot_alias() {
let _guard = ENV_MUTEX.lock().expect("env mutex poisoned");
// SAFETY: Under ENV_MUTEX.
unsafe {
std::env::set_var("LLM_BACKEND", "github-copilot");
std::env::set_var("GITHUB_COPILOT_TOKEN", "gho_test_token");
std::env::set_var(
"GITHUB_COPILOT_EXTRA_HEADERS",
"Copilot-Integration-Id:custom-chat,X-Test:enabled",
);
}
let settings = Settings::default();
let cfg = LlmConfig::resolve(&settings).expect("resolve should succeed");
assert_eq!(cfg.backend, "github_copilot");
let provider = cfg.provider.expect("provider config should be present");
assert_eq!(provider.provider_id, "github_copilot");
assert_eq!(provider.base_url, "https://api.githubcopilot.com");
assert_eq!(provider.model, "gpt-4o");
assert!(
provider
.extra_headers
.iter()
.any(|(key, value)| { key == "Copilot-Integration-Id" && value == "custom-chat" })
);
assert!(
provider
.extra_headers
.iter()
.any(|(key, value)| key == "User-Agent" && value == "GitHubCopilotChat/0.26.7")
);
assert!(
provider
.extra_headers
.iter()
.any(|(key, value)| key == "X-Test" && value == "enabled")
);
// SAFETY: Under ENV_MUTEX.
unsafe {
std::env::remove_var("LLM_BACKEND");
std::env::remove_var("GITHUB_COPILOT_TOKEN");
std::env::remove_var("GITHUB_COPILOT_EXTRA_HEADERS");
}
}
#[test]
fn nearai_backend_has_no_registry_provider() {
let _guard = ENV_MUTEX.lock().expect("env mutex poisoned");
@@ -900,19 +1039,19 @@ mod tests {
// SAFETY: Under ENV_MUTEX.
unsafe {
std::env::set_var("LLM_BACKEND", "openai_compatible");
std::env::set_var("LLM_BASE_URL", "http://env-url/v1");
std::env::set_var("LLM_BASE_URL", "http://localhost:8000/v1");
}
let settings = Settings {
llm_backend: Some("openai_compatible".to_string()),
openai_compatible_base_url: Some("http://settings-url/v1".to_string()),
openai_compatible_base_url: Some("http://localhost:9000/v1".to_string()),
..Default::default()
};
let cfg = LlmConfig::resolve(&settings).expect("resolve should succeed");
let provider = cfg.provider.expect("should have provider config");
assert_eq!(
provider.base_url, "http://env-url/v1",
provider.base_url, "http://localhost:8000/v1",
"env var should take priority over settings"
);
@@ -924,7 +1063,7 @@ mod tests {
let cfg = LlmConfig::resolve(&settings).expect("resolve should succeed");
let provider = cfg.provider.expect("should have provider config");
assert_eq!(
provider.base_url, "http://settings-url/v1",
provider.base_url, "http://localhost:9000/v1",
"settings should take priority over registry default"
);
@@ -1220,13 +1359,24 @@ mod tests {
cfg.backend, "myprovider",
"DB setting should override LLM_BACKEND env var"
);
// SAFETY: Under ENV_MUTEX.
unsafe {
std::env::remove_var("LLM_BACKEND");
}
}
// ── OpenAI Codex tests ──────────────────────────────────────────
/// Clear all openai-codex-related env vars.
fn clear_openai_codex_env() {
// SAFETY: Only called under ENV_MUTEX in tests.
unsafe {
std::env::remove_var("LLM_BACKEND");
std::env::remove_var("OPENAI_CODEX_MODEL");
std::env::remove_var("OPENAI_MODEL");
}
}
#[test]
fn builtin_override_model_used_when_no_selected_model() {
let _guard = ENV_MUTEX.lock().expect("env mutex poisoned");
@@ -1258,6 +1408,26 @@ mod tests {
);
}
#[test]
fn openai_codex_resolves_config() {
let _guard = ENV_MUTEX.lock().expect("env mutex poisoned");
clear_openai_codex_env();
let settings = Settings {
llm_backend: Some("openai_codex".to_string()),
..Default::default()
};
let cfg = LlmConfig::resolve(&settings).expect("resolve should succeed");
assert_eq!(cfg.backend, "openai_codex");
let codex = cfg.openai_codex.expect("codex config should be present");
assert_eq!(codex.model, "gpt-5.3-codex"); // default
assert!(
cfg.provider.is_none(),
"codex should not use registry provider"
);
}
#[test]
fn selected_model_takes_priority_over_builtin_override_model() {
let _guard = ENV_MUTEX.lock().expect("env mutex poisoned");
@@ -1290,6 +1460,30 @@ mod tests {
);
}
#[test]
fn openai_codex_model_env_resolution() {
let _guard = ENV_MUTEX.lock().expect("env mutex poisoned");
clear_openai_codex_env();
// SAFETY: Under ENV_MUTEX.
unsafe {
std::env::set_var("OPENAI_CODEX_MODEL", "o3-pro");
}
let settings = Settings {
llm_backend: Some("openai_codex".to_string()),
..Default::default()
};
let cfg = LlmConfig::resolve(&settings).expect("resolve should succeed");
let codex = cfg.openai_codex.expect("codex config should be present");
assert_eq!(codex.model, "o3-pro");
// SAFETY: Under ENV_MUTEX.
unsafe {
std::env::remove_var("OPENAI_CODEX_MODEL");
}
}
#[test]
fn builtin_override_api_key_used_when_no_env_var() {
let _guard = ENV_MUTEX.lock().expect("env mutex poisoned");
@@ -1326,4 +1520,103 @@ mod tests {
"builtin override api_key should be used when env var is absent"
);
}
#[test]
fn openai_codex_falls_back_to_openai_model() {
let _guard = ENV_MUTEX.lock().expect("env mutex poisoned");
clear_openai_codex_env();
// SAFETY: Under ENV_MUTEX.
unsafe {
std::env::set_var("OPENAI_MODEL", "gpt-4o");
}
let settings = Settings {
llm_backend: Some("openai_codex".to_string()),
..Default::default()
};
let cfg = LlmConfig::resolve(&settings).expect("resolve should succeed");
let codex = cfg.openai_codex.expect("codex config should be present");
assert_eq!(codex.model, "gpt-4o");
// SAFETY: Under ENV_MUTEX.
unsafe {
std::env::remove_var("OPENAI_MODEL");
}
}
#[test]
fn openai_codex_falls_back_to_selected_model() {
let _guard = ENV_MUTEX.lock().expect("env mutex poisoned");
clear_openai_codex_env();
let settings = Settings {
llm_backend: Some("openai_codex".to_string()),
selected_model: Some("gpt-4o-mini".to_string()),
..Default::default()
};
let cfg = LlmConfig::resolve(&settings).expect("resolve should succeed");
let codex = cfg.openai_codex.expect("codex config should be present");
assert_eq!(codex.model, "gpt-4o-mini");
}
/// Regression: SSRF validation on OPENAI_CODEX_API_URL (#1103).
#[test]
fn openai_codex_rejects_ssrf_api_url() {
let _guard = ENV_MUTEX.lock().expect("env mutex poisoned");
clear_openai_codex_env();
// SAFETY: Under ENV_MUTEX.
unsafe {
std::env::set_var(
"OPENAI_CODEX_API_URL",
"http://169.254.169.254/latest/meta-data",
);
}
let settings = Settings {
llm_backend: Some("openai_codex".to_string()),
..Default::default()
};
let err = LlmConfig::resolve(&settings).unwrap_err();
let msg = err.to_string();
assert!(
msg.contains("OPENAI_CODEX_API_URL"),
"error should reference the field name: {msg}"
);
// SAFETY: Under ENV_MUTEX.
unsafe {
std::env::remove_var("OPENAI_CODEX_API_URL");
}
}
/// Regression: SSRF validation on OPENAI_CODEX_AUTH_URL (#1103).
#[test]
fn openai_codex_rejects_ssrf_auth_url() {
let _guard = ENV_MUTEX.lock().expect("env mutex poisoned");
clear_openai_codex_env();
// SAFETY: Under ENV_MUTEX.
unsafe {
std::env::set_var("OPENAI_CODEX_AUTH_URL", "http://10.0.0.1");
}
let settings = Settings {
llm_backend: Some("openai_codex".to_string()),
..Default::default()
};
let err = LlmConfig::resolve(&settings).unwrap_err();
let msg = err.to_string();
assert!(
msg.contains("OPENAI_CODEX_AUTH_URL"),
"error should reference the field name: {msg}"
);
// SAFETY: Under ENV_MUTEX.
unsafe {
std::env::remove_var("OPENAI_CODEX_AUTH_URL");
}
}
}
+19 -4
View File
@@ -24,6 +24,7 @@ mod skills;
mod transcription;
mod tunnel;
mod wasm;
mod workspace;
use std::collections::HashMap;
use std::sync::{LazyLock, Mutex, Once};
@@ -53,8 +54,9 @@ pub use self::skills::SkillsConfig;
pub use self::transcription::TranscriptionConfig;
pub use self::tunnel::TunnelConfig;
pub use self::wasm::WasmConfig;
pub use self::workspace::WorkspaceConfig;
pub use crate::llm::config::{
BedrockConfig, CacheRetention, LlmConfig, NearAiConfig, OAUTH_PLACEHOLDER,
BedrockConfig, CacheRetention, LlmConfig, NearAiConfig, OAUTH_PLACEHOLDER, OpenAiCodexConfig,
RegistryProviderConfig,
};
pub use crate::llm::session::SessionConfig;
@@ -98,6 +100,7 @@ pub struct Config {
pub skills: SkillsConfig,
pub transcription: TranscriptionConfig,
pub search: WorkspaceSearchConfig,
pub workspace: WorkspaceConfig,
pub observability: crate::observability::ObservabilityConfig,
/// Channel-relay integration (Slack via external relay service).
/// Present only when both `CHANNEL_RELAY_URL` and `CHANNEL_RELAY_API_KEY` are set.
@@ -175,6 +178,9 @@ impl Config {
},
transcription: TranscriptionConfig::default(),
search: WorkspaceSearchConfig::default(),
workspace: WorkspaceConfig {
memory_layers: vec![],
},
observability: crate::observability::ObservabilityConfig::default(),
relay: None,
}
@@ -305,13 +311,21 @@ impl Config {
async fn build(settings: &Settings) -> Result<Self, ConfigError> {
let owner_id = resolve_owner_id(settings)?;
let tunnel = TunnelConfig::resolve(settings)?;
let channels = ChannelsConfig::resolve(settings, &owner_id)?;
let workspace_user_id = channels
.gateway
.as_ref()
.map(|gw| gw.user_id.clone())
.unwrap_or_else(|| "default".to_string());
Ok(Self {
owner_id: owner_id.clone(),
database: DatabaseConfig::resolve()?,
llm: LlmConfig::resolve(settings)?,
embeddings: EmbeddingsConfig::resolve(settings)?,
tunnel: TunnelConfig::resolve(settings)?,
channels: ChannelsConfig::resolve(settings, &owner_id)?,
tunnel,
channels,
agent: AgentConfig::resolve(settings)?,
safety: resolve_safety_config(settings)?,
wasm: WasmConfig::resolve(settings)?,
@@ -325,6 +339,7 @@ impl Config {
skills: SkillsConfig::resolve()?,
transcription: TranscriptionConfig::resolve(settings)?,
search: WorkspaceSearchConfig::resolve()?,
workspace: WorkspaceConfig::resolve(&workspace_user_id)?,
observability: crate::observability::ObservabilityConfig {
backend: std::env::var("OBSERVABILITY_BACKEND").unwrap_or_else(|_| "none".into()),
},
@@ -377,7 +392,7 @@ pub(crate) fn resolve_owner_id(settings: &Settings) -> Result<String, ConfigErro
/// are read by `optional_env()` before falling back to `std::env::var()`,
/// so explicit env vars always win.
///
/// Also loads tokens from OS credential stores (macOS Keychain, Linux
/// Also loads tokens from OS credential stores (macOS Keychain / Linux
/// credentials files) which don't require the secrets DB.
pub async fn inject_llm_keys_from_secrets(
secrets: &dyn crate::secrets::SecretsStore,
+208
View File
@@ -0,0 +1,208 @@
use crate::config::helpers::optional_env;
use crate::error::ConfigError;
use crate::workspace::layer::MemoryLayer;
/// Workspace memory configuration.
///
/// Controls memory layer definitions for privacy-aware writes.
/// Layers are parsed from the `MEMORY_LAYERS` env var (JSON array)
/// or default to a single private layer scoped to the gateway user.
#[derive(Debug, Clone)]
pub struct WorkspaceConfig {
pub memory_layers: Vec<MemoryLayer>,
}
impl WorkspaceConfig {
pub(crate) fn resolve(user_id: &str) -> Result<Self, ConfigError> {
let memory_layers: Vec<MemoryLayer> = match optional_env("MEMORY_LAYERS")? {
Some(json_str) => {
serde_json::from_str(&json_str).map_err(|e| ConfigError::InvalidValue {
key: "MEMORY_LAYERS".to_string(),
message: format!("must be valid JSON array of layer objects: {e}"),
})?
}
None => MemoryLayer::default_for_user(user_id),
};
// Validate layer names and scopes
for layer in &memory_layers {
if layer.name.trim().is_empty() {
return Err(ConfigError::InvalidValue {
key: "MEMORY_LAYERS".to_string(),
message: "layer name must not be empty".to_string(),
});
}
if layer.name.len() > 64 {
return Err(ConfigError::InvalidValue {
key: "MEMORY_LAYERS".to_string(),
message: format!("layer name '{}' exceeds 64 characters", layer.name),
});
}
if !layer
.name
.chars()
.all(|c| c.is_alphanumeric() || c == '_' || c == '-')
{
return Err(ConfigError::InvalidValue {
key: "MEMORY_LAYERS".to_string(),
message: format!(
"layer name '{}' contains invalid characters (only alphanumeric, _, - allowed)",
layer.name
),
});
}
if layer.scope.trim().is_empty() {
return Err(ConfigError::InvalidValue {
key: "MEMORY_LAYERS".to_string(),
message: format!("layer '{}' has an empty scope", layer.name),
});
}
}
// Check for duplicate layer names
{
let mut seen = std::collections::HashSet::new();
for layer in &memory_layers {
if !seen.insert(&layer.name) {
return Err(ConfigError::InvalidValue {
key: "MEMORY_LAYERS".to_string(),
message: format!("duplicate layer name '{}'", layer.name),
});
}
}
}
Ok(Self { memory_layers })
}
}
#[cfg(test)]
mod tests {
use super::*;
use std::sync::Mutex;
// Serialize env-var-dependent tests to avoid races.
static ENV_LOCK: Mutex<()> = Mutex::new(());
fn with_env(key: &str, val: Option<&str>, f: impl FnOnce()) {
let _guard = ENV_LOCK.lock().unwrap();
let prev = std::env::var(key).ok();
match val {
Some(v) => unsafe { std::env::set_var(key, v) },
None => unsafe { std::env::remove_var(key) },
}
f();
match prev {
Some(v) => unsafe { std::env::set_var(key, v) },
None => unsafe { std::env::remove_var(key) },
}
}
#[test]
fn valid_json_parses_correctly() {
let json = r#"[{"name":"private","scope":"alice","writable":true,"sensitivity":"private"},{"name":"shared","scope":"shared","writable":true,"sensitivity":"shared"}]"#;
with_env("MEMORY_LAYERS", Some(json), || {
let config = WorkspaceConfig::resolve("alice").expect("should parse");
assert_eq!(config.memory_layers.len(), 2);
assert_eq!(config.memory_layers[0].name, "private");
assert_eq!(config.memory_layers[1].name, "shared");
});
}
#[test]
fn invalid_json_returns_error() {
with_env("MEMORY_LAYERS", Some("not json"), || {
let result = WorkspaceConfig::resolve("alice");
assert!(result.is_err(), "invalid JSON should fail");
let err = result.unwrap_err().to_string();
assert!(
err.contains("valid JSON"),
"error should mention JSON: {err}"
);
});
}
#[test]
fn empty_layer_name_returns_error() {
let json = r#"[{"name":"","scope":"alice"}]"#;
with_env("MEMORY_LAYERS", Some(json), || {
let result = WorkspaceConfig::resolve("alice");
assert!(result.is_err(), "empty layer name should fail");
let err = result.unwrap_err().to_string();
assert!(err.contains("empty"), "error should mention empty: {err}");
});
}
#[test]
fn layer_name_exceeding_64_chars_returns_error() {
let long_name = "a".repeat(65);
let json = format!(r#"[{{"name":"{long_name}","scope":"alice"}}]"#);
with_env("MEMORY_LAYERS", Some(&json), || {
let result = WorkspaceConfig::resolve("alice");
assert!(result.is_err(), "long layer name should fail");
let err = result.unwrap_err().to_string();
assert!(
err.contains("exceeds 64"),
"error should mention 64 chars: {err}"
);
});
}
#[test]
fn layer_name_with_invalid_chars_returns_error() {
for bad_name in ["has space", "has@at", "has.dot", "has/slash"] {
let json = format!(r#"[{{"name":"{bad_name}","scope":"alice"}}]"#);
with_env("MEMORY_LAYERS", Some(&json), || {
let result = WorkspaceConfig::resolve("alice");
assert!(
result.is_err(),
"layer name '{bad_name}' should fail validation"
);
let err = result.unwrap_err().to_string();
assert!(
err.contains("invalid characters"),
"error for '{bad_name}' should mention invalid characters: {err}"
);
});
}
}
#[test]
fn empty_scope_returns_error() {
let json = r#"[{"name":"private","scope":""}]"#;
with_env("MEMORY_LAYERS", Some(json), || {
let result = WorkspaceConfig::resolve("alice");
assert!(result.is_err(), "empty scope should fail");
let err = result.unwrap_err().to_string();
assert!(
err.contains("empty scope"),
"error should mention empty scope: {err}"
);
});
}
#[test]
fn duplicate_layer_names_returns_error() {
let json = r#"[{"name":"private","scope":"alice"},{"name":"private","scope":"bob"}]"#;
with_env("MEMORY_LAYERS", Some(json), || {
let result = WorkspaceConfig::resolve("alice");
assert!(result.is_err(), "duplicate names should fail");
let err = result.unwrap_err().to_string();
assert!(
err.contains("duplicate"),
"error should mention duplicate: {err}"
);
});
}
#[test]
fn missing_env_defaults_to_single_private_layer() {
with_env("MEMORY_LAYERS", None, || {
let config = WorkspaceConfig::resolve("alice").expect("should default");
assert_eq!(config.memory_layers.len(), 1);
assert_eq!(config.memory_layers[0].name, "private");
assert_eq!(config.memory_layers[0].scope, "alice");
assert!(config.memory_layers[0].writable);
});
}
}
+28
View File
@@ -477,6 +477,34 @@ impl RoutineStore for LibSqlBackend {
Ok(())
}
async fn get_webhook_routine_by_path(
&self,
path: &str,
) -> Result<Option<Routine>, DatabaseError> {
let conn = self.connect().await?;
let mut rows = conn
.query(
&format!(
"SELECT {} FROM routines WHERE enabled = 1 AND trigger_type = 'webhook' \
AND (json_extract(trigger_config, '$.path') = ?1 \
OR (json_extract(trigger_config, '$.path') IS NULL AND CAST(id AS TEXT) = ?1))",
ROUTINE_COLUMNS
),
params![path],
)
.await
.map_err(|e| DatabaseError::Query(e.to_string()))?;
match rows
.next()
.await
.map_err(|e| DatabaseError::Query(e.to_string()))?
{
Some(row) => Ok(Some(row_to_routine_libsql(&row)?)),
None => Ok(None),
}
}
async fn list_dispatched_routine_runs(&self) -> Result<Vec<RoutineRun>, DatabaseError> {
let conn = self.connect().await?;
let mut rows = conn
+4
View File
@@ -525,6 +525,10 @@ pub trait RoutineStore: Send + Sync {
run_id: Uuid,
job_id: Uuid,
) -> Result<(), DatabaseError>;
async fn get_webhook_routine_by_path(
&self,
path: &str,
) -> Result<Option<Routine>, DatabaseError>;
/// List routine runs that were dispatched as full_job but have not yet
/// been finalized (status='running' with a linked job_id).
+7
View File
@@ -504,6 +504,13 @@ impl RoutineStore for PgBackend {
self.store.link_routine_run_to_job(run_id, job_id).await
}
async fn get_webhook_routine_by_path(
&self,
path: &str,
) -> Result<Option<Routine>, DatabaseError> {
self.store.get_webhook_routine_by_path(path).await
}
async fn list_dispatched_routine_runs(&self) -> Result<Vec<RoutineRun>, DatabaseError> {
self.store.list_dispatched_routine_runs().await
}
+18
View File
@@ -168,6 +168,9 @@ pub enum ToolError {
#[error("Tool {name} requires authentication")]
AuthRequired { name: String },
#[error("Tool {name} is not available for autonomous execution: {reason}")]
AutonomousUnavailable { name: String, reason: String },
#[error("Tool {name} is rate limited, retry after {retry_after:?}")]
RateLimited {
name: String,
@@ -301,6 +304,18 @@ pub enum WorkspaceError {
#[error("I/O error: {reason}")]
IoError { reason: String },
#[error("Not found: {path}")]
NotFound { path: String },
#[error("Layer not found: {name}")]
LayerNotFound { name: String },
#[error("Layer '{name}' is read-only")]
LayerReadOnly { name: String },
#[error("Cannot write sensitive content: no private layer available for redirect")]
PrivacyRedirectFailed,
#[error("Write rejected for '{path}': prompt injection detected ({reason})")]
InjectionRejected { path: String, reason: String },
}
@@ -373,6 +388,9 @@ pub enum RoutineError {
#[error("Not authorized to trigger routine {id}")]
NotAuthorized { id: Uuid },
#[error("Routine {name} is in cooldown period")]
Cooldown { name: String },
#[error("Routine {name} at max concurrent runs")]
MaxConcurrent { name: String },
+56
View File
@@ -463,6 +463,37 @@ fn sanitize_url_for_logging(url: &str) -> String {
}
impl ExtensionManager {
pub fn owner_id(&self) -> &str {
&self.user_id
}
pub async fn active_tool_names(&self) -> HashSet<String> {
let mut names = HashSet::new();
match self.list(None, false).await {
Ok(extensions) => {
for extension in extensions {
match extension.kind {
ExtensionKind::WasmTool if extension.active => {
names.insert(extension.name);
}
ExtensionKind::McpServer if extension.active => {
names.extend(extension.tools);
}
_ => {}
}
}
}
Err(err) => {
tracing::warn!(
owner_id = %self.user_id,
"Failed to list active extensions while resolving autonomous tool scope: {}",
err
);
}
}
names
}
#[allow(clippy::too_many_arguments)]
pub fn new(
mcp_session_manager: Arc<McpSessionManager>,
@@ -906,6 +937,31 @@ impl ExtensionManager {
&self.secrets
}
/// Inject a pre-created MCP client (from startup loading) into the manager.
///
/// Startup-loaded MCP clients register their tools in `ToolRegistry` but are
/// otherwise dropped. This method stores the client so that `list()` reports
/// accurate "connected" status and reconnection/session management works.
pub(crate) async fn inject_mcp_client(
&self,
name: String,
client: Arc<crate::tools::mcp::McpClient>,
) {
if name.is_empty() {
tracing::warn!("inject_mcp_client called with empty name; ignoring");
return;
}
if let Err(e) = Self::validate_extension_name(&name) {
tracing::warn!(
error = %e,
name = %name,
"inject_mcp_client called with invalid name; ignoring"
);
return;
}
self.mcp_clients.write().await.insert(name, client);
}
/// Register channel names that were loaded at startup.
/// Called after WASM channels are loaded so `list()` reports accurate active status.
pub async fn set_active_channels(&self, names: Vec<String>) {
+16
View File
@@ -1105,6 +1105,22 @@ impl Store {
rows.iter().map(row_to_routine).collect()
}
/// Find an enabled webhook routine by its configured path (or fallback to ID).
pub async fn get_webhook_routine_by_path(
&self,
path: &str,
) -> Result<Option<Routine>, DatabaseError> {
let conn = self.conn().await?;
let row = conn
.query_opt(
"SELECT * FROM routines WHERE enabled AND trigger_type = 'webhook' \
AND (trigger_config->>'path' = $1 OR (trigger_config->>'path' IS NULL AND id::text = $1))",
&[&path],
)
.await?;
row.as_ref().map(row_to_routine).transpose()
}
/// List all enabled cron routines whose next_fire_at <= now.
pub async fn list_due_cron_routines(&self) -> Result<Vec<Routine>, DatabaseError> {
let conn = self.conn().await?;
+45 -1
View File
@@ -13,6 +13,9 @@ Multi-provider LLM integration with circuit breaker, retry, failover, and respon
| `nearai_chat.rs` | NEAR AI Chat Completions provider (dual auth: session token or API key) |
| `codex_auth.rs` | Reads Codex CLI `auth.json`, extracts tokens, refreshes ChatGPT OAuth access tokens |
| `codex_chatgpt.rs` | Custom Responses API provider for Codex ChatGPT backend (`/backend-api/codex`) |
| `openai_codex_provider.rs` | OpenAI Codex Responses API client (SSE streaming, JWT auth, subscription billing) |
| `openai_codex_session.rs` | OAuth 2.0 session manager for OpenAI Codex (device code flow, token persistence) |
| `token_refreshing.rs` | Token-refreshing `LlmProvider` decorator for OpenAI Codex (pre-emptive refresh, zero-cost billing) |
| `reasoning.rs` | `Reasoning` struct, `ReasoningContext`, `RespondResult`, `ActionPlan`, `ToolSelection`; thinking-tag stripping; `SILENT_REPLY_TOKEN` |
| `session.rs` | NEAR AI session token management with disk + DB persistence, OAuth login flow |
| `circuit_breaker.rs` | Circuit breaker: Closed → Open → HalfOpen state machine |
@@ -34,10 +37,12 @@ Set via `LLM_BACKEND` env var:
| `nearai` (default) | NEAR AI Chat Completions | `NEARAI_SESSION_TOKEN` or `NEARAI_API_KEY` |
| `openai` | OpenAI | `OPENAI_API_KEY` |
| `anthropic` | Anthropic | `ANTHROPIC_API_KEY` |
| `github_copilot` | GitHub Copilot Chat API | `GITHUB_COPILOT_TOKEN`, `GITHUB_COPILOT_MODEL` |
| `ollama` | Ollama local | `OLLAMA_BASE_URL` |
| `openai_compatible` | Any OpenAI-compatible endpoint | `LLM_BASE_URL`, `LLM_API_KEY`, `LLM_MODEL` |
| `tinfoil` | Tinfoil TEE inference | `TINFOIL_API_KEY`, `TINFOIL_MODEL` |
| `bedrock` | AWS Bedrock (requires `--features bedrock`) | `BEDROCK_REGION`, `BEDROCK_MODEL`, `AWS_PROFILE` |
| `openai_codex` | OpenAI Codex (ChatGPT subscription) | `OPENAI_CODEX_MODEL`, `OPENAI_CODEX_CLIENT_ID` |
Codex auth reuse:
- Set `LLM_USE_CODEX_AUTH=true` to load credentials from `~/.codex/auth.json` (override with `CODEX_AUTH_PATH`).
@@ -56,6 +61,27 @@ Uses the native Converse API via `aws-sdk-bedrockruntime` (`bedrock.rs`). Requir
- `BEDROCK_MODEL` — Required model ID (e.g., `anthropic.claude-opus-4-6-v1`)
- `BEDROCK_CROSS_REGION` — Optional cross-region inference prefix (`us`, `eu`, `apac`, `global`)
## GitHub Copilot Provider Notes
`github_copilot` uses a dedicated `GithubCopilotProvider` (`github_copilot.rs`) with
direct HTTP via `reqwest::Client`. It cannot use `RigAdapter` because the Copilot API
requires a two-step authentication flow: a long-lived GitHub OAuth token is exchanged
for a short-lived Copilot session token via `api.github.com/copilot_internal/v2/token`.
The session token is cached and auto-refreshed before expiry by `CopilotTokenManager`
in `github_copilot_auth.rs`.
The API endpoint is `https://api.githubcopilot.com/chat/completions` (OpenAI Chat
Completions format). Token source: `GITHUB_COPILOT_TOKEN` env var, or the
`oauth_token` from your IDE sign-in flow (`~/.config/github-copilot/apps.json`).
The setup wizard supports GitHub device login or manual token paste.
**Known risk:** The device login flow uses the VS Code Copilot OAuth client ID
(`Iv1.b507a08c87ecfe98`) and injects VS Code identity headers (`User-Agent`,
`Editor-Version`, `Editor-Plugin-Version`, `Copilot-Integration-Id`). GitHub could
rotate this client ID at any time. If GitHub publishes an official third-party client
ID, migrate to it immediately. Advanced users can override headers via
`GITHUB_COPILOT_EXTRA_HEADERS`.
## NEAR AI Provider Gotchas
**Dual auth modes:**
@@ -148,9 +174,27 @@ To add a new provider:
Set `LLM_EXTRA_HEADERS=Key:Value,Key2:Value2` to inject headers into every request. Useful for OpenRouter attribution (`HTTP-Referer`, `X-Title`). Invalid header names/values are skipped with a warning (not a fatal error).
## OpenAI Codex Provider
Uses the Responses API at `chatgpt.com/backend-api/codex/responses` with ChatGPT subscription OAuth tokens (zero API cost — billing through subscription).
**Auth flow:** Device code OAuth via `auth.openai.com/api/accounts/deviceauth/*` endpoints. On first run, displays a code for the user to enter at a URL. Tokens are persisted to `~/.ironclaw/openai_codex_session.json` (mode 0600) and auto-refreshed before expiry.
**Provider chain:** `OpenAiCodexProvider``TokenRefreshingProvider` (pre-emptive refresh + retry on 401) → standard decorator chain. The `TokenRefreshingProvider` intercepts `AuthFailed`/`SessionExpired` errors, refreshes the OAuth token, and retries once.
**Key differences from other providers:**
- Uses Responses API (not Chat Completions) — SSE streaming with different event types
- System messages are sent as `instructions` field, not in `input` array
- Tool schemas are normalized via `normalize_schema_strict()` for OpenAI strict mode
- `cost_per_token()` returns `(0, 0)` — subscription-based billing
- `set_model()` returns error — model is fixed at construction time
- Image attachments are silently dropped with a warning log
**Env vars:** `OPENAI_CODEX_MODEL` (default: `gpt-5.3-codex`), `OPENAI_CODEX_CLIENT_ID`, `OPENAI_CODEX_AUTH_URL`, `OPENAI_CODEX_API_URL`.
## Provider Chain Construction
`build_provider_chain()` in `mod.rs` is the single source of truth for assembling decorators. The chain is:
`build_provider_chain()` in `mod.rs` is the single source of truth for assembling decorators. It creates the base provider (dispatching to `create_openai_codex_provider()` for codex, `create_llm_provider()` for everything else), then applies all decorators inline:
```
Raw provider
+34
View File
@@ -0,0 +1,34 @@
//! Shared test helpers for OpenAI Codex provider tests.
#![cfg(test)]
use crate::config::OpenAiCodexConfig;
/// Build a minimal JWT for testing (header.payload.signature).
pub(crate) fn make_test_jwt(account_id: &str) -> String {
use base64::Engine;
let engine = base64::engine::general_purpose::URL_SAFE_NO_PAD;
let header = engine.encode(b"{\"alg\":\"RS256\",\"typ\":\"JWT\"}");
let payload_json = serde_json::json!({
"sub": "user123",
"https://api.openai.com/auth": {
"chatgpt_account_id": account_id,
},
});
let payload = engine.encode(payload_json.to_string().as_bytes());
let sig = engine.encode(b"fake-signature");
format!("{header}.{payload}.{sig}")
}
/// Build a test `OpenAiCodexConfig` with a given session path.
pub(crate) fn test_codex_config(session_path: std::path::PathBuf) -> OpenAiCodexConfig {
OpenAiCodexConfig {
model: "gpt-5.3-codex".to_string(),
auth_endpoint: "https://auth.openai.com".to_string(),
api_base_url: "https://chatgpt.com/backend-api/codex".to_string(),
client_id: "test_client_id".to_string(),
session_path,
token_refresh_margin_secs: 300,
}
}
+35 -2
View File
@@ -9,6 +9,7 @@ use std::path::PathBuf;
use secrecy::SecretString;
use crate::bootstrap::ironclaw_base_dir;
use crate::llm::registry::ProviderProtocol;
use crate::llm::session::SessionConfig;
@@ -102,6 +103,36 @@ pub struct RegistryProviderConfig {
pub unsupported_params: Vec<String>,
}
/// Configuration for OpenAI Codex (ChatGPT subscription OAuth).
#[derive(Debug, Clone)]
pub struct OpenAiCodexConfig {
/// Model to use (default: "gpt-5.3-codex").
pub model: String,
/// OAuth authorization server (default: "https://auth.openai.com").
pub auth_endpoint: String,
/// Responses API base URL (default: "https://chatgpt.com/backend-api/codex").
pub api_base_url: String,
/// OAuth client ID (default: OpenAI's public Codex client).
pub client_id: String,
/// Path to session file (default: ~/.ironclaw/openai_codex_session.json).
pub session_path: PathBuf,
/// Seconds before expiry to proactively refresh (default: 300).
pub token_refresh_margin_secs: u64,
}
impl Default for OpenAiCodexConfig {
fn default() -> Self {
Self {
model: "gpt-5.3-codex".to_string(),
auth_endpoint: "https://auth.openai.com".to_string(),
api_base_url: "https://chatgpt.com/backend-api/codex".to_string(),
client_id: "app_EMoamEEZ73f0CkXaXp7hrann".to_string(),
session_path: ironclaw_base_dir().join("openai_codex_session.json"),
token_refresh_margin_secs: 300,
}
}
}
/// Configuration for AWS Bedrock (native Converse API).
#[derive(Debug, Clone)]
pub struct BedrockConfig {
@@ -134,6 +165,8 @@ pub struct LlmConfig {
pub provider: Option<RegistryProviderConfig>,
/// AWS Bedrock config (populated when backend=bedrock, requires --features bedrock).
pub bedrock: Option<BedrockConfig>,
/// OpenAI Codex config (populated when backend=openai_codex).
pub openai_codex: Option<OpenAiCodexConfig>,
/// HTTP request timeout in seconds for LLM API calls.
/// Default: 120. Increase for local LLMs (Ollama, vLLM, LM Studio) that
/// need more time for prompt evaluation on consumer hardware.
@@ -213,8 +246,8 @@ impl NearAiConfig {
} else {
"https://private.near.ai"
};
let base_url =
std::env::var("NEARAI_BASE_URL").unwrap_or_else(|_| default_base.to_string());
let base_url = crate::config::helpers::env_or_override("NEARAI_BASE_URL")
.unwrap_or_else(|| default_base.to_string());
Self {
model: String::new(),
+712
View File
@@ -0,0 +1,712 @@
//! GitHub Copilot provider (direct HTTP with token exchange).
//!
//! The GitHub Copilot API at `api.githubcopilot.com` speaks OpenAI Chat
//! Completions format but requires a two-step authentication flow:
//! 1. A long-lived GitHub OAuth token (from device login or IDE sign-in)
//! 2. A short-lived Copilot session token (exchanged via GitHub API)
//!
//! The standard OpenAI rig-core client sends `Authorization: Bearer <token>`
//! with the raw OAuth token, which gets rejected with "Authorization header
//! is badly formatted". This provider handles the token exchange transparently.
use std::collections::HashSet;
use std::sync::Arc;
use async_trait::async_trait;
use reqwest::Client;
use rust_decimal::Decimal;
use secrecy::ExposeSecret;
use serde::{Deserialize, Serialize};
use crate::llm::config::RegistryProviderConfig;
use crate::llm::costs;
use crate::llm::error::LlmError;
use crate::llm::github_copilot_auth::CopilotTokenManager;
use crate::llm::provider::{
ChatMessage, CompletionRequest, CompletionResponse, ContentPart, FinishReason, LlmProvider,
Role, ToolCall, ToolCompletionRequest, ToolCompletionResponse,
strip_unsupported_completion_params, strip_unsupported_tool_params,
};
/// GitHub Copilot provider with automatic token exchange.
pub struct GithubCopilotProvider {
client: Client,
token_manager: Arc<CopilotTokenManager>,
model: String,
base_url: String,
active_model: std::sync::RwLock<String>,
extra_headers: Vec<(String, String)>,
/// Parameter names that this provider does not support.
unsupported_params: HashSet<String>,
}
impl GithubCopilotProvider {
pub fn new(
config: &RegistryProviderConfig,
request_timeout_secs: u64,
) -> Result<Self, LlmError> {
let oauth_token = config
.api_key
.as_ref()
.map(|k| k.expose_secret().to_string())
.ok_or_else(|| {
tracing::error!("No API key configured for github_copilot — check GITHUB_COPILOT_TOKEN env var or secrets store");
LlmError::AuthFailed {
provider: "github_copilot".to_string(),
}
})?;
let client = Client::builder()
.timeout(std::time::Duration::from_secs(request_timeout_secs))
.build()
.map_err(|e| LlmError::RequestFailed {
provider: "github_copilot".to_string(),
reason: format!("Failed to build HTTP client: {e}"),
})?;
let token_manager = Arc::new(CopilotTokenManager::new(client.clone(), oauth_token));
let base_url = if config.base_url.is_empty() {
"https://api.githubcopilot.com".to_string()
} else {
config.base_url.clone()
};
let active_model = std::sync::RwLock::new(config.model.clone());
let unsupported_params: HashSet<String> =
config.unsupported_params.iter().cloned().collect();
Ok(Self {
client,
token_manager,
model: config.model.clone(),
base_url,
active_model,
extra_headers: config.extra_headers.clone(),
unsupported_params,
})
}
fn api_url(&self) -> String {
let base = self.base_url.trim_end_matches('/');
format!("{base}/chat/completions")
}
/// Strip unsupported fields from a `CompletionRequest` in place.
fn strip_unsupported_completion_params(&self, req: &mut CompletionRequest) {
strip_unsupported_completion_params(&self.unsupported_params, req);
}
/// Strip unsupported fields from a `ToolCompletionRequest` in place.
fn strip_unsupported_tool_params(&self, req: &mut ToolCompletionRequest) {
strip_unsupported_tool_params(&self.unsupported_params, req);
}
async fn send_request<R: for<'de> Deserialize<'de>>(
&self,
body: &impl Serialize,
) -> Result<R, LlmError> {
let url = self.api_url();
// Map token exchange failures to RequestFailed (retryable) rather than
// AuthFailed (non-retryable), since transient network errors during
// exchange should be retried by RetryProvider.
let token = self.token_manager.get_token().await.map_err(|e| {
tracing::warn!(error = %e, "Copilot: token exchange failed");
LlmError::RequestFailed {
provider: "github_copilot".to_string(),
reason: format!("Token exchange failed: {e}"),
}
})?;
let mut request = self
.client
.post(&url)
.bearer_auth(token.expose_secret())
.header("Content-Type", "application/json");
// Inject Copilot identity headers
for (key, value) in &self.extra_headers {
request = request.header(key.as_str(), value.as_str());
}
let response = request.json(body).send().await.map_err(|e| {
tracing::warn!(error = %e, "Copilot: HTTP request failed");
LlmError::RequestFailed {
provider: "github_copilot".to_string(),
reason: e.to_string(),
}
})?;
let status = response.status();
if !status.is_success() {
// Use shared retry-after parser (supports HTTP-date, default 60s)
let retry_after = Some(crate::llm::retry::parse_retry_after(
response.headers().get(reqwest::header::RETRY_AFTER),
));
let response_text = response
.text()
.await
.unwrap_or_else(|e| format!("(failed to read error body: {e})"));
tracing::warn!(
status = %status,
body = %crate::agent::truncate_for_preview(&response_text, 256),
"Copilot: API error response"
);
if status.as_u16() == 401 {
// Invalidate the cached session token and retry once with a
// fresh exchange — stale tokens are the most common 401 cause.
tracing::warn!("Copilot: 401 Unauthorized — invalidating session token, retrying");
self.token_manager.invalidate().await;
let fresh = self.token_manager.get_token().await.map_err(|e| {
tracing::warn!(error = %e, "Copilot: re-exchange after 401 failed");
LlmError::RequestFailed {
provider: "github_copilot".to_string(),
reason: format!("Token re-exchange after 401 failed: {e}"),
}
})?;
let mut retry_req = self
.client
.post(&url)
.bearer_auth(fresh.expose_secret())
.header("Content-Type", "application/json");
for (key, value) in &self.extra_headers {
retry_req = retry_req.header(key.as_str(), value.as_str());
}
let retry =
retry_req
.json(body)
.send()
.await
.map_err(|e| LlmError::RequestFailed {
provider: "github_copilot".to_string(),
reason: format!("Retry after 401 failed: {e}"),
})?;
if retry.status().is_success() {
let text = retry.text().await.map_err(|e| LlmError::RequestFailed {
provider: "github_copilot".to_string(),
reason: format!("Failed to read retry response body: {e}"),
})?;
return serde_json::from_str(&text).map_err(|e| {
let truncated = crate::agent::truncate_for_preview(&text, 512);
LlmError::InvalidResponse {
provider: "github_copilot".to_string(),
reason: format!("JSON parse error: {e}. Raw: {truncated}"),
}
});
}
let retry_status = retry.status();
tracing::warn!(
status = %retry_status,
"Copilot: 401 retry also failed"
);
return Err(LlmError::AuthFailed {
provider: "github_copilot".to_string(),
});
}
if status.as_u16() == 429 {
tracing::warn!(retry_after = ?retry_after, "Copilot: rate limited");
return Err(LlmError::RateLimited {
provider: "github_copilot".to_string(),
retry_after,
});
}
let truncated = crate::agent::truncate_for_preview(&response_text, 512);
return Err(LlmError::RequestFailed {
provider: "github_copilot".to_string(),
reason: format!("HTTP {status}: {truncated}"),
});
}
let response_text = response.text().await.map_err(|e| LlmError::RequestFailed {
provider: "github_copilot".to_string(),
reason: format!("Failed to read response body: {e}"),
})?;
serde_json::from_str(&response_text).map_err(|e| {
let truncated = crate::agent::truncate_for_preview(&response_text, 512);
tracing::warn!(
error = %e,
body = %truncated,
"Copilot: failed to parse response JSON"
);
LlmError::InvalidResponse {
provider: "github_copilot".to_string(),
reason: format!("JSON parse error: {e}. Raw: {truncated}"),
}
})
}
}
#[async_trait]
impl LlmProvider for GithubCopilotProvider {
async fn complete(&self, mut req: CompletionRequest) -> Result<CompletionResponse, LlmError> {
let model = req.model.take().unwrap_or_else(|| self.active_model_name());
self.strip_unsupported_completion_params(&mut req);
let messages = convert_messages(req.messages);
let request = OpenAiRequest {
model,
messages,
max_tokens: req.max_tokens,
temperature: req.temperature,
stop: req.stop_sequences,
tools: None,
tool_choice: None,
};
let response: OpenAiResponse = self.send_request(&request).await?;
let choice =
response
.choices
.into_iter()
.next()
.ok_or_else(|| LlmError::InvalidResponse {
provider: "github_copilot".to_string(),
reason: "No choices in response".to_string(),
})?;
let (content, _tool_calls) = extract_choice_content(&choice);
let finish_reason = match choice.finish_reason.as_deref() {
Some("stop") => FinishReason::Stop,
Some("length") => FinishReason::Length,
Some("tool_calls") => FinishReason::ToolUse,
Some("content_filter") => FinishReason::ContentFilter,
_ => FinishReason::Unknown,
};
Ok(CompletionResponse {
content: content.unwrap_or_default(),
finish_reason,
input_tokens: response
.usage
.as_ref()
.map(|u| u.prompt_tokens)
.unwrap_or(0),
output_tokens: response
.usage
.as_ref()
.map(|u| u.completion_tokens)
.unwrap_or(0),
cache_creation_input_tokens: 0,
cache_read_input_tokens: 0,
})
}
async fn complete_with_tools(
&self,
mut req: ToolCompletionRequest,
) -> Result<ToolCompletionResponse, LlmError> {
let model = req.model.take().unwrap_or_else(|| self.active_model_name());
self.strip_unsupported_tool_params(&mut req);
let messages = convert_messages(req.messages);
let tools: Vec<OpenAiTool> = req
.tools
.into_iter()
.map(|t| OpenAiTool {
tool_type: "function".to_string(),
function: OpenAiFunction {
name: t.name,
description: t.description,
parameters: t.parameters,
},
})
.collect();
let tool_choice = req.tool_choice.map(|tc| match tc.as_str() {
"auto" | "required" | "none" => serde_json::Value::String(tc),
specific => serde_json::json!({
"type": "function",
"function": {"name": specific}
}),
});
let request = OpenAiRequest {
model,
messages,
max_tokens: req.max_tokens,
temperature: req.temperature,
stop: req.stop_sequences,
tools: if tools.is_empty() { None } else { Some(tools) },
tool_choice,
};
let response: OpenAiResponse = self.send_request(&request).await?;
let choice =
response
.choices
.into_iter()
.next()
.ok_or_else(|| LlmError::InvalidResponse {
provider: "github_copilot".to_string(),
reason: "No choices in response".to_string(),
})?;
let (content, tool_calls) = extract_choice_content(&choice);
let finish_reason = match choice.finish_reason.as_deref() {
Some("stop") => FinishReason::Stop,
Some("length") => FinishReason::Length,
Some("tool_calls") => FinishReason::ToolUse,
Some("content_filter") => FinishReason::ContentFilter,
_ => {
if !tool_calls.is_empty() {
FinishReason::ToolUse
} else {
FinishReason::Unknown
}
}
};
Ok(ToolCompletionResponse {
content,
tool_calls,
finish_reason,
input_tokens: response
.usage
.as_ref()
.map(|u| u.prompt_tokens)
.unwrap_or(0),
output_tokens: response
.usage
.as_ref()
.map(|u| u.completion_tokens)
.unwrap_or(0),
cache_creation_input_tokens: 0,
cache_read_input_tokens: 0,
})
}
fn model_name(&self) -> &str {
&self.model
}
fn cost_per_token(&self) -> (Decimal, Decimal) {
let model = self.active_model_name();
costs::model_cost(&model).unwrap_or_else(costs::default_cost)
}
fn active_model_name(&self) -> String {
match self.active_model.read() {
Ok(guard) => guard.clone(),
Err(poisoned) => poisoned.into_inner().clone(),
}
}
fn set_model(&self, model: &str) -> Result<(), LlmError> {
match self.active_model.write() {
Ok(mut guard) => {
*guard = model.to_string();
}
Err(poisoned) => {
*poisoned.into_inner() = model.to_string();
}
}
Ok(())
}
}
// --- OpenAI Chat Completions API types ---
#[derive(Debug, Serialize)]
struct OpenAiRequest {
model: String,
messages: Vec<OpenAiMessage>,
#[serde(skip_serializing_if = "Option::is_none")]
max_tokens: Option<u32>,
#[serde(skip_serializing_if = "Option::is_none")]
temperature: Option<f32>,
#[serde(skip_serializing_if = "Option::is_none")]
stop: Option<Vec<String>>,
#[serde(skip_serializing_if = "Option::is_none")]
tools: Option<Vec<OpenAiTool>>,
#[serde(skip_serializing_if = "Option::is_none")]
tool_choice: Option<serde_json::Value>,
}
#[derive(Debug, Serialize)]
struct OpenAiMessage {
role: String,
#[serde(skip_serializing_if = "Option::is_none")]
content: Option<OpenAiContent>,
#[serde(skip_serializing_if = "Option::is_none")]
tool_calls: Option<Vec<OpenAiToolCall>>,
#[serde(skip_serializing_if = "Option::is_none")]
tool_call_id: Option<String>,
#[serde(skip_serializing_if = "Option::is_none")]
name: Option<String>,
}
/// OpenAI content can be a plain string or an array of parts (for multimodal).
#[derive(Debug, Serialize)]
#[serde(untagged)]
enum OpenAiContent {
Text(String),
Parts(Vec<OpenAiContentPart>),
}
#[derive(Debug, Serialize)]
#[serde(tag = "type")]
enum OpenAiContentPart {
#[serde(rename = "text")]
Text { text: String },
#[serde(rename = "image_url")]
ImageUrl { image_url: OpenAiImageUrl },
}
#[derive(Debug, Serialize)]
struct OpenAiImageUrl {
url: String,
}
#[derive(Debug, Serialize)]
struct OpenAiToolCall {
id: String,
#[serde(rename = "type")]
call_type: String,
function: OpenAiToolCallFunction,
}
#[derive(Debug, Serialize)]
struct OpenAiToolCallFunction {
name: String,
arguments: String,
}
#[derive(Debug, Serialize)]
struct OpenAiTool {
#[serde(rename = "type")]
tool_type: String,
function: OpenAiFunction,
}
#[derive(Debug, Serialize)]
struct OpenAiFunction {
name: String,
description: String,
parameters: serde_json::Value,
}
#[derive(Debug, Deserialize)]
struct OpenAiResponse {
choices: Vec<OpenAiChoice>,
#[serde(default)]
usage: Option<OpenAiUsage>,
}
#[derive(Debug, Deserialize)]
struct OpenAiChoice {
message: OpenAiResponseMessage,
#[serde(default)]
finish_reason: Option<String>,
}
#[derive(Debug, Deserialize)]
struct OpenAiResponseMessage {
#[serde(default)]
content: Option<String>,
#[serde(default)]
tool_calls: Option<Vec<OpenAiResponseToolCall>>,
}
#[derive(Debug, Deserialize)]
struct OpenAiResponseToolCall {
id: String,
function: OpenAiResponseFunction,
}
#[derive(Debug, Deserialize)]
struct OpenAiResponseFunction {
name: String,
arguments: String,
}
#[derive(Debug, Deserialize)]
struct OpenAiUsage {
#[serde(default)]
prompt_tokens: u32,
#[serde(default)]
completion_tokens: u32,
}
/// Convert IronClaw messages to OpenAI Chat Completions format.
fn convert_messages(messages: Vec<ChatMessage>) -> Vec<OpenAiMessage> {
messages
.into_iter()
.map(|msg| match msg.role {
Role::System => OpenAiMessage {
role: "system".to_string(),
content: Some(OpenAiContent::Text(msg.content)),
tool_calls: None,
tool_call_id: None,
name: None,
},
Role::User => {
let content = if msg.content_parts.is_empty() {
Some(OpenAiContent::Text(msg.content))
} else {
let mut parts = Vec::with_capacity(1 + msg.content_parts.len());
if !msg.content.is_empty() {
parts.push(OpenAiContentPart::Text { text: msg.content });
}
for part in msg.content_parts {
match part {
ContentPart::Text { text } => {
parts.push(OpenAiContentPart::Text { text });
}
ContentPart::ImageUrl { image_url } => {
parts.push(OpenAiContentPart::ImageUrl {
image_url: OpenAiImageUrl { url: image_url.url },
});
}
}
}
Some(OpenAiContent::Parts(parts))
};
OpenAiMessage {
role: "user".to_string(),
content,
tool_calls: None,
tool_call_id: None,
name: None,
}
}
Role::Assistant => {
let tool_calls = msg.tool_calls.map(|calls| {
calls
.into_iter()
.map(|tc| OpenAiToolCall {
id: tc.id,
call_type: "function".to_string(),
function: OpenAiToolCallFunction {
name: tc.name,
arguments: tc.arguments.to_string(),
},
})
.collect()
});
let content = if msg.content.is_empty() {
None
} else {
Some(OpenAiContent::Text(msg.content))
};
OpenAiMessage {
role: "assistant".to_string(),
content,
tool_calls,
tool_call_id: None,
name: None,
}
}
Role::Tool => OpenAiMessage {
role: "tool".to_string(),
content: Some(OpenAiContent::Text(msg.content)),
tool_calls: None,
tool_call_id: msg.tool_call_id,
name: msg.name,
},
})
.collect()
}
/// Extract text and tool calls from an OpenAI response choice.
fn extract_choice_content(choice: &OpenAiChoice) -> (Option<String>, Vec<ToolCall>) {
let content = choice.message.content.clone();
let tool_calls = choice
.message
.tool_calls
.as_ref()
.map(|calls| {
calls
.iter()
.map(|tc| ToolCall {
id: tc.id.clone(),
name: tc.function.name.clone(),
arguments: serde_json::from_str(&tc.function.arguments)
.unwrap_or(serde_json::Value::Object(serde_json::Map::new())),
})
.collect()
})
.unwrap_or_default();
(content, tool_calls)
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_convert_messages_basic() {
let messages = vec![
ChatMessage::system("You are helpful."),
ChatMessage::user("Hello"),
ChatMessage::assistant("Hi there!"),
];
let converted = convert_messages(messages);
assert_eq!(converted.len(), 3);
assert_eq!(converted[0].role, "system");
assert_eq!(converted[1].role, "user");
assert_eq!(converted[2].role, "assistant");
}
#[test]
fn test_convert_messages_tool_calls() {
let tool_calls = vec![ToolCall {
id: "call_1".to_string(),
name: "search".to_string(),
arguments: serde_json::json!({"q": "test"}),
}];
let messages = vec![
ChatMessage::user("Search"),
ChatMessage::assistant_with_tool_calls(Some("Searching...".to_string()), tool_calls),
ChatMessage::tool_result("call_1", "search", "found it"),
];
let converted = convert_messages(messages);
assert_eq!(converted.len(), 3);
assert!(converted[1].tool_calls.is_some());
assert_eq!(converted[2].role, "tool");
assert_eq!(converted[2].tool_call_id, Some("call_1".to_string()));
}
#[test]
fn test_extract_choice_text_only() {
let choice = OpenAiChoice {
message: OpenAiResponseMessage {
content: Some("Hello!".to_string()),
tool_calls: None,
},
finish_reason: Some("stop".to_string()),
};
let (content, tool_calls) = extract_choice_content(&choice);
assert_eq!(content, Some("Hello!".to_string()));
assert!(tool_calls.is_empty());
}
#[test]
fn test_extract_choice_with_tool_calls() {
let choice = OpenAiChoice {
message: OpenAiResponseMessage {
content: Some("Let me search.".to_string()),
tool_calls: Some(vec![OpenAiResponseToolCall {
id: "call_1".to_string(),
function: OpenAiResponseFunction {
name: "search".to_string(),
arguments: r#"{"q":"test"}"#.to_string(),
},
}]),
},
finish_reason: Some("tool_calls".to_string()),
};
let (content, tool_calls) = extract_choice_content(&choice);
assert_eq!(content, Some("Let me search.".to_string()));
assert_eq!(tool_calls.len(), 1);
assert_eq!(tool_calls[0].name, "search");
assert_eq!(tool_calls[0].arguments["q"], "test");
}
}
+740
View File
@@ -0,0 +1,740 @@
use std::time::Duration;
use secrecy::{ExposeSecret, SecretString};
use serde::Deserialize;
use tokio::sync::RwLock;
// ─── Risk: hardcoded VS Code Copilot identity ───────────────────────────────
//
// The client ID and editor identity headers below are extracted from the
// VS Code Copilot Chat extension. This is the *only* publicly documented
// way to access the Copilot completions API with a personal GitHub token.
//
// **Known risks:**
// • GitHub may rotate or revoke this client ID at any time, which would
// break authentication for all IronClaw users until the constant is
// updated and a new release is shipped.
// • Using another product's client ID may violate GitHub's Terms of
// Service. Maintainers should seek explicit guidance from GitHub
// before shipping this to a wide audience.
// • The editor version strings (`vscode/1.99.3`, `copilot-chat/0.26.7`)
// will become stale and could eventually be rejected by the API.
//
// **Mitigation:** If GitHub publishes an official Copilot API client ID or
// an OAuth app registration flow for third-party tools, migrate to it
// immediately.
// ─────────────────────────────────────────────────────────────────────────────
pub const GITHUB_COPILOT_CLIENT_ID: &str = "Iv1.b507a08c87ecfe98";
pub const GITHUB_COPILOT_SCOPE: &str = "read:user";
pub const GITHUB_COPILOT_DEVICE_CODE_URL: &str = "https://github.com/login/device/code";
pub const GITHUB_COPILOT_ACCESS_TOKEN_URL: &str = "https://github.com/login/oauth/access_token";
pub const GITHUB_COPILOT_MODELS_URL: &str = "https://api.githubcopilot.com/models";
pub const GITHUB_COPILOT_TOKEN_URL: &str = "https://api.github.com/copilot_internal/v2/token";
pub const GITHUB_COPILOT_USER_AGENT: &str = "GitHubCopilotChat/0.26.7";
pub const GITHUB_COPILOT_EDITOR_VERSION: &str = "vscode/1.99.3";
pub const GITHUB_COPILOT_EDITOR_PLUGIN_VERSION: &str = "copilot-chat/0.26.7";
pub const GITHUB_COPILOT_INTEGRATION_ID: &str = "vscode-chat";
/// Buffer before token expiry to trigger a refresh (5 minutes).
const TOKEN_REFRESH_BUFFER_SECS: u64 = 300;
#[derive(Debug, Clone, Deserialize)]
pub struct DeviceCodeResponse {
pub device_code: String,
pub user_code: String,
pub verification_uri: String,
pub expires_in: u64,
#[serde(default = "default_poll_interval_secs")]
pub interval: u64,
}
#[derive(Debug, Clone, Deserialize)]
struct AccessTokenResponse {
access_token: Option<String>,
error: Option<String>,
error_description: Option<String>,
}
#[derive(Debug, thiserror::Error)]
pub enum GithubCopilotAuthError {
#[error("failed to start device login: {0}")]
DeviceCodeRequest(String),
#[error("failed to poll device login: {0}")]
TokenPolling(String),
#[error("device login was denied")]
AccessDenied,
#[error("device login expired before authorization completed")]
Expired,
#[error("github copilot token validation failed: {0}")]
Validation(String),
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub enum DevicePollingStatus {
Pending,
SlowDown,
Authorized(String),
}
pub fn default_headers() -> Vec<(String, String)> {
vec![
(
"User-Agent".to_string(),
GITHUB_COPILOT_USER_AGENT.to_string(),
),
(
"Editor-Version".to_string(),
GITHUB_COPILOT_EDITOR_VERSION.to_string(),
),
(
"Editor-Plugin-Version".to_string(),
GITHUB_COPILOT_EDITOR_PLUGIN_VERSION.to_string(),
),
(
"Copilot-Integration-Id".to_string(),
GITHUB_COPILOT_INTEGRATION_ID.to_string(),
),
]
}
pub fn default_poll_interval_secs() -> u64 {
5
}
pub async fn request_device_code(
client: &reqwest::Client,
) -> Result<DeviceCodeResponse, GithubCopilotAuthError> {
let response = client
.post(GITHUB_COPILOT_DEVICE_CODE_URL)
.header(reqwest::header::ACCEPT, "application/json")
.header(reqwest::header::USER_AGENT, GITHUB_COPILOT_USER_AGENT)
.form(&[
("client_id", GITHUB_COPILOT_CLIENT_ID),
("scope", GITHUB_COPILOT_SCOPE),
])
.send()
.await
.map_err(|e| {
tracing::warn!(
error = %e,
is_timeout = e.is_timeout(),
is_connect = e.is_connect(),
url = %GITHUB_COPILOT_DEVICE_CODE_URL,
"Copilot: device code request failed"
);
GithubCopilotAuthError::DeviceCodeRequest(format_reqwest_error(&e))
})?;
if !response.status().is_success() {
let status = response.status();
let body = response.text().await.unwrap_or_default();
tracing::warn!(
status = %status,
body = %truncate_for_error(&body),
"Copilot: device code endpoint returned error"
);
return Err(GithubCopilotAuthError::DeviceCodeRequest(format!(
"HTTP {status}: {}",
truncate_for_error(&body)
)));
}
let device = response
.json::<DeviceCodeResponse>()
.await
.map_err(|e| GithubCopilotAuthError::DeviceCodeRequest(e.to_string()))?;
Ok(device)
}
pub async fn poll_for_access_token(
client: &reqwest::Client,
device_code: &str,
) -> Result<DevicePollingStatus, GithubCopilotAuthError> {
let response = client
.post(GITHUB_COPILOT_ACCESS_TOKEN_URL)
.header(reqwest::header::ACCEPT, "application/json")
.header(reqwest::header::USER_AGENT, GITHUB_COPILOT_USER_AGENT)
.form(&[
("client_id", GITHUB_COPILOT_CLIENT_ID),
("device_code", device_code),
("grant_type", "urn:ietf:params:oauth:grant-type:device_code"),
])
.send()
.await
.map_err(|e| {
tracing::warn!(
error = %e,
is_timeout = e.is_timeout(),
is_connect = e.is_connect(),
url = %GITHUB_COPILOT_ACCESS_TOKEN_URL,
"Copilot: poll request failed"
);
GithubCopilotAuthError::TokenPolling(format_reqwest_error(&e))
})?;
if !response.status().is_success() {
let status = response.status();
let body = response.text().await.unwrap_or_default();
tracing::warn!(
status = %status,
body = %truncate_for_error(&body),
"Copilot: poll endpoint returned error"
);
return Err(GithubCopilotAuthError::TokenPolling(format!(
"HTTP {status}: {}",
truncate_for_error(&body)
)));
}
let body = response
.json::<AccessTokenResponse>()
.await
.map_err(|e| GithubCopilotAuthError::TokenPolling(e.to_string()))?;
if let Some(token) = body.access_token {
return Ok(DevicePollingStatus::Authorized(token));
}
match body.error.as_deref() {
Some("authorization_pending") | None => Ok(DevicePollingStatus::Pending),
Some("slow_down") => {
tracing::debug!("Copilot: GitHub requested slow_down, increasing poll interval");
Ok(DevicePollingStatus::SlowDown)
}
Some("access_denied") => {
tracing::warn!("Copilot: device login was denied by user");
Err(GithubCopilotAuthError::AccessDenied)
}
Some("expired_token") => {
tracing::warn!("Copilot: device code expired before authorization");
Err(GithubCopilotAuthError::Expired)
}
Some(other) => {
let desc = body
.error_description
.filter(|description| !description.is_empty())
.unwrap_or_else(|| other.to_string());
tracing::warn!(error = %other, description = %desc, "Copilot: unexpected poll error");
Err(GithubCopilotAuthError::TokenPolling(desc))
}
}
}
/// Maximum consecutive transient poll failures before giving up.
const MAX_POLL_FAILURES: u32 = 5;
pub async fn wait_for_device_login(
client: &reqwest::Client,
device: &DeviceCodeResponse,
) -> Result<String, GithubCopilotAuthError> {
let expires_at = std::time::Instant::now()
.checked_add(Duration::from_secs(device.expires_in))
.ok_or(GithubCopilotAuthError::Expired)?;
let mut poll_interval = device.interval.max(1);
let mut consecutive_failures: u32 = 0;
loop {
if std::time::Instant::now() >= expires_at {
tracing::warn!("Copilot: device login expired");
return Err(GithubCopilotAuthError::Expired);
}
tokio::time::sleep(Duration::from_secs(poll_interval)).await;
match poll_for_access_token(client, &device.device_code).await {
Ok(DevicePollingStatus::Pending) => {
consecutive_failures = 0;
}
Ok(DevicePollingStatus::SlowDown) => {
consecutive_failures = 0;
poll_interval = poll_interval.saturating_add(5);
}
Ok(DevicePollingStatus::Authorized(token)) => {
return Ok(token);
}
// Definitive failures — propagate immediately
Err(GithubCopilotAuthError::AccessDenied) => {
return Err(GithubCopilotAuthError::AccessDenied);
}
Err(GithubCopilotAuthError::Expired) => {
return Err(GithubCopilotAuthError::Expired);
}
// Transient failures — retry with backoff
Err(e) => {
consecutive_failures += 1;
tracing::warn!(
error = %e,
attempt = consecutive_failures,
max = MAX_POLL_FAILURES,
"Copilot: transient poll failure, will retry"
);
if consecutive_failures >= MAX_POLL_FAILURES {
tracing::error!(
error = %e,
"Copilot: too many consecutive poll failures, giving up"
);
return Err(e);
}
// Back off on transient errors
poll_interval = (poll_interval + 2).min(30);
}
}
}
}
/// Validate a GitHub OAuth token by performing the Copilot token exchange.
///
/// This exchanges the raw OAuth token for a Copilot session token (proving the
/// token is valid and the user has Copilot access), then verifies the session
/// token works against the models endpoint.
pub async fn validate_token(
client: &reqwest::Client,
token: &str,
) -> Result<(), GithubCopilotAuthError> {
// Step 1: Exchange the OAuth token for a Copilot session token.
// This validates both that the OAuth token is valid and that the user
// has an active Copilot subscription.
let session = exchange_copilot_token(client, token).await?;
// Step 2: Verify the session token works against the models endpoint.
let mut request = client
.get(GITHUB_COPILOT_MODELS_URL)
.bearer_auth(&session.token)
.timeout(Duration::from_secs(15));
for (key, value) in default_headers() {
request = request.header(&key, value);
}
let response = request.send().await.map_err(|e| {
tracing::warn!(
error = %e,
is_timeout = e.is_timeout(),
is_connect = e.is_connect(),
"Copilot: models endpoint request failed"
);
GithubCopilotAuthError::Validation(format_reqwest_error(&e))
})?;
if response.status().is_success() {
return Ok(());
}
let status = response.status();
let body = response.text().await.unwrap_or_default();
tracing::warn!(
status = %status,
body = %truncate_for_error(&body),
"Copilot: models endpoint returned error during validation"
);
Err(GithubCopilotAuthError::Validation(format!(
"HTTP {status}: {}",
truncate_for_error(&body)
)))
}
/// Response from the Copilot token exchange endpoint.
///
/// The `token` field is an HMAC-signed session token (not a JWT) used as
/// `Authorization: Bearer <token>` for requests to `api.githubcopilot.com`.
#[derive(Debug, Clone, Deserialize)]
pub struct CopilotTokenResponse {
/// The Copilot session token (HMAC-signed, not a JWT).
pub token: String,
/// Unix timestamp (seconds) when this token expires.
pub expires_at: u64,
}
/// Exchange a GitHub OAuth token for a Copilot API session token.
///
/// Calls `GET https://api.github.com/copilot_internal/v2/token` with the
/// GitHub OAuth token in `Authorization: token <oauth_token>` format.
/// Returns a short-lived session token for `api.githubcopilot.com`.
pub async fn exchange_copilot_token(
client: &reqwest::Client,
oauth_token: &str,
) -> Result<CopilotTokenResponse, GithubCopilotAuthError> {
let token_trimmed = oauth_token.trim();
let mut request = client
.get(GITHUB_COPILOT_TOKEN_URL)
.header(reqwest::header::ACCEPT, "application/json")
// GitHub Copilot uses `token` auth scheme, not `Bearer`
.header(
reqwest::header::AUTHORIZATION,
format!("token {token_trimmed}"),
)
.timeout(Duration::from_secs(15));
for (key, value) in default_headers() {
request = request.header(&key, value);
}
let response = request.send().await.map_err(|e| {
tracing::warn!(
error = %e,
is_timeout = e.is_timeout(),
is_connect = e.is_connect(),
"Copilot: token exchange HTTP request failed"
);
GithubCopilotAuthError::Validation(format_reqwest_error(&e))
})?;
if !response.status().is_success() {
let status = response.status();
let body = response.text().await.unwrap_or_default();
tracing::warn!(
status = %status,
body = %truncate_for_error(&body),
"Copilot: token exchange endpoint returned error"
);
return Err(GithubCopilotAuthError::Validation(format!(
"Copilot token exchange failed: HTTP {status}: {}",
truncate_for_error(&body)
)));
}
let token_response = response.json::<CopilotTokenResponse>().await.map_err(|e| {
tracing::warn!(error = %e, "Copilot: failed to parse token exchange response");
GithubCopilotAuthError::Validation(e.to_string())
})?;
Ok(token_response)
}
/// Manages a cached Copilot API session token with automatic refresh.
///
/// The GitHub Copilot API requires a two-step authentication:
/// 1. A long-lived GitHub OAuth token (from device login or IDE sign-in)
/// 2. A short-lived Copilot session token (exchanged via `/copilot_internal/v2/token`)
///
/// This manager caches the session token and refreshes it automatically
/// before it expires (with a 5-minute buffer).
pub struct CopilotTokenManager {
client: reqwest::Client,
oauth_token: SecretString,
cached: RwLock<Option<CachedCopilotToken>>,
}
#[derive(Clone)]
struct CachedCopilotToken {
token: SecretString,
expires_at: u64,
}
fn unix_now() -> u64 {
std::time::SystemTime::now()
.duration_since(std::time::UNIX_EPOCH)
.unwrap_or_default()
.as_secs()
}
impl CopilotTokenManager {
/// Create a new token manager with the given GitHub OAuth token.
pub fn new(client: reqwest::Client, oauth_token: String) -> Self {
Self {
client,
oauth_token: SecretString::from(oauth_token),
cached: RwLock::new(None),
}
}
/// Get a valid Copilot session token, refreshing if needed.
///
/// Returns the cached token if it has more than 5 minutes remaining,
/// otherwise exchanges the OAuth token for a fresh session token.
pub async fn get_token(&self) -> Result<SecretString, GithubCopilotAuthError> {
// Fast path: check if cached token is still valid under read lock.
{
let guard = self.cached.read().await;
if let Some(ref cached) = *guard {
let now = unix_now();
if cached.expires_at > now + TOKEN_REFRESH_BUFFER_SECS {
return Ok(cached.token.clone());
}
tracing::debug!(
expires_at = cached.expires_at,
now = now,
"Copilot: cached session token expired or expiring soon, refreshing"
);
}
}
// Slow path: acquire write lock and re-check (another caller may have
// already refreshed while we waited for the lock).
let mut guard = self.cached.write().await;
if let Some(ref cached) = *guard {
let now = unix_now();
if cached.expires_at > now + TOKEN_REFRESH_BUFFER_SECS {
return Ok(cached.token.clone());
}
}
let response =
exchange_copilot_token(&self.client, self.oauth_token.expose_secret()).await?;
let token = SecretString::from(response.token);
let expires_at = response.expires_at;
*guard = Some(CachedCopilotToken {
token: token.clone(),
expires_at,
});
tracing::debug!(expires_at = expires_at, "Copilot session token refreshed");
Ok(token)
}
/// Invalidate the cached session token.
///
/// Called when the API returns 401, so the next `get_token()` call
/// will perform a fresh token exchange instead of reusing the stale token.
pub async fn invalidate(&self) {
let mut guard = self.cached.write().await;
*guard = None;
tracing::debug!("Copilot session token invalidated");
}
}
fn truncate_for_error(body: &str) -> String {
const LIMIT: usize = 200;
if body.len() <= LIMIT {
return body.to_string();
}
let end = crate::util::floor_char_boundary(body, LIMIT);
format!("{}...", &body[..end])
}
/// Format a reqwest error with its full causal chain for debugging.
///
/// `reqwest::Error::to_string()` often just says "error sending request"
/// without the underlying cause (timeout, DNS, TLS, connection refused).
/// This walks the `source()` chain to surface the real problem.
fn format_reqwest_error(e: &reqwest::Error) -> String {
use std::error::Error;
let mut msg = e.to_string();
let mut source = e.source();
while let Some(cause) = source {
msg.push_str(&format!(": {cause}"));
source = cause.source();
}
msg
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn default_headers_include_required_identity_headers() {
let headers = default_headers();
assert!(headers.iter().any(|(key, value)| {
key == "Copilot-Integration-Id" && value == GITHUB_COPILOT_INTEGRATION_ID
}));
assert!(
headers
.iter()
.any(|(key, value)| key == "Editor-Version"
&& value == GITHUB_COPILOT_EDITOR_VERSION)
);
assert!(
headers
.iter()
.any(|(key, value)| key == "User-Agent" && value == GITHUB_COPILOT_USER_AGENT)
);
}
#[test]
fn truncate_for_error_preserves_utf8_boundaries() {
let long = "日本語".repeat(100);
let truncated = truncate_for_error(&long);
assert!(truncated.ends_with("..."));
assert!(truncated.is_char_boundary(truncated.len() - 3));
}
#[test]
fn truncate_for_error_short_strings_unchanged() {
let short = "hello";
assert_eq!(truncate_for_error(short), "hello");
}
// --- poll_for_access_token response parsing ---
fn parse_access_token_body(json: &str) -> AccessTokenResponse {
serde_json::from_str(json).expect("valid JSON")
}
#[test]
fn parse_authorization_pending_response() {
let body: AccessTokenResponse =
parse_access_token_body(r#"{"error": "authorization_pending"}"#);
assert!(body.access_token.is_none());
assert_eq!(body.error.as_deref(), Some("authorization_pending"));
}
#[test]
fn parse_slow_down_response() {
let body: AccessTokenResponse = parse_access_token_body(r#"{"error": "slow_down"}"#);
assert_eq!(body.error.as_deref(), Some("slow_down"));
}
#[test]
fn parse_access_denied_response() {
let body: AccessTokenResponse = parse_access_token_body(r#"{"error": "access_denied"}"#);
assert_eq!(body.error.as_deref(), Some("access_denied"));
}
#[test]
fn parse_expired_token_response() {
let body: AccessTokenResponse = parse_access_token_body(r#"{"error": "expired_token"}"#);
assert_eq!(body.error.as_deref(), Some("expired_token"));
}
#[test]
fn parse_successful_token_response() {
let body: AccessTokenResponse =
parse_access_token_body(r#"{"access_token": "ghu_abc123"}"#);
assert_eq!(body.access_token.as_deref(), Some("ghu_abc123"));
assert!(body.error.is_none());
}
#[test]
fn parse_error_with_description() {
let body: AccessTokenResponse = parse_access_token_body(
r#"{"error": "bad_verification_code", "error_description": "The code has expired"}"#,
);
assert_eq!(body.error.as_deref(), Some("bad_verification_code"));
assert_eq!(
body.error_description.as_deref(),
Some("The code has expired")
);
}
#[test]
fn parse_device_code_response_with_defaults() {
let json = r#"{
"device_code": "dc_123",
"user_code": "ABCD-1234",
"verification_uri": "https://github.com/login/device",
"expires_in": 900
}"#;
let resp: DeviceCodeResponse = serde_json::from_str(json).expect("valid JSON");
assert_eq!(resp.device_code, "dc_123");
assert_eq!(resp.user_code, "ABCD-1234");
assert_eq!(resp.interval, 5); // default_poll_interval_secs
assert_eq!(resp.expires_in, 900);
}
#[test]
fn parse_device_code_response_with_custom_interval() {
let json = r#"{
"device_code": "dc_456",
"user_code": "EFGH-5678",
"verification_uri": "https://github.com/login/device",
"expires_in": 600,
"interval": 10
}"#;
let resp: DeviceCodeResponse = serde_json::from_str(json).expect("valid JSON");
assert_eq!(resp.interval, 10);
}
// --- CopilotTokenManager ---
#[tokio::test]
async fn token_manager_caches_token_and_returns_same_value() {
// Pre-populate the cache with a token that expires far in the future.
let client = reqwest::Client::new();
let manager = CopilotTokenManager::new(client, "unused_oauth".to_string());
let far_future = unix_now() + 3600;
{
let mut guard = manager.cached.write().await;
*guard = Some(CachedCopilotToken {
token: SecretString::from("cached_session_token".to_string()),
expires_at: far_future,
});
}
let token = manager.get_token().await.expect("should return cached");
assert_eq!(token.expose_secret(), "cached_session_token");
// A second call should return the same cached token.
let token2 = manager.get_token().await.expect("should return cached");
assert_eq!(token2.expose_secret(), "cached_session_token");
}
#[tokio::test]
async fn token_manager_invalidation_clears_cache() {
let client = reqwest::Client::new();
let manager = CopilotTokenManager::new(client, "unused_oauth".to_string());
let far_future = unix_now() + 3600;
{
let mut guard = manager.cached.write().await;
*guard = Some(CachedCopilotToken {
token: SecretString::from("old_token".to_string()),
expires_at: far_future,
});
}
manager.invalidate().await;
let guard = manager.cached.read().await;
assert!(guard.is_none(), "cache should be empty after invalidation");
}
#[tokio::test]
async fn token_manager_expired_token_triggers_refresh_path() {
let client = reqwest::Client::new();
let manager = CopilotTokenManager::new(client, "unused_oauth".to_string());
// Set a token that is already expired (expires_at in the past).
{
let mut guard = manager.cached.write().await;
*guard = Some(CachedCopilotToken {
token: SecretString::from("stale_token".to_string()),
expires_at: 1, // way in the past
});
}
// get_token will try the slow path (token exchange) which will fail
// because we have no real server, but this proves the cached stale
// token is NOT returned.
let result = manager.get_token().await;
assert!(
result.is_err(),
"expired cached token should trigger exchange, which fails without a server"
);
}
#[tokio::test]
async fn token_manager_within_buffer_triggers_refresh() {
let client = reqwest::Client::new();
let manager = CopilotTokenManager::new(client, "unused_oauth".to_string());
// Set a token that expires within the refresh buffer window.
let expires_soon = unix_now() + TOKEN_REFRESH_BUFFER_SECS - 10;
{
let mut guard = manager.cached.write().await;
*guard = Some(CachedCopilotToken {
token: SecretString::from("expiring_soon".to_string()),
expires_at: expires_soon,
});
}
let result = manager.get_token().await;
assert!(
result.is_err(),
"token within buffer should trigger exchange"
);
}
// --- CopilotTokenResponse parsing ---
#[test]
fn parse_copilot_token_response() {
let json = r#"{"token": "tid=abc;exp=999;sku=123;sig=xyz", "expires_at": 1700000000}"#;
let resp: CopilotTokenResponse = serde_json::from_str(json).expect("valid JSON");
assert!(resp.token.starts_with("tid="));
assert_eq!(resp.expires_at, 1700000000);
}
}
+79 -2
View File
@@ -18,8 +18,12 @@ pub mod config;
pub mod costs;
pub mod error;
pub mod failover;
mod github_copilot;
pub(crate) mod github_copilot_auth;
mod nearai_chat;
pub mod oauth_helpers;
pub mod openai_codex_provider;
pub mod openai_codex_session;
mod provider;
mod reasoning;
pub mod recording;
@@ -29,6 +33,10 @@ pub mod retry;
mod rig_adapter;
pub mod session;
pub mod smart_routing;
mod token_refreshing;
#[cfg(test)]
mod codex_test_helpers;
pub mod image_models;
pub mod models;
@@ -37,12 +45,14 @@ pub mod vision_models;
pub use circuit_breaker::{CircuitBreakerConfig, CircuitBreakerProvider};
pub use config::{
BedrockConfig, CacheRetention, LlmConfig, NearAiConfig, OAUTH_PLACEHOLDER,
BedrockConfig, CacheRetention, LlmConfig, NearAiConfig, OAUTH_PLACEHOLDER, OpenAiCodexConfig,
RegistryProviderConfig,
};
pub use error::LlmError;
pub use failover::{CooldownConfig, FailoverProvider};
pub use nearai_chat::{DEFAULT_MODEL, ModelInfo, NearAiChatProvider, default_models};
pub use openai_codex_provider::OpenAiCodexProvider;
pub use openai_codex_session::{OpenAiCodexSession, OpenAiCodexSessionManager};
pub use provider::{
ChatMessage, CompletionRequest, CompletionResponse, ContentPart, FinishReason, ImageUrl,
LlmProvider, ModelMetadata, Role, ToolCall, ToolCompletionRequest, ToolCompletionResponse,
@@ -59,6 +69,7 @@ pub use retry::{RetryConfig, RetryProvider};
pub use rig_adapter::RigAdapter;
pub use session::{SessionConfig, SessionManager, create_session_manager};
pub use smart_routing::{SmartRoutingConfig, SmartRoutingProvider, TaskComplexity};
pub use token_refreshing::TokenRefreshingProvider;
use std::sync::Arc;
@@ -99,6 +110,15 @@ pub async fn create_llm_provider(
}
}
if config.backend == "openai_codex" {
return Err(LlmError::RequestFailed {
provider: "openai_codex".to_string(),
reason:
"OpenAI Codex uses a dedicated factory path. Use build_provider_chain() instead of create_llm_provider()."
.to_string(),
});
}
let reg_config = config
.provider
.as_ref()
@@ -155,6 +175,17 @@ fn create_registry_provider(
ProviderProtocol::OpenAiCompletions => create_openai_compat_from_registry(config),
ProviderProtocol::Anthropic => create_anthropic_from_registry(config),
ProviderProtocol::Ollama => create_ollama_from_registry(config),
ProviderProtocol::GithubCopilot => {
let provider =
github_copilot::GithubCopilotProvider::new(config, request_timeout_secs)?;
tracing::debug!(
provider = %config.provider_id,
model = %config.model,
base_url = %config.base_url,
"Using GitHub Copilot provider (token exchange)"
);
Ok(Arc::new(provider))
}
}
}
@@ -376,6 +407,47 @@ fn create_ollama_from_registry(
Ok(Arc::new(adapter))
}
/// Create an OpenAI Codex provider with OAuth authentication.
///
/// This is async because it needs to ensure authentication before
/// creating the provider (which requires a valid Bearer token).
///
/// Uses the Responses API (`chatgpt.com/backend-api/codex/responses`)
/// instead of the Chat Completions API, matching OpenClaw's approach.
async fn create_openai_codex_provider(
config: &LlmConfig,
) -> Result<Arc<dyn LlmProvider>, LlmError> {
let codex = config
.openai_codex
.as_ref()
.ok_or_else(|| LlmError::AuthFailed {
provider: "openai_codex".to_string(),
})?;
let session_mgr = Arc::new(OpenAiCodexSessionManager::new(codex.clone())?);
session_mgr.ensure_authenticated().await?;
let token = session_mgr.get_access_token().await?;
let provider = Arc::new(OpenAiCodexProvider::new(
&codex.model,
&codex.api_base_url,
token.expose_secret(),
config.request_timeout_secs,
)?);
tracing::info!(
"Using OpenAI Codex (Responses API, model: {}, base: {})",
codex.model,
codex.api_base_url,
);
Ok(Arc::new(TokenRefreshingProvider::new(
provider,
session_mgr,
)))
}
/// Create a cheap/fast LLM provider for lightweight tasks (heartbeat, routing, evaluation).
///
/// Resolution order:
@@ -462,7 +534,11 @@ pub async fn build_provider_chain(
),
LlmError,
> {
let llm = create_llm_provider(config, session.clone()).await?;
let llm: Arc<dyn LlmProvider> = if config.backend == "openai_codex" {
create_openai_codex_provider(config).await?
} else {
create_llm_provider(config, session.clone()).await?
};
tracing::debug!("LLM provider initialized: {}", llm.model_name());
// 1. Retry
@@ -634,6 +710,7 @@ mod tests {
request_timeout_secs: 120,
cheap_model: None,
smart_routing_cascade: true,
openai_codex: None,
}
}
+3 -2
View File
@@ -332,8 +332,8 @@ pub(crate) async fn fetch_openai_compatible_models(
/// Uses [`NearAiConfig::for_model_discovery()`] to construct a minimal NEAR AI
/// config, then wraps it in an `LlmConfig` with session config for auth.
pub(crate) fn build_nearai_model_fetch_config() -> crate::config::LlmConfig {
let auth_base_url =
std::env::var("NEARAI_AUTH_URL").unwrap_or_else(|_| "https://private.near.ai".to_string());
let auth_base_url = crate::config::helpers::env_or_override("NEARAI_AUTH_URL")
.unwrap_or_else(|| "https://private.near.ai".to_string());
crate::config::LlmConfig {
backend: "nearai".to_string(),
@@ -347,5 +347,6 @@ pub(crate) fn build_nearai_model_fetch_config() -> crate::config::LlmConfig {
request_timeout_secs: 120,
cheap_model: None,
smart_routing_cascade: false,
openai_codex: None,
}
}
+2 -2
View File
@@ -390,7 +390,7 @@ mod tests {
#[allow(clippy::await_holding_lock)]
#[tokio::test]
async fn bind_rejects_wildcard_ipv4() {
let _guard = ENV_MUTEX.lock().unwrap_or_else(|e| e.into_inner());
let _guard = ENV_MUTEX.lock().expect("env mutex poisoned");
let original = std::env::var("OAUTH_CALLBACK_HOST").ok();
// SAFETY: Under ENV_MUTEX, no concurrent env access.
unsafe { std::env::set_var("OAUTH_CALLBACK_HOST", "0.0.0.0") };
@@ -414,7 +414,7 @@ mod tests {
#[allow(clippy::await_holding_lock)]
#[tokio::test]
async fn bind_rejects_wildcard_ipv6() {
let _guard = ENV_MUTEX.lock().unwrap_or_else(|e| e.into_inner());
let _guard = ENV_MUTEX.lock().expect("env mutex poisoned");
let original = std::env::var("OAUTH_CALLBACK_HOST").ok();
// SAFETY: Under ENV_MUTEX, no concurrent env access.
unsafe { std::env::set_var("OAUTH_CALLBACK_HOST", "::") };
File diff suppressed because it is too large Load Diff
+731
View File
@@ -0,0 +1,731 @@
//! OAuth 2.0 session manager for OpenAI Codex (ChatGPT subscription).
//!
//! Supports two auth flows:
//! - **Device Code** (primary): Works on headless servers, no browser needed.
//! - **Browser PKCE** (fallback): Standard OAuth for local machines.
//!
//! Tokens are persisted to `~/.ironclaw/openai_codex_session.json` and
//! auto-refreshed before expiry.
use chrono::{DateTime, Utc};
use reqwest::Client;
use reqwest::header::{HeaderMap, HeaderValue, USER_AGENT};
use secrecy::SecretString;
use serde::{Deserialize, Serialize};
use tokio::sync::{Mutex, RwLock};
use crate::config::OpenAiCodexConfig;
use crate::error::LlmError;
/// Persisted OAuth session data.
///
/// Note: `Debug` is manually implemented to redact tokens.
#[derive(Serialize, Deserialize)]
pub struct OpenAiCodexSession {
pub(crate) access_token: String,
pub(crate) refresh_token: String,
pub(crate) expires_at: DateTime<Utc>,
pub(crate) created_at: DateTime<Utc>,
}
impl std::fmt::Debug for OpenAiCodexSession {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("OpenAiCodexSession")
.field("access_token", &"[REDACTED]")
.field("refresh_token", &"[REDACTED]")
.field("expires_at", &self.expires_at)
.field("created_at", &self.created_at)
.finish()
}
}
/// Request body for the device code usercode endpoint.
#[derive(Debug, Serialize)]
struct UserCodeRequest {
client_id: String,
}
/// Response from the device code usercode endpoint.
#[derive(Debug, Deserialize)]
struct UserCodeResponse {
/// Unique ID for this device auth session.
device_auth_id: String,
/// Code the user enters in their browser.
user_code: String,
/// URL where the user enters the code (may not be present).
#[serde(default = "default_verification_uri")]
verification_uri: String,
/// Polling interval in seconds (OpenAI sends this as a string).
#[serde(
default = "default_interval",
deserialize_with = "deserialize_string_or_u64"
)]
interval: u64,
/// Expiry timestamp (OpenAI sends `expires_at` as ISO-8601).
#[serde(default)]
expires_at: Option<String>,
/// Seconds until the device code expires (standard field, may not be present).
#[serde(default)]
expires_in: Option<u64>,
}
fn default_verification_uri() -> String {
"https://auth.openai.com/codex/device".to_string()
}
fn default_interval() -> u64 {
5
}
/// Deserialize a value that may be either a string or a number as u64.
fn deserialize_string_or_u64<'de, D>(deserializer: D) -> Result<u64, D::Error>
where
D: serde::Deserializer<'de>,
{
use serde::de;
struct StringOrU64;
impl<'de> de::Visitor<'de> for StringOrU64 {
type Value = u64;
fn expecting(&self, formatter: &mut std::fmt::Formatter) -> std::fmt::Result {
formatter.write_str("a string or integer")
}
fn visit_u64<E: de::Error>(self, v: u64) -> Result<u64, E> {
Ok(v)
}
fn visit_str<E: de::Error>(self, v: &str) -> Result<u64, E> {
v.parse().map_err(de::Error::custom)
}
}
deserializer.deserialize_any(StringOrU64)
}
impl UserCodeResponse {
/// Get the expiry duration in seconds, from either `expires_in` or `expires_at`.
fn expires_in_secs(&self) -> u64 {
if let Some(secs) = self.expires_in {
return secs;
}
if let Some(ref ts) = self.expires_at
&& let Ok(dt) = chrono::DateTime::parse_from_rfc3339(ts)
{
let remaining = dt.signed_duration_since(Utc::now()).num_seconds();
return remaining.max(0) as u64;
}
900 // default 15 minutes
}
}
/// Request body for polling the device auth token endpoint.
#[derive(Debug, Serialize)]
struct DeviceTokenPollRequest {
device_auth_id: String,
user_code: String,
}
/// Successful response from the device auth token endpoint.
/// Returns an authorization code + PKCE pair for the final token exchange.
#[derive(Debug, Deserialize)]
struct DeviceAuthCodeResponse {
authorization_code: String,
#[allow(dead_code)]
code_challenge: String,
code_verifier: String,
}
/// Response from the final OAuth token exchange.
#[derive(Debug, Deserialize)]
struct TokenResponse {
access_token: String,
#[serde(default)]
refresh_token: String,
#[serde(default)]
expires_in: u64,
#[serde(default)]
#[allow(dead_code)]
token_type: String,
}
/// Manages OpenAI Codex OAuth sessions with persistence and auto-refresh.
pub struct OpenAiCodexSessionManager {
config: OpenAiCodexConfig,
client: Client,
session: RwLock<Option<OpenAiCodexSession>>,
renewal_lock: Mutex<()>,
}
impl OpenAiCodexSessionManager {
/// Create a new session manager. Tries to load existing session from disk.
///
/// # Errors
///
/// Returns `LlmError` if the HTTP client cannot be constructed.
pub fn new(config: OpenAiCodexConfig) -> Result<Self, LlmError> {
let mut headers = HeaderMap::new();
headers.insert(
USER_AGENT,
HeaderValue::from_static(concat!("ironclaw/", env!("CARGO_PKG_VERSION"))),
);
let client = Client::builder()
.default_headers(headers)
.timeout(std::time::Duration::from_secs(30))
.build()
.map_err(|e| LlmError::RequestFailed {
provider: "openai_codex".into(),
reason: format!("HTTP client build failed: {e}"),
})?;
let mgr = Self {
config,
client,
session: RwLock::new(None),
renewal_lock: Mutex::new(()),
};
// Try synchronous load from disk during construction
if let Ok(data) = std::fs::read_to_string(&mgr.config.session_path)
&& let Ok(session) = serde_json::from_str::<OpenAiCodexSession>(&data)
&& let Ok(mut guard) = mgr.session.try_write()
{
*guard = Some(session);
tracing::info!(
"Loaded OpenAI Codex session from {}",
mgr.config.session_path.display()
);
}
Ok(mgr)
}
/// Check if we have a session (may be expired).
pub async fn has_session(&self) -> bool {
self.session.read().await.is_some()
}
/// Check if the current access token needs refreshing.
pub async fn needs_refresh(&self) -> bool {
let guard = self.session.read().await;
match guard.as_ref() {
None => true,
Some(s) => {
let margin =
chrono::Duration::seconds(self.config.token_refresh_margin_secs as i64);
Utc::now() + margin >= s.expires_at
}
}
}
/// Get the current access token, refreshing if needed.
///
/// If the token is within the refresh margin, silently refreshes first.
/// If no session exists, returns an AuthFailed error.
pub async fn get_access_token(&self) -> Result<SecretString, LlmError> {
if self.needs_refresh().await {
let has_refresh = self
.session
.read()
.await
.as_ref()
.map(|s| !s.refresh_token.is_empty())
.unwrap_or(false);
if has_refresh {
self.refresh_tokens().await?;
} else {
return Err(LlmError::AuthFailed {
provider: "openai_codex".to_string(),
});
}
}
let guard = self.session.read().await;
guard
.as_ref()
.map(|s| SecretString::from(s.access_token.clone()))
.ok_or_else(|| LlmError::AuthFailed {
provider: "openai_codex".to_string(),
})
}
/// Ensure we have a valid session. Loads from disk, refreshes, or prompts login.
pub async fn ensure_authenticated(&self) -> Result<(), LlmError> {
// Try loading from disk if we don't have a session
if !self.has_session().await {
let _ = self.load_session().await;
}
if !self.has_session().await {
// No session at all -- need to authenticate
return self.device_code_login().await;
}
if self.needs_refresh().await {
// Try refresh; if it fails, re-authenticate
match self.refresh_tokens().await {
Ok(()) => Ok(()),
Err(e) => {
tracing::info!("Token refresh failed ({}), re-authenticating...", e);
self.device_code_login().await
}
}
} else {
Ok(())
}
}
/// Run OpenAI's device code auth flow.
///
/// Uses OpenAI's custom `/api/accounts/deviceauth/*` endpoints (not the standard
/// Auth0 `/oauth/device/code` which is behind Cloudflare managed challenge).
///
/// Flow:
/// 1. POST `/api/accounts/deviceauth/usercode` → get device_auth_id + user_code
/// 2. Poll POST `/api/accounts/deviceauth/token` → get authorization_code + PKCE
/// 3. Exchange via POST `/oauth/token` → get access_token + refresh_token
pub async fn device_code_login(&self) -> Result<(), LlmError> {
let _guard = self.renewal_lock.lock().await;
let auth_base = format!("{}/api/accounts", self.config.auth_endpoint);
// Step 1: Request device code
let usercode_url = format!("{}/deviceauth/usercode", auth_base);
let resp = self
.client
.post(&usercode_url)
.json(&UserCodeRequest {
client_id: self.config.client_id.clone(),
})
.send()
.await
.map_err(|e| LlmError::SessionRenewalFailed {
provider: "openai_codex".to_string(),
reason: format!("Device code request failed: {}", e),
})?;
if !resp.status().is_success() {
let status = resp.status();
let body = resp.text().await.unwrap_or_default();
return Err(LlmError::SessionRenewalFailed {
provider: "openai_codex".to_string(),
reason: format!("Device code request failed: HTTP {} -- {}", status, body),
});
}
let body_text = resp
.text()
.await
.map_err(|e| LlmError::SessionRenewalFailed {
provider: "openai_codex".to_string(),
reason: format!("Failed to read device code response: {}", e),
})?;
tracing::debug!("Device code response received ({} bytes)", body_text.len());
let device: UserCodeResponse =
serde_json::from_str(&body_text).map_err(|e| LlmError::SessionRenewalFailed {
provider: "openai_codex".to_string(),
reason: format!(
"Failed to parse device code response: {} ({} bytes)",
e,
body_text.len()
),
})?;
// Step 2: Display code to user
println!();
println!("===========================================================");
println!(" OpenAI Codex Authentication ");
println!("===========================================================");
println!();
println!(" 1. Open this URL in any browser:");
println!(" {}", device.verification_uri);
println!();
println!(" 2. Enter this code:");
println!();
println!(" [ {} ]", device.user_code);
println!();
let expires_secs = device.expires_in_secs();
println!(
" Waiting for authorization... (expires in {} min)",
expires_secs / 60
);
println!("===========================================================");
println!();
// Step 3: Poll for authorization code
let poll_url = format!("{}/deviceauth/token", auth_base);
let mut interval = std::time::Duration::from_secs(device.interval.max(5));
let deadline = tokio::time::Instant::now() + std::time::Duration::from_secs(expires_secs);
let auth_code = loop {
tokio::time::sleep(interval).await;
if tokio::time::Instant::now() >= deadline {
return Err(LlmError::SessionRenewalFailed {
provider: "openai_codex".to_string(),
reason: "Device code authorization timed out".to_string(),
});
}
let resp = self
.client
.post(&poll_url)
.json(&DeviceTokenPollRequest {
device_auth_id: device.device_auth_id.clone(),
user_code: device.user_code.clone(),
})
.send()
.await
.map_err(|e| LlmError::SessionRenewalFailed {
provider: "openai_codex".to_string(),
reason: format!("Token poll request failed: {}", e),
})?;
let status = resp.status();
if status.is_success() {
let code_resp: DeviceAuthCodeResponse =
resp.json()
.await
.map_err(|e| LlmError::SessionRenewalFailed {
provider: "openai_codex".to_string(),
reason: format!("Failed to parse auth code response: {}", e),
})?;
break code_resp;
}
// 403 = authorization_pending, keep polling
// 404 = device code not found / not enabled
if status == reqwest::StatusCode::FORBIDDEN {
continue;
}
if status == reqwest::StatusCode::NOT_FOUND {
return Err(LlmError::SessionRenewalFailed {
provider: "openai_codex".to_string(),
reason: "Device code login is not enabled. Please check your OpenAI account settings.".to_string(),
});
}
// Slow down on 429, cap at 60s to avoid unbounded growth
if status == reqwest::StatusCode::TOO_MANY_REQUESTS {
interval = (interval + std::time::Duration::from_secs(5))
.min(std::time::Duration::from_secs(60));
continue;
}
let body = resp.text().await.unwrap_or_default();
return Err(LlmError::SessionRenewalFailed {
provider: "openai_codex".to_string(),
reason: format!("Device auth poll failed: HTTP {} -- {}", status, body),
});
};
// Step 4: Exchange authorization code for tokens (form-encoded, per Auth0 spec)
let token_url = format!("{}/oauth/token", self.config.auth_endpoint);
let resp = self
.client
.post(&token_url)
.form(&[
("grant_type", "authorization_code"),
("code", &auth_code.authorization_code),
("code_verifier", &auth_code.code_verifier),
("client_id", &self.config.client_id),
(
"redirect_uri",
&format!("{}/deviceauth/callback", self.config.auth_endpoint),
),
])
.send()
.await
.map_err(|e| LlmError::SessionRenewalFailed {
provider: "openai_codex".to_string(),
reason: format!("Token exchange failed: {}", e),
})?;
if !resp.status().is_success() {
let status = resp.status();
let body = resp.text().await.unwrap_or_default();
return Err(LlmError::SessionRenewalFailed {
provider: "openai_codex".to_string(),
reason: format!("Token exchange failed: HTTP {} -- {}", status, body),
});
}
let token_resp: TokenResponse =
resp.json()
.await
.map_err(|e| LlmError::SessionRenewalFailed {
provider: "openai_codex".to_string(),
reason: format!("Failed to parse token response: {}", e),
})?;
let session = OpenAiCodexSession {
access_token: token_resp.access_token,
refresh_token: token_resp.refresh_token,
expires_at: Utc::now()
+ chrono::Duration::seconds(if token_resp.expires_in > 0 {
token_resp.expires_in
} else {
tracing::warn!("Token response has expires_in=0, defaulting to 3600s");
3600
} as i64),
created_at: Utc::now(),
};
self.save_session(&session).await?;
self.set_session(session).await;
println!();
println!("Authentication successful!");
println!();
Ok(())
}
/// Refresh the access token using the refresh token.
pub async fn refresh_tokens(&self) -> Result<(), LlmError> {
let _guard = self.renewal_lock.lock().await;
// Double-check: another task may have refreshed while we waited on the lock
if !self.needs_refresh().await {
return Ok(());
}
let refresh_token = {
let guard = self.session.read().await;
guard
.as_ref()
.map(|s| s.refresh_token.clone())
.ok_or_else(|| LlmError::AuthFailed {
provider: "openai_codex".to_string(),
})?
};
let token_url = format!("{}/oauth/token", self.config.auth_endpoint);
let resp = self
.client
.post(&token_url)
.form(&[
("grant_type", "refresh_token"),
("refresh_token", refresh_token.as_str()),
("client_id", self.config.client_id.as_str()),
])
.send()
.await
.map_err(|e| LlmError::SessionRenewalFailed {
provider: "openai_codex".to_string(),
reason: format!("Token refresh request failed: {}", e),
})?;
if !resp.status().is_success() {
let status = resp.status();
let body = resp.text().await.unwrap_or_default();
return Err(LlmError::SessionRenewalFailed {
provider: "openai_codex".to_string(),
reason: format!("Token refresh failed: HTTP {} -- {}", status, body),
});
}
let token_resp: TokenResponse =
resp.json()
.await
.map_err(|e| LlmError::SessionRenewalFailed {
provider: "openai_codex".to_string(),
reason: format!("Failed to parse refresh response: {}", e),
})?;
let session = OpenAiCodexSession {
access_token: token_resp.access_token,
refresh_token: token_resp.refresh_token,
expires_at: Utc::now()
+ chrono::Duration::seconds(if token_resp.expires_in > 0 {
token_resp.expires_in
} else {
tracing::warn!("Token response has expires_in=0, defaulting to 3600s");
3600
} as i64),
created_at: Utc::now(),
};
self.save_session(&session).await?;
self.set_session(session).await;
tracing::debug!("OpenAI Codex token refreshed successfully");
Ok(())
}
/// Save session data to disk with restrictive permissions.
pub async fn save_session(&self, session: &OpenAiCodexSession) -> Result<(), LlmError> {
if let Some(parent) = self.config.session_path.parent() {
tokio::fs::create_dir_all(parent).await.map_err(|e| {
LlmError::Io(std::io::Error::new(
e.kind(),
format!("Failed to create session directory: {}", e),
))
})?;
}
let json =
serde_json::to_string_pretty(session).map_err(|e| LlmError::SessionRenewalFailed {
provider: "openai_codex".to_string(),
reason: format!("Failed to serialize session: {}", e),
})?;
tokio::fs::write(&self.config.session_path, &json)
.await
.map_err(|e| {
LlmError::Io(std::io::Error::new(
e.kind(),
format!("Failed to write session file: {}", e),
))
})?;
#[cfg(unix)]
{
use std::os::unix::fs::PermissionsExt;
let perms = std::fs::Permissions::from_mode(0o600);
tokio::fs::set_permissions(&self.config.session_path, perms)
.await
.map_err(|e| {
LlmError::Io(std::io::Error::new(
e.kind(),
format!("Failed to set permissions: {}", e),
))
})?;
}
Ok(())
}
/// Load session from disk.
pub async fn load_session(&self) -> Result<(), LlmError> {
let data = tokio::fs::read_to_string(&self.config.session_path)
.await
.map_err(|e| {
LlmError::Io(std::io::Error::new(
e.kind(),
format!("Failed to read session file: {}", e),
))
})?;
let session: OpenAiCodexSession =
serde_json::from_str(&data).map_err(|e| LlmError::SessionRenewalFailed {
provider: "openai_codex".to_string(),
reason: format!("Failed to parse session file: {}", e),
})?;
let mut guard = self.session.write().await;
*guard = Some(session);
tracing::info!(
"Loaded OpenAI Codex session from {}",
self.config.session_path.display()
);
Ok(())
}
/// Set session directly (for testing or after auth).
pub async fn set_session(&self, session: OpenAiCodexSession) {
let mut guard = self.session.write().await;
*guard = Some(session);
}
/// Handle a 401 response by refreshing, or re-authenticating.
pub async fn handle_auth_failure(&self) -> Result<(), LlmError> {
match self.refresh_tokens().await {
Ok(()) => Ok(()),
Err(_) => self.device_code_login().await,
}
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::llm::codex_test_helpers::test_codex_config as test_config;
use tempfile::tempdir;
#[tokio::test]
async fn test_save_and_load_session() {
let dir = tempdir().unwrap();
let path = dir.path().join("session.json");
let config = test_config(path.clone());
let mgr = OpenAiCodexSessionManager::new(config).unwrap();
// No session initially
assert!(!mgr.has_session().await);
// Save a session
let session = OpenAiCodexSession {
access_token: "access_abc".to_string(),
refresh_token: "refresh_xyz".to_string(),
expires_at: chrono::Utc::now() + chrono::Duration::hours(1),
created_at: chrono::Utc::now(),
};
mgr.save_session(&session).await.unwrap();
mgr.set_session(session).await;
assert!(mgr.has_session().await);
// Load from disk in a new manager
let config2 = test_config(path);
let mgr2 = OpenAiCodexSessionManager::new(config2).unwrap();
mgr2.load_session().await.unwrap();
assert!(mgr2.has_session().await);
}
#[tokio::test]
async fn test_needs_refresh_when_near_expiry() {
let dir = tempdir().unwrap();
let config = test_config(dir.path().join("session.json"));
let mgr = OpenAiCodexSessionManager::new(config).unwrap();
// Token expiring in 2 minutes (margin is 300s = 5 min)
let session = OpenAiCodexSession {
access_token: "access_abc".to_string(),
refresh_token: "refresh_xyz".to_string(),
expires_at: chrono::Utc::now() + chrono::Duration::minutes(2),
created_at: chrono::Utc::now(),
};
mgr.set_session(session).await;
assert!(mgr.needs_refresh().await);
}
#[test]
fn device_code_parse_error_redacts_body() {
// Regression: the parse error used to include raw body_text which could
// contain sensitive auth data. Now it only shows byte count.
let body_text = r#"{"secret_token":"sk-12345","error":"unexpected"}"#;
let err: Result<UserCodeResponse, _> = serde_json::from_str(body_text);
assert!(err.is_err());
let e = err.unwrap_err();
let error_msg = format!(
"Failed to parse device code response: {} ({} bytes)",
e,
body_text.len()
);
assert!(
!error_msg.contains("sk-12345"),
"error message must not contain raw body: {error_msg}"
);
assert!(
error_msg.contains("bytes"),
"error message should show byte count"
);
}
#[tokio::test]
async fn test_no_refresh_when_fresh() {
let dir = tempdir().unwrap();
let config = test_config(dir.path().join("session.json"));
let mgr = OpenAiCodexSessionManager::new(config).unwrap();
// Token expiring in 30 minutes (margin is 300s = 5 min)
let session = OpenAiCodexSession {
access_token: "access_abc".to_string(),
refresh_token: "refresh_xyz".to_string(),
expires_at: chrono::Utc::now() + chrono::Duration::minutes(30),
created_at: chrono::Utc::now(),
};
mgr.set_session(session).await;
assert!(!mgr.needs_refresh().await);
}
}
+2
View File
@@ -37,6 +37,8 @@ pub enum ProviderProtocol {
Anthropic,
/// Ollama API (OpenAI-ish, no API key required).
Ollama,
/// GitHub Copilot API (OpenAI-compatible with token exchange).
GithubCopilot,
}
/// How the setup wizard should collect credentials for this provider.
+1 -1
View File
@@ -132,7 +132,7 @@ fn round_f32_to_f64(val: f32) -> f64 {
///
/// This is applied as a clone-and-transform at the provider boundary so the
/// original tool definitions remain unchanged for other providers.
fn normalize_schema_strict(schema: &JsonValue) -> JsonValue {
pub(crate) fn normalize_schema_strict(schema: &JsonValue) -> JsonValue {
let mut schema = schema.clone();
normalize_schema_recursive(&mut schema);
schema
+191
View File
@@ -0,0 +1,191 @@
//! Token-refreshing LlmProvider decorator for OpenAI Codex.
//!
//! Wraps an `OpenAiCodexProvider` and:
//! - Pre-emptively refreshes the OAuth access token before each call if near expiry
//! - Updates the inner provider's token after refresh (no client rebuild needed)
//! - Retries once on `AuthFailed` / `SessionExpired` after refreshing
//! - Overrides `cost_per_token()` to return (0, 0) since billing is through subscription
use std::sync::Arc;
use async_trait::async_trait;
use rust_decimal::Decimal;
use secrecy::ExposeSecret;
use crate::error::LlmError;
use crate::llm::openai_codex_provider::OpenAiCodexProvider;
use crate::llm::openai_codex_session::OpenAiCodexSessionManager;
use crate::llm::provider::{
CompletionRequest, CompletionResponse, LlmProvider, ModelMetadata, ToolCompletionRequest,
ToolCompletionResponse,
};
/// Decorator that refreshes OAuth tokens before API calls and reports zero cost.
///
/// The inner `OpenAiCodexProvider` manages its own token state, so after a
/// refresh we just call `update_token()` -- no client rebuild is needed.
pub struct TokenRefreshingProvider {
inner: Arc<OpenAiCodexProvider>,
session: Arc<OpenAiCodexSessionManager>,
}
impl TokenRefreshingProvider {
pub fn new(inner: Arc<OpenAiCodexProvider>, session: Arc<OpenAiCodexSessionManager>) -> Self {
Self { inner, session }
}
/// Push a fresh token from the session manager into the inner provider.
async fn update_inner_token(&self) -> Result<(), LlmError> {
let token = self.session.get_access_token().await?;
self.inner.update_token(token.expose_secret()).await?;
tracing::debug!("Updated inner provider token after refresh");
Ok(())
}
/// Best-effort pre-emptive token refresh before an API call.
///
/// If refresh fails (e.g., no refresh token), we log and continue so the
/// actual request still fires and the retry-on-auth-failure path can kick in.
async fn ensure_fresh_token(&self) {
if self.session.needs_refresh().await {
match self.session.refresh_tokens().await {
Ok(()) => {
if let Err(e) = self.update_inner_token().await {
tracing::warn!(
"Pre-emptive token update failed: {e}, will retry on auth failure"
);
}
}
Err(e) => {
tracing::warn!(
"Pre-emptive token refresh failed: {e}, will retry on auth failure"
);
}
}
}
}
}
#[async_trait]
impl LlmProvider for TokenRefreshingProvider {
fn model_name(&self) -> &str {
self.inner.model_name()
}
fn cost_per_token(&self) -> (Decimal, Decimal) {
(Decimal::ZERO, Decimal::ZERO)
}
async fn complete(&self, request: CompletionRequest) -> Result<CompletionResponse, LlmError> {
self.ensure_fresh_token().await;
match self.inner.complete(request.clone()).await {
Err(LlmError::AuthFailed { .. } | LlmError::SessionExpired { .. }) => {
tracing::info!("Auth failure during complete(), refreshing and retrying once");
self.session.handle_auth_failure().await?;
self.update_inner_token().await?;
self.inner.complete(request).await
}
other => other,
}
}
async fn complete_with_tools(
&self,
request: ToolCompletionRequest,
) -> Result<ToolCompletionResponse, LlmError> {
self.ensure_fresh_token().await;
match self.inner.complete_with_tools(request.clone()).await {
Err(LlmError::AuthFailed { .. } | LlmError::SessionExpired { .. }) => {
tracing::info!(
"Auth failure during complete_with_tools(), refreshing and retrying once"
);
self.session.handle_auth_failure().await?;
self.update_inner_token().await?;
self.inner.complete_with_tools(request).await
}
other => other,
}
}
async fn list_models(&self) -> Result<Vec<String>, LlmError> {
self.ensure_fresh_token().await;
self.inner.list_models().await
}
async fn model_metadata(&self) -> Result<ModelMetadata, LlmError> {
self.ensure_fresh_token().await;
self.inner.model_metadata().await
}
fn active_model_name(&self) -> String {
self.inner.model_name().to_string()
}
fn effective_model_name(&self, requested_model: Option<&str>) -> String {
self.inner.effective_model_name(requested_model)
}
fn set_model(&self, model: &str) -> Result<(), LlmError> {
self.inner.set_model(model)
}
fn calculate_cost(&self, _input_tokens: u32, _output_tokens: u32) -> Decimal {
Decimal::ZERO
}
fn cache_write_multiplier(&self) -> Decimal {
self.inner.cache_write_multiplier()
}
fn cache_read_discount(&self) -> Decimal {
self.inner.cache_read_discount()
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::llm::codex_test_helpers::{make_test_jwt, test_codex_config};
use crate::llm::openai_codex_session::OpenAiCodexSessionManager;
use tempfile::tempdir;
fn make_provider_and_session() -> (TokenRefreshingProvider, tempfile::TempDir) {
let dir = tempdir().unwrap();
let config = test_codex_config(dir.path().join("session.json"));
let jwt = make_test_jwt("acct_test");
let inner = Arc::new(
OpenAiCodexProvider::new(&config.model, &config.api_base_url, &jwt, 300)
.expect("provider creation should succeed"),
);
let session = Arc::new(OpenAiCodexSessionManager::new(config).unwrap());
(TokenRefreshingProvider::new(inner, session), dir)
}
#[test]
fn test_model_name_delegates() {
let (provider, _dir) = make_provider_and_session();
assert_eq!(provider.model_name(), "gpt-5.3-codex");
}
#[test]
fn test_cost_per_token_zero() {
let (provider, _dir) = make_provider_and_session();
let (input, output) = provider.cost_per_token();
assert_eq!(input, Decimal::ZERO);
assert_eq!(output, Decimal::ZERO);
}
#[test]
fn test_calculate_cost_zero() {
let (provider, _dir) = make_provider_and_session();
assert_eq!(provider.calculate_cost(1000, 500), Decimal::ZERO);
}
#[test]
fn test_active_model_name_delegates() {
let (provider, _dir) = make_provider_and_session();
assert_eq!(provider.active_model_name(), "gpt-5.3-codex");
}
}
+41
View File
@@ -139,6 +139,47 @@ async fn async_main() -> anyhow::Result<()> {
)
.await;
}
Some(Command::Login { openai_codex }) => {
init_cli_tracing();
if *openai_codex {
// Resolve codex config so OPENAI_CODEX_* env overrides are
// honoured even when LLM_BACKEND isn't set to openai_codex.
let codex_config = {
let config = Config::from_env()
.await
.map_err(|e| anyhow::anyhow!("{}", e))?;
config.llm.openai_codex.unwrap_or_else(|| {
use ironclaw::llm::OpenAiCodexConfig;
let mut cfg = OpenAiCodexConfig::default();
if let Ok(v) = std::env::var("OPENAI_CODEX_AUTH_URL") {
cfg.auth_endpoint = v;
}
if let Ok(v) = std::env::var("OPENAI_CODEX_API_URL") {
cfg.api_base_url = v;
}
if let Ok(v) = std::env::var("OPENAI_CODEX_CLIENT_ID") {
cfg.client_id = v;
}
if let Ok(v) = std::env::var("OPENAI_CODEX_SESSION_PATH") {
cfg.session_path = std::path::PathBuf::from(v);
}
cfg
})
};
let mgr = ironclaw::llm::OpenAiCodexSessionManager::new(codex_config)
.map_err(|e| anyhow::anyhow!("{}", e))?;
mgr.device_code_login()
.await
.map_err(|e| anyhow::anyhow!("{}", e))?;
println!(
"OpenAI Codex authentication complete. Set LLM_BACKEND=openai_codex to use it."
);
} else {
println!("Specify a provider to authenticate with:");
println!(" ironclaw login --openai-codex (ChatGPT subscription)");
}
return Ok(());
}
Some(Command::Onboard {
skip_auth,
channels_only,
+1 -1
View File
@@ -99,7 +99,7 @@ pub struct Settings {
pub secrets_master_key_hex: Option<String>,
// === Step 3: Inference Provider ===
/// LLM backend: "nearai", "anthropic", "openai", "ollama", "openai_compatible", "tinfoil", "bedrock".
/// LLM backend: "nearai", "anthropic", "openai", "github_copilot", "ollama", "openai_compatible", "tinfoil", "bedrock".
#[serde(default)]
pub llm_backend: Option<String>,
+30 -21
View File
@@ -216,8 +216,9 @@ env-var mode or skipped secrets.
|----------|-------------|-------------|---------|
| NEAR AI Chat | Browser OAuth or session token | - | `NEARAI_SESSION_TOKEN` |
| NEAR AI Cloud | API key | `llm_nearai_api_key` | `NEARAI_API_KEY` |
| Anthropic | API key | `anthropic_api_key` | `ANTHROPIC_API_KEY` |
| OpenAI | API key | `openai_api_key` | `OPENAI_API_KEY` |
| Anthropic | API key | `llm_anthropic_api_key` | `ANTHROPIC_API_KEY` |
| OpenAI | API key | `llm_openai_api_key` | `OPENAI_API_KEY` |
| GitHub Copilot | OAuth token | `llm_github_copilot_token` | `GITHUB_COPILOT_TOKEN` |
| Ollama | None | - | - |
| OpenRouter | API key | `llm_openrouter_api_key` | `OPENROUTER_API_KEY` |
| OpenAI-compatible | Optional API key | `llm_compatible_api_key` | `LLM_API_KEY` |
@@ -240,6 +241,12 @@ with its own secret name and env var. It is **not** stored as `openai_compatible
5. Preserve `selected_model` on a same-backend re-run; clear it only when
switching to a different backend
**GitHub Copilot** (`setup_github_copilot`):
- Offers **GitHub device login** (recommended) or manual token paste
- Device login uses the VS Code Copilot OAuth client and stores the resulting token as `llm_github_copilot_token`
- Validates the token against `https://api.githubcopilot.com/models` before saving
- Injects `GITHUB_COPILOT_TOKEN` into the config overlay for immediate provider use
**NEAR AI** (`setup_nearai`):
- Calls `session_manager.ensure_authenticated()` which shows the auth menu:
- Options 1-2 (GitHub/Google): browser OAuth → **NEAR AI Chat** mode
@@ -406,26 +413,24 @@ Contains only the settings needed BEFORE database connection. Written by
```env
DATABASE_BACKEND="libsql"
LIBSQL_PATH="/Users/name/.ironclaw/ironclaw.db"
LLM_BACKEND="openai_compatible"
LLM_BASE_URL="http://my-vllm:8000/v1"
SECRETS_MASTER_KEY="..." # only if env key source selected
ONBOARD_COMPLETED="true"
```
Or for PostgreSQL + NEAR AI:
Or for PostgreSQL:
```env
DATABASE_BACKEND="postgres"
DATABASE_URL="postgres://user:pass@localhost/ironclaw"
LLM_BACKEND="nearai"
```
Or for Ollama:
```env
LLM_BACKEND="ollama"
OLLAMA_BASE_URL="http://localhost:11434"
SECRETS_MASTER_KEY="..."
ONBOARD_COMPLETED="true"
```
**Why separate?** Chicken-and-egg: you need `DATABASE_BACKEND` to know
which database to connect to, and `LLM_BACKEND` to know whether to
attempt NEAR AI session auth -- neither can be stored in the database.
which database to connect to, and `SECRETS_MASTER_KEY` to decrypt the
secrets store — neither can be stored in the database. LLM settings
(`LLM_BACKEND`, base URLs, model names) are persisted to the DB via
`persist_settings()` and loaded after connection. API keys are stored
encrypted in the secrets DB.
**Layer 2: Database settings table** (everything else)
@@ -487,16 +492,20 @@ Final step of the wizard:
4. Print configuration summary
```
Bootstrap vars written to `~/.ironclaw/.env`:
Bootstrap vars written to `~/.ironclaw/.env` (only true chicken-and-egg vars
that are needed before the DB is connected):
- `DATABASE_BACKEND` (always)
- `DATABASE_URL` (if postgres)
- `LIBSQL_PATH` (if libsql)
- `LIBSQL_URL` (if turso sync)
- `LLM_BACKEND` (always, when set)
- `LLM_BASE_URL` (if openai_compatible)
- `OLLAMA_BASE_URL` (if ollama)
- `NEARAI_API_KEY` (if API key auth path)
- `SECRETS_MASTER_KEY` (if env key source selected in Step 2)
- `ONBOARD_COMPLETED` (always, "true")
- Channel/sandbox vars: `CLAUDE_CODE_ENABLED`, `SIGNAL_HTTP_URL`, `SIGNAL_ACCOUNT`, etc. (channel init may precede DB)
LLM settings (`LLM_BACKEND`, `LLM_BASE_URL`, model, API keys) are persisted
to the DB via `persist_settings()` and loaded by `Config::from_db_with_toml()`
after connection. API keys are stored encrypted in the secrets DB and injected
via `inject_llm_keys_from_secrets()`.
**Invariant:** Both Layer 1 and Layer 2 must be written. If the database
write fails, the wizard returns an error and the `.env` file is not written.
@@ -528,7 +537,7 @@ pub struct Settings {
pub secrets_master_key_source: KeySource, // Keychain | Env | None
// Step 3: Inference
pub llm_backend: Option<String>, // "nearai" | "anthropic" | "openai" | "ollama" | "openai_compatible" | "bedrock"
pub llm_backend: Option<String>, // "nearai" | "anthropic" | "openai" | "github_copilot" | "ollama" | "openai_compatible" | "bedrock"
pub ollama_base_url: Option<String>,
pub openai_compatible_base_url: Option<String>,
@@ -586,7 +595,7 @@ in the database `secrets` table. The wizard writes secrets like:
```
telegram_bot_token → encrypted bot token
telegram_webhook_secret → encrypted webhook HMAC secret
anthropic_api_key → encrypted API key
llm_anthropic_api_key → encrypted API key
```
---
+211 -72
View File
@@ -3,7 +3,7 @@
//! The wizard guides users through:
//! 1. Database connection
//! 2. Security (secrets master key)
//! 3. Inference provider (NEAR AI, Anthropic, OpenAI, Ollama, OpenAI-compatible)
//! 3. Inference provider (NEAR AI, Anthropic, OpenAI, GitHub Copilot, OpenAI Codex, Ollama, OpenAI-compatible)
//! 4. Model selection
//! 5. Embeddings
//! 6. Channel configuration
@@ -1083,8 +1083,10 @@ impl SetupWizard {
print_info(&format!("Current provider: {}", display));
println!();
let is_known =
current == "nearai" || current == "bedrock" || registry.is_known(&current);
let is_known = current == "nearai"
|| current == "bedrock"
|| current == "openai_codex"
|| registry.is_known(&current);
if is_known && confirm("Keep current provider?", true).map_err(SetupError::Io)? {
if current == "bedrock" {
@@ -1093,6 +1095,10 @@ impl SetupWizard {
print_info("Keeping existing AWS Bedrock configuration.");
return Ok(());
}
if current == "openai_codex" {
print_info("Keeping existing OpenAI Codex configuration.");
return Ok(());
}
return self.run_provider_setup(&current, &registry).await;
}
@@ -1107,7 +1113,7 @@ impl SetupWizard {
print_info("Select your inference provider:");
println!();
// Build menu: NearAI first, then all registry providers with setup hints, then Bedrock
// Build menu: NearAI first, then OpenAI Codex, then registry providers, then Bedrock
let selectable = registry.selectable();
let mut options: Vec<String> = Vec::with_capacity(2 + selectable.len());
let mut provider_ids: Vec<String> = Vec::with_capacity(2 + selectable.len());
@@ -1115,6 +1121,9 @@ impl SetupWizard {
options.push("NEAR AI - multi-model access via NEAR account".to_string());
provider_ids.push("nearai".to_string());
options.push("OpenAI Codex - ChatGPT subscription (Plus/Pro/Max)".to_string());
provider_ids.push("openai_codex".to_string());
for def in &selectable {
let label = format!(
"{:<17}- {}",
@@ -1158,6 +1167,10 @@ impl SetupWizard {
return self.setup_nearai().await;
}
if provider_id == "openai_codex" {
return self.setup_openai_codex().await;
}
let def = registry
.find(provider_id)
.ok_or_else(|| SetupError::Config(format!("Unknown provider: {}", provider_id)))?;
@@ -1178,6 +1191,10 @@ impl SetupWizard {
return self.setup_anthropic().await;
}
if provider_id == "github_copilot" {
return self.setup_github_copilot().await;
}
match setup {
crate::llm::registry::SetupHint::ApiKey {
secret_name,
@@ -1340,6 +1357,100 @@ impl SetupWizard {
}
}
async fn setup_github_copilot(&mut self) -> Result<(), SetupError> {
print_info("GitHub Copilot authentication:");
let options = &[
"GitHub device login (recommended)",
"Paste an existing token (from IDE or personal access token)",
];
let choice = select_one("Auth method:", options).map_err(SetupError::Io)?;
match choice {
0 => self.setup_github_copilot_device_login().await,
_ => self.setup_github_copilot_paste_token().await,
}
}
async fn setup_github_copilot_paste_token(&mut self) -> Result<(), SetupError> {
self.set_llm_backend_preserving_model("github_copilot");
print_info("Paste your GitHub token (requires an active Copilot subscription).");
print_info("Sources: `gh auth token`, or the oauth_token field in");
print_info("~/.config/github-copilot/apps.json (VS Code) or ~/.config/gh/hosts.yml.");
let token_secret = secret_input("GitHub Copilot token").map_err(SetupError::Io)?;
let token = token_secret.expose_secret().trim().to_string();
if token.is_empty() {
return Err(SetupError::Auth("No token provided".to_string()));
}
let client = reqwest::Client::builder()
.timeout(std::time::Duration::from_secs(15))
.build()
.map_err(|e| SetupError::Auth(format!("Failed to create HTTP client: {e}")))?;
self.save_github_copilot_token(&client, &token).await
}
async fn setup_github_copilot_device_login(&mut self) -> Result<(), SetupError> {
self.set_llm_backend_preserving_model("github_copilot");
let client = reqwest::Client::builder()
.timeout(std::time::Duration::from_secs(15))
.build()
.map_err(|e| SetupError::Auth(format!("Failed to create HTTP client: {e}")))?;
let device = crate::llm::github_copilot_auth::request_device_code(&client)
.await
.map_err(|e| SetupError::Auth(e.to_string()))?;
print_info("Authorize IronClaw with GitHub Copilot in your browser.");
print_info(&format!("Verification URL: {}", device.verification_uri));
print_info(&format!("One-time code: {}", device.user_code));
if let Err(e) = open::that(&device.verification_uri) {
tracing::debug!(
url = %device.verification_uri,
error = %e,
"Failed to open GitHub Copilot device login URL"
);
print_info("Open the URL above manually if your browser did not launch.");
} else {
print_info("Opened your browser to GitHub device login.");
}
print_info("Waiting for GitHub authorization...");
let token = crate::llm::github_copilot_auth::wait_for_device_login(&client, &device)
.await
.map_err(|e| SetupError::Auth(e.to_string()))?;
self.save_github_copilot_token(&client, &token).await
}
async fn save_github_copilot_token(
&mut self,
client: &reqwest::Client,
token: &str,
) -> Result<(), SetupError> {
crate::llm::github_copilot_auth::validate_token(client, token)
.await
.map_err(|e| SetupError::Auth(e.to_string()))?;
if let Ok(ctx) = self.init_secrets_context().await {
let key = SecretString::from(token.to_string());
ctx.save_secret("llm_github_copilot_token", &key)
.await
.map_err(|e| SetupError::Config(format!("Failed to save GitHub token: {e}")))?;
print_success("GitHub Copilot token encrypted and saved");
} else {
print_info("Secrets not available. Set GITHUB_COPILOT_TOKEN in your environment.");
}
crate::config::inject_single_var("GITHUB_COPILOT_TOKEN", token);
self.llm_api_key = Some(SecretString::from(token.to_string()));
print_success("GitHub Copilot configured");
Ok(())
}
/// Anthropic OAuth setup: extract token from `claude login` credentials.
async fn setup_anthropic_oauth(&mut self) -> Result<(), SetupError> {
self.set_llm_backend_preserving_model("anthropic");
@@ -1490,6 +1601,29 @@ impl SetupWizard {
Ok(())
}
/// OpenAI Codex (ChatGPT subscription) setup: device code OAuth flow.
async fn setup_openai_codex(&mut self) -> Result<(), SetupError> {
self.settings.llm_backend = Some("openai_codex".to_string());
if self.settings.selected_model.is_some() {
self.settings.selected_model = None;
}
use crate::config::OpenAiCodexConfig;
use crate::llm::OpenAiCodexSessionManager;
let config = OpenAiCodexConfig::default();
let mgr = OpenAiCodexSessionManager::new(config).map_err(|e| {
SetupError::Config(format!("OpenAI Codex session manager init failed: {}", e))
})?;
mgr.device_code_login().await.map_err(|e| {
SetupError::Config(format!("OpenAI Codex authentication failed: {}", e))
})?;
print_success("OpenAI Codex configured (ChatGPT subscription)");
Ok(())
}
/// Generic Ollama-style setup: just needs a base URL, no API key.
fn setup_ollama_generic(
&mut self,
@@ -2618,16 +2752,17 @@ impl SetupWizard {
/// Write bootstrap environment variables to `~/.ironclaw/.env`.
///
/// These are the chicken-and-egg settings needed before the database is
/// connected (DATABASE_BACKEND, DATABASE_URL, LLM_BACKEND, etc.).
/// Only true chicken-and-egg settings are written here — things needed
/// before the database is connected: `DATABASE_BACKEND`, `DATABASE_URL`,
/// `LIBSQL_PATH`, `SECRETS_MASTER_KEY`, `ONBOARD_COMPLETED`, and
/// channel config vars (Signal, Claude Code sandbox).
///
/// **Credentials are NOT written here.** API keys and OAuth tokens live
/// only in the encrypted secrets DB. `LlmConfig::resolve()` defers
/// gracefully when credentials are missing during early startup, and the
/// re-resolution in `AppBuilder::build_all()` fills them in after
/// `inject_llm_keys_from_secrets()` loads from encrypted storage.
/// **LLM settings and credentials are NOT written here.** `LLM_BACKEND`,
/// base URLs, and model names are persisted to the DB via
/// `persist_settings()` and loaded by `Config::from_db_with_toml()`.
/// API keys live only in the encrypted secrets DB and are injected via
/// `inject_llm_keys_from_secrets()` after DB init.
fn write_bootstrap_env(&self) -> Result<(), SetupError> {
let registry = crate::llm::ProviderRegistry::load();
let mut env_vars: Vec<(String, String)> = Vec::new();
if let Some(ref backend) = self.settings.database_backend {
@@ -2643,66 +2778,6 @@ impl SetupWizard {
env_vars.push(("LIBSQL_URL".to_string(), url.clone()));
}
// LLM bootstrap vars: same chicken-and-egg problem as DATABASE_BACKEND.
// Config::from_env() needs the backend before the DB is connected.
if let Some(ref backend) = self.settings.llm_backend {
env_vars.push(("LLM_BACKEND".to_string(), backend.clone()));
}
if let Some(ref url) = self.settings.openai_compatible_base_url {
env_vars.push(("LLM_BASE_URL".to_string(), url.clone()));
}
if let Some(ref url) = self.settings.ollama_base_url {
env_vars.push(("OLLAMA_BASE_URL".to_string(), url.clone()));
}
if let Some(ref region) = self.settings.bedrock_region {
env_vars.push(("BEDROCK_REGION".to_string(), region.clone()));
}
if self.settings.llm_backend.as_deref() == Some("bedrock") {
if let Some(ref model) = self.settings.selected_model {
env_vars.push(("BEDROCK_MODEL".to_string(), model.clone()));
}
if let Some(ref cross) = self.settings.bedrock_cross_region {
env_vars.push(("BEDROCK_CROSS_REGION".to_string(), cross.clone()));
}
if let Some(ref profile) = self.settings.bedrock_profile {
env_vars.push(("AWS_PROFILE".to_string(), profile.clone()));
}
}
// Model name: same chicken-and-egg — Config::from_env() resolves the
// model before the DB is connected, so we must persist it to .env.
// Write the backend-specific env var so the correct resolution path
// picks it up (looked up from the provider registry).
// Bedrock model is already written above as BEDROCK_MODEL, skip here.
if self.settings.llm_backend.as_deref() != Some("bedrock")
&& let Some(ref model) = self.settings.selected_model
{
let backend_str = self.settings.llm_backend.as_deref().unwrap_or("nearai");
let model_env = registry.model_env_var(backend_str);
env_vars.push((model_env.to_string(), model.clone()));
}
// Also write provider-specific base URL env var if the provider
// defines one (e.g., GROQ doesn't need LLM_BASE_URL since its
// default is compiled in, but it doesn't hurt to be explicit).
if let Some(ref backend) = self.settings.llm_backend
&& let Some(def) = registry.find(backend)
&& let Some(ref base_url_env) = def.base_url_env
&& let Some(ref base_url) = def.default_base_url
&& base_url_env != "LLM_BASE_URL"
&& base_url_env != "OLLAMA_BASE_URL"
{
env_vars.push((base_url_env.clone(), base_url.clone()));
}
// Preserve NEARAI_API_KEY if present (set by API key auth flow
// via the thread-safe runtime env overlay).
if let Some(api_key) = crate::config::helpers::env_or_override("NEARAI_API_KEY")
&& !api_key.is_empty()
{
env_vars.push(("NEARAI_API_KEY".to_string(), api_key));
}
// Secrets master key (env var mode): write to .env so it's available
// on next startup before the DB is connected.
if let Some(ref key_hex) = self.settings.secrets_master_key_hex {
@@ -2963,6 +3038,7 @@ impl SetupWizard {
"ollama" => "Ollama",
"openai_compatible" => "OpenAI-compatible",
"bedrock" => "AWS Bedrock",
"openai_codex" => "OpenAI Codex",
other => other,
};
println!(" Provider: {}", display);
@@ -3530,6 +3606,36 @@ mod tests {
);
}
#[test]
fn test_github_copilot_setup_preserves_model_for_same_backend() {
let mut wizard = SetupWizard::new();
wizard.settings.llm_backend = Some("github_copilot".to_string());
wizard.settings.selected_model = Some("gpt-4o".to_string());
wizard.set_llm_backend_preserving_model("github_copilot");
assert_eq!(wizard.settings.selected_model.as_deref(), Some("gpt-4o"));
assert_eq!(
wizard.settings.llm_backend.as_deref(),
Some("github_copilot")
);
}
#[test]
fn test_github_copilot_setup_clears_stale_model_on_switch() {
let mut wizard = SetupWizard::new();
wizard.settings.llm_backend = Some("openai".to_string());
wizard.settings.selected_model = Some("gpt-5".to_string());
wizard.set_llm_backend_preserving_model("github_copilot");
assert!(wizard.settings.selected_model.is_none());
assert_eq!(
wizard.settings.llm_backend.as_deref(),
Some("github_copilot")
);
}
#[test]
fn test_is_openai_chat_model_includes_gpt5_and_filters_non_chat_variants() {
assert!(is_openai_chat_model("gpt-5"));
@@ -3887,6 +3993,39 @@ mod tests {
);
}
/// Regression: API key set via inject_single_var (the path used by
/// setup_api_key_provider during onboarding) must be picked up by
/// for_model_discovery() so model listing uses cloud-api auth
/// instead of falling back to session-token auth.
#[test]
fn test_model_discovery_picks_up_injected_var() {
use secrecy::ExposeSecret;
let _lock = ENV_MUTEX.lock().unwrap();
let _guard = EnvGuard::clear("NEARAI_API_KEY");
let _guard2 = EnvGuard::clear("NEARAI_BASE_URL");
crate::config::inject_single_var("NEARAI_API_KEY", "injected-wizard-key");
let config = build_nearai_model_fetch_config();
// Clean up: empty values are treated as unset by env_or_override()
// at every layer (real env, runtime overrides, INJECTED_VARS).
crate::config::inject_single_var("NEARAI_API_KEY", "");
assert!(
config.nearai.api_key.is_some(),
"for_model_discovery must read NEARAI_API_KEY from inject_single_var overlay"
);
assert_eq!(
config.nearai.api_key.as_ref().unwrap().expose_secret(),
"injected-wizard-key"
);
assert_eq!(
config.nearai.base_url, "https://cloud-api.near.ai",
"API key from overlay must select cloud-api base URL"
);
}
/// Regression: API key set via set_runtime_env (interactive api_key_login
/// path) must be picked up by build_nearai_model_fetch_config so that
/// model listing doesn't fall back to session-token auth and re-trigger
+210
View File
@@ -0,0 +1,210 @@
use std::collections::HashSet;
use std::sync::Arc;
use crate::extensions::ExtensionManager;
use super::ToolRegistry;
pub const AUTONOMOUS_TOOL_DENYLIST: &[&str] = &[
"routine_create",
"routine_update",
"routine_delete",
"routine_fire",
"event_emit",
"create_job",
"job_prompt",
"restart",
"tool_install",
"tool_auth",
"tool_activate",
"tool_remove",
"tool_upgrade",
"skill_install",
"skill_remove",
"secret_list",
"secret_delete",
];
pub fn is_autonomous_tool_denylisted(tool_name: &str) -> bool {
AUTONOMOUS_TOOL_DENYLIST.contains(&tool_name)
}
pub fn autonomous_unavailable_message(tool_name: &str, owner_id: &str) -> String {
if is_autonomous_tool_denylisted(tool_name) {
format!("Tool '{tool_name}' is not available in autonomous jobs or routines")
} else {
format!("Tool '{tool_name}' is not currently available for owner '{owner_id}'")
}
}
pub fn autonomous_unavailable_error(tool_name: &str, owner_id: &str) -> crate::error::ToolError {
crate::error::ToolError::AutonomousUnavailable {
name: tool_name.to_string(),
reason: autonomous_unavailable_message(tool_name, owner_id),
}
}
pub async fn autonomous_allowed_tool_names(
tools: &Arc<ToolRegistry>,
extension_manager: Option<&Arc<ExtensionManager>>,
owner_id: &str,
) -> HashSet<String> {
let mut allowed = tools.builtin_tool_names().await;
allowed.retain(|name| !is_autonomous_tool_denylisted(name));
if let Some(extension_manager) = extension_manager
&& extension_manager.owner_id() == owner_id
{
allowed.extend(
extension_manager
.active_tool_names()
.await
.into_iter()
.filter(|name| !is_autonomous_tool_denylisted(name)),
);
}
allowed
}
#[cfg(test)]
mod tests {
use std::path::Path;
use std::time::Duration;
use async_trait::async_trait;
use secrecy::SecretString;
use super::*;
use crate::context::JobContext;
use crate::extensions::ExtensionManager;
use crate::hooks::HookRegistry;
use crate::secrets::{InMemorySecretsStore, SecretsCrypto, SecretsStore};
use crate::tools::mcp::{McpProcessManager, McpSessionManager};
use crate::tools::{Tool, ToolError, ToolOutput};
struct FakeTool {
name: &'static str,
}
#[async_trait]
impl Tool for FakeTool {
fn name(&self) -> &str {
self.name
}
fn description(&self) -> &str {
"test tool"
}
fn parameters_schema(&self) -> serde_json::Value {
serde_json::json!({
"type": "object",
"properties": {},
})
}
async fn execute(
&self,
_params: serde_json::Value,
_ctx: &JobContext,
) -> Result<ToolOutput, ToolError> {
Ok(ToolOutput::text("ok", Duration::from_millis(1)))
}
}
async fn write_test_extension_wasm(tools_dir: &Path, name: &str) {
tokio::fs::create_dir_all(tools_dir)
.await
.expect("create test tools dir");
tokio::fs::write(tools_dir.join(format!("{name}.wasm")), b"\0asm")
.await
.expect("write wasm marker");
}
fn make_extension_manager(
tools: Arc<ToolRegistry>,
tools_dir: &Path,
owner_id: &str,
) -> Arc<ExtensionManager> {
let crypto = Arc::new(
SecretsCrypto::new(SecretString::from(
"0123456789abcdef0123456789abcdef0123456789abcdef0123456789abcdef",
))
.expect("test crypto"),
);
let secrets: Arc<dyn SecretsStore + Send + Sync> =
Arc::new(InMemorySecretsStore::new(crypto));
Arc::new(ExtensionManager::new(
Arc::new(McpSessionManager::new()),
Arc::new(McpProcessManager::new()),
secrets,
tools,
Some(Arc::new(HookRegistry::default())),
None,
tools_dir.to_path_buf(),
tools_dir.join("channels"),
None,
owner_id.to_string(),
None,
Vec::new(),
))
}
#[tokio::test]
async fn autonomous_scope_keeps_allowed_builtins_and_blocks_denylisted_builtins() {
let tools = Arc::new(ToolRegistry::new());
tools.register_sync(Arc::new(FakeTool { name: "echo" }));
tools.register_sync(Arc::new(FakeTool { name: "restart" }));
let allowed = autonomous_allowed_tool_names(&tools, None, "default").await;
assert!(allowed.contains("echo"));
assert!(!allowed.contains("restart"));
}
#[tokio::test]
async fn autonomous_scope_includes_active_extension_tools_for_matching_owner() {
let temp_dir = tempfile::tempdir().expect("tempdir");
let tools_dir = temp_dir.path().join("wasm-tools");
let tools = Arc::new(ToolRegistry::new());
tools
.register(Arc::new(FakeTool { name: "owner_gate" }))
.await;
write_test_extension_wasm(&tools_dir, "owner_gate").await;
let manager = make_extension_manager(tools.clone(), &tools_dir, "default");
let allowed = autonomous_allowed_tool_names(&tools, Some(&manager), "default").await;
assert!(allowed.contains("owner_gate"));
}
#[tokio::test]
async fn autonomous_scope_excludes_inactive_extension_tools() {
let temp_dir = tempfile::tempdir().expect("tempdir");
let tools_dir = temp_dir.path().join("wasm-tools");
let tools = Arc::new(ToolRegistry::new());
let manager = make_extension_manager(tools.clone(), &tools_dir, "default");
let allowed = autonomous_allowed_tool_names(&tools, Some(&manager), "default").await;
assert!(!allowed.contains("owner_gate"));
}
#[tokio::test]
async fn autonomous_scope_excludes_active_extension_tools_for_other_owner() {
let temp_dir = tempfile::tempdir().expect("tempdir");
let tools_dir = temp_dir.path().join("wasm-tools");
let tools = Arc::new(ToolRegistry::new());
tools
.register(Arc::new(FakeTool { name: "owner_gate" }))
.await;
write_test_extension_wasm(&tools_dir, "owner_gate").await;
let manager = make_extension_manager(tools.clone(), &tools_dir, "someone-else");
let allowed = autonomous_allowed_tool_names(&tools, Some(&manager), "default").await;
assert!(!allowed.contains("owner_gate"));
}
}
+77 -48
View File
@@ -194,6 +194,15 @@ impl Tool for MemoryWriteTool {
"type": "boolean",
"description": "If true, append to existing content. If false, replace entirely.",
"default": true
},
"layer": {
"type": "string",
"description": "Memory layer to write to (e.g. 'private', 'household', 'finance'). When omitted, writes to the workspace's default scope."
},
"force": {
"type": "boolean",
"description": "Skip privacy classification and write directly to the specified layer without redirect. Use when you're certain the content belongs in the target layer.",
"default": false
}
},
"required": ["content"]
@@ -256,67 +265,86 @@ impl Tool for MemoryWriteTool {
.and_then(|v| v.as_bool())
.unwrap_or(true);
// Prompt injection scanning for system-prompt files is handled by
// Workspace::write() / Workspace::append() — no need to duplicate here.
let layer = params.get("layer").and_then(|v| v.as_str());
let force = params
.get("force")
.and_then(|v| v.as_bool())
.unwrap_or(false);
let path = match target {
"memory" => {
if append {
self.workspace
.append_memory(content)
.await
.map_err(map_write_err)?;
} else {
self.workspace
.write(paths::MEMORY, content)
.await
.map_err(map_write_err)?;
}
paths::MEMORY.to_string()
}
// Resolve the target to a workspace path
let resolved_path = match target {
"memory" => paths::MEMORY.to_string(),
"daily_log" => {
let tz = crate::timezone::parse_timezone(&ctx.user_timezone)
.unwrap_or(chrono_tz::Tz::UTC);
let now = chrono::Utc::now().with_timezone(&tz);
format!("daily/{}.md", now.format("%Y-%m-%d"))
}
"heartbeat" => paths::HEARTBEAT.to_string(),
path => path.to_string(),
};
// When a layer is specified, route through layer-aware methods for ALL targets.
// Otherwise, use default workspace methods (which include injection scanning).
let layer_result = if let Some(layer_name) = layer {
let result = if append {
self.workspace
.append_daily_log_tz(content, tz)
.append_to_layer(layer_name, &resolved_path, content, force)
.await
.map_err(map_write_err)?
}
"heartbeat" => {
if append {
} else {
self.workspace
.write_to_layer(layer_name, &resolved_path, content, force)
.await
.map_err(map_write_err)?
};
Some((result.actual_layer, result.redirected))
} else {
// No layer specified — use default workspace methods.
// Prompt injection scanning for system-prompt files is handled by
// Workspace::write() / Workspace::append().
match target {
"memory" => {
if append {
self.workspace
.append_memory(content)
.await
.map_err(map_write_err)?;
} else {
self.workspace
.write(paths::MEMORY, content)
.await
.map_err(map_write_err)?;
}
}
"daily_log" => {
let tz = crate::timezone::parse_timezone(&ctx.user_timezone)
.unwrap_or(chrono_tz::Tz::UTC);
self.workspace
.append(paths::HEARTBEAT, content)
.await
.map_err(map_write_err)?;
} else {
self.workspace
.write(paths::HEARTBEAT, content)
.append_daily_log_tz(content, tz)
.await
.map_err(map_write_err)?;
}
paths::HEARTBEAT.to_string()
}
path => {
if append {
self.workspace
.append(path, content)
.await
.map_err(map_write_err)?;
} else {
self.workspace
.write(path, content)
.await
.map_err(map_write_err)?;
_ => {
if append {
self.workspace
.append(&resolved_path, content)
.await
.map_err(map_write_err)?;
} else {
self.workspace
.write(&resolved_path, content)
.await
.map_err(map_write_err)?;
}
}
path.to_string()
}
None
};
// Sync derived identity documents when the profile is written.
// Normalize the path to match Workspace::normalize_path(): trim, strip
// leading/trailing slashes, collapse all consecutive slashes.
let normalized_path = {
let trimmed = path.trim().trim_matches('/');
let trimmed = resolved_path.trim().trim_matches('/');
let mut result = String::new();
let mut last_was_slash = false;
for c in trimmed.chars() {
@@ -339,9 +367,6 @@ impl Tool for MemoryWriteTool {
tracing::info!("profile write: synced USER.md + assistant-directives.md");
synced_docs.extend_from_slice(&[paths::USER, paths::ASSISTANT_DIRECTIVES]);
// Persist the onboarding-completed flag and set the
// in-memory safety net so BOOTSTRAP.md injection stops
// even if the LLM forgets to delete it.
self.workspace.mark_bootstrap_completed();
let toml_path = crate::settings::Settings::default_toml_path();
if let Ok(Some(mut settings)) = crate::settings::Settings::load_toml(&toml_path)
@@ -364,10 +389,14 @@ impl Tool for MemoryWriteTool {
let mut output = serde_json::json!({
"status": "written",
"path": path,
"path": resolved_path,
"append": append,
"content_length": content.len(),
});
if let Some((actual_layer, redirected)) = layer_result {
output["layer"] = serde_json::Value::String(actual_layer);
output["redirected"] = serde_json::Value::Bool(redirected);
}
if !synced_docs.is_empty() {
output["synced"] = serde_json::json!(synced_docs);
}
+88 -355
View File
@@ -19,9 +19,8 @@ use serde_json::{Map, Value};
use uuid::Uuid;
use crate::agent::routine::{
FullJobPermissionDefaultMode, FullJobPermissionMode, NotifyConfig, Routine, RoutineAction,
RoutineGuardrails, Trigger, load_full_job_permission_settings, next_cron_fire,
normalize_cron_expression, normalize_tool_names,
NotifyConfig, Routine, RoutineAction, RoutineGuardrails, Trigger, next_cron_fire,
normalize_cron_expression,
};
use crate::agent::routine_engine::RoutineEngine;
use crate::context::JobContext;
@@ -48,6 +47,10 @@ enum NormalizedTriggerRequest {
event_type: String,
filters: HashMap<String, String>,
},
Webhook {
path: Option<String>,
secret: Option<String>,
},
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
@@ -56,21 +59,12 @@ enum NormalizedExecutionMode {
FullJob,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
enum RequestedFullJobPermissionMode {
Explicit,
InheritOwner,
CopyOwner,
}
#[derive(Debug, Clone, PartialEq, Eq)]
struct NormalizedExecutionRequest {
mode: NormalizedExecutionMode,
context_paths: Vec<String>,
use_tools: bool,
max_tool_rounds: u32,
tool_permissions: Vec<String>,
permission_mode: Option<RequestedFullJobPermissionMode>,
}
#[derive(Debug, Clone, PartialEq, Eq)]
@@ -154,16 +148,6 @@ fn execution_properties() -> Value {
"maximum": crate::agent::routine::MAX_TOOL_ROUNDS_LIMIT,
"default": 3,
"description": "Only applies when execution.mode='lightweight' and use_tools=true. Runtime-capped to prevent loops."
},
"tool_permissions": {
"type": "array",
"items": { "type": "string" },
"description": "Only applies when execution.mode='full_job'. These tools are pre-authorized for Always-approval checks."
},
"permission_mode": {
"type": "string",
"enum": ["inherit_owner", "explicit", "copy_owner"],
"description": "Only applies when execution.mode='full_job'. 'inherit_owner' uses the owner defaults at run time, 'explicit' uses only tool_permissions, and 'copy_owner' snapshots the current owner allowlist into tool_permissions."
}
})
}
@@ -336,22 +320,12 @@ fn lightweight_execution_variant() -> Value {
fn full_job_execution_variant() -> Value {
serde_json::json!({
"type": "object",
"description": "Full-job execution. Uses owner-scoped permission defaults plus tool_permissions and ignores lightweight-only fields such as use_tools, max_tool_rounds, and context_paths.",
"description": "Full-job execution. Uses the owner's live autonomous tool scope and ignores lightweight-only fields such as use_tools, max_tool_rounds, and context_paths.",
"properties": {
"mode": {
"type": "string",
"enum": ["full_job"],
"description": "Full-job execution mode."
},
"tool_permissions": {
"type": "array",
"items": { "type": "string" },
"description": "Tools pre-authorized for Always-approval checks."
},
"permission_mode": {
"type": "string",
"enum": ["inherit_owner", "explicit", "copy_owner"],
"description": "When omitted, new routines use the owner default. 'copy_owner' snapshots the current owner allowlist into this routine."
}
},
"required": ["mode"]
@@ -369,7 +343,7 @@ fn execution_discovery_schema() -> Value {
],
"examples": [
{ "mode": "lightweight", "use_tools": true, "max_tool_rounds": 3 },
{ "mode": "full_job", "permission_mode": "inherit_owner", "tool_permissions": ["message", "http"] }
{ "mode": "full_job" }
]
})
}
@@ -418,9 +392,7 @@ fn routine_create_examples() -> Vec<Value> {
"filters": { "repository": "nearai/ironclaw" }
},
"execution": {
"mode": "full_job",
"permission_mode": "inherit_owner",
"tool_permissions": ["message"]
"mode": "full_job"
}
}),
]
@@ -433,7 +405,7 @@ fn routine_create_tool_summary() -> ToolDiscoverySummary {
"request.kind='cron' requires request.schedule.".into(),
"request.kind='message_event' requires request.pattern.".into(),
"request.kind='system_event' requires request.source and request.event_type.".into(),
"execution.mode='full_job' uses permission_mode and tool_permissions, and ignores use_tools, max_tool_rounds, and context_paths.".into(),
"execution.mode='full_job' uses the owner's live autonomous tool scope and ignores use_tools, max_tool_rounds, and context_paths.".into(),
],
notes: vec![
"Omitting execution defaults to lightweight mode.".into(),
@@ -590,22 +562,6 @@ fn routine_create_schema(include_compatibility_aliases: bool) -> Value {
"description": "Compatibility alias for execution.max_tool_rounds."
}),
);
properties.insert(
"tool_permissions".to_string(),
serde_json::json!({
"type": "array",
"items": { "type": "string" },
"description": "Compatibility alias for execution.tool_permissions."
}),
);
properties.insert(
"permission_mode".to_string(),
serde_json::json!({
"type": "string",
"enum": ["inherit_owner", "explicit", "copy_owner"],
"description": "Compatibility alias for execution.permission_mode."
}),
);
properties.insert(
"notify_channel".to_string(),
serde_json::json!({
@@ -684,16 +640,6 @@ pub(crate) fn routine_update_parameters_schema() -> Value {
"description": {
"type": "string",
"description": "New description"
},
"tool_permissions": {
"type": "array",
"items": { "type": "string" },
"description": "Updated Always-approval tool allowlist for full_job routines only."
},
"permission_mode": {
"type": "string",
"enum": ["inherit_owner", "explicit", "copy_owner"],
"description": "Updated permission mode for full_job routines only. 'copy_owner' snapshots the current owner allowlist into the routine and persists as explicit."
}
},
"required": ["name"]
@@ -739,27 +685,6 @@ fn u64_field(params: &Value, group: &str, field: &str, aliases: &[&str]) -> Opti
}
fn string_array_field(params: &Value, group: &str, field: &str, aliases: &[&str]) -> Vec<String> {
normalize_tool_names(
nested_object(params, group)
.and_then(|obj| obj.get(field))
.and_then(Value::as_array)
.or_else(|| {
aliases
.iter()
.find_map(|alias| params.get(*alias).and_then(Value::as_array))
})
.into_iter()
.flatten()
.filter_map(|value| value.as_str().map(String::from)),
)
}
fn optional_string_array_field(
params: &Value,
group: &str,
field: &str,
aliases: &[&str],
) -> Option<Vec<String>> {
nested_object(params, group)
.and_then(|obj| obj.get(field))
.and_then(Value::as_array)
@@ -769,11 +694,21 @@ fn optional_string_array_field(
.find_map(|alias| params.get(*alias).and_then(Value::as_array))
})
.map(|arr| {
normalize_tool_names(
arr.iter()
.filter_map(|value| value.as_str().map(String::from)),
)
let mut seen = std::collections::HashSet::new();
arr.iter()
.filter_map(Value::as_str)
.map(str::trim)
.filter(|value| !value.is_empty())
.filter_map(|value| {
if seen.insert(value.to_string()) {
Some(value.to_string())
} else {
None
}
})
.collect()
})
.unwrap_or_default()
}
fn object_field(
@@ -896,6 +831,11 @@ fn parse_routine_trigger(params: &Value) -> Result<NormalizedTriggerRequest, Too
filters,
})
}
"webhook" => {
let path = string_field(params, "request", "path", &["webhook_path"]);
let secret = string_field(params, "request", "secret", &["webhook_secret"]);
Ok(NormalizedTriggerRequest::Webhook { path, secret })
}
other => Err(ToolError::InvalidParameters(format!(
"unknown request.kind: {other}"
))),
@@ -912,20 +852,6 @@ fn parse_execution_mode(value: Option<String>) -> Result<NormalizedExecutionMode
}
}
fn parse_requested_full_job_permission_mode(
value: Option<String>,
) -> Result<Option<RequestedFullJobPermissionMode>, ToolError> {
match value.as_deref() {
None => Ok(None),
Some("explicit") => Ok(Some(RequestedFullJobPermissionMode::Explicit)),
Some("inherit_owner") => Ok(Some(RequestedFullJobPermissionMode::InheritOwner)),
Some("copy_owner") => Ok(Some(RequestedFullJobPermissionMode::CopyOwner)),
Some(other) => Err(ToolError::InvalidParameters(format!(
"unknown full_job permission_mode: {other}"
))),
}
}
fn parse_routine_execution(params: &Value) -> Result<NormalizedExecutionRequest, ToolError> {
let mode = parse_execution_mode(string_field(params, "execution", "mode", &["action_type"]))?;
let context_paths =
@@ -935,26 +861,12 @@ fn parse_routine_execution(params: &Value) -> Result<NormalizedExecutionRequest,
.unwrap_or(3)
.clamp(1, crate::agent::routine::MAX_TOOL_ROUNDS_LIMIT as u64)
as u32;
let tool_permissions = string_array_field(
params,
"execution",
"tool_permissions",
&["tool_permissions"],
);
let permission_mode = parse_requested_full_job_permission_mode(string_field(
params,
"execution",
"permission_mode",
&["permission_mode"],
))?;
Ok(NormalizedExecutionRequest {
mode,
context_paths,
use_tools,
max_tool_rounds,
tool_permissions,
permission_mode,
})
}
@@ -1012,92 +924,31 @@ fn build_routine_trigger(trigger: &NormalizedTriggerRequest) -> Trigger {
event_type: event_type.clone(),
filters: filters.clone(),
},
NormalizedTriggerRequest::Webhook { path, secret } => Trigger::Webhook {
path: path.clone(),
secret: secret.clone(),
},
}
}
async fn build_routine_action(
store: &dyn Database,
user_id: &str,
fn build_routine_action(
name: &str,
prompt: &str,
execution: &NormalizedExecutionRequest,
) -> Result<RoutineAction, ToolError> {
) -> RoutineAction {
match execution.mode {
NormalizedExecutionMode::Lightweight => Ok(RoutineAction::Lightweight {
NormalizedExecutionMode::Lightweight => RoutineAction::Lightweight {
prompt: prompt.to_string(),
context_paths: execution.context_paths.clone(),
max_tokens: 4096,
use_tools: execution.use_tools,
max_tool_rounds: execution.max_tool_rounds,
}),
NormalizedExecutionMode::FullJob => {
let mut owner_settings = None;
let requested_mode = match execution.permission_mode {
Some(mode) => mode,
None => {
let settings = load_full_job_permission_settings(store, user_id)
.await
.map_err(|e| {
ToolError::ExecutionFailed(format!(
"failed to load routine permission settings: {e}"
))
})?;
let mode = match settings.default_mode {
FullJobPermissionDefaultMode::Explicit => {
RequestedFullJobPermissionMode::Explicit
}
FullJobPermissionDefaultMode::InheritOwner => {
RequestedFullJobPermissionMode::InheritOwner
}
FullJobPermissionDefaultMode::CopyOwner => {
RequestedFullJobPermissionMode::CopyOwner
}
};
owner_settings = Some(settings);
mode
}
};
let (permission_mode, tool_permissions) = match requested_mode {
RequestedFullJobPermissionMode::Explicit => (
FullJobPermissionMode::Explicit,
execution.tool_permissions.clone(),
),
RequestedFullJobPermissionMode::InheritOwner => (
FullJobPermissionMode::InheritOwner,
execution.tool_permissions.clone(),
),
RequestedFullJobPermissionMode::CopyOwner => {
let owner_allowed_tools = match owner_settings {
Some(settings) => settings.owner_allowed_tools,
None => {
load_full_job_permission_settings(store, user_id)
.await
.map_err(|e| {
ToolError::ExecutionFailed(format!(
"failed to load routine permission settings: {e}"
))
})?
.owner_allowed_tools
}
};
(
FullJobPermissionMode::Explicit,
normalize_tool_names(
owner_allowed_tools
.into_iter()
.chain(execution.tool_permissions.iter().cloned()),
),
)
}
};
Ok(RoutineAction::FullJob {
title: name.to_string(),
description: prompt.to_string(),
max_iterations: 10,
tool_permissions,
permission_mode,
})
}
},
NormalizedExecutionMode::FullJob => RoutineAction::FullJob {
title: name.to_string(),
description: prompt.to_string(),
max_iterations: 10,
},
}
}
@@ -1108,13 +959,6 @@ fn routine_requests_full_job(params: &Value) -> bool {
)
}
fn routine_permission_fields_present(params: &Value) -> bool {
nested_object(params, "execution").is_some_and(|execution| {
execution.contains_key("tool_permissions") || execution.contains_key("permission_mode")
}) || params.get("tool_permissions").is_some()
|| params.get("permission_mode").is_some()
}
fn event_emit_schema(include_source_alias: bool) -> Value {
let mut schema = serde_json::json!({
"type": "object",
@@ -1241,14 +1085,8 @@ impl Tool for RoutineCreateTool {
let start = std::time::Instant::now();
let normalized = parse_routine_create_request(&params)?;
let trigger = build_routine_trigger(&normalized.trigger);
let action = build_routine_action(
self.store.as_ref(),
&ctx.user_id,
&normalized.name,
&normalized.prompt,
&normalized.execution,
)
.await?;
let action =
build_routine_action(&normalized.name, &normalized.prompt, &normalized.execution);
// Compute next fire time for cron
let next_fire = if let Trigger::Cron {
@@ -1412,22 +1250,13 @@ impl Tool for RoutineUpdateTool {
fn description(&self) -> &str {
"Update an existing routine. Can change prompt, description, enabled state, cron schedule/timezone, \
or full_job permission settings. Pass the routine name and only the fields you want to change. \
This does not convert trigger types."
Pass the routine name and only the fields you want to change. This does not convert trigger types."
}
fn parameters_schema(&self) -> serde_json::Value {
routine_update_parameters_schema()
}
fn requires_approval(&self, params: &serde_json::Value) -> ApprovalRequirement {
if routine_permission_fields_present(params) {
ApprovalRequirement::UnlessAutoApproved
} else {
ApprovalRequirement::Never
}
}
async fn execute(
&self,
params: serde_json::Value,
@@ -1460,72 +1289,6 @@ impl Tool for RoutineUpdateTool {
}
}
let requested_permission_mode = parse_requested_full_job_permission_mode(string_field(
&params,
"execution",
"permission_mode",
&["permission_mode"],
))?;
let requested_tool_permissions = optional_string_array_field(
&params,
"execution",
"tool_permissions",
&["tool_permissions"],
);
let updates_permissions =
requested_permission_mode.is_some() || requested_tool_permissions.is_some();
if updates_permissions {
match &mut routine.action {
RoutineAction::FullJob {
tool_permissions,
permission_mode,
..
} => {
let next_tool_permissions =
requested_tool_permissions.unwrap_or_else(|| tool_permissions.clone());
match requested_permission_mode {
Some(RequestedFullJobPermissionMode::Explicit) => {
*permission_mode = FullJobPermissionMode::Explicit;
*tool_permissions = next_tool_permissions;
}
Some(RequestedFullJobPermissionMode::InheritOwner) => {
*permission_mode = FullJobPermissionMode::InheritOwner;
*tool_permissions = next_tool_permissions;
}
Some(RequestedFullJobPermissionMode::CopyOwner) => {
let owner_settings = load_full_job_permission_settings(
self.store.as_ref(),
&ctx.user_id,
)
.await
.map_err(|e| {
ToolError::ExecutionFailed(format!(
"failed to load routine permission settings: {e}"
))
})?;
*permission_mode = FullJobPermissionMode::Explicit;
*tool_permissions = normalize_tool_names(
owner_settings
.owner_allowed_tools
.into_iter()
.chain(next_tool_permissions),
);
}
None => {
*tool_permissions = next_tool_permissions;
}
}
}
RoutineAction::Lightweight { .. } => {
return Err(ToolError::InvalidParameters(
"permission_mode and tool_permissions can only be updated for full_job routines"
.to_string(),
));
}
}
}
// Validate timezone param if provided
let new_timezone = params
.get("timezone")
@@ -1936,8 +1699,6 @@ mod tests {
"context_paths",
"use_tools",
"max_tool_rounds",
"tool_permissions",
"permission_mode",
"notify_channel",
"notify_user",
"cooldown_secs",
@@ -2036,8 +1797,7 @@ mod tests {
"timezone": "UTC"
},
"execution": {
"mode": "full_job",
"tool_permissions": ["message", "http"]
"mode": "full_job"
},
"delivery": {
"channel": "telegram",
@@ -2062,11 +1822,6 @@ mod tests {
matches!(parsed.execution.mode, NormalizedExecutionMode::FullJob),
"expected full_job execution mode",
);
assert_eq!(
parsed.execution.tool_permissions,
vec!["message".to_string(), "http".to_string()],
);
assert_eq!(parsed.execution.permission_mode, None);
assert_eq!(parsed.delivery.channel.as_deref(), Some("telegram"));
assert_eq!(parsed.delivery.user.as_deref(), Some("ops-team"));
assert_eq!(parsed.cooldown_secs, 30);
@@ -2108,6 +1863,37 @@ mod tests {
);
}
#[test]
fn parses_context_paths_with_trim_drop_empty_and_stable_dedupe() {
let params = serde_json::json!({
"name": "deploy-watch",
"prompt": "Look for deploy requests.",
"request": {
"kind": "manual"
},
"execution": {
"context_paths": [
" context/deploy.md ",
"",
" ",
"context/deploy.md",
"context/notes.md"
]
}
});
let parsed =
parse_routine_create_request(&params).expect("parse context_paths normalization");
assert_eq!(
parsed.execution.context_paths,
vec![
"context/deploy.md".to_string(),
"context/notes.md".to_string()
],
);
}
#[test]
fn parses_grouped_system_event_request() {
let params = serde_json::json!({
@@ -2187,7 +1973,6 @@ mod tests {
"event_pattern": "hello",
"event_channel": "telegram",
"action_type": "full_job",
"tool_permissions": ["message"],
"notify_channel": "telegram",
"notify_user": "123"
});
@@ -2206,10 +1991,6 @@ mod tests {
matches!(parsed.execution.mode, NormalizedExecutionMode::FullJob),
"expected full_job execution mode",
);
assert_eq!(
parsed.execution.tool_permissions,
vec!["message".to_string()],
);
assert_eq!(parsed.delivery.channel.as_deref(), Some("telegram"));
assert_eq!(parsed.delivery.user.as_deref(), Some("123"));
}
@@ -2396,9 +2177,8 @@ mod tests {
.and_then(Value::as_object)
.expect("full_job properties");
assert!(
full_job_props.contains_key("tool_permissions")
&& full_job_props.contains_key("permission_mode"),
"full_job variant should expose permission fields",
full_job_props.len() == 1 && full_job_props.contains_key("mode"),
"full_job variant should only expose the execution mode",
);
}
@@ -2503,8 +2283,6 @@ mod tests {
"schedule",
"timezone",
"description",
"tool_permissions",
"permission_mode",
] {
let _ = schema_property(&schema, field);
}
@@ -2587,71 +2365,26 @@ mod tests {
);
}
#[cfg(feature = "libsql")]
#[tokio::test]
async fn build_full_job_action_defaults_to_inherit_owner_for_new_routines() {
let (db, _tmp) = crate::testing::test_db().await;
#[test]
fn build_full_job_action_uses_live_owner_scope_defaults() {
let execution = NormalizedExecutionRequest {
mode: NormalizedExecutionMode::FullJob,
context_paths: Vec::new(),
use_tools: false,
max_tool_rounds: 3,
tool_permissions: vec!["shell".to_string()],
permission_mode: None,
};
let action =
build_routine_action(db.as_ref(), "default", "issue-1316", "Run it", &execution)
.await
.expect("build action");
let action = build_routine_action("issue-1316", "Run it", &execution);
assert!(matches!(
action,
RoutineAction::FullJob {
permission_mode: FullJobPermissionMode::InheritOwner,
tool_permissions,
..
} if tool_permissions == vec!["shell".to_string()]
));
}
#[cfg(feature = "libsql")]
#[tokio::test]
async fn build_full_job_action_copy_owner_snapshots_allowlist() {
let (db, _tmp) = crate::testing::test_db().await;
db.set_setting(
"default",
crate::agent::routine::FULL_JOB_OWNER_ALLOWED_TOOLS_SETTING_KEY,
&serde_json::json!(["http", "shell"]),
)
.await
.expect("set owner allowlist");
let execution = NormalizedExecutionRequest {
mode: NormalizedExecutionMode::FullJob,
context_paths: Vec::new(),
use_tools: false,
max_tool_rounds: 3,
tool_permissions: vec!["message".to_string(), "shell".to_string()],
permission_mode: Some(RequestedFullJobPermissionMode::CopyOwner),
};
let action =
build_routine_action(db.as_ref(), "default", "issue-1316", "Run it", &execution)
.await
.expect("build action");
assert!(matches!(
action,
RoutineAction::FullJob {
permission_mode: FullJobPermissionMode::Explicit,
tool_permissions,
..
} if tool_permissions
== vec![
"http".to_string(),
"shell".to_string(),
"message".to_string(),
]
title,
description,
max_iterations,
} if title == "issue-1316"
&& description == "Run it"
&& max_iterations == 10
));
}
}
+5
View File
@@ -7,6 +7,7 @@
//! - Delegate tasks to other services
//! - Build new software and tools
mod autonomy;
pub mod builder;
pub mod builtin;
mod coercion;
@@ -20,6 +21,10 @@ pub mod wasm;
mod registry;
mod tool;
pub use autonomy::{
AUTONOMOUS_TOOL_DENYLIST, autonomous_allowed_tool_names, autonomous_unavailable_error,
autonomous_unavailable_message, is_autonomous_tool_denylisted,
};
pub use builder::{
BuildPhase, BuildRequirement, BuildResult, BuildSoftwareTool, BuilderConfig, Language,
LlmSoftwareBuilder, SoftwareBuilder, SoftwareType, Template, TemplateEngine, TemplateType,
+43 -8
View File
@@ -83,7 +83,7 @@ const PROTECTED_TOOL_NAMES: &[&str] = &[
/// Registry of available tools.
pub struct ToolRegistry {
tools: RwLock<HashMap<String, Arc<dyn Tool>>>,
/// Tracks which names were registered as built-in (protected from shadowing).
/// Tracks which names were registered via the built-in startup path.
builtin_names: RwLock<std::collections::HashSet<String>>,
/// Shared credential registry populated by WASM tools, consumed by HTTP tool.
credential_registry: Option<Arc<SharedCredentialRegistry>>,
@@ -138,10 +138,12 @@ impl ToolRegistry {
&self.rate_limiter
}
/// Register a tool. Rejects dynamic tools that try to shadow a built-in name.
/// Register a tool. Rejects dynamic tools that try to shadow a protected built-in name.
pub async fn register(&self, tool: Arc<dyn Tool>) {
let name = tool.name().to_string();
if self.builtin_names.read().await.contains(&name) {
if PROTECTED_TOOL_NAMES.contains(&name.as_str())
&& self.builtin_names.read().await.contains(&name)
{
tracing::warn!(
tool = %name,
"Rejected tool registration: would shadow a built-in tool"
@@ -157,10 +159,7 @@ impl ToolRegistry {
let name = tool.name().to_string();
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()
{
if let Ok(mut builtins) = self.builtin_names.try_write() {
builtins.insert(name.clone());
}
tracing::debug!("Registered tool: {}", name);
@@ -210,6 +209,11 @@ impl ToolRegistry {
self.tools.read().await.values().cloned().collect()
}
/// Get the set of built-in tool names currently registered.
pub async fn builtin_tool_names(&self) -> std::collections::HashSet<String> {
self.builtin_names.read().await.clone()
}
/// Get tool definitions for LLM function calling.
pub async fn tool_definitions(&self) -> Vec<ToolDefinition> {
let mut defs: Vec<ToolDefinition> = self
@@ -888,7 +892,7 @@ mod tests {
#[tokio::test]
async fn test_builtin_tool_cannot_be_shadowed() {
let registry = ToolRegistry::new();
// Register echo as built-in (uses register_sync which marks protected names)
// Register echo as built-in (uses register_sync and echo is protected).
registry.register_sync(Arc::new(EchoTool));
assert!(registry.has("echo").await);
@@ -935,6 +939,37 @@ mod tests {
assert_ne!(desc, "EVIL SHADOW");
}
#[tokio::test]
async fn test_builtin_tool_names_include_non_protected_sync_tools() {
struct NonProtectedBuiltin;
#[async_trait::async_trait]
impl Tool for NonProtectedBuiltin {
fn name(&self) -> &str {
"owner_gate"
}
fn description(&self) -> &str {
"test builtin"
}
fn parameters_schema(&self) -> serde_json::Value {
serde_json::json!({})
}
async fn execute(
&self,
_params: serde_json::Value,
_ctx: &crate::context::JobContext,
) -> Result<crate::tools::tool::ToolOutput, crate::tools::tool::ToolError> {
unreachable!()
}
}
let registry = ToolRegistry::new();
registry.register_sync(Arc::new(NonProtectedBuiltin));
let builtins = registry.builtin_tool_names().await;
assert!(builtins.contains("owner_gate"));
}
#[tokio::test(flavor = "multi_thread", worker_threads = 4)]
async fn concurrent_register_and_read_no_panic() {
use std::sync::Arc as StdArc;
+18 -22
View File
@@ -28,30 +28,29 @@ impl ApprovalRequirement {
}
}
/// Approval context for autonomous tool execution (routines, background jobs).
/// Precomputed autonomous tool scope for background jobs and routines.
///
/// Interactive sessions don't use this type — they rely on session-level
/// auto-approve lists managed by the UI. This enum models only the autonomous
/// case where no interactive user is present.
/// Interactive sessions don't use this type — they still rely on
/// `requires_approval()` and session-level approval state.
#[derive(Debug, Clone)]
pub enum ApprovalContext {
/// Autonomous job with no interactive user. `UnlessAutoApproved` tools are
/// pre-approved. `Always` tools are blocked unless listed in `allowed_tools`.
/// Autonomous job with no interactive user. Only tools in `allowed_tools`
/// may run; interactive approval requirements are ignored.
Autonomous {
/// Tool names that are pre-authorized even for `Always` approval.
/// Tool names that may run autonomously for this job/run.
allowed_tools: std::collections::HashSet<String>,
},
}
impl ApprovalContext {
/// Create an autonomous context with no extra tool permissions.
/// Create an autonomous context with no allowed tools.
pub fn autonomous() -> Self {
Self::Autonomous {
allowed_tools: std::collections::HashSet::new(),
}
}
/// Create an autonomous context with specific tools pre-authorized.
/// Create an autonomous context with specific allowed tools.
pub fn autonomous_with_tools(tools: impl IntoIterator<Item = String>) -> Self {
Self::Autonomous {
allowed_tools: tools.into_iter().collect(),
@@ -59,13 +58,9 @@ impl ApprovalContext {
}
/// Check whether a tool invocation is blocked in this context.
pub fn is_blocked(&self, tool_name: &str, requirement: ApprovalRequirement) -> bool {
pub fn is_blocked(&self, tool_name: &str, _requirement: ApprovalRequirement) -> bool {
match self {
Self::Autonomous { allowed_tools } => match requirement {
ApprovalRequirement::Never => false,
ApprovalRequirement::UnlessAutoApproved => false,
ApprovalRequirement::Always => !allowed_tools.contains(tool_name),
},
Self::Autonomous { allowed_tools } => !allowed_tools.contains(tool_name),
}
}
@@ -889,26 +884,27 @@ mod tests {
}
#[test]
fn test_approval_context_autonomous_allows_unless_auto_approved() {
fn test_approval_context_autonomous_blocks_tools_not_in_scope() {
let ctx = ApprovalContext::autonomous();
assert!(!ctx.is_blocked("shell", ApprovalRequirement::Never));
assert!(!ctx.is_blocked("shell", ApprovalRequirement::UnlessAutoApproved));
assert!(ctx.is_blocked("shell", ApprovalRequirement::Never));
assert!(ctx.is_blocked("shell", ApprovalRequirement::UnlessAutoApproved));
assert!(ctx.is_blocked("shell", ApprovalRequirement::Always));
}
#[test]
fn test_approval_context_autonomous_with_tools_allows_always() {
fn test_approval_context_autonomous_with_tools_allows_registered_name() {
let ctx =
ApprovalContext::autonomous_with_tools(["shell".to_string(), "message".to_string()]);
assert!(!ctx.is_blocked("shell", ApprovalRequirement::Never));
assert!(!ctx.is_blocked("shell", ApprovalRequirement::Always));
assert!(!ctx.is_blocked("message", ApprovalRequirement::Always));
assert!(ctx.is_blocked("http", ApprovalRequirement::Always));
}
#[test]
fn test_approval_context_never_is_not_blocked() {
fn test_approval_context_blocks_never_when_not_in_scope() {
let ctx = ApprovalContext::autonomous();
assert!(!ctx.is_blocked("any_tool", ApprovalRequirement::Never));
assert!(ctx.is_blocked("any_tool", ApprovalRequirement::Never));
}
#[test]
@@ -946,7 +942,7 @@ mod tests {
"other",
ApprovalRequirement::Always
));
assert!(!ApprovalContext::is_blocked_or_default(
assert!(ApprovalContext::is_blocked_or_default(
&ctx,
"any",
ApprovalRequirement::UnlessAutoApproved
+49 -18
View File
@@ -30,7 +30,9 @@ use crate::llm::{
use crate::safety::SafetyLayer;
use crate::tools::execute::process_tool_result;
use crate::tools::rate_limiter::RateLimitResult;
use crate::tools::{ApprovalContext, ToolRegistry, prepare_tool_params, redact_params};
use crate::tools::{
ApprovalContext, ToolRegistry, autonomous_unavailable_error, prepare_tool_params, redact_params,
};
/// Shared dependencies for worker execution.
///
@@ -486,22 +488,20 @@ Report when the job is complete or if you encounter issues you cannot resolve."#
let normalized_params = prepare_tool_params(tool.as_ref(), params);
// Fetch job context early so we have the real user_id for approval, hooks,
// and rate limiting decisions.
let mut job_ctx = deps.context_manager.get_context(job_id).await?;
// Propagate http_interceptor for trace recording/replay
if job_ctx.http_interceptor.is_none() {
job_ctx.http_interceptor = deps.http_interceptor.clone();
}
// Check approval: use context-aware check if available, else block all non-Never tools
let requirement = tool.requires_approval(&normalized_params);
let blocked =
ApprovalContext::is_blocked_or_default(&deps.approval_context, tool_name, requirement);
if blocked {
return Err(crate::error::ToolError::AuthRequired {
name: tool_name.to_string(),
}
.into());
}
// Fetch job context early so we have the real user_id for hooks and rate limiting
let mut job_ctx = deps.context_manager.get_context(job_id).await?;
// Propagate http_interceptor for trace recording/replay
if job_ctx.http_interceptor.is_none() {
job_ctx.http_interceptor = deps.http_interceptor.clone();
return Err(autonomous_unavailable_error(tool_name, &job_ctx.user_id).into());
}
// Check per-tool rate limit before running hooks or executing (cheaper check first)
@@ -761,12 +761,12 @@ Report when the job is complete or if you encounter issues you cannot resolve."#
);
reason_ctx.messages.push(message);
match &result {
match result {
Ok(raw_output) => {
let sanitized = self
.deps
.safety
.sanitize_tool_output(&selection.tool_name, raw_output);
.sanitize_tool_output(&selection.tool_name, &raw_output);
self.log_event(
"tool_result",
serde_json::json!({
@@ -807,7 +807,14 @@ Report when the job is complete or if you encounter issues you cannot resolve."#
}),
);
Ok(())
if matches!(
&e,
Error::Tool(crate::error::ToolError::AutonomousUnavailable { .. })
) {
Err(e)
} else {
Ok(())
}
}
}
}
@@ -1802,7 +1809,7 @@ mod tests {
}
#[tokio::test]
async fn test_approval_context_unblocks_unless_auto_approved() {
async fn test_approval_context_requires_explicit_allowed_tool_names() {
let worker_blocked = make_worker_with_approval(vec![Arc::new(ApprovalTool)], None).await;
let result = worker_blocked
.execute_tool("needs_approval", &serde_json::json!({}))
@@ -1815,13 +1822,18 @@ mod tests {
let worker_allowed = make_worker_with_approval(
vec![Arc::new(ApprovalTool)],
Some(crate::tools::ApprovalContext::autonomous()),
Some(crate::tools::ApprovalContext::autonomous_with_tools([
"needs_approval".to_string(),
])),
)
.await;
let result = worker_allowed
.execute_tool("needs_approval", &serde_json::json!({}))
.await;
assert!(result.is_ok(), "Should be allowed with autonomous context"); // safety: test
assert!(
result.is_ok(),
"Should be allowed when the tool is in the autonomous scope"
); // safety: test
}
#[tokio::test]
@@ -1857,6 +1869,25 @@ mod tests {
);
}
#[tokio::test]
async fn test_approval_context_returns_structured_autonomous_unavailable_error() {
let worker = make_worker_with_approval(
vec![Arc::new(AlwaysApprovalTool)],
Some(crate::tools::ApprovalContext::autonomous()),
)
.await;
let result = worker
.execute_tool("always_approval", &serde_json::json!({}))
.await;
assert!(matches!(
result,
Err(Error::Tool(crate::error::ToolError::AutonomousUnavailable { name, .. }))
if name == "always_approval"
));
}
#[tokio::test]
async fn test_token_budget_exceeded_fails_job() {
let worker = make_worker(vec![]).await;
+158
View File
@@ -0,0 +1,158 @@
use serde::Deserialize;
/// Sensitivity level for a memory layer.
#[derive(Debug, Clone, Default, PartialEq, Eq, Deserialize)]
#[serde(rename_all = "snake_case")]
pub enum LayerSensitivity {
#[default]
Private,
Shared,
}
/// A named memory layer with read/write permissions and a scope.
///
/// Layers map to synthetic `user_id` values in the workspace tables.
/// The `scope` field is the user_id used for DB queries on this layer.
#[derive(Debug, Clone, Deserialize)]
pub struct MemoryLayer {
pub name: String,
pub scope: String,
#[serde(default = "default_true")]
pub writable: bool,
#[serde(default)]
pub sensitivity: LayerSensitivity,
}
fn default_true() -> bool {
true
}
impl MemoryLayer {
/// Build the default layer set: a single private layer for the given user_id.
pub fn default_for_user(user_id: &str) -> Vec<MemoryLayer> {
vec![MemoryLayer {
name: "private".to_string(),
scope: user_id.to_string(),
writable: true,
sensitivity: LayerSensitivity::Private,
}]
}
/// Extract read scopes (all layer scope values).
pub fn read_scopes(layers: &[MemoryLayer]) -> Vec<String> {
layers.iter().map(|l| l.scope.clone()).collect()
}
/// Extract writable scopes only.
pub fn writable_scopes(layers: &[MemoryLayer]) -> Vec<String> {
layers
.iter()
.filter(|l| l.writable)
.map(|l| l.scope.clone())
.collect()
}
/// Find a layer by name. Returns None if not found.
pub fn find<'a>(layers: &'a [MemoryLayer], name: &str) -> Option<&'a MemoryLayer> {
layers.iter().find(|l| l.name == name)
}
/// Find the private layer (first layer with Private sensitivity).
pub fn private_layer(layers: &[MemoryLayer]) -> Option<&MemoryLayer> {
layers
.iter()
.find(|l| l.sensitivity == LayerSensitivity::Private)
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn default_for_user_creates_single_private_layer() {
let layers = MemoryLayer::default_for_user("alice");
assert_eq!(layers.len(), 1);
assert_eq!(layers[0].name, "private");
assert_eq!(layers[0].scope, "alice");
assert!(layers[0].writable);
assert_eq!(layers[0].sensitivity, LayerSensitivity::Private);
}
#[test]
fn read_scopes_collects_all() {
let layers = vec![
MemoryLayer {
name: "private".into(),
scope: "alice".into(),
writable: true,
sensitivity: LayerSensitivity::Private,
},
MemoryLayer {
name: "shared".into(),
scope: "shared".into(),
writable: true,
sensitivity: LayerSensitivity::Shared,
},
MemoryLayer {
name: "reports".into(),
scope: "reports".into(),
writable: false,
sensitivity: LayerSensitivity::Shared,
},
];
let scopes = MemoryLayer::read_scopes(&layers);
assert_eq!(scopes, vec!["alice", "shared", "reports"]);
}
#[test]
fn writable_scopes_filters_read_only() {
let layers = vec![
MemoryLayer {
name: "private".into(),
scope: "alice".into(),
writable: true,
sensitivity: LayerSensitivity::Private,
},
MemoryLayer {
name: "reports".into(),
scope: "reports".into(),
writable: false,
sensitivity: LayerSensitivity::Shared,
},
];
let scopes = MemoryLayer::writable_scopes(&layers);
assert_eq!(scopes, vec!["alice"]);
}
#[test]
fn find_returns_matching_layer() {
let layers = MemoryLayer::default_for_user("alice");
assert!(MemoryLayer::find(&layers, "private").is_some());
assert!(MemoryLayer::find(&layers, "shared").is_none());
}
#[test]
fn deserialize_from_json() {
let json = serde_json::json!({
"name": "shared",
"scope": "shared",
"writable": true,
"sensitivity": "shared"
});
let layer: MemoryLayer = serde_json::from_value(json).unwrap();
assert_eq!(layer.name, "shared");
assert_eq!(layer.sensitivity, LayerSensitivity::Shared);
}
#[test]
fn deserialize_defaults() {
let json = serde_json::json!({
"name": "private",
"scope": "alice"
});
let layer: MemoryLayer = serde_json::from_value(json).unwrap();
assert!(layer.writable); // default true
assert_eq!(layer.sensitivity, LayerSensitivity::Private); // default
}
}
+178 -3
View File
@@ -45,6 +45,8 @@ mod document;
mod embedding_cache;
mod embeddings;
pub mod hygiene;
pub mod layer;
pub mod privacy;
#[cfg(feature = "postgres")]
mod repository;
mod search;
@@ -61,6 +63,17 @@ pub use search::{
FusionStrategy, RankedResult, SearchConfig, SearchResult, fuse_results, reciprocal_rank_fusion,
};
/// Result of a layer-aware write operation.
///
/// Contains the written document plus metadata about whether the write
/// was redirected to a different layer (e.g., sensitive content redirected
/// from shared to private).
pub struct WriteResult {
pub document: MemoryDocument,
pub redirected: bool,
pub actual_layer: String,
}
use std::sync::Arc;
use chrono::{NaiveDate, Utc};
@@ -344,20 +357,29 @@ pub struct Workspace {
bootstrap_completed: std::sync::atomic::AtomicBool,
/// Default search configuration applied to all queries.
search_defaults: SearchConfig,
/// Memory layers this workspace has access to.
memory_layers: Vec<crate::workspace::layer::MemoryLayer>,
/// Optional privacy classifier for shared layer writes.
/// When None, writes go exactly where requested — no silent redirect.
privacy_classifier: Option<Arc<dyn crate::workspace::privacy::PrivacyClassifier>>,
}
impl Workspace {
/// Create a new workspace backed by a PostgreSQL connection pool.
#[cfg(feature = "postgres")]
pub fn new(user_id: impl Into<String>, pool: Pool) -> Self {
let user_id_str = user_id.into();
let memory_layers = crate::workspace::layer::MemoryLayer::default_for_user(&user_id_str);
Self {
user_id: user_id.into(),
user_id: user_id_str,
agent_id: None,
storage: WorkspaceStorage::Repo(Repository::new(pool)),
embeddings: None,
bootstrap_pending: std::sync::atomic::AtomicBool::new(false),
bootstrap_completed: std::sync::atomic::AtomicBool::new(false),
search_defaults: SearchConfig::default(),
memory_layers,
privacy_classifier: None,
}
}
@@ -365,14 +387,18 @@ impl Workspace {
///
/// Use this for libSQL or any other backend that implements the Database trait.
pub fn new_with_db(user_id: impl Into<String>, db: Arc<dyn crate::db::Database>) -> Self {
let user_id_str = user_id.into();
let memory_layers = crate::workspace::layer::MemoryLayer::default_for_user(&user_id_str);
Self {
user_id: user_id.into(),
user_id: user_id_str,
agent_id: None,
storage: WorkspaceStorage::Db(db),
embeddings: None,
bootstrap_pending: std::sync::atomic::AtomicBool::new(false),
bootstrap_completed: std::sync::atomic::AtomicBool::new(false),
search_defaults: SearchConfig::default(),
memory_layers,
privacy_classifier: None,
}
}
@@ -444,6 +470,32 @@ impl Workspace {
self
}
/// Configure memory layers for this workspace.
///
/// Also updates read_user_ids to include all layer scopes.
pub fn with_memory_layers(mut self, layers: Vec<crate::workspace::layer::MemoryLayer>) -> Self {
self.memory_layers = layers;
self
}
/// Set a privacy classifier for shared layer writes.
///
/// When set, writes to shared layers are checked against the classifier
/// and redirected to the private layer if sensitive content is detected.
/// When unset (the default), writes go exactly where requested.
pub fn with_privacy_classifier(
mut self,
classifier: Arc<dyn crate::workspace::privacy::PrivacyClassifier>,
) -> Self {
self.privacy_classifier = Some(classifier);
self
}
/// Get the configured memory layers.
pub fn memory_layers(&self) -> &[crate::workspace::layer::MemoryLayer] {
&self.memory_layers
}
/// Get the user ID.
pub fn user_id(&self) -> &str {
&self.user_id
@@ -501,7 +553,9 @@ impl Workspace {
/// Append content to a file.
///
/// Creates the file if it doesn't exist.
/// Adds a newline separator between existing and new content.
/// Uses a single `\n` separator (suitable for log-style entries).
/// For semantic separation (e.g., memory entries), use `append_memory()`
/// which uses `\n\n`.
pub async fn append(&self, path: &str, content: &str) -> Result<(), WorkspaceError> {
let path = normalize_path(path);
let doc = self
@@ -526,6 +580,127 @@ impl Workspace {
Ok(())
}
/// Resolve the target scope for a layer write, optionally applying privacy guards.
///
/// Validates that the layer exists and is writable. When a privacy classifier
/// is configured on the workspace AND `force` is false, checks shared-layer
/// writes for sensitive content and redirects to the private layer.
///
/// By default no classifier is set — writes go exactly where requested.
/// This is intentional: the LLM chooses the correct layer via system prompt
/// guidance, and a regex classifier can't improve on that decision without
/// unacceptable false positive rates in household contexts (e.g., "doctor",
/// "therapy", phone numbers). Operators who want a safety net can configure
/// one via `with_privacy_classifier()`.
///
/// # Multi-tenant safety (Issue #59)
///
/// Layer scopes are currently used directly as `user_id` for DB operations.
/// In a multi-tenant deployment, an operator could configure a scope that
/// collides with another user's ID, granting write access to their data.
/// Future work should namespace or validate scopes to prevent this.
///
/// Returns `(scope, actual_layer_name, redirected)`.
fn resolve_layer_target(
&self,
layer_name: &str,
content: &str,
force: bool,
) -> Result<(String, String, bool), WorkspaceError> {
use crate::workspace::layer::{LayerSensitivity, MemoryLayer};
let layer = MemoryLayer::find(&self.memory_layers, layer_name).ok_or_else(|| {
WorkspaceError::LayerNotFound {
name: layer_name.to_string(),
}
})?;
if !layer.writable {
return Err(WorkspaceError::LayerReadOnly {
name: layer_name.to_string(),
});
}
if !force
&& layer.sensitivity == LayerSensitivity::Shared
&& let Some(ref classifier) = self.privacy_classifier
&& classifier.classify(content).is_sensitive
{
tracing::warn!(
layer = layer_name,
"Redirected sensitive content to private layer"
);
let private = MemoryLayer::private_layer(&self.memory_layers)
.ok_or(WorkspaceError::PrivacyRedirectFailed)?;
if !private.writable {
return Err(WorkspaceError::PrivacyRedirectFailed);
}
return Ok((private.scope.clone(), private.name.clone(), true));
}
Ok((layer.scope.clone(), layer_name.to_string(), false))
}
/// Write to a specific memory layer.
///
/// Checks that the layer exists and is writable. Uses the layer's scope
/// as the user_id for the database write. For shared layers, sensitive
/// content is automatically redirected to the private layer unless
/// `force` is set.
pub async fn write_to_layer(
&self,
layer_name: &str,
path: &str,
content: &str,
force: bool,
) -> Result<WriteResult, WorkspaceError> {
let (scope, actual_layer, redirected) =
self.resolve_layer_target(layer_name, content, force)?;
let path = normalize_path(path);
let doc = self
.storage
.get_or_create_document_by_path(&scope, self.agent_id, &path)
.await?;
self.storage.update_document(doc.id, content).await?;
self.reindex_document(doc.id).await?;
let document = self.storage.get_document_by_id(doc.id).await?;
Ok(WriteResult {
document,
redirected,
actual_layer,
})
}
/// Write to a layer, with append semantics.
pub async fn append_to_layer(
&self,
layer_name: &str,
path: &str,
content: &str,
force: bool,
) -> Result<WriteResult, WorkspaceError> {
let (scope, actual_layer, redirected) =
self.resolve_layer_target(layer_name, content, force)?;
let path = normalize_path(path);
let doc = self
.storage
.get_or_create_document_by_path(&scope, self.agent_id, &path)
.await?;
let new_content = if doc.content.is_empty() {
content.to_string()
} else {
format!("{}\n\n{}", doc.content, content)
};
self.storage.update_document(doc.id, &new_content).await?;
self.reindex_document(doc.id).await?;
let document = self.storage.get_document_by_id(doc.id).await?;
Ok(WriteResult {
document,
redirected,
actual_layer,
})
}
/// Check if a file exists.
pub async fn exists(&self, path: &str) -> Result<bool, WorkspaceError> {
let path = normalize_path(path);
+276
View File
@@ -0,0 +1,276 @@
use regex::Regex;
/// Result of privacy classification, including confidence level.
///
/// Confidence enables downstream callers to apply thresholds (e.g., only
/// redirect above 0.8) and supports future upgrade to LLM-based classifiers
/// that produce probabilistic scores.
#[derive(Debug, Clone)]
pub struct SensitivityResult {
pub is_sensitive: bool,
pub confidence: f32,
}
/// Classifies content as potentially sensitive for privacy purposes.
///
/// Used to guard writes to shared memory layers -- if content is flagged
/// as sensitive, it can be redirected to the private layer instead.
pub trait PrivacyClassifier: Send + Sync {
/// Classify content and return sensitivity with confidence score.
fn classify(&self, content: &str) -> SensitivityResult;
}
/// Pattern-based privacy classifier using regex matching.
///
/// Default patterns target hard PII (SSN, credit card numbers) where silent
/// redirect is clearly correct. Ambiguous terms (health vocabulary, contact
/// info) are intentionally excluded — they cause false positives in household
/// contexts and silently redirect content the user intended to share.
///
/// Operators who need broader coverage should use `ConfigurablePrivacyClassifier`
/// with domain-specific patterns.
pub struct PatternPrivacyClassifier {
patterns: Vec<Regex>,
}
impl PatternPrivacyClassifier {
pub fn new() -> Result<Self, regex::Error> {
let pattern_strs = [
// SSN — always PII
r"\b\d{3}-\d{2}-\d{4}\b",
// Credit card (basic) — always PII
r"\b\d{4}[\s-]?\d{4}[\s-]?\d{4}[\s-]?\d{4}\b",
// Credentials and auth tokens — high-confidence PII
r"(?i)\b(password|passwd|api[_-]?key|auth[_-]?token|secret[_-]?key)\b",
];
let patterns = pattern_strs
.iter()
.map(|p| Regex::new(p))
.collect::<Result<Vec<_>, _>>()?;
Ok(Self { patterns })
}
}
impl PrivacyClassifier for PatternPrivacyClassifier {
fn classify(&self, content: &str) -> SensitivityResult {
let is_sensitive = self.patterns.iter().any(|p| p.is_match(content));
SensitivityResult {
is_sensitive,
// Regex is binary — matched or not. Always full confidence.
confidence: if is_sensitive { 1.0 } else { 0.0 },
}
}
}
/// User-configurable privacy classifier.
///
/// Accepts custom regex patterns at construction time, allowing operators
/// to tune sensitivity for their use case (e.g., drop health terms that
/// cause false positives, add domain-specific patterns).
///
/// ```
/// use ironclaw::workspace::privacy::ConfigurablePrivacyClassifier;
/// use ironclaw::workspace::privacy::PrivacyClassifier;
///
/// let classifier = ConfigurablePrivacyClassifier::new(vec![
/// r"\b\d{3}-\d{2}-\d{4}\b".into(), // SSN only
/// ]).unwrap();
/// assert!(classifier.classify("SSN: 123-45-6789").is_sensitive);
/// assert!(!classifier.classify("saw the doctor today").is_sensitive);
/// ```
pub struct ConfigurablePrivacyClassifier {
patterns: Vec<Regex>,
}
impl ConfigurablePrivacyClassifier {
/// Create a classifier from user-supplied regex strings.
///
/// Returns an error if any pattern fails to compile.
pub fn new(pattern_strs: Vec<String>) -> Result<Self, regex::Error> {
let patterns = pattern_strs
.iter()
.map(|p| Regex::new(p))
.collect::<Result<Vec<_>, _>>()?;
Ok(Self { patterns })
}
}
impl PrivacyClassifier for ConfigurablePrivacyClassifier {
fn classify(&self, content: &str) -> SensitivityResult {
let is_sensitive = self.patterns.iter().any(|p| p.is_match(content));
SensitivityResult {
is_sensitive,
confidence: if is_sensitive { 1.0 } else { 0.0 },
}
}
}
#[cfg(test)]
mod tests {
use super::*;
fn classifier() -> PatternPrivacyClassifier {
PatternPrivacyClassifier::new().unwrap()
}
// Hard PII — must always trigger
#[test]
fn detects_ssn() {
let result = classifier().classify("My SSN is 123-45-6789");
assert!(result.is_sensitive);
assert_eq!(result.confidence, 1.0);
}
#[test]
fn detects_credit_card() {
let result = classifier().classify("Card: 4111 1111 1111 1111");
assert!(result.is_sensitive);
assert_eq!(result.confidence, 1.0);
}
#[test]
fn detects_password() {
assert!(classifier().classify("my password is hunter2").is_sensitive);
}
#[test]
fn detects_api_key() {
assert!(
classifier()
.classify("set the api_key to sk-1234")
.is_sensitive
);
}
// Household content — must NOT trigger (previous false positives)
#[test]
fn allows_normal_household_content() {
let result = classifier().classify("We need to buy groceries for dinner Saturday");
assert!(!result.is_sensitive);
assert_eq!(result.confidence, 0.0);
}
#[test]
fn allows_doctor_mention() {
assert!(
!classifier()
.classify("the doctor's office called about Saturday")
.is_sensitive
);
}
#[test]
fn allows_email_address() {
assert!(
!classifier()
.classify("email [email protected] about the leak")
.is_sensitive
);
}
#[test]
fn allows_phone_number() {
assert!(
!classifier()
.classify("call the restaurant at 555-123-4567")
.is_sensitive
);
}
#[test]
fn allows_medical_terms_in_context() {
assert!(
!classifier()
.classify("Started new medication for anxiety")
.is_sensitive
);
}
#[test]
fn configurable_with_custom_patterns() {
let c = ConfigurablePrivacyClassifier::new(vec![
r"\b\d{3}-\d{2}-\d{4}\b".into(), // SSN only
])
.unwrap();
assert!(c.classify("SSN: 123-45-6789").is_sensitive);
// Health terms no longer trigger with SSN-only config
assert!(!c.classify("saw the doctor today").is_sensitive);
}
#[test]
fn configurable_rejects_bad_regex() {
let result = ConfigurablePrivacyClassifier::new(vec!["[invalid".into()]);
assert!(result.is_err());
}
#[test]
fn configurable_empty_patterns_allows_everything() {
let c = ConfigurablePrivacyClassifier::new(vec![]).unwrap();
assert!(!c.classify("My SSN is 123-45-6789").is_sensitive);
}
// Format variants
#[test]
fn detects_credit_card_no_separators() {
assert!(
classifier()
.classify("card 4111111111111111 on file")
.is_sensitive
);
}
#[test]
fn detects_credit_card_with_dashes() {
assert!(
classifier()
.classify("Card: 4111-1111-1111-1111")
.is_sensitive
);
}
#[test]
fn detects_ssn_bare() {
assert!(classifier().classify("123-45-6789").is_sensitive);
}
#[test]
fn detects_auth_token_keyword() {
assert!(
classifier()
.classify("set auth_token to abc123")
.is_sensitive
);
}
#[test]
fn detects_secret_key_keyword() {
assert!(
classifier()
.classify("the secret_key is sk-prod-xyz")
.is_sensitive
);
}
#[test]
fn detects_pii_in_longer_document() {
let content = "Meeting notes from Thursday.\n\
Discussed budget and timeline.\n\
SSN is 999-88-7777 for the insurance form.\n\
Action items: follow up with vendor.";
assert!(classifier().classify(content).is_sensitive);
}
#[test]
fn empty_string_is_not_sensitive() {
assert!(!classifier().classify("").is_sensitive);
}
#[test]
fn partial_ssn_not_sensitive() {
assert!(
!classifier()
.classify("code 123-45 in the system")
.is_sensitive
);
}
}
+1
View File
@@ -56,6 +56,7 @@ fn bootstrap_env_round_trips_llm_backend() {
for backend in &[
"nearai",
"anthropic",
"github_copilot",
"ollama",
"openai_compatible",
"tinfoil",
+1 -4
View File
@@ -15,8 +15,7 @@ mod tests {
use uuid::Uuid;
use ironclaw::agent::routine::{
FullJobPermissionMode, Routine, RoutineAction, RoutineGuardrails, RoutineRun, RunStatus,
Trigger,
Routine, RoutineAction, RoutineGuardrails, RoutineRun, RunStatus, Trigger,
};
use ironclaw::context::{JobContext, JobState};
use ironclaw::db::Database;
@@ -46,8 +45,6 @@ mod tests {
title: "Test job".to_string(),
description: "Test description".to_string(),
max_iterations: 5,
tool_permissions: vec![],
permission_mode: FullJobPermissionMode::Explicit,
},
guardrails: RoutineGuardrails {
cooldown: std::time::Duration::from_secs(0),
+82 -36
View File
@@ -4,7 +4,6 @@ These tests exercise the explicit owner model across:
- the web gateway chat UI
- the owner-scoped HTTP webhook channel
- routine tools / routines tab
- job creation via routine execution / jobs tab
"""
import asyncio
@@ -13,7 +12,13 @@ import uuid
import httpx
from helpers import SEL, AUTH_TOKEN, signed_http_webhook_headers
from helpers import (
AUTH_TOKEN,
SEL,
api_get,
api_post,
signed_http_webhook_headers,
)
async def _send_and_get_response(
@@ -58,13 +63,14 @@ async def _post_http_webhook(
content: str,
sender_id: str,
thread_id: str,
) -> str:
wait_for_response: bool = True,
) -> str | None:
"""Send a signed request to the owner-scoped HTTP webhook channel."""
payload = {
"user_id": sender_id,
"thread_id": thread_id,
"content": content,
"wait_for_response": True,
"wait_for_response": wait_for_response,
}
body = json.dumps(payload).encode("utf-8")
@@ -81,8 +87,9 @@ async def _post_http_webhook(
)
data = response.json()
assert data["status"] == "accepted", f"Unexpected webhook response: {data}"
assert data["response"], f"Expected synchronous response body, got: {data}"
return data["response"]
if wait_for_response:
assert data["response"], f"Expected synchronous response body, got: {data}"
return data.get("response")
async def _open_tab(page, tab: str) -> None:
@@ -112,22 +119,60 @@ async def _wait_for_routine(base_url: str, name: str, timeout: float = 20.0) ->
raise AssertionError(f"Routine '{name}' was not created within {timeout}s")
async def _wait_for_job(base_url: str, title: str, timeout: float = 30.0) -> dict:
"""Poll the jobs API until the named job exists."""
async with httpx.AsyncClient() as client:
for _ in range(int(timeout * 2)):
response = await client.get(
f"{base_url}/api/jobs",
headers={"Authorization": f"Bearer {AUTH_TOKEN}"},
timeout=10,
)
response.raise_for_status()
jobs = response.json()["jobs"]
for job in jobs:
if job["title"] == title:
return job
await _poll_sleep()
raise AssertionError(f"Job '{title}' was not created within {timeout}s")
async def _wait_for_http_thread(base_url: str, title_fragment: str, timeout: float = 20.0) -> str:
"""Poll the chat thread list until the matching HTTP thread is visible."""
for _ in range(int(timeout * 2)):
response = await api_get(base_url, "/api/chat/threads", timeout=10)
response.raise_for_status()
threads = response.json()["threads"]
for thread in threads:
if thread.get("channel") != "http":
continue
if title_fragment in (thread.get("title") or ""):
return thread["id"]
await _poll_sleep()
raise AssertionError(
f"HTTP thread containing '{title_fragment}' was not visible within {timeout}s"
)
async def _wait_for_pending_approval(
base_url: str,
thread_id: str,
timeout: float = 20.0,
) -> dict:
"""Poll chat history until the thread exposes a pending approval payload."""
for _ in range(int(timeout * 2)):
response = await api_get(
base_url,
f"/api/chat/history?thread_id={thread_id}",
timeout=10,
)
response.raise_for_status()
pending = response.json().get("pending_approval")
if pending:
return pending
await _poll_sleep()
raise AssertionError(f"Thread '{thread_id}' did not expose a pending approval")
async def _approve_pending_request(base_url: str, thread_id: str, request_id: str) -> None:
"""Approve a pending tool request through the web gateway API."""
response = await api_post(
base_url,
"/api/chat/approval",
json={
"request_id": request_id,
"action": "approve",
"thread_id": thread_id,
},
timeout=10,
)
assert response.status_code == 202, (
f"Approval submission failed: {response.status_code} {response.text[:400]}"
)
data = response.json()
assert data["status"] == "accepted", f"Unexpected approval response: {data}"
async def _poll_sleep() -> None:
@@ -194,33 +239,34 @@ async def test_web_created_routine_is_listed_from_http_channel_across_senders(
assert routine_name in second_sender_text, second_sender_text
async def test_http_created_full_job_routine_can_be_run_from_web_and_shows_in_jobs(
async def test_http_created_full_job_routine_is_visible_in_web_after_approval(
page,
ironclaw_server,
http_channel_server,
):
"""A full-job routine created via HTTP can be run from the web UI and create a job."""
"""A full-job routine created via HTTP appears in the web owner UI after approval."""
routine_name = f"owner-job-{uuid.uuid4().hex[:8]}"
response_text = await _post_http_webhook(
await _post_http_webhook(
http_channel_server,
content=f"create full-job owner routine {routine_name}",
sender_id="http-job-sender",
thread_id="owner-job-thread",
wait_for_response=False,
)
assert routine_name in response_text
await _wait_for_routine(ironclaw_server, routine_name)
thread_id = await _wait_for_http_thread(ironclaw_server, routine_name)
pending = await _wait_for_pending_approval(ironclaw_server, thread_id)
assert pending["tool_name"] == "routine_create"
await _approve_pending_request(
ironclaw_server,
thread_id,
pending["request_id"],
)
routine = await _wait_for_routine(ironclaw_server, routine_name)
assert routine["action_type"] == "full_job"
await _open_tab(page, "routines")
routine_row = page.locator(SEL["routine_row"]).filter(has_text=routine_name).first
await routine_row.wait_for(state="visible", timeout=15000)
await routine_row.locator('button[data-action="trigger-routine"]').click()
await _wait_for_job(ironclaw_server, routine_name, timeout=45.0)
await _open_tab(page, "jobs")
await page.locator(SEL["job_row"]).filter(has_text=routine_name).first.wait_for(
state="visible",
timeout=20000,
)
+3 -20
View File
@@ -10,7 +10,7 @@ mod support;
mod tests {
use std::time::Duration;
use ironclaw::agent::routine::{FullJobPermissionMode, RoutineAction, Trigger};
use ironclaw::agent::routine::{RoutineAction, Trigger};
use crate::support::test_rig::TestRigBuilder;
use crate::support::trace_llm::LlmTrace;
@@ -356,15 +356,8 @@ mod tests {
}
match &routine.action {
RoutineAction::FullJob {
description,
tool_permissions,
permission_mode,
..
} => {
RoutineAction::FullJob { description, .. } => {
assert!(description.contains("Summarize the new issue"));
assert_eq!(tool_permissions, &vec!["shell".to_string()]);
assert_eq!(permission_mode, &FullJobPermissionMode::InheritOwner);
}
other => panic!("expected full_job action, got {other:?}"),
}
@@ -412,18 +405,8 @@ mod tests {
}
match &routine.action {
RoutineAction::FullJob {
description,
tool_permissions,
permission_mode,
..
} => {
RoutineAction::FullJob { description, .. } => {
assert!(description.contains("Prepare the morning digest"));
assert_eq!(
tool_permissions,
&vec!["message".to_string(), "http".to_string()]
);
assert_eq!(permission_mode, &FullJobPermissionMode::InheritOwner);
}
other => panic!("expected full_job action, got {other:?}"),
}
+380 -88
View File
@@ -8,27 +8,33 @@ mod support;
#[cfg(feature = "libsql")]
mod tests {
use std::path::Path;
use std::sync::Arc;
use std::time::Duration;
use chrono::Utc;
use libsql::params;
use secrecy::SecretString;
use uuid::Uuid;
use ironclaw::agent::routine::{
FullJobPermissionMode, NotifyConfig, Routine, RoutineAction, RoutineGuardrails, RoutineRun,
RunStatus, Trigger,
NotifyConfig, Routine, RoutineAction, RoutineGuardrails, RoutineRun, RunStatus, Trigger,
};
use ironclaw::agent::routine_engine::RoutineEngine;
use ironclaw::agent::{HeartbeatConfig, HeartbeatRunner, SandboxReadiness, Scheduler};
use ironclaw::agent::{
HeartbeatConfig, HeartbeatRunner, SandboxReadiness, Scheduler, SchedulerDeps,
};
use ironclaw::channels::IncomingMessage;
use ironclaw::config::{AgentConfig, RoutineConfig, SafetyConfig};
use ironclaw::context::{ContextManager, JobContext};
use ironclaw::db::{Database, libsql::LibSqlBackend};
use ironclaw::extensions::ExtensionManager;
use ironclaw::hooks::HookRegistry;
use ironclaw::llm::LlmProvider;
use ironclaw::safety::SafetyLayer;
use ironclaw::secrets::{InMemorySecretsStore, SecretsCrypto, SecretsStore};
use ironclaw::tools::builtin::routine::RoutineUpdateTool;
use ironclaw::tools::mcp::{McpProcessManager, McpSessionManager};
use ironclaw::tools::{ApprovalRequirement, Tool, ToolError, ToolOutput, ToolRegistry};
use ironclaw::workspace::Workspace;
use ironclaw::workspace::hygiene::HygieneConfig;
@@ -165,11 +171,7 @@ mod tests {
}
}
fn make_full_job_routine(
name: &str,
permission_mode: FullJobPermissionMode,
tool_permissions: Vec<String>,
) -> Routine {
fn make_full_job_routine(name: &str) -> Routine {
Routine {
id: Uuid::new_v4(),
name: name.to_string(),
@@ -181,8 +183,6 @@ mod tests {
title: name.to_string(),
description: "Use the owner-gated tool when permitted.".to_string(),
max_iterations: 3,
tool_permissions,
permission_mode,
},
guardrails: RoutineGuardrails {
cooldown: Duration::from_secs(0),
@@ -234,27 +234,112 @@ mod tests {
LlmTrace::single_turn("test-owner-gate", "run owner gate", steps)
}
async fn setup_owner_gate_engine(db: Arc<dyn Database>, trace: LlmTrace) -> Arc<RoutineEngine> {
fn owner_gate_lightweight_trace() -> LlmTrace {
LlmTrace::single_turn(
"test-owner-gate-lightweight",
"run owner gate",
vec![
TraceStep {
request_hint: None,
response: TraceResponse::ToolCalls {
tool_calls: vec![TraceToolCall {
id: "call_owner_gate".to_string(),
name: "owner_gate".to_string(),
arguments: serde_json::json!({}),
}],
input_tokens: 40,
output_tokens: 10,
},
expected_tool_results: vec![],
},
TraceStep {
request_hint: None,
response: TraceResponse::Text {
content: "ROUTINE_OK".to_string(),
input_tokens: 20,
output_tokens: 5,
},
expected_tool_results: vec![],
},
],
)
}
async fn write_test_extension_wasm(tools_dir: &Path, name: &str) {
tokio::fs::create_dir_all(tools_dir)
.await
.expect("create test wasm tools dir");
tokio::fs::write(tools_dir.join(format!("{name}.wasm")), b"\0asm")
.await
.expect("write test wasm tool marker");
}
fn make_test_extension_manager(
tools: Arc<ToolRegistry>,
tools_dir: &Path,
owner_id: &str,
) -> Arc<ExtensionManager> {
let crypto = Arc::new(
SecretsCrypto::new(SecretString::from(
"0123456789abcdef0123456789abcdef0123456789abcdef0123456789abcdef",
))
.expect("test crypto"),
);
let secrets: Arc<dyn SecretsStore + Send + Sync> =
Arc::new(InMemorySecretsStore::new(crypto));
Arc::new(ExtensionManager::new(
Arc::new(McpSessionManager::new()),
Arc::new(McpProcessManager::new()),
secrets,
tools,
None,
None,
tools_dir.to_path_buf(),
tools_dir.join("channels"),
None,
owner_id.to_string(),
None,
Vec::new(),
))
}
async fn setup_owner_gate_engine(
db: Arc<dyn Database>,
trace: LlmTrace,
tools_dir: &Path,
extension_owner_id: Option<&str>,
activate_owner_gate: bool,
) -> Arc<RoutineEngine> {
let ws = create_workspace(&db);
let (notify_tx, _rx) = tokio::sync::mpsc::channel(16);
let registry = Arc::new(ToolRegistry::new());
registry
.register(Arc::new(OwnerGateTool { store: db.clone() }))
.await;
if extension_owner_id.is_some() {
registry
.register(Arc::new(OwnerGateTool { store: db.clone() }))
.await;
}
if activate_owner_gate {
write_test_extension_wasm(tools_dir, "owner_gate").await;
}
let safety = Arc::new(SafetyLayer::new(&SafetyConfig {
max_output_length: 100_000,
injection_check_enabled: false,
}));
let llm: Arc<dyn LlmProvider> = Arc::new(TraceLlm::from_trace(trace));
let extension_manager = extension_owner_id
.map(|owner_id| make_test_extension_manager(registry.clone(), tools_dir, owner_id));
let scheduler = Arc::new(Scheduler::new(
AgentConfig::for_testing(),
Arc::new(ContextManager::new(5)),
llm.clone(),
safety.clone(),
registry.clone(),
Some(db.clone()),
Arc::new(HookRegistry::new()),
SchedulerDeps {
tools: registry.clone(),
extension_manager: extension_manager.clone(),
store: Some(db.clone()),
hooks: Arc::new(HookRegistry::new()),
},
));
Arc::new(RoutineEngine::new(
@@ -264,9 +349,10 @@ mod tests {
ws,
notify_tx,
Some(scheduler),
extension_manager,
registry,
safety,
SandboxReadiness::DisabledByConfig,
SandboxReadiness::Available,
))
}
@@ -303,6 +389,28 @@ mod tests {
}
}
async fn wait_for_any_run_completion(db: &Arc<dyn Database>, routine_id: Uuid) -> RoutineRun {
let deadline = std::time::Instant::now() + Duration::from_secs(10);
loop {
let runs = db
.list_routine_runs(routine_id, 10)
.await
.expect("list_routine_runs");
if let Some(run) = runs
.into_iter()
.find(|run| run.status != RunStatus::Running)
{
return run;
}
assert!(
std::time::Instant::now() < deadline,
"timed out waiting for any routine run for {routine_id} to complete"
);
tokio::time::sleep(Duration::from_millis(100)).await;
}
}
// -----------------------------------------------------------------------
// Test 1: cron_routine_fires
// -----------------------------------------------------------------------
@@ -345,6 +453,7 @@ mod tests {
ws,
notify_tx,
None,
None,
tools,
safety,
SandboxReadiness::DisabledByConfig,
@@ -423,6 +532,7 @@ mod tests {
ws,
notify_tx,
None,
None,
tools,
safety,
SandboxReadiness::DisabledByConfig,
@@ -517,6 +627,7 @@ mod tests {
ws,
notify_tx,
None,
None,
tools,
safety,
SandboxReadiness::DisabledByConfig,
@@ -625,6 +736,7 @@ mod tests {
ws,
notify_tx,
None,
None,
tools,
safety,
SandboxReadiness::DisabledByConfig,
@@ -767,6 +879,7 @@ mod tests {
ws,
notify_tx,
None,
None,
tools,
safety,
SandboxReadiness::DisabledByConfig,
@@ -953,6 +1066,7 @@ mod tests {
ws,
notify_tx,
None,
None,
tools,
safety,
SandboxReadiness::DisabledByConfig,
@@ -1083,6 +1197,7 @@ mod tests {
ws,
notify_tx,
None, // no scheduler — rejected before dispatch
None,
tools,
safety,
SandboxReadiness::DisabledByConfig,
@@ -1100,8 +1215,6 @@ mod tests {
title: "t".to_string(),
description: "d".to_string(),
max_iterations: 3,
tool_permissions: vec![],
permission_mode: ironclaw::agent::routine::FullJobPermissionMode::Explicit,
},
guardrails: RoutineGuardrails {
cooldown: Duration::from_secs(0),
@@ -1192,6 +1305,7 @@ mod tests {
ws,
notify_tx,
None,
None,
tools,
safety,
SandboxReadiness::DisabledByConfig,
@@ -1250,28 +1364,27 @@ mod tests {
}
// -----------------------------------------------------------------------
// Test: inherit_owner full_job routines can use owner-gated tools
// Test: lightweight manual routines use the owner's active extension tools
// -----------------------------------------------------------------------
#[tokio::test]
async fn full_job_inherit_owner_uses_owner_allowlist() {
let (backend, _tmp) = create_test_backend().await;
async fn lightweight_manual_routine_uses_active_owner_extension_tool() {
let (backend, tmp) = create_test_backend().await;
let db: Arc<dyn Database> = backend;
let engine = setup_owner_gate_engine(db.clone(), owner_gate_trace(true)).await;
db.set_setting(
"default",
ironclaw::agent::routine::FULL_JOB_OWNER_ALLOWED_TOOLS_SETTING_KEY,
&serde_json::json!(["owner_gate"]),
let tools_dir = tmp.path().join("wasm-tools");
let engine = setup_owner_gate_engine(
db.clone(),
owner_gate_lightweight_trace(),
tools_dir.as_path(),
Some("default"),
true,
)
.await
.expect("set owner allowlist");
.await;
let routine = make_full_job_routine(
"inherit-owner-allowed",
FullJobPermissionMode::InheritOwner,
vec![],
);
let mut routine = make_routine("manual-owner-gate", Trigger::Manual, "Use owner_gate.");
if let RoutineAction::Lightweight { use_tools, .. } = &mut routine.action {
*use_tools = true;
}
db.create_routine(&routine).await.expect("create_routine");
let run_id = engine
@@ -1285,20 +1398,141 @@ mod tests {
}
// -----------------------------------------------------------------------
// Test: inherit_owner full_job routines stay blocked without owner allowlist
// Test: full_job cron routines use the owner's active extension tools
// -----------------------------------------------------------------------
#[tokio::test]
async fn full_job_inherit_owner_blocks_without_owner_allowlist() {
let (backend, _tmp) = create_test_backend().await;
async fn full_job_cron_routine_uses_active_owner_extension_tool() {
let (backend, tmp) = create_test_backend().await;
let db: Arc<dyn Database> = backend;
let engine = setup_owner_gate_engine(db.clone(), owner_gate_trace(false)).await;
let tools_dir = tmp.path().join("wasm-tools");
let engine = setup_owner_gate_engine(
db.clone(),
owner_gate_trace(true),
tools_dir.as_path(),
Some("default"),
true,
)
.await;
let routine = make_full_job_routine(
"inherit-owner-blocked",
FullJobPermissionMode::InheritOwner,
vec![],
let mut routine = make_full_job_routine("cron-owner-gate");
routine.trigger = Trigger::Cron {
schedule: "* * * * *".to_string(),
timezone: None,
};
routine.next_fire_at = Some(Utc::now() - chrono::Duration::minutes(1));
db.create_routine(&routine).await.expect("create_routine");
engine.check_cron_triggers().await;
let run = wait_for_any_run_completion(&db, routine.id).await;
assert_eq!(run.status, RunStatus::Ok);
assert_eq!(owner_gate_count(&db).await, 1);
}
// -----------------------------------------------------------------------
// Test: lightweight event routines use the owner's active extension tools
// -----------------------------------------------------------------------
#[tokio::test]
async fn lightweight_event_routine_uses_active_owner_extension_tool() {
let (backend, tmp) = create_test_backend().await;
let db: Arc<dyn Database> = backend;
let tools_dir = tmp.path().join("wasm-tools");
let engine = setup_owner_gate_engine(
db.clone(),
owner_gate_lightweight_trace(),
tools_dir.as_path(),
Some("default"),
true,
)
.await;
let mut routine = make_routine(
"event-owner-gate",
Trigger::Event {
channel: None,
pattern: "owner-gate".to_string(),
},
"Use owner_gate.",
);
if let RoutineAction::Lightweight { use_tools, .. } = &mut routine.action {
*use_tools = true;
}
db.create_routine(&routine).await.expect("create_routine");
engine.refresh_event_cache().await;
let fired = engine
.check_event_triggers("default", "test", "owner-gate")
.await;
assert_eq!(fired, 1, "expected one matching event routine");
let run = wait_for_any_run_completion(&db, routine.id).await;
assert_eq!(run.status, RunStatus::Ok);
assert_eq!(owner_gate_count(&db).await, 1);
}
// -----------------------------------------------------------------------
// Test: full_job system-event routines use the owner's active extension tools
// -----------------------------------------------------------------------
#[tokio::test]
async fn full_job_system_event_routine_uses_active_owner_extension_tool() {
let (backend, tmp) = create_test_backend().await;
let db: Arc<dyn Database> = backend;
let tools_dir = tmp.path().join("wasm-tools");
let engine = setup_owner_gate_engine(
db.clone(),
owner_gate_trace(true),
tools_dir.as_path(),
Some("default"),
true,
)
.await;
let mut routine = make_full_job_routine("system-owner-gate");
routine.trigger = Trigger::SystemEvent {
source: "github".to_string(),
event_type: "issue.opened".to_string(),
filters: std::collections::HashMap::new(),
};
db.create_routine(&routine).await.expect("create_routine");
engine.refresh_event_cache().await;
let fired = engine
.emit_system_event(
"github",
"issue.opened",
&serde_json::json!({"issue_number": 7}),
Some("default"),
)
.await;
assert_eq!(fired, 1, "expected one matching system_event routine");
let run = wait_for_any_run_completion(&db, routine.id).await;
assert_eq!(run.status, RunStatus::Ok);
assert_eq!(owner_gate_count(&db).await, 1);
}
// -----------------------------------------------------------------------
// Test: autonomous runs fail loudly when an extension tool is inactive
// -----------------------------------------------------------------------
#[tokio::test]
async fn full_job_blocks_without_active_owner_extension_tool() {
let (backend, tmp) = create_test_backend().await;
let db: Arc<dyn Database> = backend;
let tools_dir = tmp.path().join("wasm-tools");
let engine = setup_owner_gate_engine(
db.clone(),
owner_gate_trace(false),
tools_dir.as_path(),
Some("default"),
false,
)
.await;
let routine = make_full_job_routine("inactive-owner-gate");
db.create_routine(&routine).await.expect("create_routine");
let run_id = engine
@@ -1309,27 +1543,67 @@ mod tests {
assert_eq!(run.status, RunStatus::Failed);
assert_eq!(owner_gate_count(&db).await, 0);
let failure_reason = db
.get_agent_job_failure_reason(run.job_id.expect("linked job id"))
.await
.expect("load job failure reason")
.expect("missing job failure reason");
assert!(
failure_reason.contains("owner_gate"),
"expected missing-tool failure reason, got {failure_reason}"
);
}
// -----------------------------------------------------------------------
// Test: legacy full_job routines remain explicit until updated
// Test: extension tools activated for another owner are not inherited
// -----------------------------------------------------------------------
#[tokio::test]
async fn legacy_full_job_stays_explicit_until_updated() {
let (backend, _tmp) = create_test_backend().await;
async fn full_job_blocks_when_extension_belongs_to_another_owner() {
let (backend, tmp) = create_test_backend().await;
let db: Arc<dyn Database> = backend;
let tools_dir = tmp.path().join("wasm-tools");
let engine = setup_owner_gate_engine(
db.clone(),
owner_gate_trace(false),
tools_dir.as_path(),
Some("someone-else"),
true,
)
.await;
let routine = make_full_job_routine("other-owner-gate");
db.create_routine(&routine).await.expect("create_routine");
let run_id = engine
.fire_manual(routine.id, None)
.await
.expect("fire manual");
let run = wait_for_run_completion(&db, routine.id, run_id).await;
assert_eq!(run.status, RunStatus::Failed);
assert_eq!(owner_gate_count(&db).await, 0);
let failure_reason = db
.get_agent_job_failure_reason(run.job_id.expect("linked job id"))
.await
.expect("load job failure reason")
.expect("missing job failure reason");
assert!(
failure_reason.contains("owner_gate"),
"expected owner-mismatch failure reason, got {failure_reason}"
);
}
// -----------------------------------------------------------------------
// Test: legacy permission fields are ignored on read and removed on rewrite
// -----------------------------------------------------------------------
#[tokio::test]
async fn legacy_full_job_permission_fields_are_ignored_and_removed_on_update() {
let (backend, tmp) = create_test_backend().await;
let db: Arc<dyn Database> = backend.clone();
db.set_setting(
"default",
ironclaw::agent::routine::FULL_JOB_OWNER_ALLOWED_TOOLS_SETTING_KEY,
&serde_json::json!(["owner_gate"]),
)
.await
.expect("set owner allowlist");
let legacy_routine =
make_full_job_routine("legacy-full-job", FullJobPermissionMode::Explicit, vec![]);
let legacy_routine = make_full_job_routine("legacy-full-job");
db.create_routine(&legacy_routine)
.await
.expect("create_routine");
@@ -1342,59 +1616,77 @@ mod tests {
"title": legacy_routine.name,
"description": "Use the owner-gated tool when permitted.",
"max_iterations": 3,
"tool_permissions": [],
"tool_permissions": ["owner_gate"],
"permission_mode": "inherit_owner",
})
.to_string(),
legacy_routine.id.to_string(),
],
)
.await
.expect("strip permission_mode from action_config");
.expect("inject legacy permission fields into action_config");
let blocked_engine = setup_owner_gate_engine(db.clone(), owner_gate_trace(false)).await;
let first_run_id = blocked_engine
.fire_manual(legacy_routine.id, None)
let loaded = db
.get_routine(legacy_routine.id)
.await
.expect("fire manual legacy routine");
let first_run = wait_for_run_completion(&db, legacy_routine.id, first_run_id).await;
.expect("get_routine")
.expect("routine should still exist");
assert!(matches!(
loaded.action,
RoutineAction::FullJob {
ref title,
ref description,
max_iterations,
} if title == "legacy-full-job"
&& description == "Use the owner-gated tool when permitted."
&& max_iterations == 3
));
assert_eq!(first_run.status, RunStatus::Failed);
assert_eq!(owner_gate_count(&db).await, 0);
let update_tool = RoutineUpdateTool::new(db.clone(), blocked_engine.clone());
let tools_dir = tmp.path().join("wasm-tools");
let engine = setup_owner_gate_engine(
db.clone(),
owner_gate_trace(false),
tools_dir.as_path(),
None,
false,
)
.await;
let update_tool = RoutineUpdateTool::new(db.clone(), engine);
let update_ctx = JobContext::with_user("default", "update", "update legacy routine");
update_tool
.execute(
serde_json::json!({
"name": legacy_routine.name,
"permission_mode": "inherit_owner",
"prompt": "Updated legacy description",
}),
&update_ctx,
)
.await
.expect("routine_update should succeed");
let updated = db
.get_routine(legacy_routine.id)
let mut rows = conn
.query(
"SELECT action_config FROM routines WHERE id = ?1",
params![legacy_routine.id.to_string()],
)
.await
.expect("get_routine")
.expect("routine should still exist");
assert!(matches!(
updated.action,
RoutineAction::FullJob {
permission_mode: FullJobPermissionMode::InheritOwner,
..
}
));
let allowed_engine = setup_owner_gate_engine(db.clone(), owner_gate_trace(true)).await;
let second_run_id = allowed_engine
.fire_manual(legacy_routine.id, None)
.expect("select updated action_config");
let row = rows
.next()
.await
.expect("fire manual updated routine");
let second_run = wait_for_run_completion(&db, legacy_routine.id, second_run_id).await;
.expect("next row")
.expect("updated routine row");
let action_config_raw: String = row.get(0).expect("action_config text");
let action_config: serde_json::Value =
serde_json::from_str(&action_config_raw).expect("parse updated action_config");
assert_eq!(second_run.status, RunStatus::Ok);
assert_eq!(owner_gate_count(&db).await, 1);
assert_eq!(
action_config,
serde_json::json!({
"title": "legacy-full-job",
"description": "Updated legacy description",
"max_iterations": 3,
})
);
}
}
+8 -36
View File
@@ -15,7 +15,7 @@ mod tests {
use chrono::Utc;
use ironclaw::agent::routine::{
FullJobPermissionMode, NotifyConfig, Routine, RoutineAction, RoutineGuardrails, Trigger,
NotifyConfig, Routine, RoutineAction, RoutineGuardrails, Trigger,
};
use uuid::Uuid;
@@ -266,7 +266,7 @@ mod tests {
}
#[tokio::test]
async fn routines_detail_exposes_full_job_permission_resolution() {
async fn routines_detail_omits_legacy_full_job_permission_surface() {
let mock = MockOpenAiServerBuilder::new()
.with_default_response(MockOpenAiResponse::Text("ack".to_string()))
.start()
@@ -276,25 +276,6 @@ mod tests {
GatewayWorkflowHarness::start_openai_compatible(&mock.openai_base_url(), "mock-model")
.await;
harness
.db
.set_setting(
&harness.user_id,
ironclaw::agent::routine::FULL_JOB_OWNER_ALLOWED_TOOLS_SETTING_KEY,
&serde_json::json!(["shell", "http"]),
)
.await
.expect("set owner allowlist");
harness
.db
.set_setting(
&harness.user_id,
ironclaw::agent::routine::FULL_JOB_DEFAULT_PERMISSION_MODE_SETTING_KEY,
&serde_json::json!("copy_owner"),
)
.await
.expect("set owner default mode");
let routine = Routine {
id: Uuid::new_v4(),
name: "wf-full-job-permissions".to_string(),
@@ -306,8 +287,6 @@ mod tests {
title: "permission-detail".to_string(),
description: "Check effective permission detail".to_string(),
max_iterations: 3,
tool_permissions: vec!["message".to_string()],
permission_mode: FullJobPermissionMode::InheritOwner,
},
guardrails: RoutineGuardrails {
cooldown: Duration::from_secs(0),
@@ -346,21 +325,14 @@ mod tests {
.await
.expect("invalid detail response");
assert_eq!(
detail["full_job_permissions"]["permission_mode"].as_str(),
Some("inherit_owner")
assert!(
detail.get("full_job_permissions").is_none(),
"detail response should not expose legacy permission fields: {detail}"
);
assert_eq!(detail["action"]["type"].as_str(), Some("full_job"));
assert_eq!(
detail["full_job_permissions"]["default_permission_mode"].as_str(),
Some("copy_owner")
);
assert_eq!(
detail["full_job_permissions"]["owner_allowed_tools"],
serde_json::json!(["shell", "http"])
);
assert_eq!(
detail["full_job_permissions"]["effective_tool_permissions"],
serde_json::json!(["shell", "http", "message"])
detail["action"]["description"].as_str(),
Some("Check effective permission detail")
);
harness.shutdown().await;
+360
View File
@@ -0,0 +1,360 @@
#![cfg(feature = "libsql")]
//! Integration tests for layered memory using file-backed libSQL.
use std::sync::Arc;
use ironclaw::db::Database;
use ironclaw::db::libsql::LibSqlBackend;
use ironclaw::workspace::Workspace;
use ironclaw::workspace::layer::{LayerSensitivity, MemoryLayer};
use ironclaw::workspace::privacy::PatternPrivacyClassifier;
async fn setup() -> (Arc<dyn Database>, tempfile::TempDir) {
let dir = tempfile::tempdir().expect("create temp dir");
let db_path = dir.path().join("test.db");
let backend = LibSqlBackend::new_local(&db_path).await.expect("create db");
backend.run_migrations().await.expect("run migrations");
let db: Arc<dyn Database> = Arc::new(backend);
(db, dir)
}
fn test_layers() -> Vec<MemoryLayer> {
vec![
MemoryLayer {
name: "private".into(),
scope: "alice".into(),
writable: true,
sensitivity: LayerSensitivity::Private,
},
MemoryLayer {
name: "shared".into(),
scope: "shared".into(),
writable: true,
sensitivity: LayerSensitivity::Shared,
},
MemoryLayer {
name: "reports".into(),
scope: "reports".into(),
writable: false,
sensitivity: LayerSensitivity::Shared,
},
]
}
#[tokio::test]
async fn write_to_private_layer() {
let (db, _dir) = setup().await;
let ws = Workspace::new_with_db("alice", db).with_memory_layers(test_layers());
let result = ws
.write_to_layer("private", "notes/test.md", "Private note", false)
.await
.expect("write should succeed");
assert_eq!(result.document.content, "Private note");
assert!(!result.redirected);
assert_eq!(result.actual_layer, "private");
}
#[tokio::test]
async fn write_to_shared_layer() {
let (db, _dir) = setup().await;
let ws = Workspace::new_with_db("alice", db).with_memory_layers(test_layers());
let result = ws
.write_to_layer("shared", "plans/dinner.md", "Dinner Saturday at 6", false)
.await
.expect("write should succeed");
assert_eq!(result.document.content, "Dinner Saturday at 6");
assert!(!result.redirected);
assert_eq!(result.actual_layer, "shared");
}
#[tokio::test]
async fn write_to_read_only_layer_fails() {
let (db, _dir) = setup().await;
let ws = Workspace::new_with_db("alice", db).with_memory_layers(test_layers());
let result = ws
.write_to_layer("reports", "notes/budget.md", "Some budget note", false)
.await;
assert!(result.is_err());
}
#[tokio::test]
async fn write_to_unknown_layer_fails() {
let (db, _dir) = setup().await;
let ws = Workspace::new_with_db("alice", db).with_memory_layers(test_layers());
let result = ws
.write_to_layer("nonexistent", "notes/test.md", "content", false)
.await;
assert!(result.is_err());
}
#[tokio::test]
async fn no_redirect_without_classifier() {
let (db, _dir) = setup().await;
let ws = Workspace::new_with_db("alice", db).with_memory_layers(test_layers());
// Without a classifier, PII goes exactly where requested
let result = ws
.write_to_layer("shared", "notes/pii.md", "My SSN is 123-45-6789", false)
.await
.expect("write should succeed");
assert!(!result.redirected);
assert_eq!(result.actual_layer, "shared");
}
#[tokio::test]
async fn sensitive_content_redirected_to_private() {
let (db, _dir) = setup().await;
let db_clone = db.clone();
let ws = Workspace::new_with_db("alice", db)
.with_memory_layers(test_layers())
.with_privacy_classifier(Arc::new(PatternPrivacyClassifier::new().unwrap()));
// Write content containing hard PII to shared layer -- should be redirected
let result = ws
.write_to_layer("shared", "notes/pii.md", "My SSN is 123-45-6789", false)
.await
.expect("write should succeed (redirected)");
// WriteResult should indicate redirect to private layer
assert!(result.redirected, "Should be redirected");
assert_eq!(result.actual_layer, "private");
assert_eq!(result.document.content, "My SSN is 123-45-6789");
// Content should be in the private scope (alice), not the shared scope
let private_doc = ws.read("notes/pii.md").await;
assert!(
private_doc.is_ok(),
"Should find content in private scope (alice)"
);
assert_eq!(private_doc.unwrap().content, "My SSN is 123-45-6789");
// Verify content is NOT in the shared scope (same DB, different user_id)
let ws_shared = Workspace::new_with_db("shared", db_clone);
let shared_doc = ws_shared.read("notes/pii.md").await;
assert!(
shared_doc.is_err(),
"Should NOT find content in shared scope"
);
}
#[tokio::test]
async fn default_write_still_works() {
let (db, _dir) = setup().await;
let ws = Workspace::new_with_db("alice", db).with_memory_layers(test_layers());
// Regular write (no layer) should still work
let doc = ws
.write("notes/test.md", "Regular note")
.await
.expect("write should succeed");
assert_eq!(doc.content, "Regular note");
}
#[tokio::test]
async fn append_to_layer_works() {
let (db, _dir) = setup().await;
let ws = Workspace::new_with_db("alice", db).with_memory_layers(test_layers());
// Write initial content to a layer
ws.write_to_layer("private", "notes/log.md", "Entry one", false)
.await
.expect("initial write should succeed");
// Append to the same layer path
let result = ws
.append_to_layer("private", "notes/log.md", "Entry two", false)
.await
.expect("append should succeed");
// Content should be concatenated with double newline
assert!(
result.document.content.contains("Entry one"),
"Should contain first entry"
);
assert!(
result.document.content.contains("Entry two"),
"Should contain second entry"
);
}
#[tokio::test]
async fn sensitive_content_fails_without_private_layer() {
let (db, _dir) = setup().await;
// Workspace with classifier but only shared layers (no private layer for redirect)
let shared_only_layers = vec![MemoryLayer {
name: "shared".into(),
scope: "shared".into(),
writable: true,
sensitivity: LayerSensitivity::Shared,
}];
let ws = Workspace::new_with_db("alice", db)
.with_memory_layers(shared_only_layers)
.with_privacy_classifier(Arc::new(PatternPrivacyClassifier::new().unwrap()));
// Writing PII content should fail (no private layer to redirect to)
let result = ws
.write_to_layer("shared", "notes/pii.md", "My SSN is 123-45-6789", false)
.await;
assert!(
result.is_err(),
"Should fail when no private layer available for redirect"
);
}
#[tokio::test]
async fn append_sensitive_to_shared_redirects() {
let (db, _dir) = setup().await;
let ws = Workspace::new_with_db("alice", db)
.with_memory_layers(test_layers())
.with_privacy_classifier(Arc::new(PatternPrivacyClassifier::new().unwrap()));
// Append PII content to shared layer -- should be redirected
let result = ws
.append_to_layer(
"shared",
"notes/pii.md",
"Card number is 4111 1111 1111 1111",
false,
)
.await
.expect("append should succeed (redirected)");
assert!(result.redirected, "Should be redirected");
assert_eq!(result.actual_layer, "private");
assert!(result.document.content.contains("4111"));
}
#[tokio::test]
async fn force_skips_privacy_redirect() {
let (db, _dir) = setup().await;
let ws = Workspace::new_with_db("alice", db)
.with_memory_layers(test_layers())
.with_privacy_classifier(Arc::new(PatternPrivacyClassifier::new().unwrap()));
// PII content with force=true should stay in shared layer
let result = ws
.write_to_layer("shared", "notes/pii.md", "My SSN is 123-45-6789", true)
.await
.expect("write should succeed without redirect");
assert!(
!result.redirected,
"Should NOT be redirected with force=true"
);
assert_eq!(result.actual_layer, "shared");
}
#[tokio::test]
async fn search_finds_private_layer_content() {
let (db, _dir) = setup().await;
let ws = Workspace::new_with_db("alice", db).with_memory_layers(test_layers());
// Write to the private layer (scope = "alice" = user_id)
ws.write_to_layer(
"private",
"notes/private.md",
"My private thought about waffles",
false,
)
.await
.unwrap();
// Search should find content in the primary scope
let results = ws.search("waffles", 10).await.unwrap();
assert!(
!results.is_empty(),
"Should find results in the private layer"
);
}
#[tokio::test]
async fn write_to_private_invisible_from_shared_scope() {
let (db, _dir) = setup().await;
let db_clone = db.clone();
let ws = Workspace::new_with_db("alice", db).with_memory_layers(test_layers());
ws.write_to_layer("private", "notes/secret.md", "Private data", false)
.await
.expect("write should succeed");
let ws_shared = Workspace::new_with_db("shared", db_clone);
let result = ws_shared.read("notes/secret.md").await;
assert!(
result.is_err(),
"Shared scope must not read private layer content"
);
}
#[tokio::test]
async fn write_to_shared_invisible_from_private_scope() {
let (db, _dir) = setup().await;
let db_clone = db.clone();
let ws = Workspace::new_with_db("alice", db).with_memory_layers(test_layers());
ws.write_to_layer("shared", "plans/visible.md", "Shared plan", false)
.await
.expect("write should succeed");
let ws_alice = Workspace::new_with_db("alice", db_clone);
let result = ws_alice.read("plans/visible.md").await;
assert!(
result.is_err(),
"Private scope must not read shared layer content without multi-scope"
);
}
#[tokio::test]
async fn write_empty_path_to_layer() {
let (db, _dir) = setup().await;
let ws = Workspace::new_with_db("alice", db).with_memory_layers(test_layers());
let result = ws.write_to_layer("private", "", "content", false).await;
// normalize_path("") returns "" — the write succeeds with an empty-string path
assert!(result.is_ok(), "write with empty path should succeed");
let write_result = result.unwrap();
assert_eq!(write_result.document.content, "content");
assert!(!write_result.redirected);
assert_eq!(write_result.actual_layer, "private");
}
#[tokio::test]
async fn overwrite_existing_content_in_layer() {
let (db, _dir) = setup().await;
let ws = Workspace::new_with_db("alice", db).with_memory_layers(test_layers());
ws.write_to_layer("private", "notes/evolving.md", "Version 1", false)
.await
.expect("first write");
let result = ws
.write_to_layer("private", "notes/evolving.md", "Version 2", false)
.await
.expect("overwrite should succeed");
assert_eq!(result.document.content, "Version 2");
assert!(!result.redirected);
}
#[tokio::test]
async fn sensitive_write_to_private_layer_not_redirected() {
let (db, _dir) = setup().await;
let ws = Workspace::new_with_db("alice", db)
.with_memory_layers(test_layers())
.with_privacy_classifier(Arc::new(PatternPrivacyClassifier::new().unwrap()));
let result = ws
.write_to_layer("private", "notes/pii.md", "My SSN is 123-45-6789", false)
.await
.expect("write to private should succeed");
assert!(
!result.redirected,
"Private layer writes should not redirect"
);
assert_eq!(result.actual_layer, "private");
}
+2
View File
@@ -210,6 +210,7 @@ async fn start_test_server_with_provider(
skill_catalog: None,
chat_rate_limiter: ironclaw::channels::web::server::RateLimiter::new(30, 60),
oauth_rate_limiter: ironclaw::channels::web::server::RateLimiter::new(10, 60),
webhook_rate_limiter: ironclaw::channels::web::server::RateLimiter::new(10, 60),
registry_entries: Vec::new(),
cost_guard: None,
routine_engine: Arc::new(tokio::sync::RwLock::new(None)),
@@ -702,6 +703,7 @@ async fn test_no_llm_provider_returns_503() {
skill_catalog: None,
chat_rate_limiter: ironclaw::channels::web::server::RateLimiter::new(30, 60),
oauth_rate_limiter: ironclaw::channels::web::server::RateLimiter::new(10, 60),
webhook_rate_limiter: ironclaw::channels::web::server::RateLimiter::new(10, 60),
registry_entries: Vec::new(),
cost_guard: None,
routine_engine: Arc::new(tokio::sync::RwLock::new(None)),
@@ -230,6 +230,7 @@ impl GatewayWorkflowHarness {
skill_catalog: components.skill_catalog.clone(),
chat_rate_limiter: RateLimiter::new(120, 60),
oauth_rate_limiter: RateLimiter::new(10, 60),
webhook_rate_limiter: RateLimiter::new(10, 60),
registry_entries: Vec::new(),
cost_guard: Some(Arc::clone(&components.cost_guard)),
routine_engine: Arc::clone(&routine_slot),
+1
View File
@@ -591,6 +591,7 @@ impl TestRigBuilder {
Arc::clone(ws),
notify_tx,
None,
None,
components.tools.clone(),
components.safety.clone(),
ironclaw::agent::SandboxReadiness::Available, // tests don't use real Docker
+1
View File
@@ -58,6 +58,7 @@ async fn start_test_server() -> (
skill_catalog: None,
chat_rate_limiter: ironclaw::channels::web::server::RateLimiter::new(30, 60),
oauth_rate_limiter: ironclaw::channels::web::server::RateLimiter::new(10, 60),
webhook_rate_limiter: ironclaw::channels::web::server::RateLimiter::new(10, 60),
registry_entries: Vec::new(),
cost_guard: None,
routine_engine: Arc::new(tokio::sync::RwLock::new(None)),
+13 -4
View File
@@ -42,10 +42,10 @@ impl exports::near::agent::tool::Guest for WebSearchTool {
}
fn description() -> String {
"Search the web using Brave Search. Returns titles, URLs, descriptions, and \
publication dates for matching web pages. Supports filtering by country, \
language, and freshness. Authentication is handled via the 'brave_api_key' \
secret injected by the host."
"Search the web using Brave Search. Returns titles, URLs, descriptions, \
publication dates, and thumbnail images for matching web pages. Supports \
filtering by country, language, and freshness. Authentication is handled \
via the 'brave_api_key' secret injected by the host."
.to_string()
}
}
@@ -76,6 +76,12 @@ struct BraveSearchResult {
url: Option<String>,
description: Option<String>,
age: Option<String>,
thumbnail: Option<BraveThumbnail>,
}
#[derive(Debug, Deserialize)]
struct BraveThumbnail {
src: Option<String>,
}
fn execute_inner(params: &str) -> Result<String, String> {
@@ -198,6 +204,9 @@ fn execute_inner(params: &str) -> Result<String, String> {
if let Some(age) = r.age {
entry["published"] = serde_json::json!(age);
}
if let Some(thumb) = r.thumbnail.and_then(|t| t.src) {
entry["thumbnail"] = serde_json::json!(thumb);
}
// Extract hostname for site_name.
if let Some(host) = extract_hostname(&url) {
entry["site_name"] = serde_json::json!(host);