Merge remote-tracking branch 'origin/main' into fix/setup-audit-fixes

Resolve conflicts between main's simplified config (no bootstrap param,
env-only DatabaseConfig) and our branch's typed ChannelSetupError.

- config.rs: take main's simpler resolve() signatures (no bootstrap)
- main.rs: remove dead check_onboard_needed block and CACHED_KEYCHAIN_KEY ref
- channels.rs: keep ChannelSetupError types, restore settings params from main
- wizard.rs: pass &self.settings to setup_telegram, use ? with From impl
- settings.rs: fix test_llm_backend_round_trip (use std::fs::write, tempfile::tempdir)

Co-Authored-By: Claude Opus 4.6 <[email protected]>
This commit is contained in:
Illia Polosukhin
2026-02-14 14:05:54 -08:00
co-authored by Claude Opus 4.6
48 changed files with 3760 additions and 2499 deletions
Generated
+3 -11
View File
@@ -2210,11 +2210,11 @@ dependencies = [
"hyper 1.8.1",
"hyper-util",
"rustls",
"rustls-native-certs",
"rustls-pki-types",
"tokio",
"tokio-rustls",
"tower-service",
"webpki-roots",
]
[[package]]
@@ -2548,6 +2548,7 @@ dependencies = [
"tower-http 0.6.8",
"tracing",
"tracing-subscriber",
"url",
"urlencoding",
"uuid",
"wasmparser 0.220.1",
@@ -4033,6 +4034,7 @@ dependencies = [
"pin-project-lite",
"quinn",
"rustls",
"rustls-native-certs",
"rustls-pki-types",
"serde",
"serde_json",
@@ -4050,7 +4052,6 @@ dependencies = [
"wasm-bindgen-futures",
"wasm-streams",
"web-sys",
"webpki-roots",
]
[[package]]
@@ -6237,15 +6238,6 @@ dependencies = [
"wasm-bindgen",
]
[[package]]
name = "webpki-roots"
version = "1.0.5"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "12bed680863276c63889429bfd6cab3b99943659923822de1c8a39c49e4d722c"
dependencies = [
"rustls-pki-types",
]
[[package]]
name = "which"
version = "4.4.2"
+3 -2
View File
@@ -22,7 +22,7 @@ tokio-stream = { version = "0.1", features = ["sync"] }
futures = "0.3"
# HTTP client
reqwest = { version = "0.12", default-features = false, features = ["json", "rustls-tls", "stream"] }
reqwest = { version = "0.12", default-features = false, features = ["json", "rustls-tls-native-roots", "stream"] }
# Serialization
serde = { version = "1", features = ["derive"] }
@@ -84,7 +84,8 @@ fs4 = "0.6"
# Secrecy for sensitive values
secrecy = { version = "0.10", features = ["serde"] }
# URL encoding for OAuth flow
# URL parsing and encoding
url = "2"
urlencoding = "2"
# Open URLs in browser
+3 -3
View File
@@ -133,7 +133,7 @@ This document tracks feature parity between IronClaw (Rust implementation) and O
|---------|----------|----------|-------|
| Pi agent runtime | ✅ | | IronClaw uses custom runtime |
| RPC-based execution | ✅ | ✅ | Orchestrator/worker pattern |
| Multi-provider failover | ✅ | | Provider fallback chains |
| Multi-provider failover | ✅ | | `FailoverProvider` tries providers sequentially on retryable errors |
| Per-sender sessions | ✅ | ✅ | |
| Global sessions | ✅ | ❌ | Optional shared context |
| Session pruning | ✅ | ❌ | Auto cleanup old sessions |
@@ -173,7 +173,7 @@ This document tracks feature parity between IronClaw (Rust implementation) and O
| Feature | OpenClaw | IronClaw | Notes |
|---------|----------|----------|-------|
| Auto-discovery | ✅ | ❌ | |
| Failover chains | ✅ | | Provider fallback |
| Failover chains | ✅ | | `FailoverProvider` with configurable `fallback_model` |
| Cooldown management | ✅ | ❌ | Skip failed providers |
| Per-session model override | ✅ | ✅ | Model selector in TUI |
| Model selection UI | ✅ | ✅ | TUI keyboard shortcut |
@@ -419,7 +419,7 @@ This document tracks feature parity between IronClaw (Rust implementation) and O
- ❌ Slack channel (real implementation)
- ✅ Telegram channel (WASM, DM pairing, caption, /start)
- ❌ WhatsApp channel
- Multi-provider failover
- Multi-provider failover (`FailoverProvider` with retryable error classification)
- ❌ Hooks system (beforeInbound, beforeToolCall, etc.)
### P2 - Medium Priority
+36 -36
View File
@@ -181,42 +181,42 @@ External content passes through multiple security layers:
## Architecture
```
┌────────────────────────────────────────────────────────────────────
│ Channels
│ ┌──────┐ ┌──────┐ ┌─────────────┐ ┌─────────────┐
│ │ REPL │ │ HTTP │ │WASM Channels│ │ Web Gateway │
│ └──┬───┘ └──┬───┘ └──────┬──────┘ │ (SSE + WS) │
│ │ │ │ └──────┬──────┘
│ └─────────┴──────────────┴────────────────┘
│ │
│ ┌─────────▼─────────┐
│ │ Agent Loop │ Intent routing
│ └────┬─────────────┘
│ │ │
│ ┌──────────▼───┐ ┌──▼──────────────┐
│ │ Scheduler │ │ Routines Engine │
│ │(parallel jobs)│ │(cron, event, wh) │
│ └──────┬───────┘ └────────┬─────────┘
│ │ │
│ ┌─────────────┼───────────────────┘
│ │ │
│ ┌───▼────┐ ┌────▼────────────────┐
│ │ Local │ │ Orchestrator │
│ │Workers │ │ ┌───────────────┐ │
│ │(in-proc)│ │ │ Docker Sandbox│ │
│ └───┬────┘ │ │ Containers │ │
│ │ │ │ ┌───────────┐ │ │
│ │ │ │ │Worker / CC│ │ │
│ │ │ │ └───────────┘ │ │
│ │ │ └───────────────┘ │
│ │ └─────────┬───────────┘
│ └──────────────────┤
│ │
│ ┌───────────▼──────────┐
│ │ Tool Registry │
│ │ Built-in, MCP, WASM │
│ └──────────────────────┘
└────────────────────────────────────────────────────────────────────
┌────────────────────────────────────────────────────────────────┐
│ Channels │
│ ┌──────┐ ┌──────┐ ┌─────────────┐ ┌─────────────┐ │
│ │ REPL │ │ HTTP │ │WASM Channels│ │ Web Gateway │ │
│ └──┬───┘ └──┬───┘ └──────┬──────┘ │ (SSE + WS) │ │
│ │ │ │ └──────┬──────┘ │
│ └─────────┴──────────────┴────────────────┘ │
│ │ │
│ ┌─────────▼─────────┐ │
│ │ Agent Loop │ Intent routing │
│ └────┬─────────────┘ │
│ │ │ │
│ ┌──────────▼───┐ ┌──▼──────────────┐ │
│ │ Scheduler │ │ Routines Engine │ │
│ │(parallel jobs)│ │(cron, event, wh) │ │
│ └──────┬───────┘ └────────┬─────────┘ │
│ │ │ │
│ ┌─────────────┼───────────────────┘ │
│ │ │ │
│ ┌───▼────┐ ┌────▼────────────────┐ │
│ │ Local │ │ Orchestrator │ │
│ │Workers │ │ ┌───────────────┐ │ │
│ │(in-proc)│ │ │ Docker Sandbox│ │ │
│ └───┬────┘ │ │ Containers │ │ │
│ │ │ │ ┌───────────┐ │ │ │
│ │ │ │ │Worker / CC│ │ │ │
│ │ │ │ └───────────┘ │ │ │
│ │ │ └───────────────┘ │ │
│ │ └─────────┬───────────┘ │
│ └──────────────────┤ │
│ │ │
│ ┌───────────▼──────────┐ │
│ │ Tool Registry │ │
│ │ Built-in, MCP, WASM │ │
│ └──────────────────────┘ │
└────────────────────────────────────────────────────────────────┘
```
### Core Components
-1
View File
@@ -81,7 +81,6 @@ async fn main() -> anyhow::Result<()> {
let session = create_session_manager(SessionConfig {
auth_base_url: config.llm.nearai.auth_base_url.clone(),
session_path: config.llm.nearai.session_path.clone(),
..Default::default()
})
.await;
let llm = create_llm_provider(&config.llm, session)?;
+14 -6
View File
@@ -1171,10 +1171,18 @@ impl Agent {
&& tc.name == "shell"
&& let Some(cmd) = tc
.arguments
.as_str()
.and_then(|s| serde_json::from_str::<serde_json::Value>(s).ok())
.and_then(|v| {
v.get("command").and_then(|c| c.as_str().map(String::from))
.get("command")
.and_then(|c| c.as_str().map(String::from))
.or_else(|| {
tc.arguments
.as_str()
.and_then(|s| {
serde_json::from_str::<serde_json::Value>(s).ok()
})
.and_then(|v| {
v.get("command")
.and_then(|c| c.as_str().map(String::from))
})
})
&& crate::tools::builtin::shell::requires_explicit_approval(&cmd)
{
@@ -1236,7 +1244,7 @@ impl Agent {
&message.channel,
StatusUpdate::ToolResult {
name: tc.name.clone(),
preview: truncate_for_preview(output, 200),
preview: output.clone(),
},
&message.metadata,
)
@@ -1710,7 +1718,7 @@ impl Agent {
&message.channel,
StatusUpdate::ToolResult {
name: pending.tool_name.clone(),
preview: truncate_for_preview(output, 200),
preview: output.clone(),
},
&message.metadata,
)
+25 -3
View File
@@ -299,6 +299,7 @@ Report when the job is complete or if you encounter issues you cannot resolve."#
parameters: tc.arguments.clone(),
reasoning: String::new(),
alternatives: vec![],
tool_call_id: tc.id.clone(),
};
self.process_tool_result(reason_ctx, &selection, result)
@@ -565,7 +566,7 @@ Report when the job is complete or if you encounter issues you cannot resolve."#
);
reason_ctx.messages.push(ChatMessage::tool_result(
"tool_call_id",
&selection.tool_call_id,
&selection.tool_name,
wrapped,
));
@@ -597,7 +598,7 @@ Report when the job is complete or if you encounter issues you cannot resolve."#
}
reason_ctx.messages.push(ChatMessage::tool_result(
"tool_call_id",
&selection.tool_call_id,
&selection.tool_name,
format!("Error: {}", e),
));
@@ -647,12 +648,15 @@ Report when the job is complete or if you encounter issues you cannot resolve."#
.execute_tool(&action.tool_name, &action.parameters)
.await;
// Create a synthetic ToolSelection for process_tool_result
// Create a synthetic ToolSelection for process_tool_result.
// Plan actions don't originate from an LLM tool_call response so
// there is no real tool_call_id; generate a unique one.
let selection = ToolSelection {
tool_name: action.tool_name.clone(),
parameters: action.parameters.clone(),
reasoning: action.reasoning.clone(),
alternatives: vec![],
tool_call_id: format!("plan_{}_{}", self.job_id, i),
};
// Process the result
@@ -774,8 +778,26 @@ impl From<TaskOutput> for Result<String, Error> {
#[cfg(test)]
mod tests {
use crate::llm::ToolSelection;
use crate::util::llm_signals_completion;
#[test]
fn test_tool_selection_preserves_call_id() {
let selection = ToolSelection {
tool_name: "memory_search".to_string(),
parameters: serde_json::json!({"query": "test"}),
reasoning: "Need to search memory".to_string(),
alternatives: vec![],
tool_call_id: "call_abc123".to_string(),
};
assert_eq!(selection.tool_call_id, "call_abc123");
assert_ne!(
selection.tool_call_id, "tool_call_id",
"tool_call_id must not be the hardcoded placeholder string"
);
}
#[test]
fn test_completion_positive_signals() {
assert!(llm_signals_completion("The job is complete."));
+228 -165
View File
@@ -1,147 +1,128 @@
//! Bootstrap configuration for IronClaw.
//! Bootstrap helpers for IronClaw.
//!
//! These are the only settings that MUST live on disk because they're needed
//! before the database connection is established. Everything else lives in the
//! `settings` table in PostgreSQL.
//! The only setting that truly needs disk persistence before the database is
//! available is `DATABASE_URL` (chicken-and-egg: can't connect to DB without
//! it). Everything else is auto-detected or read from env vars.
//!
//! File: `~/.ironclaw/bootstrap.json`
//! File: `~/.ironclaw/.env` (standard dotenvy format)
use std::path::PathBuf;
use serde::{Deserialize, Serialize};
use crate::settings::KeySource;
/// Minimal config needed to connect to the database and decrypt secrets.
///
/// This is the only JSON file IronClaw reads from disk at startup.
/// All other configuration lives in the `settings` table in PostgreSQL.
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct BootstrapConfig {
/// Database connection URL (postgres://...).
#[serde(default)]
pub database_url: Option<String>,
/// Database connection pool size.
#[serde(default)]
pub database_pool_size: Option<usize>,
/// Source for the secrets master key.
#[serde(default)]
pub secrets_master_key_source: KeySource,
/// Whether onboarding wizard has been completed.
#[serde(default)]
pub onboard_completed: bool,
/// Path to the IronClaw-specific `.env` file: `~/.ironclaw/.env`.
pub fn ironclaw_env_path() -> PathBuf {
dirs::home_dir()
.unwrap_or_else(|| PathBuf::from("."))
.join(".ironclaw")
.join(".env")
}
impl Default for BootstrapConfig {
fn default() -> Self {
Self {
database_url: None,
database_pool_size: None,
secrets_master_key_source: KeySource::None,
onboard_completed: false,
}
/// Load env vars from `~/.ironclaw/.env` (in addition to the standard `.env`).
///
/// Call this **after** `dotenvy::dotenv()` so that the standard `./.env`
/// takes priority over `~/.ironclaw/.env`. dotenvy never overwrites
/// existing env vars, so the effective priority is:
///
/// explicit env vars > `./.env` > `~/.ironclaw/.env`
///
/// If `~/.ironclaw/.env` doesn't exist but the legacy `bootstrap.json` does,
/// extracts `DATABASE_URL` from it and writes the `.env` file (one-time
/// upgrade from the old config format).
pub fn load_ironclaw_env() {
let path = ironclaw_env_path();
if !path.exists() {
// One-time upgrade: extract DATABASE_URL from legacy bootstrap.json
migrate_bootstrap_json_to_env(&path);
}
if path.exists() {
let _ = dotenvy::from_path(&path);
}
}
impl BootstrapConfig {
/// Default bootstrap file path: `~/.ironclaw/bootstrap.json`.
pub fn default_path() -> PathBuf {
dirs::home_dir()
.unwrap_or_else(|| PathBuf::from("."))
.join(".ironclaw")
.join("bootstrap.json")
/// If `bootstrap.json` exists, pull `database_url` out of it and write `.env`.
fn migrate_bootstrap_json_to_env(env_path: &std::path::Path) {
let ironclaw_dir = env_path
.parent()
.unwrap_or_else(|| std::path::Path::new("."));
let bootstrap_path = ironclaw_dir.join("bootstrap.json");
if !bootstrap_path.exists() {
return;
}
/// Legacy settings.json path (for migration detection).
pub fn legacy_settings_path() -> PathBuf {
dirs::home_dir()
.unwrap_or_else(|| PathBuf::from("."))
.join(".ironclaw")
.join("settings.json")
}
let content = match std::fs::read_to_string(&bootstrap_path) {
Ok(c) => c,
Err(_) => return,
};
/// Load from the default path, falling back to legacy settings.json,
/// then to defaults if neither exists.
pub fn load() -> Self {
let bootstrap_path = Self::default_path();
if bootstrap_path.exists() {
return Self::load_from(&bootstrap_path);
// Minimal parse: just grab database_url from the JSON
let parsed: serde_json::Value = match serde_json::from_str(&content) {
Ok(v) => v,
Err(_) => return,
};
if let Some(url) = parsed.get("database_url").and_then(|v| v.as_str()) {
if let Some(parent) = env_path.parent()
&& let Err(e) = std::fs::create_dir_all(parent)
{
eprintln!("Warning: failed to create {}: {}", parent.display(), e);
return;
}
// Fall back to legacy settings.json (extract just the 4 bootstrap fields)
let legacy_path = Self::legacy_settings_path();
if legacy_path.exists() {
return Self::load_from_legacy(&legacy_path);
if let Err(e) = std::fs::write(env_path, format!("DATABASE_URL=\"{}\"\n", url)) {
eprintln!("Warning: failed to migrate bootstrap.json to .env: {}", e);
return;
}
Self::default()
}
/// Load from a specific path.
pub fn load_from(path: &PathBuf) -> Self {
match std::fs::read_to_string(path) {
Ok(data) => serde_json::from_str(&data).unwrap_or_default(),
Err(_) => Self::default(),
}
}
/// Extract bootstrap fields from a legacy settings.json.
fn load_from_legacy(path: &PathBuf) -> Self {
match std::fs::read_to_string(path) {
Ok(data) => {
// The legacy Settings struct is a superset; serde will ignore extra fields.
serde_json::from_str(&data).unwrap_or_default()
}
Err(_) => Self::default(),
}
}
/// Save to the default path.
pub fn save(&self) -> std::io::Result<()> {
self.save_to(&Self::default_path())
}
/// Save to a specific path.
pub fn save_to(&self, path: &PathBuf) -> std::io::Result<()> {
if let Some(parent) = path.parent() {
std::fs::create_dir_all(parent)?;
}
let json = serde_json::to_string_pretty(self)
.map_err(|e| std::io::Error::new(std::io::ErrorKind::InvalidData, e.to_string()))?;
std::fs::write(path, json)
rename_to_migrated(&bootstrap_path);
eprintln!(
"Migrated DATABASE_URL from bootstrap.json to {}",
env_path.display()
);
}
}
/// One-time migration from disk config files to the database settings table.
/// Write `DATABASE_URL` to `~/.ironclaw/.env`.
///
/// On first boot after upgrade, checks if:
/// 1. `~/.ironclaw/settings.json` exists
/// 2. The DB settings table is empty for this user
/// Creates the parent directory if it doesn't exist.
/// The value is double-quoted so that `#` (common in URL-encoded passwords)
/// and other shell-special characters are preserved by dotenvy.
pub fn save_database_url(url: &str) -> std::io::Result<()> {
let path = ironclaw_env_path();
if let Some(parent) = path.parent() {
std::fs::create_dir_all(parent)?;
}
std::fs::write(&path, format!("DATABASE_URL=\"{}\"\n", url))
}
/// One-time migration of legacy `~/.ironclaw/settings.json` into the database.
///
/// If both conditions hold, migrates settings, MCP servers, and session data
/// to the database, writes `bootstrap.json`, and renames old files to `.migrated`.
/// Only runs when a `settings.json` exists on disk AND the DB has no settings
/// yet. After the wizard writes directly to the DB, this path is only hit by
/// users upgrading from the old disk-only configuration.
///
/// After syncing, renames `settings.json` to `.migrated` so it won't trigger again.
pub async fn migrate_disk_to_db(
store: &dyn crate::db::Database,
user_id: &str,
) -> Result<(), MigrationError> {
let legacy_settings_path = BootstrapConfig::legacy_settings_path();
let ironclaw_dir = dirs::home_dir()
.unwrap_or_else(|| PathBuf::from("."))
.join(".ironclaw");
let legacy_settings_path = ironclaw_dir.join("settings.json");
if !legacy_settings_path.exists() {
tracing::debug!("No legacy settings.json found, skipping disk-to-DB migration");
return Ok(());
}
// Only migrate if DB is empty for this user
// If DB already has settings, this is not a first boot, the wizard already
// wrote directly to the DB. Just clean up the stale file.
let has_settings = store.has_settings(user_id).await.map_err(|e| {
MigrationError::Database(format!("Failed to check existing settings: {}", e))
})?;
if has_settings {
tracing::debug!(
"DB already has settings for user '{}', skipping migration",
user_id
);
tracing::info!("DB already has settings, renaming stale settings.json");
rename_to_migrated(&legacy_settings_path);
return Ok(());
}
@@ -160,22 +141,14 @@ pub async fn migrate_disk_to_db(
tracing::info!("Migrated {} settings to database", db_map.len());
}
// 2. Write bootstrap.json with the 4 essential fields
let bootstrap = BootstrapConfig {
database_url: settings.database_url.clone(),
database_pool_size: settings.database_pool_size,
secrets_master_key_source: settings.secrets_master_key_source,
onboard_completed: settings.onboard_completed,
};
bootstrap
.save()
.map_err(|e| MigrationError::Io(format!("Failed to write bootstrap.json: {}", e)))?;
tracing::info!("Wrote bootstrap.json");
// 2. Write DATABASE_URL to ~/.ironclaw/.env
if let Some(ref url) = settings.database_url {
save_database_url(url)
.map_err(|e| MigrationError::Io(format!("Failed to write .env: {}", e)))?;
tracing::info!("Wrote DATABASE_URL to {}", ironclaw_env_path().display());
}
// 3. Migrate mcp-servers.json if it exists
let ironclaw_dir = dirs::home_dir()
.unwrap_or_else(|| PathBuf::from("."))
.join(".ironclaw");
let mcp_path = ironclaw_dir.join("mcp-servers.json");
if mcp_path.exists() {
match std::fs::read_to_string(&mcp_path) {
@@ -236,12 +209,19 @@ pub async fn migrate_disk_to_db(
// 5. Rename settings.json to .migrated (don't delete, safety net)
rename_to_migrated(&legacy_settings_path);
// 6. Clean up old bootstrap.json if it exists (superseded by .env)
let old_bootstrap = ironclaw_dir.join("bootstrap.json");
if old_bootstrap.exists() {
rename_to_migrated(&old_bootstrap);
tracing::info!("Renamed old bootstrap.json to .migrated");
}
tracing::info!("Disk-to-DB migration complete");
Ok(())
}
/// Rename a file to `<name>.migrated` as a safety net.
fn rename_to_migrated(path: &PathBuf) {
fn rename_to_migrated(path: &std::path::Path) {
let mut migrated = path.as_os_str().to_owned();
migrated.push(".migrated");
if let Err(e) = std::fs::rename(path, &migrated) {
@@ -264,62 +244,145 @@ mod tests {
use tempfile::tempdir;
#[test]
fn test_bootstrap_save_load() {
fn test_save_and_load_database_url() {
let dir = tempdir().unwrap();
let path = dir.path().join("bootstrap.json");
let env_path = dir.path().join(".env");
let config = BootstrapConfig {
database_url: Some("postgres://localhost/test".to_string()),
database_pool_size: Some(5),
secrets_master_key_source: KeySource::Keychain,
onboard_completed: true,
};
// Write in the quoted format that save_database_url uses
let url = "postgres://localhost:5432/ironclaw_test";
std::fs::write(&env_path, format!("DATABASE_URL=\"{}\"\n", url)).unwrap();
config.save_to(&path).unwrap();
let loaded = BootstrapConfig::load_from(&path);
// Verify the content is a valid dotenv line (quoted)
let content = std::fs::read_to_string(&env_path).unwrap();
assert_eq!(
loaded.database_url,
Some("postgres://localhost/test".to_string())
content,
"DATABASE_URL=\"postgres://localhost:5432/ironclaw_test\"\n"
);
assert_eq!(loaded.database_pool_size, Some(5));
assert_eq!(loaded.secrets_master_key_source, KeySource::Keychain);
assert!(loaded.onboard_completed);
// Verify dotenvy can parse it (strips quotes automatically)
let parsed: Vec<(String, String)> = dotenvy::from_path_iter(&env_path)
.unwrap()
.filter_map(|r| r.ok())
.collect();
assert_eq!(parsed.len(), 1);
assert_eq!(parsed[0].0, "DATABASE_URL");
assert_eq!(parsed[0].1, url);
}
#[test]
fn test_bootstrap_from_legacy_settings() {
fn test_save_database_url_with_hash_in_password() {
let dir = tempdir().unwrap();
let path = dir.path().join("settings.json");
let env_path = dir.path().join(".env");
// Write a legacy settings.json with many extra fields
let legacy = serde_json::json!({
"database_url": "postgres://localhost/ironclaw",
"database_pool_size": 10,
// URLs with # in the password are common (URL-encoded special chars).
// Without quoting, dotenvy treats # as a comment delimiter.
let url = "postgres://user:p%23ss@localhost:5432/ironclaw";
std::fs::write(&env_path, format!("DATABASE_URL=\"{}\"\n", url)).unwrap();
let parsed: Vec<(String, String)> = dotenvy::from_path_iter(&env_path)
.unwrap()
.filter_map(|r| r.ok())
.collect();
assert_eq!(parsed.len(), 1);
assert_eq!(parsed[0].0, "DATABASE_URL");
assert_eq!(parsed[0].1, url);
}
#[test]
fn test_save_database_url_creates_parent_dirs() {
let dir = tempdir().unwrap();
let nested = dir.path().join("deep").join("nested");
let env_path = nested.join(".env");
// Parent doesn't exist yet
assert!(!nested.exists());
// The global function uses a fixed path, so we test the logic directly
std::fs::create_dir_all(&nested).unwrap();
std::fs::write(&env_path, "DATABASE_URL=postgres://test\n").unwrap();
assert!(env_path.exists());
let content = std::fs::read_to_string(&env_path).unwrap();
assert!(content.contains("DATABASE_URL=postgres://test"));
}
#[test]
fn test_ironclaw_env_path() {
let path = ironclaw_env_path();
assert!(path.ends_with(".ironclaw/.env"));
}
#[test]
fn test_migrate_bootstrap_json_to_env() {
let dir = tempdir().unwrap();
let env_path = dir.path().join(".env");
let bootstrap_path = dir.path().join("bootstrap.json");
// Write a legacy bootstrap.json
let bootstrap_json = serde_json::json!({
"database_url": "postgres://localhost/ironclaw_upgrade",
"database_pool_size": 5,
"secrets_master_key_source": "keychain",
"onboard_completed": true,
"selected_model": "claude-3-5-sonnet",
"agent": { "name": "testbot", "max_parallel_jobs": 3 },
"heartbeat": { "enabled": true }
"onboard_completed": true
});
std::fs::write(&path, serde_json::to_string_pretty(&legacy).unwrap()).unwrap();
std::fs::write(
&bootstrap_path,
serde_json::to_string_pretty(&bootstrap_json).unwrap(),
)
.unwrap();
let config = BootstrapConfig::load_from_legacy(&path);
assert!(!env_path.exists());
assert!(bootstrap_path.exists());
// Run the migration
migrate_bootstrap_json_to_env(&env_path);
// .env should now exist with DATABASE_URL
assert!(env_path.exists());
let content = std::fs::read_to_string(&env_path).unwrap();
assert_eq!(
config.database_url,
Some("postgres://localhost/ironclaw".to_string())
content,
"DATABASE_URL=\"postgres://localhost/ironclaw_upgrade\"\n"
);
assert_eq!(config.database_pool_size, Some(10));
assert_eq!(config.secrets_master_key_source, KeySource::Keychain);
assert!(config.onboard_completed);
// bootstrap.json should be renamed to .migrated
assert!(!bootstrap_path.exists());
assert!(dir.path().join("bootstrap.json.migrated").exists());
}
#[test]
fn test_bootstrap_defaults() {
let config = BootstrapConfig::default();
assert!(config.database_url.is_none());
assert!(config.database_pool_size.is_none());
assert_eq!(config.secrets_master_key_source, KeySource::None);
assert!(!config.onboard_completed);
fn test_migrate_bootstrap_json_no_database_url() {
let dir = tempdir().unwrap();
let env_path = dir.path().join(".env");
let bootstrap_path = dir.path().join("bootstrap.json");
// bootstrap.json with no database_url
let bootstrap_json = serde_json::json!({
"onboard_completed": false
});
std::fs::write(
&bootstrap_path,
serde_json::to_string_pretty(&bootstrap_json).unwrap(),
)
.unwrap();
migrate_bootstrap_json_to_env(&env_path);
// .env should NOT be created
assert!(!env_path.exists());
// bootstrap.json should remain (no migration happened)
assert!(bootstrap_path.exists());
}
#[test]
fn test_migrate_bootstrap_json_missing() {
let dir = tempdir().unwrap();
let env_path = dir.path().join(".env");
// No bootstrap.json at all
migrate_bootstrap_json_to_env(&env_path);
// Nothing should happen
assert!(!env_path.exists());
}
}
+9 -5
View File
@@ -37,6 +37,9 @@ use crate::agent::truncate_for_preview;
use crate::channels::{Channel, IncomingMessage, MessageStream, OutgoingResponse, StatusUpdate};
use crate::error::ChannelError;
/// Max characters for tool result previews in the terminal.
const CLI_TOOL_RESULT_MAX: usize = 200;
/// Max characters for thinking/status messages in the terminal.
const CLI_STATUS_MAX: usize = 200;
@@ -265,7 +268,7 @@ impl Channel for ReplChannel {
std::thread::spawn(move || {
// Single message mode: send it and return
if let Some(msg) = single_message {
let incoming = IncomingMessage::new("repl", "user", &msg);
let incoming = IncomingMessage::new("repl", "default", &msg);
let _ = tx.blocking_send(incoming);
return;
}
@@ -333,21 +336,21 @@ impl Channel for ReplChannel {
_ => {}
}
let msg = IncomingMessage::new("repl", "user", line);
let msg = IncomingMessage::new("repl", "default", line);
if tx.blocking_send(msg).is_err() {
break;
}
}
Err(ReadlineError::Interrupted) => {
// Ctrl+C: send /interrupt
let msg = IncomingMessage::new("repl", "user", "/interrupt");
let msg = IncomingMessage::new("repl", "default", "/interrupt");
if tx.blocking_send(msg).is_err() {
break;
}
}
Err(ReadlineError::Eof) => {
// Ctrl+D: send /quit so the agent loop runs graceful shutdown
let msg = IncomingMessage::new("repl", "user", "/quit");
let msg = IncomingMessage::new("repl", "default", "/quit");
let _ = tx.blocking_send(msg);
break;
}
@@ -418,7 +421,8 @@ impl Channel for ReplChannel {
}
}
StatusUpdate::ToolResult { name: _, preview } => {
eprintln!(" \x1b[90m{preview}\x1b[0m");
let display = truncate_for_preview(&preview, CLI_TOOL_RESULT_MAX);
eprintln!(" \x1b[90m{display}\x1b[0m");
}
StatusUpdate::StreamChunk(chunk) => {
// Print separator on the false-to-true transition
+82 -12
View File
@@ -76,6 +76,9 @@ struct ChannelStoreData {
credentials: HashMap<String, String>,
/// Pairing store for DM pairing (guest access control).
pairing_store: Arc<PairingStore>,
/// Dedicated tokio runtime for HTTP requests, lazily initialized.
/// Reused across multiple `http_request` calls within one execution.
http_runtime: Option<tokio::runtime::Runtime>,
}
impl ChannelStoreData {
@@ -96,6 +99,7 @@ impl ChannelStoreData {
table: ResourceTable::new(),
credentials,
pairing_store,
http_runtime: None,
}
}
@@ -283,10 +287,25 @@ impl near::agent::channel_host::Host for ChannelStoreData {
.map(|h| h.max_response_bytes)
.unwrap_or(10 * 1024 * 1024);
// Make the HTTP request using blocking I/O
// We're already in a spawn_blocking context, so we can use block_on
let result = tokio::runtime::Handle::current().block_on(async {
let client = reqwest::Client::new();
// Make the HTTP request using a dedicated single-threaded runtime.
// We're inside spawn_blocking, so we can't rely on the main runtime's
// I/O driver (it may be busy with WASM compilation or other startup work).
// A dedicated runtime gives us our own I/O driver and avoids contention.
// The runtime is lazily created and reused across calls within one execution.
if self.http_runtime.is_none() {
self.http_runtime = Some(
tokio::runtime::Builder::new_current_thread()
.enable_all()
.build()
.map_err(|e| format!("Failed to create HTTP runtime: {e}"))?,
);
}
let rt = self.http_runtime.as_ref().expect("just initialized");
let result = rt.block_on(async {
let client = reqwest::Client::builder()
.connect_timeout(std::time::Duration::from_secs(10))
.build()
.map_err(|e| format!("Failed to build HTTP client: {e}"))?;
let mut request = match method.to_uppercase().as_str() {
"GET" => client.get(&url),
@@ -308,9 +327,9 @@ impl near::agent::channel_host::Host for ChannelStoreData {
request = request.body(body_bytes);
}
// Send request with caller-specified timeout (default 30s).
// Cap at callback_timeout to prevent outliving the host wrapper.
let timeout = std::time::Duration::from_millis(timeout_ms.unwrap_or(30_000) as u64);
// Send request with caller-specified timeout (default 30s, max 5min).
let timeout_ms = timeout_ms.unwrap_or(30_000).min(300_000) as u64;
let timeout = std::time::Duration::from_millis(timeout_ms);
let response = request.timeout(timeout).send().await.map_err(|e| {
// Walk the full error chain so we get the actual root cause
// (DNS, TLS, connection refused, etc.) instead of just
@@ -795,7 +814,21 @@ impl WasmChannel {
.await;
match result {
Ok(Ok((config, _host_state))) => {
Ok(Ok((config, mut host_state))) => {
// Surface WASM guest logs (errors/warnings from webhook setup, etc.)
for entry in host_state.take_logs() {
match entry.level {
crate::tools::wasm::LogLevel::Error => {
tracing::error!(channel = %self.name, "{}", entry.message);
}
crate::tools::wasm::LogLevel::Warn => {
tracing::warn!(channel = %self.name, "{}", entry.message);
}
_ => {
tracing::debug!(channel = %self.name, "{}", entry.message);
}
}
}
tracing::info!(
channel = %self.name,
display_name = %config.display_name,
@@ -2615,15 +2648,52 @@ mod tests {
assert_eq!(store.redact_credentials(input), input);
}
/// Verify that the block_on-inside-spawn_blocking pattern used by the WASM
/// channel HTTP host function doesn't deadlock or panic.
/// Verify that WASM HTTP host functions work using a dedicated
/// current-thread runtime inside spawn_blocking.
#[tokio::test]
async fn test_block_on_inside_spawn_blocking_does_not_deadlock() {
async fn test_dedicated_runtime_inside_spawn_blocking() {
let result = tokio::task::spawn_blocking(|| {
tokio::runtime::Handle::current().block_on(async { 42 })
let rt = tokio::runtime::Builder::new_current_thread()
.enable_all()
.build()
.expect("failed to build runtime");
rt.block_on(async { 42 })
})
.await
.expect("spawn_blocking panicked");
assert_eq!(result, 42);
}
/// Verify a real HTTP request works using the dedicated-runtime pattern.
/// This catches DNS, TLS, and I/O driver issues that trivial tests miss.
#[tokio::test]
#[ignore] // requires network
async fn test_dedicated_runtime_real_http() {
let result = tokio::task::spawn_blocking(|| {
let rt = tokio::runtime::Builder::new_current_thread()
.enable_all()
.build()
.expect("failed to build runtime");
rt.block_on(async {
let client = reqwest::Client::builder()
.connect_timeout(std::time::Duration::from_secs(10))
.build()
.expect("failed to build client");
let resp = client
.get("https://api.telegram.org/bot000/getMe")
.timeout(std::time::Duration::from_secs(10))
.send()
.await;
match resp {
Ok(r) => r.status().as_u16(),
Err(e) if e.is_timeout() => panic!("request timed out: {e}"),
Err(e) => panic!("unexpected error: {e}"),
}
})
})
.await
.expect("spawn_blocking panicked");
// 404 because "000" is not a valid bot token
assert_eq!(result, 404);
}
}
+24 -56
View File
@@ -48,8 +48,6 @@ pub enum ConfigCommand {
/// Connects to the database to read/write settings. Falls back to disk
/// if the database is not available.
pub async fn run_config_command(cmd: ConfigCommand) -> anyhow::Result<()> {
let _ = dotenvy::dotenv();
// Try to connect to the DB for settings access
let db: Option<Arc<dyn crate::db::Database>> = match connect_db().await {
Ok(d) => Some(d),
@@ -92,7 +90,7 @@ async fn load_settings(store: Option<&dyn crate::db::Database>) -> Settings {
_ => {}
}
}
Settings::load()
Settings::default()
}
/// List all settings.
@@ -155,19 +153,17 @@ async fn set_setting(
.set(path, value)
.map_err(|e| anyhow::anyhow!("{}", e))?;
// Save to DB if available, otherwise disk
if let Some(store) = store {
let json_value = match serde_json::from_str::<serde_json::Value>(value) {
Ok(v) => v,
Err(_) => serde_json::Value::String(value.to_string()),
};
store
.set_setting(DEFAULT_USER_ID, path, &json_value)
.await
.map_err(|e| anyhow::anyhow!("Failed to save to database: {}", e))?;
} else {
settings.save()?;
}
let store = store.ok_or_else(|| {
anyhow::anyhow!("Database connection required to save settings. Check DATABASE_URL.")
})?;
let json_value = match serde_json::from_str::<serde_json::Value>(value) {
Ok(v) => v,
Err(_) => serde_json::Value::String(value.to_string()),
};
store
.set_setting(DEFAULT_USER_ID, path, &json_value)
.await
.map_err(|e| anyhow::anyhow!("Failed to save to database: {}", e))?;
println!("Set {} = {}", path, value);
Ok(())
@@ -180,17 +176,13 @@ async fn reset_setting(store: Option<&dyn crate::db::Database>, path: &str) -> a
.get(path)
.ok_or_else(|| anyhow::anyhow!("Unknown setting: {}", path))?;
// Delete from DB (falling back to default) or reset on disk
if let Some(store) = store {
store
.delete_setting(DEFAULT_USER_ID, path)
.await
.map_err(|e| anyhow::anyhow!("Failed to delete setting from database: {}", e))?;
} else {
let mut settings = Settings::load();
settings.reset(path).map_err(|e| anyhow::anyhow!("{}", e))?;
settings.save()?;
}
let store = store.ok_or_else(|| {
anyhow::anyhow!("Database connection required to reset settings. Check DATABASE_URL.")
})?;
store
.delete_setting(DEFAULT_USER_ID, path)
.await
.map_err(|e| anyhow::anyhow!("Failed to delete setting from database: {}", e))?;
println!("Reset {} to default: {}", path, default_value);
Ok(())
@@ -200,37 +192,13 @@ async fn reset_setting(store: Option<&dyn crate::db::Database>, path: &str) -> a
fn show_path(has_db: bool) -> anyhow::Result<()> {
if has_db {
println!("Settings stored in: database (settings table)");
println!(
"Bootstrap config: {}",
crate::bootstrap::BootstrapConfig::default_path().display()
);
} else {
let path = Settings::default_path();
println!("Settings stored in: {} (disk fallback)", path.display());
if path.exists() {
let metadata = std::fs::metadata(&path)?;
println!(" Size: {} bytes", metadata.len());
if let Ok(modified) = metadata.modified() {
use std::time::SystemTime;
let duration = SystemTime::now()
.duration_since(modified)
.unwrap_or_default();
let secs = duration.as_secs();
if secs < 60 {
println!(" Modified: {} seconds ago", secs);
} else if secs < 3600 {
println!(" Modified: {} minutes ago", secs / 60);
} else if secs < 86400 {
println!(" Modified: {} hours ago", secs / 3600);
} else {
println!(" Modified: {} days ago", secs / 86400);
}
}
} else {
println!(" (does not exist, using defaults)");
}
println!("Settings stored in: PostgreSQL (not connected, using defaults)");
}
println!(
"Env config: {}",
crate::bootstrap::ironclaw_env_path().display()
);
Ok(())
}
+1
View File
@@ -12,6 +12,7 @@
mod config;
mod mcp;
pub mod memory;
pub mod oauth_defaults;
mod pairing;
pub mod status;
mod tool;
+343
View File
@@ -0,0 +1,343 @@
//! Shared OAuth infrastructure: built-in credentials, callback server, landing pages.
//!
//! Every OAuth flow in the codebase (WASM tool auth, MCP server auth, NEAR AI login)
//! uses the same callback port, landing page, and listener logic from this module.
//!
//! # Built-in Credentials
//!
//! Many CLI tools (gcloud, rclone, gdrive) ship with default OAuth credentials
//! so users don't need to register their own OAuth app. Google explicitly
//! documents that client_secret for "Desktop App" / "Installed App" types
//! is NOT actually secret.
//!
//! Default credentials are hardcoded below. They can be overridden at:
//!
//! - **Compile time**: Set IRONCLAW_GOOGLE_CLIENT_ID / IRONCLAW_GOOGLE_CLIENT_SECRET
//! env vars before building to replace the hardcoded defaults.
//! - **Runtime**: Users can set GOOGLE_OAUTH_CLIENT_ID / GOOGLE_OAUTH_CLIENT_SECRET
//! env vars, which take priority over built-in defaults.
use std::time::Duration;
use tokio::io::{AsyncBufReadExt, AsyncWriteExt, BufReader};
use tokio::net::TcpListener;
// ── Built-in credentials ────────────────────────────────────────────────
pub struct OAuthCredentials {
pub client_id: &'static str,
pub client_secret: &'static str,
}
/// Google OAuth "Desktop App" credentials, shared across all Google tools.
/// Compile-time env vars override the hardcoded defaults below.
const GOOGLE_CLIENT_ID: &str = match option_env!("IRONCLAW_GOOGLE_CLIENT_ID") {
Some(v) => v,
None => "564604149681-efo25d43rs85v0tibdepsmdv5dsrhhr0.apps.googleusercontent.com",
};
const GOOGLE_CLIENT_SECRET: &str = match option_env!("IRONCLAW_GOOGLE_CLIENT_SECRET") {
Some(v) => v,
None => "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",
};
/// Returns built-in OAuth credentials for a provider, keyed by secret_name.
///
/// The secret_name comes from the tool's capabilities.json `auth.secret_name` field.
/// Returns `None` if no built-in credentials are configured for that provider.
pub fn builtin_credentials(secret_name: &str) -> Option<OAuthCredentials> {
match secret_name {
"google_oauth_token" => Some(OAuthCredentials {
client_id: GOOGLE_CLIENT_ID,
client_secret: GOOGLE_CLIENT_SECRET,
}),
_ => None,
}
}
// ── Shared callback server ──────────────────────────────────────────────
/// Fixed port for all OAuth callbacks.
///
/// Every redirect URI registered with providers must use this port:
/// `http://localhost:9876/callback` (or `/auth/callback` for NEAR AI).
pub const OAUTH_CALLBACK_PORT: u16 = 9876;
/// Error from the OAuth callback listener.
#[derive(Debug, thiserror::Error)]
pub enum OAuthCallbackError {
#[error("Port {0} is in use (another auth flow running?): {1}")]
PortInUse(u16, String),
#[error("Authorization denied by user")]
Denied,
#[error("Timed out waiting for authorization")]
Timeout,
#[error("IO error: {0}")]
Io(String),
}
/// Bind the OAuth callback listener on the fixed port.
///
/// Tries IPv6 loopback (`[::1]`) first so that `http://localhost:…` redirects
/// work on systems where `localhost` resolves to `::1`. Falls back to IPv4
/// (`127.0.0.1`) only if IPv6 fails for a reason other than `AddrInUse`
/// (e.g., IPv6 not supported on the host). If the port is already occupied
/// on IPv6, the port is occupied period, so we fail immediately.
pub async fn bind_callback_listener() -> Result<TcpListener, OAuthCallbackError> {
let ipv6_addr = format!("[::1]:{}", OAUTH_CALLBACK_PORT);
match TcpListener::bind(&ipv6_addr).await {
Ok(listener) => return Ok(listener),
Err(e) if e.kind() == std::io::ErrorKind::AddrInUse => {
return Err(OAuthCallbackError::PortInUse(
OAUTH_CALLBACK_PORT,
e.to_string(),
));
}
Err(_) => {
// IPv6 not available on this host, fall back to IPv4
}
}
TcpListener::bind(format!("127.0.0.1:{}", OAUTH_CALLBACK_PORT))
.await
.map_err(|e| {
if e.kind() == std::io::ErrorKind::AddrInUse {
OAuthCallbackError::PortInUse(OAUTH_CALLBACK_PORT, e.to_string())
} else {
OAuthCallbackError::Io(e.to_string())
}
})
}
/// Wait for an OAuth callback and extract a query parameter value.
///
/// Listens for a GET request matching `path_prefix` (e.g., "/callback" or "/auth/callback"),
/// extracts the value of `param_name` (e.g., "code" or "token"), and shows a branded
/// landing page using `display_name` (e.g., "Google", "Notion", "NEAR AI").
///
/// Times out after 5 minutes.
pub async fn wait_for_callback(
listener: TcpListener,
path_prefix: &str,
param_name: &str,
display_name: &str,
) -> Result<String, OAuthCallbackError> {
let path_prefix = path_prefix.to_string();
let param_name = param_name.to_string();
let display_name = display_name.to_string();
tokio::time::timeout(Duration::from_secs(300), async move {
loop {
let (mut socket, _) = listener
.accept()
.await
.map_err(|e| OAuthCallbackError::Io(e.to_string()))?;
let mut reader = BufReader::new(&mut socket);
let mut request_line = String::new();
reader
.read_line(&mut request_line)
.await
.map_err(|e| OAuthCallbackError::Io(e.to_string()))?;
if let Some(path) = request_line.split_whitespace().nth(1)
&& path.starts_with(&path_prefix)
&& let Some(query) = path.split('?').nth(1)
{
// Check for error first
if query.contains("error=") {
let html = landing_html(&display_name, false);
let response = format!(
"HTTP/1.1 400 Bad Request\r\n\
Content-Type: text/html; charset=utf-8\r\n\
Connection: close\r\n\
\r\n\
{}",
html
);
let _ = socket.write_all(response.as_bytes()).await;
return Err(OAuthCallbackError::Denied);
}
// Look for the target parameter
for param in query.split('&') {
let parts: Vec<&str> = param.splitn(2, '=').collect();
if parts.len() == 2 && parts[0] == param_name {
let value = urlencoding::decode(parts[1])
.unwrap_or_else(|_| parts[1].into())
.into_owned();
let html = landing_html(&display_name, true);
let response = format!(
"HTTP/1.1 200 OK\r\n\
Content-Type: text/html; charset=utf-8\r\n\
Connection: close\r\n\
\r\n\
{}",
html
);
let _ = socket.write_all(response.as_bytes()).await;
let _ = socket.shutdown().await;
return Ok(value);
}
}
}
// Not the callback we're looking for
let response = "HTTP/1.1 404 Not Found\r\nConnection: close\r\n\r\n";
let _ = socket.write_all(response.as_bytes()).await;
}
})
.await
.map_err(|_| OAuthCallbackError::Timeout)?
}
/// Escape a string for safe interpolation into HTML content.
fn html_escape(s: &str) -> String {
let mut out = String::with_capacity(s.len());
for c in s.chars() {
match c {
'&' => out.push_str("&amp;"),
'<' => out.push_str("&lt;"),
'>' => out.push_str("&gt;"),
'"' => out.push_str("&quot;"),
'\'' => out.push_str("&#x27;"),
_ => out.push(c),
}
}
out
}
/// HTML landing page shown in the browser after an OAuth redirect.
pub fn landing_html(provider_name: &str, success: bool) -> String {
let safe_name = html_escape(provider_name);
let (icon, heading, subtitle, accent) = if success {
(
r##"<div style="width:64px;height:64px;border-radius:50%;background:#22c55e;display:flex;align-items:center;justify-content:center;margin:0 auto 24px">
<svg width="32" height="32" viewBox="0 0 24 24" fill="none" stroke="#fff" stroke-width="3" stroke-linecap="round" stroke-linejoin="round"><polyline points="20 6 9 17 4 12"/></svg>
</div>"##,
format!("{} Connected", safe_name),
"You can close this window and return to your terminal.",
"#22c55e",
)
} else {
(
r##"<div style="width:64px;height:64px;border-radius:50%;background:#ef4444;display:flex;align-items:center;justify-content:center;margin:0 auto 24px">
<svg width="32" height="32" viewBox="0 0 24 24" fill="none" stroke="#fff" stroke-width="3" stroke-linecap="round" stroke-linejoin="round"><line x1="18" y1="6" x2="6" y2="18"/><line x1="6" y1="6" x2="18" y2="18"/></svg>
</div>"##,
"Authorization Failed".to_string(),
"The request was denied. You can close this window and try again.",
"#ef4444",
)
};
format!(
r#"<!DOCTYPE html>
<html lang="en">
<head>
<meta charset="utf-8">
<meta name="viewport" content="width=device-width,initial-scale=1">
<title>IronClaw - {heading}</title>
<style>
* {{ margin:0; padding:0; box-sizing:border-box }}
body {{
font-family: -apple-system, BlinkMacSystemFont, "Segoe UI", Roboto, Helvetica, Arial, sans-serif;
background: #0a0a0a;
color: #e5e5e5;
display: flex;
justify-content: center;
align-items: center;
min-height: 100vh;
}}
.card {{
text-align: center;
padding: 48px 40px;
max-width: 420px;
border: 1px solid #262626;
border-radius: 16px;
background: #141414;
}}
h1 {{
font-size: 22px;
font-weight: 600;
margin-bottom: 8px;
color: #fafafa;
}}
p {{
font-size: 14px;
color: #a3a3a3;
line-height: 1.5;
}}
.accent {{ color: {accent}; }}
.brand {{
margin-top: 32px;
font-size: 12px;
color: #525252;
letter-spacing: 0.5px;
text-transform: uppercase;
}}
</style>
</head>
<body>
<div class="card">
{icon}
<h1>{heading}</h1>
<p>{subtitle}</p>
<div class="brand">IronClaw</div>
</div>
</body>
</html>"#,
heading = heading,
icon = icon,
subtitle = subtitle,
accent = accent,
)
}
#[cfg(test)]
mod tests {
use crate::cli::oauth_defaults::{builtin_credentials, landing_html};
#[test]
fn test_unknown_provider_returns_none() {
assert!(builtin_credentials("unknown_token").is_none());
}
#[test]
fn test_google_returns_based_on_compile_env() {
let creds = builtin_credentials("google_oauth_token");
assert!(creds.is_some());
let creds = creds.unwrap();
assert!(!creds.client_id.is_empty());
assert!(!creds.client_secret.is_empty());
}
#[test]
fn test_landing_html_success_contains_key_elements() {
let html = landing_html("Google", true);
assert!(html.contains("Google Connected"));
assert!(html.contains("charset"));
assert!(html.contains("IronClaw"));
assert!(html.contains("#22c55e")); // green accent
assert!(!html.contains("Failed"));
}
#[test]
fn test_landing_html_escapes_provider_name() {
let html = landing_html("<script>alert(1)</script>", true);
assert!(!html.contains("<script>"));
assert!(html.contains("&lt;script&gt;"));
}
#[test]
fn test_landing_html_error_contains_key_elements() {
let html = landing_html("Notion", false);
assert!(html.contains("Authorization Failed"));
assert!(html.contains("charset"));
assert!(html.contains("IronClaw"));
assert!(html.contains("#ef4444")); // red accent
assert!(!html.contains("Connected"));
}
}
+15 -17
View File
@@ -9,7 +9,7 @@ use crate::settings::Settings;
/// Run the status command, printing system health info.
pub async fn run_status_command() -> anyhow::Result<()> {
let settings = Settings::load();
let settings = Settings::default();
println!("IronClaw Status");
println!("===============\n");
@@ -22,10 +22,9 @@ pub async fn run_status_command() -> anyhow::Result<()> {
);
// Database
let db_url_set = settings.database_url.is_some() || std::env::var("DATABASE_URL").is_ok();
let db_url_set = std::env::var("DATABASE_URL").is_ok();
print!(" Database: ");
if db_url_set {
// Try to connect
match check_database().await {
Ok(()) => println!("connected"),
Err(e) => println!("error ({})", e),
@@ -43,13 +42,14 @@ pub async fn run_status_command() -> anyhow::Result<()> {
println!("not found (run `ironclaw onboard`)");
}
// Secrets
// Secrets (auto-detect: env var or keychain)
print!(" Secrets: ");
let secrets_configured = settings.secrets_master_key_source != crate::settings::KeySource::None
|| std::env::var("SECRETS_MASTER_KEY").is_ok()
|| crate::secrets::keychain::has_master_key().await;
if secrets_configured {
println!("configured ({:?})", settings.secrets_master_key_source);
let has_env_key = std::env::var("SECRETS_MASTER_KEY").is_ok();
let has_keychain = crate::secrets::keychain::has_master_key().await;
if has_env_key {
println!("configured (env)");
} else if has_keychain {
println!("configured (keychain)");
} else {
println!("not configured");
}
@@ -129,20 +129,18 @@ pub async fn run_status_command() -> anyhow::Result<()> {
Err(_) => println!("none configured"),
}
// Settings path
println!("\n Settings: {}", Settings::default_path().display());
// Config path
println!(
"\n Config: {}",
crate::bootstrap::ironclaw_env_path().display()
);
Ok(())
}
#[cfg(feature = "postgres")]
async fn check_database() -> anyhow::Result<()> {
let _ = dotenvy::dotenv();
let settings = Settings::load();
let url = std::env::var("DATABASE_URL")
.ok()
.or(settings.database_url)
.ok_or_else(|| anyhow::anyhow!("no URL"))?;
let url = std::env::var("DATABASE_URL").map_err(|_| anyhow::anyhow!("DATABASE_URL not set"))?;
let config: deadpool_postgres::Config = deadpool_postgres::Config {
url: Some(url),
+116 -91
View File
@@ -829,20 +829,73 @@ async fn auth_tool(name: String, dir: Option<PathBuf>, user_id: String) -> anyho
}
// Save the token
save_token(secrets_store.as_ref(), &user_id, &auth, &token).await?;
save_token(secrets_store.as_ref(), &user_id, &auth, &token, None, None).await?;
print_success(display_name);
return Ok(());
}
// Check for OAuth configuration
if let Some(ref oauth) = auth.oauth {
return auth_tool_oauth(secrets_store.as_ref(), &user_id, &auth, oauth).await;
// For providers with shared tokens (e.g., all Google tools share google_oauth_token),
// combine scopes from all installed tools so one auth covers everything.
let combined = combine_provider_scopes(&tools_dir, &auth.secret_name, oauth).await;
if combined.scopes.len() > oauth.scopes.len() {
let extra = combined.scopes.len() - oauth.scopes.len();
println!(
" Including scopes from {} other installed tool(s) sharing this credential.",
extra
);
println!();
}
return auth_tool_oauth(secrets_store.as_ref(), &user_id, &auth, &combined).await;
}
// Fall back to manual entry
auth_tool_manual(secrets_store.as_ref(), &user_id, &auth).await
}
/// Scan the tools directory for all capabilities files sharing the same secret_name
/// and combine their OAuth scopes. This way, authing any Google tool requests scopes
/// for ALL installed Google tools, so one login covers everything.
async fn combine_provider_scopes(
tools_dir: &Path,
secret_name: &str,
base_oauth: &crate::tools::wasm::OAuthConfigSchema,
) -> crate::tools::wasm::OAuthConfigSchema {
let mut all_scopes: std::collections::HashSet<String> =
base_oauth.scopes.iter().cloned().collect();
if let Ok(mut entries) = tokio::fs::read_dir(tools_dir).await {
while let Ok(Some(entry)) = entries.next_entry().await {
let path = entry.path();
if path.extension().and_then(|e| e.to_str()) != Some("json") {
continue;
}
let name = path
.file_name()
.and_then(|n| n.to_str())
.unwrap_or_default();
if !name.ends_with(".capabilities.json") {
continue;
}
if let Ok(content) = tokio::fs::read_to_string(&path).await
&& let Ok(caps) = CapabilitiesFile::from_json(&content)
&& let Some(auth) = &caps.auth
&& auth.secret_name == secret_name
&& let Some(oauth) = &auth.oauth
{
all_scopes.extend(oauth.scopes.iter().cloned());
}
}
}
let mut combined = base_oauth.clone();
combined.scopes = all_scopes.into_iter().collect();
combined.scopes.sort(); // deterministic ordering
combined
}
/// OAuth browser-based login flow.
async fn auth_tool_oauth(
store: &(dyn SecretsStore + Send + Sync),
@@ -853,12 +906,14 @@ async fn auth_tool_oauth(
use base64::{Engine, engine::general_purpose::URL_SAFE_NO_PAD};
use rand::RngCore;
use sha2::{Digest, Sha256};
use tokio::io::{AsyncBufReadExt, AsyncWriteExt, BufReader};
use tokio::net::TcpListener;
use crate::cli::oauth_defaults::{self, OAUTH_CALLBACK_PORT};
let display_name = auth.display_name.as_deref().unwrap_or(&auth.secret_name);
// Get client_id from config or env
// Get client_id: capabilities file > runtime env var > built-in defaults
let builtin = oauth_defaults::builtin_credentials(&auth.secret_name);
let client_id = oauth
.client_id
.clone()
@@ -868,41 +923,32 @@ async fn auth_tool_oauth(
.as_ref()
.and_then(|env| std::env::var(env).ok())
})
.or_else(|| builtin.as_ref().map(|c| c.client_id.to_string()))
.ok_or_else(|| {
anyhow::anyhow!(
"OAuth client_id not configured.\n\
Set it in the capabilities file or via environment variable."
Set {} env var, or build with IRONCLAW_GOOGLE_CLIENT_ID.",
oauth.client_id_env.as_deref().unwrap_or("the client_id")
)
})?;
// Get client_secret if provided
let client_secret = oauth.client_secret.clone().or_else(|| {
oauth
.client_secret_env
.as_ref()
.and_then(|env| std::env::var(env).ok())
});
// Get client_secret: capabilities file > runtime env var > built-in defaults
let client_secret = oauth
.client_secret
.clone()
.or_else(|| {
oauth
.client_secret_env
.as_ref()
.and_then(|env| std::env::var(env).ok())
})
.or_else(|| builtin.as_ref().map(|c| c.client_secret.to_string()));
println!(" Starting OAuth authentication...");
println!();
// Find an available port for the callback
let mut listener = None;
let mut port = 0;
for p in 9876..=9886 {
match TcpListener::bind(format!("127.0.0.1:{}", p)).await {
Ok(l) => {
listener = Some(l);
port = p;
break;
}
Err(_) => continue,
}
}
let listener = listener.ok_or_else(|| anyhow::anyhow!("Could not find available port"))?;
let redirect_uri = format!("http://localhost:{}/callback", port);
let listener = oauth_defaults::bind_callback_listener().await?;
let redirect_uri = format!("http://localhost:{}/callback", OAUTH_CALLBACK_PORT);
// Generate PKCE verifier and challenge
let (code_verifier, code_challenge) = if oauth.use_pkce {
@@ -961,63 +1007,8 @@ async fn auth_tool_oauth(
println!(" Waiting for authorization...");
// Wait for callback with timeout
let timeout = std::time::Duration::from_secs(300);
let code = tokio::time::timeout(timeout, async {
loop {
let (mut socket, _) = listener.accept().await?;
let mut reader = BufReader::new(&mut socket);
let mut request_line = String::new();
reader.read_line(&mut request_line).await?;
// Parse GET /callback?code=xxx HTTP/1.1
if let Some(path) = request_line.split_whitespace().nth(1)
&& path.starts_with("/callback")
&& let Some(query) = path.split('?').nth(1) {
for param in query.split('&') {
let parts: Vec<&str> = param.splitn(2, '=').collect();
if parts.len() == 2 && parts[0] == "code" {
let code = urlencoding::decode(parts[1])
.unwrap_or_else(|_| parts[1].into())
.into_owned();
// Send success response
let response = format!(
"HTTP/1.1 200 OK\r\n\
Content-Type: text/html\r\n\
\r\n\
<!DOCTYPE html><html><body style=\"font-family: sans-serif; \
display: flex; justify-content: center; align-items: center; \
height: 100vh; margin: 0; background: #191919; color: white;\">\
<div style=\"text-align: center;\">\
<h1>✓ {} Connected!</h1>\
<p>You can close this window.</p>\
</div></body></html>",
display_name
);
let _ = socket.write_all(response.as_bytes()).await;
let _ = socket.shutdown().await;
return Ok::<_, anyhow::Error>(code);
}
}
// Check for error
if query.contains("error=") {
let response =
"HTTP/1.1 400 Bad Request\r\n\r\nAuthorization denied";
let _ = socket.write_all(response.as_bytes()).await;
return Err(anyhow::anyhow!("Authorization denied by user"));
}
}
let response = "HTTP/1.1 404 Not Found\r\n\r\n";
let _ = socket.write_all(response.as_bytes()).await;
}
})
.await
.map_err(|_| anyhow::anyhow!("Timed out waiting for authorization"))??;
let code =
oauth_defaults::wait_for_callback(listener, "/callback", "code", display_name).await?;
println!();
println!(" Exchanging code for token...");
@@ -1068,8 +1059,19 @@ async fn auth_tool_oauth(
)
})?;
// Save the token
save_token(store, user_id, auth, access_token).await?;
let refresh_token = token_data.get("refresh_token").and_then(|v| v.as_str());
let expires_in = token_data.get("expires_in").and_then(|v| v.as_u64());
// Save the token (with refresh token and expiry if provided)
save_token(
store,
user_id,
auth,
access_token,
refresh_token,
expires_in,
)
.await?;
// Extract any additional info for display
let workspace_name = token_data
@@ -1171,8 +1173,8 @@ async fn auth_tool_manual(
}
}
// Save the token
save_token(store, user_id, auth, &token).await?;
// Save the token (manual path: no refresh token or expiry)
save_token(store, user_id, auth, &token, None, None).await?;
print_success(display_name);
Ok(())
}
@@ -1263,11 +1265,16 @@ async fn validate_token(
}
/// Save token to secrets store.
///
/// Optionally stores a refresh token (as `{secret_name}_refresh_token`) and
/// sets `expires_at` on the access token so the runtime can auto-refresh.
async fn save_token(
store: &(dyn SecretsStore + Send + Sync),
user_id: &str,
auth: &crate::tools::wasm::AuthCapabilitySchema,
token: &str,
refresh_token: Option<&str>,
expires_in: Option<u64>,
) -> anyhow::Result<()> {
let mut params = CreateSecretParams::new(&auth.secret_name, token);
@@ -1275,11 +1282,29 @@ async fn save_token(
params = params.with_provider(provider);
}
if let Some(secs) = expires_in {
let expires_at = chrono::Utc::now() + chrono::Duration::seconds(secs as i64);
params = params.with_expiry(expires_at);
}
store
.create(user_id, params)
.await
.map_err(|e| anyhow::anyhow!("Failed to save token: {}", e))?;
// Store refresh token separately (no expiry, it's long-lived)
if let Some(rt) = refresh_token {
let refresh_name = format!("{}_refresh_token", auth.secret_name);
let mut refresh_params = CreateSecretParams::new(&refresh_name, rt);
if let Some(ref provider) = auth.provider {
refresh_params = refresh_params.with_provider(provider);
}
store
.create(user_id, refresh_params)
.await
.map_err(|e| anyhow::anyhow!("Failed to save refresh token: {}", e))?;
}
Ok(())
}
+47 -88
View File
@@ -1,9 +1,9 @@
//! Configuration for IronClaw.
//!
//! Settings are loaded with priority: env var > database > default.
//! The database replaces the old `settings.json` file for all settings
//! except the 4 bootstrap fields (database_url, pool_size, secrets key
//! source, onboard_completed) which live in `~/.ironclaw/bootstrap.json`.
//! `DATABASE_URL` lives in `~/.ironclaw/.env` (loaded via dotenvy early
//! in startup). Everything else comes from env vars, the DB settings
//! table, or auto-detection.
use std::path::PathBuf;
use std::time::Duration;
@@ -40,9 +40,9 @@ impl Config {
pub async fn from_db(
store: &dyn crate::db::Database,
user_id: &str,
bootstrap: &crate::bootstrap::BootstrapConfig,
) -> Result<Self, ConfigError> {
let _ = dotenvy::dotenv();
crate::bootstrap::load_ironclaw_env();
// Load all settings from DB into a Settings struct
let db_settings = match store.get_all_settings(user_id).await {
@@ -53,7 +53,7 @@ impl Config {
}
};
Self::build(bootstrap, &db_settings).await
Self::build(&db_settings).await
}
/// Load configuration from environment variables only (no database).
@@ -61,20 +61,20 @@ impl Config {
/// Used during early startup before the database is connected,
/// and by CLI commands that don't have DB access.
/// Falls back to legacy `settings.json` on disk if present.
///
/// Loads both `./.env` (standard, higher priority) and `~/.ironclaw/.env`
/// (lower priority) via dotenvy, which never overwrites existing vars.
pub async fn from_env() -> Result<Self, ConfigError> {
let _ = dotenvy::dotenv();
let bootstrap = crate::bootstrap::BootstrapConfig::load();
crate::bootstrap::load_ironclaw_env();
let settings = Settings::load();
Self::build(&bootstrap, &settings).await
Self::build(&settings).await
}
/// Build config from bootstrap + settings (shared by from_env and from_db).
async fn build(
bootstrap: &crate::bootstrap::BootstrapConfig,
settings: &Settings,
) -> Result<Self, ConfigError> {
/// Build config from settings (shared by from_env and from_db).
async fn build(settings: &Settings) -> Result<Self, ConfigError> {
Ok(Self {
database: DatabaseConfig::resolve(bootstrap, settings)?,
database: DatabaseConfig::resolve()?,
llm: LlmConfig::resolve(settings)?,
embeddings: EmbeddingsConfig::resolve(settings)?,
tunnel: TunnelConfig::resolve(settings)?,
@@ -82,7 +82,7 @@ impl Config {
agent: AgentConfig::resolve(settings)?,
safety: SafetyConfig::resolve()?,
wasm: WasmConfig::resolve()?,
secrets: SecretsConfig::resolve(bootstrap).await?,
secrets: SecretsConfig::resolve().await?,
builder: BuilderModeConfig::resolve()?,
heartbeat: HeartbeatConfig::resolve(settings)?,
routines: RoutineConfig::resolve()?,
@@ -179,36 +179,20 @@ pub struct DatabaseConfig {
}
impl DatabaseConfig {
fn resolve(
bootstrap: &crate::bootstrap::BootstrapConfig,
settings: &Settings,
) -> Result<Self, ConfigError> {
// Priority: env var > settings > default
fn resolve() -> Result<Self, ConfigError> {
let backend: DatabaseBackend = if let Some(b) = optional_env("DATABASE_BACKEND")? {
b.parse().map_err(|e| ConfigError::InvalidValue {
key: "DATABASE_BACKEND".to_string(),
message: e,
})?
} else if let Some(ref b) = settings.database_backend {
match b.parse() {
Ok(backend) => backend,
Err(e) => {
tracing::warn!(
"Invalid database_backend '{}' in settings: {}. Using default.",
b,
e
);
DatabaseBackend::default()
}
}
} else {
DatabaseBackend::default()
};
// PostgreSQL URL is required only when using the postgres backend.
// For libsql backend, default to an empty placeholder.
// DATABASE_URL is loaded from ~/.ironclaw/.env via dotenvy early in startup.
let url = optional_env("DATABASE_URL")?
.or_else(|| bootstrap.database_url.clone())
.or_else(|| {
if backend == DatabaseBackend::LibSql {
Some("unused://libsql".to_string())
@@ -221,29 +205,17 @@ impl DatabaseConfig {
hint: "Run 'ironclaw onboard' or set DATABASE_URL environment variable".to_string(),
})?;
let pool_size = optional_env("DATABASE_POOL_SIZE")?
.map(|s| s.parse())
.transpose()
.map_err(|e| ConfigError::InvalidValue {
key: "DATABASE_POOL_SIZE".to_string(),
message: format!("must be a positive integer: {e}"),
})?
.or(bootstrap.database_pool_size)
.unwrap_or(10);
let pool_size = parse_optional_env("DATABASE_POOL_SIZE", 10)?;
// Priority: env var > settings > default (if libsql backend)
let libsql_path = optional_env("LIBSQL_PATH")?
.map(PathBuf::from)
.or_else(|| settings.libsql_path.as_ref().map(PathBuf::from))
.or_else(|| {
if backend == DatabaseBackend::LibSql {
Some(default_libsql_path())
} else {
None
}
});
let libsql_path = optional_env("LIBSQL_PATH")?.map(PathBuf::from).or_else(|| {
if backend == DatabaseBackend::LibSql {
Some(default_libsql_path())
} else {
None
}
});
let libsql_url = optional_env("LIBSQL_URL")?.or_else(|| settings.libsql_url.clone());
let libsql_url = optional_env("LIBSQL_URL")?;
let libsql_auth_token = optional_env("LIBSQL_AUTH_TOKEN")?.map(SecretString::from);
if libsql_url.is_some() && libsql_auth_token.is_none() {
@@ -417,6 +389,15 @@ pub struct NearAiConfig {
pub api_mode: NearAiApiMode,
/// API key for cloud-api (required for chat_completions mode)
pub api_key: Option<SecretString>,
/// Optional fallback model for failover (default: None).
/// When set, a secondary provider is created with this model and wrapped
/// in a `FailoverProvider` so transient errors on the primary model
/// automatically fall through to the fallback.
pub fallback_model: Option<String>,
/// Maximum number of retries for transient errors (default: 3).
/// With the default of 3, the provider makes up to 4 total attempts
/// (1 initial + 3 retries) before giving up.
pub max_retries: u32,
}
impl LlmConfig {
@@ -473,6 +454,8 @@ impl LlmConfig {
.unwrap_or_else(default_session_path),
api_mode,
api_key: nearai_api_key,
fallback_model: optional_env("NEARAI_FALLBACK_MODEL")?,
max_retries: parse_optional_env("NEARAI_MAX_RETRIES", 3)?,
};
// Resolve provider-specific configs based on backend
@@ -891,49 +874,25 @@ impl std::fmt::Debug for SecretsConfig {
/// Avoids re-prompting the OS keychain on every `SecretsConfig::resolve()` call
/// (e.g. `Config::from_env()` then `Config::from_db()`). Thread-safe alternative
/// to caching in a process env var.
static CACHED_KEYCHAIN_KEY: std::sync::OnceLock<String> = std::sync::OnceLock::new();
impl SecretsConfig {
async fn resolve(bootstrap: &crate::bootstrap::BootstrapConfig) -> Result<Self, ConfigError> {
/// Auto-detect secrets master key from env var, then OS keychain.
///
/// Sequential probe: SECRETS_MASTER_KEY env var first, then OS keychain.
/// No saved "source" needed; just try each source in order.
async fn resolve() -> Result<Self, ConfigError> {
use crate::settings::KeySource;
let (master_key, source) = if let Some(env_key) = optional_env("SECRETS_MASTER_KEY")? {
(Some(SecretString::from(env_key)), KeySource::Env)
} else {
match bootstrap.secrets_master_key_source {
KeySource::Keychain => {
// Check process-level cache first (set on previous resolve() call)
if let Some(cached) = CACHED_KEYCHAIN_KEY.get() {
(
Some(SecretString::from(cached.clone())),
KeySource::Keychain,
)
} else {
// Try to load from OS keychain (async on Linux)
match crate::secrets::keychain::get_master_key().await {
Ok(key_bytes) => {
let key_hex: String =
key_bytes.iter().map(|b| format!("{b:02x}")).collect();
let _ = CACHED_KEYCHAIN_KEY.set(key_hex.clone());
(Some(SecretString::from(key_hex)), KeySource::Keychain)
}
Err(_) => {
tracing::warn!(
"Secrets configured for keychain but key not found. \
Run 'ironclaw onboard' to reconfigure."
);
(None, KeySource::None)
}
}
}
// Probe the OS keychain; if a key is stored, use it
match crate::secrets::keychain::get_master_key().await {
Ok(key_bytes) => {
let key_hex: String = key_bytes.iter().map(|b| format!("{:02x}", b)).collect();
(Some(SecretString::from(key_hex)), KeySource::Keychain)
}
KeySource::Env => {
tracing::warn!(
"Secrets configured for env var but SECRETS_MASTER_KEY not set."
);
(None, KeySource::None)
}
KeySource::None => (None, KeySource::None),
Err(_) => (None, KeySource::None),
}
};
+5
View File
@@ -36,6 +36,11 @@ pub struct Store {
#[cfg(feature = "postgres")]
impl Store {
/// Wrap an existing pool (useful when the caller already has a connection).
pub fn from_pool(pool: Pool) -> Self {
Self { pool }
}
/// Create a new store and connect to the database.
pub async fn new(config: &DatabaseConfig) -> Result<Self, DatabaseError> {
let mut cfg = Config::new();
+483
View File
@@ -0,0 +1,483 @@
//! Multi-provider LLM failover.
//!
//! Wraps multiple LlmProvider instances and tries each in sequence
//! until one succeeds. Transparent to callers --- same LlmProvider trait.
use std::future::Future;
use std::sync::Arc;
use std::sync::atomic::{AtomicUsize, Ordering};
use async_trait::async_trait;
use rust_decimal::Decimal;
use crate::error::LlmError;
use crate::llm::provider::{
CompletionRequest, CompletionResponse, LlmProvider, ToolCompletionRequest,
ToolCompletionResponse,
};
/// Returns `true` if the error is transient and the request should be retried
/// on the next provider in the failover chain.
///
/// Retryable: `RequestFailed`, `RateLimited`, `InvalidResponse`,
/// `SessionRenewalFailed`, `ModelNotAvailable`, `Http`, `Io`.
///
/// `ModelNotAvailable` is retryable because the next provider in the chain may
/// offer a different model, so it's worth trying.
///
/// Non-retryable errors (`AuthFailed`, `SessionExpired`, `ContextLengthExceeded`)
/// propagate immediately because a different provider won't fix them.
fn is_retryable(err: &LlmError) -> bool {
matches!(
err,
LlmError::RequestFailed { .. }
| LlmError::RateLimited { .. }
| LlmError::InvalidResponse { .. }
| LlmError::SessionRenewalFailed { .. }
// ModelNotAvailable is retryable: the next provider may offer a different model.
| LlmError::ModelNotAvailable { .. }
| LlmError::Http(_)
| LlmError::Io(_)
)
}
/// An LLM provider that wraps multiple providers and tries each in sequence
/// on transient failures.
///
/// The first provider in the list is the primary. If it fails with a retryable
/// error, the next provider is tried, and so on. Non-retryable errors
/// (e.g. `AuthFailed`, `ContextLengthExceeded`) propagate immediately.
pub struct FailoverProvider {
providers: Vec<Arc<dyn LlmProvider>>,
/// Index of the provider that last handled a request successfully.
/// Used by `model_name()` and `cost_per_token()` so downstream cost
/// tracking reflects the provider that actually served the request.
last_used: AtomicUsize,
}
impl FailoverProvider {
/// Create a new failover provider.
///
/// Returns an error if `providers` is empty.
pub fn new(providers: Vec<Arc<dyn LlmProvider>>) -> Result<Self, LlmError> {
if providers.is_empty() {
return Err(LlmError::RequestFailed {
provider: "failover".to_string(),
reason: "FailoverProvider requires at least one provider".to_string(),
});
}
Ok(Self {
providers,
last_used: AtomicUsize::new(0),
})
}
/// Try each provider in sequence until one succeeds or all fail.
async fn try_providers<T, F, Fut>(&self, mut call: F) -> Result<T, LlmError>
where
F: FnMut(Arc<dyn LlmProvider>) -> Fut,
Fut: Future<Output = Result<T, LlmError>>,
{
let mut last_error: Option<LlmError> = None;
for (i, provider) in self.providers.iter().enumerate() {
let result = call(Arc::clone(provider)).await;
match result {
Ok(response) => {
self.last_used.store(i, Ordering::Relaxed);
return Ok(response);
}
Err(err) => {
if !is_retryable(&err) {
return Err(err);
}
if i + 1 < self.providers.len() {
tracing::warn!(
provider = %provider.model_name(),
error = %err,
next_provider = %self.providers[i + 1].model_name(),
"Provider failed with retryable error, trying next provider"
);
}
last_error = Some(err);
}
}
}
// SAFETY: providers is non-empty (checked in `new`), so at least one
// iteration ran and `last_error` is `Some`.
Err(last_error.expect("providers list is non-empty"))
}
}
#[async_trait]
impl LlmProvider for FailoverProvider {
fn model_name(&self) -> &str {
self.providers[self.last_used.load(Ordering::Relaxed)].model_name()
}
fn cost_per_token(&self) -> (Decimal, Decimal) {
self.providers[self.last_used.load(Ordering::Relaxed)].cost_per_token()
}
async fn complete(&self, request: CompletionRequest) -> Result<CompletionResponse, LlmError> {
self.try_providers(|provider| {
let req = request.clone();
async move { provider.complete(req).await }
})
.await
}
async fn complete_with_tools(
&self,
request: ToolCompletionRequest,
) -> Result<ToolCompletionResponse, LlmError> {
self.try_providers(|provider| {
let req = request.clone();
async move { provider.complete_with_tools(req).await }
})
.await
}
async fn list_models(&self) -> Result<Vec<String>, LlmError> {
let mut all_models = Vec::new();
for provider in &self.providers {
match provider.list_models().await {
Ok(models) => all_models.extend(models),
Err(err) => {
tracing::warn!(
provider = %provider.model_name(),
error = %err,
"Failed to list models from provider, skipping"
);
}
}
}
all_models.sort();
all_models.dedup();
Ok(all_models)
}
}
#[cfg(test)]
mod tests {
use super::*;
use std::sync::Mutex;
use std::time::Duration;
use crate::llm::provider::{CompletionResponse, FinishReason, ToolCompletionResponse};
/// A mock LLM provider that returns a predetermined result.
struct MockProvider {
name: String,
input_cost: Decimal,
output_cost: Decimal,
complete_result: Mutex<Option<Result<CompletionResponse, LlmError>>>,
tool_complete_result: Mutex<Option<Result<ToolCompletionResponse, LlmError>>>,
}
impl MockProvider {
fn succeeding(name: &str, content: &str) -> Self {
Self {
name: name.to_string(),
input_cost: Decimal::ZERO,
output_cost: Decimal::ZERO,
complete_result: Mutex::new(Some(Ok(CompletionResponse {
content: content.to_string(),
input_tokens: 10,
output_tokens: 5,
finish_reason: FinishReason::Stop,
response_id: None,
}))),
tool_complete_result: Mutex::new(Some(Ok(ToolCompletionResponse {
content: Some(content.to_string()),
tool_calls: vec![],
input_tokens: 10,
output_tokens: 5,
finish_reason: FinishReason::Stop,
response_id: None,
}))),
}
}
fn succeeding_with_cost(
name: &str,
content: &str,
input_cost: Decimal,
output_cost: Decimal,
) -> Self {
Self {
input_cost,
output_cost,
..Self::succeeding(name, content)
}
}
fn failing_retryable(name: &str) -> Self {
Self {
name: name.to_string(),
input_cost: Decimal::ZERO,
output_cost: Decimal::ZERO,
complete_result: Mutex::new(Some(Err(LlmError::RequestFailed {
provider: name.to_string(),
reason: "server error".to_string(),
}))),
tool_complete_result: Mutex::new(Some(Err(LlmError::RequestFailed {
provider: name.to_string(),
reason: "server error".to_string(),
}))),
}
}
fn failing_non_retryable(name: &str) -> Self {
Self {
name: name.to_string(),
input_cost: Decimal::ZERO,
output_cost: Decimal::ZERO,
complete_result: Mutex::new(Some(Err(LlmError::AuthFailed {
provider: name.to_string(),
}))),
tool_complete_result: Mutex::new(Some(Err(LlmError::AuthFailed {
provider: name.to_string(),
}))),
}
}
fn failing_rate_limited(name: &str) -> Self {
Self {
name: name.to_string(),
input_cost: Decimal::ZERO,
output_cost: Decimal::ZERO,
complete_result: Mutex::new(Some(Err(LlmError::RateLimited {
provider: name.to_string(),
retry_after: Some(Duration::from_secs(30)),
}))),
tool_complete_result: Mutex::new(Some(Err(LlmError::RateLimited {
provider: name.to_string(),
retry_after: Some(Duration::from_secs(30)),
}))),
}
}
}
#[async_trait]
impl LlmProvider for MockProvider {
fn model_name(&self) -> &str {
&self.name
}
fn cost_per_token(&self) -> (Decimal, Decimal) {
(self.input_cost, self.output_cost)
}
async fn complete(
&self,
_request: CompletionRequest,
) -> Result<CompletionResponse, LlmError> {
self.complete_result
.lock()
.unwrap()
.take()
.expect("MockProvider::complete called more than once")
}
async fn complete_with_tools(
&self,
_request: ToolCompletionRequest,
) -> Result<ToolCompletionResponse, LlmError> {
self.tool_complete_result
.lock()
.unwrap()
.take()
.expect("MockProvider::complete_with_tools called more than once")
}
async fn list_models(&self) -> Result<Vec<String>, LlmError> {
Ok(vec![self.name.clone()])
}
}
fn make_request() -> CompletionRequest {
CompletionRequest::new(vec![crate::llm::ChatMessage::user("hello")])
}
fn make_tool_request() -> ToolCompletionRequest {
ToolCompletionRequest::new(vec![crate::llm::ChatMessage::user("hello")], vec![])
}
// Test 1: Primary succeeds, no failover occurs.
#[tokio::test]
async fn primary_succeeds_no_failover() {
let primary = Arc::new(MockProvider::succeeding("primary", "primary response"));
let fallback = Arc::new(MockProvider::succeeding("fallback", "fallback response"));
let failover = FailoverProvider::new(vec![primary, fallback]).unwrap();
let response = failover.complete(make_request()).await.unwrap();
assert_eq!(response.content, "primary response");
}
// Test 2: Primary fails with retryable error, fallback succeeds.
#[tokio::test]
async fn primary_fails_retryable_fallback_succeeds() {
let primary = Arc::new(MockProvider::failing_retryable("primary"));
let fallback = Arc::new(MockProvider::succeeding("fallback", "fallback response"));
let failover = FailoverProvider::new(vec![primary, fallback]).unwrap();
let response = failover.complete(make_request()).await.unwrap();
assert_eq!(response.content, "fallback response");
}
// Test 3: All providers fail, returns last error.
#[tokio::test]
async fn all_providers_fail_returns_last_error() {
let primary = Arc::new(MockProvider::failing_retryable("primary"));
let fallback = Arc::new(MockProvider::failing_retryable("fallback"));
let failover = FailoverProvider::new(vec![primary, fallback]).unwrap();
let err = failover.complete(make_request()).await.unwrap_err();
match err {
LlmError::RequestFailed { provider, .. } => {
assert_eq!(provider, "fallback");
}
other => panic!("expected RequestFailed, got: {other:?}"),
}
}
// Test 4: Non-retryable error fails immediately, no failover.
#[tokio::test]
async fn non_retryable_error_fails_immediately() {
let primary = Arc::new(MockProvider::failing_non_retryable("primary"));
let fallback = Arc::new(MockProvider::succeeding("fallback", "fallback response"));
let failover = FailoverProvider::new(vec![primary, fallback]).unwrap();
let err = failover.complete(make_request()).await.unwrap_err();
match err {
LlmError::AuthFailed { provider } => {
assert_eq!(provider, "primary");
}
other => panic!("expected AuthFailed, got: {other:?}"),
}
}
// Test 5: Three providers, first two fail (retryable), third succeeds.
#[tokio::test]
async fn three_providers_first_two_fail_third_succeeds() {
let p1 = Arc::new(MockProvider::failing_retryable("provider-1"));
let p2 = Arc::new(MockProvider::failing_rate_limited("provider-2"));
let p3 = Arc::new(MockProvider::succeeding("provider-3", "third time lucky"));
let failover = FailoverProvider::new(vec![p1, p2, p3]).unwrap();
let response = failover.complete(make_request()).await.unwrap();
assert_eq!(response.content, "third time lucky");
}
// Test: complete_with_tools follows same failover logic.
#[tokio::test]
async fn complete_with_tools_failover() {
let primary = Arc::new(MockProvider::failing_retryable("primary"));
let fallback = Arc::new(MockProvider::succeeding("fallback", "tools fallback"));
let failover = FailoverProvider::new(vec![primary, fallback]).unwrap();
let response = failover
.complete_with_tools(make_tool_request())
.await
.unwrap();
assert_eq!(response.content.as_deref(), Some("tools fallback"));
}
// Test: model_name and cost_per_token reflect the last-used provider.
#[tokio::test]
async fn model_name_and_cost_track_last_used_provider() {
let fallback_cost = Decimal::new(15, 6); // 0.000015
let primary = Arc::new(MockProvider::failing_retryable("primary-model"));
let fallback = Arc::new(MockProvider::succeeding_with_cost(
"fallback-model",
"ok",
fallback_cost,
fallback_cost,
));
let failover = FailoverProvider::new(vec![primary, fallback]).unwrap();
// Before any call, defaults to primary (index 0).
assert_eq!(failover.model_name(), "primary-model");
assert_eq!(failover.cost_per_token(), (Decimal::ZERO, Decimal::ZERO));
// After failover, should reflect the fallback provider.
let _ = failover.complete(make_request()).await.unwrap();
assert_eq!(failover.model_name(), "fallback-model");
assert_eq!(failover.cost_per_token(), (fallback_cost, fallback_cost));
}
// Test: list_models aggregates from all providers.
#[tokio::test]
async fn list_models_aggregates_all() {
let p1 = Arc::new(MockProvider::succeeding("model-a", "ok"));
let p2 = Arc::new(MockProvider::succeeding("model-b", "ok"));
let failover = FailoverProvider::new(vec![p1, p2]).unwrap();
let models = failover.list_models().await.unwrap();
assert!(models.contains(&"model-a".to_string()));
assert!(models.contains(&"model-b".to_string()));
}
// Test: is_retryable correctly classifies errors.
#[test]
fn retryable_classification() {
// Retryable
assert!(is_retryable(&LlmError::RequestFailed {
provider: "p".into(),
reason: "err".into(),
}));
assert!(is_retryable(&LlmError::RateLimited {
provider: "p".into(),
retry_after: None,
}));
assert!(is_retryable(&LlmError::InvalidResponse {
provider: "p".into(),
reason: "bad json".into(),
}));
assert!(is_retryable(&LlmError::SessionRenewalFailed {
provider: "p".into(),
reason: "timeout".into(),
}));
assert!(is_retryable(&LlmError::Io(std::io::Error::new(
std::io::ErrorKind::ConnectionReset,
"reset"
))));
assert!(is_retryable(&LlmError::ModelNotAvailable {
provider: "p".into(),
model: "m".into(),
}));
// Non-retryable
assert!(!is_retryable(&LlmError::AuthFailed {
provider: "p".into(),
}));
assert!(!is_retryable(&LlmError::SessionExpired {
provider: "p".into(),
}));
assert!(!is_retryable(&LlmError::ContextLengthExceeded {
used: 100_000,
limit: 50_000,
}));
}
// Test: empty providers list returns error (not panic).
#[test]
fn empty_providers_returns_error() {
let result = FailoverProvider::new(vec![]);
assert!(result.is_err());
}
}
+22 -12
View File
@@ -8,13 +8,16 @@
//! - **OpenAI-compatible**: Any endpoint that speaks the OpenAI API
mod costs;
pub mod failover;
mod nearai;
mod nearai_chat;
mod provider;
mod reasoning;
mod retry;
mod rig_adapter;
pub mod session;
pub use failover::FailoverProvider;
pub use nearai::{ModelInfo, NearAiProvider};
pub use nearai_chat::NearAiChatProvider;
pub use provider::{
@@ -33,7 +36,7 @@ use std::sync::Arc;
use rig::client::CompletionClient;
use secrecy::ExposeSecret;
use crate::config::{LlmBackend, LlmConfig, NearAiApiMode};
use crate::config::{LlmBackend, LlmConfig, NearAiApiMode, NearAiConfig};
use crate::error::LlmError;
/// Create an LLM provider based on configuration.
@@ -46,7 +49,7 @@ pub fn create_llm_provider(
session: Arc<SessionManager>,
) -> Result<Arc<dyn LlmProvider>, LlmError> {
match config.backend {
LlmBackend::NearAi => create_nearai_provider(config, session),
LlmBackend::NearAi => create_llm_provider_with_config(&config.nearai, session),
LlmBackend::OpenAi => create_openai_provider(config),
LlmBackend::Anthropic => create_anthropic_provider(config),
LlmBackend::Ollama => create_ollama_provider(config),
@@ -54,21 +57,28 @@ pub fn create_llm_provider(
}
}
fn create_nearai_provider(
config: &LlmConfig,
/// Create an LLM provider from a `NearAiConfig` directly.
///
/// This is useful when constructing additional providers for failover,
/// where only the model name differs from the primary config.
pub fn create_llm_provider_with_config(
config: &NearAiConfig,
session: Arc<SessionManager>,
) -> Result<Arc<dyn LlmProvider>, LlmError> {
match config.nearai.api_mode {
match config.api_mode {
NearAiApiMode::Responses => {
tracing::info!("Using NEAR AI Responses API (chat-api) with session auth");
Ok(Arc::new(NearAiProvider::new(
config.nearai.clone(),
session,
)))
tracing::info!(
model = %config.model,
"Using Responses API (chat-api) with session auth"
);
Ok(Arc::new(NearAiProvider::new(config.clone(), session)))
}
NearAiApiMode::ChatCompletions => {
tracing::info!("Using NEAR AI Chat Completions API (cloud-api) with API key auth");
Ok(Arc::new(NearAiChatProvider::new(config.nearai.clone())?))
tracing::info!(
model = %config.model,
"Using Chat Completions API (cloud-api) with API key auth"
);
Ok(Arc::new(NearAiChatProvider::new(config.clone())?))
}
}
}
+108 -66
View File
@@ -19,6 +19,7 @@ use crate::llm::provider::{
ChatMessage, CompletionRequest, CompletionResponse, FinishReason, LlmProvider, Role, ToolCall,
ToolCompletionRequest, ToolCompletionResponse,
};
use crate::llm::retry::{is_retryable_status, retry_backoff_delay};
use crate::llm::session::SessionManager;
/// Information about an available model from NEAR AI API.
@@ -270,88 +271,139 @@ impl NearAiProvider {
}
}
/// Inner request implementation without retry logic.
/// Inner request implementation with retry logic for transient errors.
///
/// Retries on HTTP 429, 500, 502, 503, 504 with exponential backoff.
/// Does not retry on client errors (400, 401, 403, 404) or parse errors.
async fn send_request_inner<T: Serialize + std::fmt::Debug, R: for<'de> Deserialize<'de>>(
&self,
path: &str,
body: &T,
) -> Result<R, LlmError> {
let url = self.api_url(path);
let token = self.session.get_token().await?;
let max_retries = self.config.max_retries;
tracing::debug!("Sending request to NEAR AI: {}", url);
tracing::debug!("Request body: {:?}", body);
for attempt in 0..=max_retries {
let token = self.session.get_token().await?;
let response = self
.client
.post(&url)
.header("Authorization", format!("Bearer {}", token.expose_secret()))
.header("Content-Type", "application/json")
.json(body)
.send()
.await
.map_err(|e| {
tracing::error!("NEAR AI request failed: {}", e);
e
})?;
tracing::debug!(
"Sending request to NEAR AI: {} (attempt {})",
url,
attempt + 1
);
tracing::debug!("Request body: {:?}", body);
let status = response.status();
let response_text = response.text().await.unwrap_or_default();
let response = self
.client
.post(&url)
.header("Authorization", format!("Bearer {}", token.expose_secret()))
.header("Content-Type", "application/json")
.json(body)
.send()
.await;
tracing::debug!("NEAR AI response status: {}", status);
tracing::debug!("NEAR AI response body: {}", response_text);
let response = match response {
Ok(r) => r,
Err(e) => {
tracing::error!("NEAR AI request failed: {}", e);
// Network errors (timeout, connection refused) are transient
if attempt < max_retries {
let delay = retry_backoff_delay(attempt);
tracing::warn!(
"NEAR AI request error (attempt {}/{}), retrying in {:?}: {}",
attempt + 1,
max_retries + 1,
delay,
e,
);
tokio::time::sleep(delay).await;
continue;
}
return Err(e.into());
}
};
if !status.is_success() {
// Check for session expiration (401 with specific message patterns)
if status.as_u16() == 401 {
let is_session_expired = response_text.to_lowercase().contains("session")
&& (response_text.to_lowercase().contains("expired")
|| response_text.to_lowercase().contains("invalid"));
let status = response.status();
let response_text = response.text().await.unwrap_or_default();
if is_session_expired {
return Err(LlmError::SessionExpired {
tracing::debug!("NEAR AI response status: {}", status);
tracing::debug!("NEAR AI response body: {}", response_text);
if !status.is_success() {
let status_code = status.as_u16();
// Check for session expiration (401 with specific message patterns)
if status_code == 401 {
let lower = response_text.to_lowercase();
let is_session_expired = lower.contains("session")
&& (lower.contains("expired") || lower.contains("invalid"));
if is_session_expired {
return Err(LlmError::SessionExpired {
provider: "nearai".to_string(),
});
}
// Generic 401 -- not retryable
return Err(LlmError::AuthFailed {
provider: "nearai".to_string(),
});
}
// Generic 401 without session expiration indication
return Err(LlmError::AuthFailed {
provider: "nearai".to_string(),
});
}
// Check if this is a transient error worth retrying
if is_retryable_status(status_code) && attempt < max_retries {
let delay = retry_backoff_delay(attempt);
tracing::warn!(
"NEAR AI returned HTTP {} (attempt {}/{}), retrying in {:?}",
status_code,
attempt + 1,
max_retries + 1,
delay,
);
tokio::time::sleep(delay).await;
continue;
}
// Try to parse as JSON error
if let Ok(error) = serde_json::from_str::<NearAiErrorResponse>(&response_text) {
if status.as_u16() == 429 {
return Err(LlmError::RateLimited {
// Non-retryable error or exhausted retries
if let Ok(error) = serde_json::from_str::<NearAiErrorResponse>(&response_text) {
if status_code == 429 {
return Err(LlmError::RateLimited {
provider: "nearai".to_string(),
retry_after: None,
});
}
return Err(LlmError::RequestFailed {
provider: "nearai".to_string(),
retry_after: None,
reason: error.error,
});
}
return Err(LlmError::RequestFailed {
provider: "nearai".to_string(),
reason: error.error,
reason: format!("HTTP {}: {}", status, response_text),
});
}
return Err(LlmError::RequestFailed {
provider: "nearai".to_string(),
reason: format!("HTTP {}: {}", status, response_text),
});
// Success -- parse the response
return match serde_json::from_str::<R>(&response_text) {
Ok(parsed) => Ok(parsed),
Err(e) => {
tracing::debug!("Response is not expected JSON format: {}", e);
tracing::debug!("Will try alternative parsing in caller");
Err(LlmError::InvalidResponse {
provider: "nearai".to_string(),
reason: format!("Parse error: {}. Raw: {}", e, response_text),
})
}
};
}
// Try to parse as our expected type
match serde_json::from_str::<R>(&response_text) {
Ok(parsed) => Ok(parsed),
Err(e) => {
tracing::debug!("Response is not expected JSON format: {}", e);
tracing::debug!("Will try alternative parsing in caller");
Err(LlmError::InvalidResponse {
provider: "nearai".to_string(),
reason: format!("Parse error: {}. Raw: {}", e, response_text),
})
}
}
// This is unreachable because the loop always returns, but the compiler
// cannot prove that. Return a generic error as a safety net.
Err(LlmError::RequestFailed {
provider: "nearai".to_string(),
reason: "retry loop exited unexpectedly".to_string(),
})
}
}
@@ -456,7 +508,7 @@ impl LlmProvider for NearAiProvider {
Err(e) => return Err(e),
};
tracing::debug!("NEAR AI response: {:?}", response);
tracing::debug!("NEAR AI response: output_items={}", response.output.len());
// Extract text from response output
// Try multiple formats since API response shape may vary
@@ -464,11 +516,6 @@ impl LlmProvider for NearAiProvider {
.output
.iter()
.filter_map(|item| {
tracing::debug!(
"Processing output item: type={}, text={:?}",
item.item_type,
item.text
);
if item.item_type == "message" {
// First check for direct text field on item
if let Some(ref text) = item.text {
@@ -479,11 +526,6 @@ impl LlmProvider for NearAiProvider {
contents
.iter()
.filter_map(|c| {
tracing::debug!(
"Content item: type={}, text={:?}",
c.content_type,
c.text
);
// Accept various content types that might contain text
match c.content_type.as_str() {
"output_text" | "text" => c.text.clone(),
+91 -38
View File
@@ -16,6 +16,7 @@ use crate::llm::provider::{
ChatMessage, CompletionRequest, CompletionResponse, FinishReason, LlmProvider, ModelMetadata,
Role, ToolCall, ToolCompletionRequest, ToolCompletionResponse,
};
use crate::llm::retry::{is_retryable_status, retry_backoff_delay};
/// NEAR AI Chat Completions API provider.
pub struct NearAiChatProvider {
@@ -62,64 +63,116 @@ impl NearAiChatProvider {
.unwrap_or_default()
}
/// Send a request to the chat completions API.
/// Send a request to the chat completions API with retry on transient errors.
///
/// Retries on HTTP 429, 500, 502, 503, 504 with exponential backoff.
/// Does not retry on client errors (400, 401, 403, 404) or parse errors.
async fn send_request<T: Serialize, R: for<'de> Deserialize<'de>>(
&self,
body: &T,
) -> Result<R, LlmError> {
let url = self.api_url("chat/completions");
let max_retries = self.config.max_retries;
tracing::debug!("Sending request to NEAR AI Chat: {}", url);
for attempt in 0..=max_retries {
tracing::debug!(
"Sending request to NEAR AI Chat: {} (attempt {})",
url,
attempt + 1,
);
if tracing::enabled!(tracing::Level::DEBUG)
&& let Ok(json) = serde_json::to_string(body)
{
tracing::debug!("NEAR AI Chat request body: {}", json);
}
if tracing::enabled!(tracing::Level::DEBUG)
&& let Ok(json) = serde_json::to_string(body)
{
tracing::debug!("NEAR AI Chat request body: {}", json);
}
let response = self
.client
.post(&url)
.header("Authorization", format!("Bearer {}", self.api_key()))
.header("Content-Type", "application/json")
.json(body)
.send()
.await
.map_err(|e| {
tracing::error!("NEAR AI Chat request failed: {}", e);
LlmError::RequestFailed {
provider: "nearai_chat".to_string(),
reason: e.to_string(),
let response = self
.client
.post(&url)
.header("Authorization", format!("Bearer {}", self.api_key()))
.header("Content-Type", "application/json")
.json(body)
.send()
.await;
let response = match response {
Ok(r) => r,
Err(e) => {
tracing::error!("NEAR AI Chat request failed: {}", e);
if attempt < max_retries {
let delay = retry_backoff_delay(attempt);
tracing::warn!(
"NEAR AI Chat request error (attempt {}/{}), retrying in {:?}: {}",
attempt + 1,
max_retries + 1,
delay,
e,
);
tokio::time::sleep(delay).await;
continue;
}
return Err(LlmError::RequestFailed {
provider: "nearai_chat".to_string(),
reason: e.to_string(),
});
}
})?;
};
let status = response.status();
let response_text = response.text().await.unwrap_or_default();
let status = response.status();
let response_text = response.text().await.unwrap_or_default();
tracing::debug!("NEAR AI Chat response status: {}", status);
tracing::debug!("NEAR AI Chat response body: {}", response_text);
tracing::debug!("NEAR AI Chat response status: {}", status);
tracing::debug!("NEAR AI Chat response body: {}", response_text);
if !status.is_success() {
if status.as_u16() == 401 {
return Err(LlmError::AuthFailed {
if !status.is_success() {
let status_code = status.as_u16();
// Auth errors are not retryable
if status_code == 401 {
return Err(LlmError::AuthFailed {
provider: "nearai_chat".to_string(),
});
}
// Transient errors: retry with backoff
if is_retryable_status(status_code) && attempt < max_retries {
let delay = retry_backoff_delay(attempt);
tracing::warn!(
"NEAR AI Chat returned HTTP {} (attempt {}/{}), retrying in {:?}",
status_code,
attempt + 1,
max_retries + 1,
delay,
);
tokio::time::sleep(delay).await;
continue;
}
// Non-retryable or exhausted retries
if status_code == 429 {
return Err(LlmError::RateLimited {
provider: "nearai_chat".to_string(),
retry_after: None,
});
}
return Err(LlmError::RequestFailed {
provider: "nearai_chat".to_string(),
reason: format!("HTTP {}: {}", status, response_text),
});
}
if status.as_u16() == 429 {
return Err(LlmError::RateLimited {
provider: "nearai_chat".to_string(),
retry_after: None,
});
}
return Err(LlmError::RequestFailed {
// Success — parse the response
return serde_json::from_str(&response_text).map_err(|e| LlmError::InvalidResponse {
provider: "nearai_chat".to_string(),
reason: format!("HTTP {}: {}", status, response_text),
reason: format!("JSON parse error: {}. Raw: {}", e, response_text),
});
}
serde_json::from_str(&response_text).map_err(|e| LlmError::InvalidResponse {
// Safety net: unreachable because the loop always returns
Err(LlmError::RequestFailed {
provider: "nearai_chat".to_string(),
reason: format!("JSON parse error: {}. Raw: {}", e, response_text),
reason: "retry loop exited unexpectedly".to_string(),
})
}
+7
View File
@@ -113,6 +113,12 @@ pub struct ToolSelection {
pub reasoning: String,
/// Alternative tools considered.
pub alternatives: Vec<String>,
/// The tool call ID from the LLM response.
///
/// OpenAI-compatible providers assign each tool call a unique ID that must
/// be echoed back in the corresponding tool result message. Without this,
/// the provider cannot match results to their originating calls.
pub tool_call_id: String,
}
/// Token usage from a single LLM call.
@@ -244,6 +250,7 @@ impl Reasoning {
parameters: tool_call.arguments,
reasoning: reasoning.clone(),
alternatives: vec![],
tool_call_id: tool_call.id,
})
.collect();
+96
View File
@@ -0,0 +1,96 @@
//! Shared retry helpers for LLM providers.
//!
//! Provides exponential backoff with jitter and retryable status classification
//! used by both `NearAiProvider` and `NearAiChatProvider`.
use std::time::Duration;
use rand::Rng;
/// Returns `true` if the HTTP status code is transient and worth retrying.
pub(crate) fn is_retryable_status(status: u16) -> bool {
matches!(status, 429 | 500 | 502 | 503 | 504)
}
/// Calculate exponential backoff delay with random jitter.
///
/// Base delay is 1 second, doubled each attempt, with +/-25% jitter.
/// - attempt 0: ~1s (0.75s - 1.25s)
/// - attempt 1: ~2s (1.5s - 2.5s)
/// - attempt 2: ~4s (3.0s - 5.0s)
pub(crate) fn retry_backoff_delay(attempt: u32) -> Duration {
let base_ms: u64 = 1000u64.saturating_mul(2u64.saturating_pow(attempt));
let jitter_range = base_ms / 4; // 25%
let jitter = if jitter_range > 0 {
let offset = rand::thread_rng().gen_range(0..=jitter_range * 2);
offset as i64 - jitter_range as i64
} else {
0
};
let delay_ms = (base_ms as i64 + jitter).max(100) as u64;
Duration::from_millis(delay_ms)
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_is_retryable_status() {
// Transient errors should be retryable
assert!(is_retryable_status(429));
assert!(is_retryable_status(500));
assert!(is_retryable_status(502));
assert!(is_retryable_status(503));
assert!(is_retryable_status(504));
// Client errors should not be retryable
assert!(!is_retryable_status(400));
assert!(!is_retryable_status(401));
assert!(!is_retryable_status(403));
assert!(!is_retryable_status(404));
assert!(!is_retryable_status(422));
// Success codes should not be retryable
assert!(!is_retryable_status(200));
assert!(!is_retryable_status(201));
}
#[test]
fn test_retry_backoff_delay_exponential_growth() {
// Run multiple samples to verify the range, accounting for jitter
for _ in 0..20 {
let d0 = retry_backoff_delay(0);
let d1 = retry_backoff_delay(1);
let d2 = retry_backoff_delay(2);
// Attempt 0: base 1000ms, jitter +/-250ms -> [750, 1250]
assert!(d0.as_millis() >= 750, "attempt 0 too low: {:?}", d0);
assert!(d0.as_millis() <= 1250, "attempt 0 too high: {:?}", d0);
// Attempt 1: base 2000ms, jitter +/-500ms -> [1500, 2500]
assert!(d1.as_millis() >= 1500, "attempt 1 too low: {:?}", d1);
assert!(d1.as_millis() <= 2500, "attempt 1 too high: {:?}", d1);
// Attempt 2: base 4000ms, jitter +/-1000ms -> [3000, 5000]
assert!(d2.as_millis() >= 3000, "attempt 2 too low: {:?}", d2);
assert!(d2.as_millis() <= 5000, "attempt 2 too high: {:?}", d2);
}
}
#[test]
fn test_retry_backoff_delay_minimum() {
// Even at attempt 0, delay should be at least 100ms (the minimum floor)
for _ in 0..20 {
let delay = retry_backoff_delay(0);
assert!(delay.as_millis() >= 100);
}
}
#[test]
fn test_retry_backoff_delay_no_overflow() {
// Very high attempt numbers should not panic from overflow
let delay = retry_backoff_delay(30);
assert!(delay.as_millis() >= 100);
}
}
+17 -160
View File
@@ -31,8 +31,6 @@ pub struct SessionConfig {
pub auth_base_url: String,
/// Path to session file (e.g., ~/.ironclaw/session.json).
pub session_path: PathBuf,
/// Port range for OAuth callback server.
pub callback_port_range: (u16, u16),
}
impl Default for SessionConfig {
@@ -40,7 +38,6 @@ impl Default for SessionConfig {
Self {
auth_base_url: "https://private.near.ai".to_string(),
session_path: default_session_path(),
callback_port_range: (9876, 9886),
}
}
}
@@ -222,38 +219,21 @@ impl SessionManager {
/// Start the OAuth login flow.
///
/// 1. Find an available port for the callback server
/// 1. Bind the fixed callback port
/// 2. Print the auth URL and attempt to open browser
/// 3. Wait for OAuth callback with session token
/// 4. Save and return the token
async fn initiate_login(&self) -> Result<(), LlmError> {
use tokio::io::{AsyncBufReadExt, AsyncWriteExt, BufReader};
use tokio::net::TcpListener;
use crate::cli::oauth_defaults::{self, OAUTH_CALLBACK_PORT};
// Find an available port
let mut listener = None;
let mut port = 0;
let listener = oauth_defaults::bind_callback_listener()
.await
.map_err(|e| LlmError::SessionRenewalFailed {
provider: "nearai".to_string(),
reason: e.to_string(),
})?;
for p in self.config.callback_port_range.0..=self.config.callback_port_range.1 {
match TcpListener::bind(format!("127.0.0.1:{}", p)).await {
Ok(l) => {
listener = Some(l);
port = p;
break;
}
Err(_) => continue,
}
}
let listener = listener.ok_or_else(|| LlmError::SessionRenewalFailed {
provider: "nearai".to_string(),
reason: format!(
"Could not find available port in range {}-{}",
self.config.callback_port_range.0, self.config.callback_port_range.1
),
})?;
let callback_url = format!("http://127.0.0.1:{}", port);
let callback_url = format!("http://127.0.0.1:{}", OAUTH_CALLBACK_PORT);
// Show auth provider menu
println!();
@@ -333,137 +313,16 @@ impl SessionManager {
println!();
println!("Waiting for authentication...");
// Wait for callback with timeout
// The API redirects to: {frontend_callback}/auth/callback?token=X&session_id=X&expires_at=X&is_new_user=X
let timeout = std::time::Duration::from_secs(300); // 5 minutes
let selected_provider = auth_provider.to_string();
let (session_token, auth_provider) = tokio::time::timeout(timeout, async move {
loop {
let (mut socket, _) = listener.accept().await.map_err(|e| {
LlmError::SessionRenewalFailed {
provider: "nearai".to_string(),
reason: format!("Failed to accept connection: {}", e),
}
// The NEAR AI API redirects to: {frontend_callback}/auth/callback?token=X&...
let session_token =
oauth_defaults::wait_for_callback(listener, "/auth/callback", "token", "NEAR AI")
.await
.map_err(|e| LlmError::SessionRenewalFailed {
provider: "nearai".to_string(),
reason: e.to_string(),
})?;
let mut reader = BufReader::new(&mut socket);
let mut request_line = String::new();
reader.read_line(&mut request_line).await.map_err(|e| {
LlmError::SessionRenewalFailed {
provider: "nearai".to_string(),
reason: format!("Failed to read request: {}", e),
}
})?;
// Parse GET /auth/callback?token=xxx&session_id=xxx&expires_at=xxx&is_new_user=xxx HTTP/1.1
if let Some(path) = request_line.split_whitespace().nth(1)
&& path.starts_with("/auth/callback") {
// Parse query parameters
if let Some(query) = path.split('?').nth(1) {
let mut token = None;
for param in query.split('&') {
let parts: Vec<&str> = param.splitn(2, '=').collect();
if parts.len() == 2 && parts[0] == "token" {
token = Some(
urlencoding::decode(parts[1])
.unwrap_or_else(|_| parts[1].into())
.into_owned(),
);
}
}
if let Some(token) = token {
// Send success response with nice styling
let response = concat!(
"HTTP/1.1 200 OK\r\n",
"Content-Type: text/html; charset=utf-8\r\n",
"Connection: close\r\n",
"\r\n",
"<!DOCTYPE html>\n",
"<html>\n",
"<head>\n",
" <meta charset=\"utf-8\">\n",
" <title>NEAR AI - Authentication Successful</title>\n",
" <style>\n",
" * { margin: 0; padding: 0; box-sizing: border-box; }\n",
" body {\n",
" font-family: -apple-system, BlinkMacSystemFont, 'Segoe UI', Roboto, sans-serif;\n",
" background: linear-gradient(135deg, #1a1a2e 0%, #16213e 100%);\n",
" min-height: 100vh;\n",
" display: flex;\n",
" align-items: center;\n",
" justify-content: center;\n",
" color: #fff;\n",
" }\n",
" .container {\n",
" text-align: center;\n",
" padding: 3rem;\n",
" background: rgba(255,255,255,0.05);\n",
" border-radius: 16px;\n",
" backdrop-filter: blur(10px);\n",
" border: 1px solid rgba(255,255,255,0.1);\n",
" max-width: 400px;\n",
" }\n",
" .checkmark {\n",
" width: 80px;\n",
" height: 80px;\n",
" background: linear-gradient(135deg, #00d9a5 0%, #00b386 100%);\n",
" border-radius: 50%;\n",
" display: flex;\n",
" align-items: center;\n",
" justify-content: center;\n",
" margin: 0 auto 1.5rem;\n",
" font-size: 40px;\n",
" }\n",
" h1 {\n",
" font-size: 1.5rem;\n",
" font-weight: 600;\n",
" margin-bottom: 0.75rem;\n",
" }\n",
" p {\n",
" color: rgba(255,255,255,0.7);\n",
" font-size: 0.95rem;\n",
" line-height: 1.5;\n",
" }\n",
" .brand {\n",
" margin-top: 2rem;\n",
" padding-top: 1.5rem;\n",
" border-top: 1px solid rgba(255,255,255,0.1);\n",
" font-size: 0.8rem;\n",
" color: rgba(255,255,255,0.4);\n",
" }\n",
" </style>\n",
"</head>\n",
"<body>\n",
" <div class=\"container\">\n",
" <div class=\"checkmark\">&#10003;</div>\n",
" <h1>Authentication Successful</h1>\n",
" <p>You can close this window and return to the terminal.</p>\n",
" <div class=\"brand\">NEAR AI Agent</div>\n",
" </div>\n",
"</body>\n",
"</html>"
);
let _ = socket.write_all(response.as_bytes()).await;
let _ = socket.shutdown().await;
return Ok::<_, LlmError>((token, Some(selected_provider.clone())));
}
}
}
// Not the callback we're looking for, send 404
let response = "HTTP/1.1 404 Not Found\r\nConnection: close\r\n\r\n";
let _ = socket.write_all(response.as_bytes()).await;
}
})
.await
.map_err(|_| LlmError::SessionRenewalFailed {
provider: "nearai".to_string(),
reason: "Authentication timed out after 5 minutes".to_string(),
})??;
let auth_provider = Some(auth_provider.to_string());
// Save the token
self.save_session(&session_token, auth_provider.as_deref())
@@ -669,7 +528,6 @@ mod tests {
let config = SessionConfig {
auth_base_url: "https://example.com".to_string(),
session_path: session_path.clone(),
callback_port_range: (9900, 9910),
};
let manager = SessionManager::new_async(config.clone()).await;
@@ -710,7 +568,6 @@ mod tests {
let config = SessionConfig {
auth_base_url: "https://example.com".to_string(),
session_path: dir.path().join("nonexistent.json"),
callback_port_range: (9900, 9910),
};
let manager = SessionManager::new_async(config).await;
+44 -43
View File
@@ -22,7 +22,10 @@ use ironclaw::{
config::Config,
context::ContextManager,
extensions::ExtensionManager,
llm::{SessionConfig, create_llm_provider, create_session_manager},
llm::{
FailoverProvider, LlmProvider, SessionConfig, create_llm_provider,
create_llm_provider_with_config, create_session_manager,
},
orchestrator::{
ContainerJobConfig, ContainerJobManager, OrchestratorApi, TokenStore,
api::OrchestratorState,
@@ -89,7 +92,6 @@ async fn main() -> anyhow::Result<()> {
.init();
// Memory commands need database (and optionally embeddings)
let _ = dotenvy::dotenv();
let config = Config::from_env()
.await
.map_err(|e| anyhow::anyhow!("{}", e))?;
@@ -98,7 +100,6 @@ async fn main() -> anyhow::Result<()> {
let session = ironclaw::llm::create_session_manager(ironclaw::llm::SessionConfig {
auth_base_url: config.llm.nearai.auth_base_url.clone(),
session_path: config.llm.nearai.session_path.clone(),
..Default::default()
})
.await;
@@ -151,7 +152,6 @@ async fn main() -> anyhow::Result<()> {
return run_pairing_command(pairing_cmd.clone()).map_err(|e| anyhow::anyhow!("{}", e));
}
Some(Command::Status) => {
let _ = dotenvy::dotenv();
tracing_subscriber::fmt()
.with_env_filter(
EnvFilter::try_from_default_env().unwrap_or_else(|_| EnvFilter::new("warn")),
@@ -242,8 +242,10 @@ async fn main() -> anyhow::Result<()> {
skip_auth,
channels_only,
}) => {
// Load .env before running onboarding wizard
// Load .env files before running onboarding wizard.
// Standard ./.env first (higher priority), then ~/.ironclaw/.env.
let _ = dotenvy::dotenv();
ironclaw::bootstrap::load_ironclaw_env();
#[cfg(any(feature = "postgres", feature = "libsql"))]
{
@@ -266,13 +268,16 @@ async fn main() -> anyhow::Result<()> {
}
}
// Load .env if present
// Load .env files early so DATABASE_URL (and any other vars) are
// available to all subsequent env-based config resolution.
// Standard ./.env first (higher priority), then ~/.ironclaw/.env.
let _ = dotenvy::dotenv();
ironclaw::bootstrap::load_ironclaw_env();
// Enhanced first-run detection
#[cfg(any(feature = "postgres", feature = "libsql"))]
if !cli.no_onboard
&& let Some(reason) = check_onboard_needed().await
&& let Some(reason) = check_onboard_needed()
{
println!("Onboarding needed: {}", reason);
println!();
@@ -280,9 +285,6 @@ async fn main() -> anyhow::Result<()> {
wizard.run().await?;
}
// Load bootstrap config (4 fields that must live on disk)
let bootstrap = ironclaw::bootstrap::BootstrapConfig::load();
// Load initial config from env + disk (before DB is available)
let mut config = match Config::from_env().await {
Ok(c) => c,
@@ -298,14 +300,10 @@ async fn main() -> anyhow::Result<()> {
Err(e) => return Err(e.into()),
};
// Keychain master key caching is handled by CACHED_KEYCHAIN_KEY OnceLock
// in SecretsConfig::resolve(), so repeated resolve() calls skip the keychain.
// Initialize session manager and authenticate before channel setup
let session_config = SessionConfig {
auth_base_url: config.llm.nearai.auth_base_url.clone(),
session_path: config.llm.nearai.session_path.clone(),
..Default::default()
};
let session = create_session_manager(session_config).await;
@@ -316,7 +314,7 @@ async fn main() -> anyhow::Result<()> {
// Initialize tracing
let env_filter = EnvFilter::try_from_default_env()
.unwrap_or_else(|_| EnvFilter::new("ironclaw=info,tower_http=debug"));
.unwrap_or_else(|_| EnvFilter::new("ironclaw=info,tower_http=warn"));
// Create log broadcaster before tracing init so the WebLogLayer can capture all events.
// This gets wired to the gateway's /api/logs/events SSE endpoint later.
@@ -424,7 +422,7 @@ async fn main() -> anyhow::Result<()> {
}
// Reload config from DB now that we have a connection.
match Config::from_db(db.as_ref(), "default", &bootstrap).await {
match Config::from_db(db.as_ref(), "default").await {
Ok(db_config) => {
config = db_config;
tracing::info!("Configuration reloaded from database");
@@ -499,7 +497,7 @@ async fn main() -> anyhow::Result<()> {
// Re-resolve LlmConfig now that env vars may have been populated
if let Some(ref db_ref) = db {
match Config::from_db(db_ref.as_ref(), "default", &bootstrap).await {
match Config::from_db(db_ref.as_ref(), "default").await {
Ok(refreshed) => {
config = refreshed;
tracing::debug!("LlmConfig re-resolved after secret injection");
@@ -515,6 +513,27 @@ async fn main() -> anyhow::Result<()> {
let llm = create_llm_provider(&config.llm, session.clone())?;
tracing::info!("LLM provider initialized: {}", llm.model_name());
// Wrap in failover if a fallback model is configured
let llm: Arc<dyn LlmProvider> =
if let Some(fallback_model) = config.llm.nearai.fallback_model.as_ref() {
if fallback_model == &config.llm.nearai.model {
tracing::warn!(
"fallback_model is the same as primary model, failover may not be effective"
);
}
let mut fallback_config = config.llm.nearai.clone();
fallback_config.model = fallback_model.clone();
let fallback = create_llm_provider_with_config(&fallback_config, session.clone())?;
tracing::info!(
primary = %llm.model_name(),
fallback = %fallback.model_name(),
"LLM failover enabled"
);
Arc::new(FailoverProvider::new(vec![llm, fallback])?)
} else {
llm
};
// Initialize safety layer
let safety = Arc::new(SafetyLayer::new(&config.safety));
tracing::info!("Safety layer initialized");
@@ -608,7 +627,10 @@ async fn main() -> anyhow::Result<()> {
// Both register into the shared ToolRegistry (RwLock-based) so concurrent writes are safe.
let wasm_tools_future = async {
if let Some(ref runtime) = wasm_tool_runtime {
let loader = WasmToolLoader::new(Arc::clone(runtime), Arc::clone(&tools));
let mut loader = WasmToolLoader::new(Arc::clone(runtime), Arc::clone(&tools));
if let Some(ref secrets) = secrets_store {
loader = loader.with_secrets_store(Arc::clone(secrets));
}
// Load installed tools from ~/.ironclaw/tools/
match loader.load_from_dir(&config.wasm.tools_dir).await {
@@ -1195,13 +1217,11 @@ async fn main() -> anyhow::Result<()> {
/// Check if onboarding is needed and return the reason.
///
/// Returns `Some(reason)` if onboarding should be triggered, `None` otherwise.
/// Called after `load_ironclaw_env()`, so DATABASE_URL from `~/.ironclaw/.env`
/// is already in the environment.
#[cfg(any(feature = "postgres", feature = "libsql"))]
async fn check_onboard_needed() -> Option<&'static str> {
let bootstrap = ironclaw::bootstrap::BootstrapConfig::load();
// Database not configured (and not in env)
let has_db = bootstrap.database_url.is_some()
|| std::env::var("DATABASE_URL").is_ok()
fn check_onboard_needed() -> Option<&'static str> {
let has_db = std::env::var("DATABASE_URL").is_ok()
|| std::env::var("LIBSQL_PATH").is_ok()
|| ironclaw::config::default_libsql_path().exists();
@@ -1209,25 +1229,6 @@ async fn check_onboard_needed() -> Option<&'static str> {
return Some("Database not configured");
}
// Secrets not configured (and not in env)
if bootstrap.secrets_master_key_source == ironclaw::settings::KeySource::None
&& std::env::var("SECRETS_MASTER_KEY").is_err()
&& !ironclaw::secrets::keychain::has_master_key().await
{
// Only require secrets setup if user hasn't explicitly disabled it
// For now, we don't require it for first run
}
// First run (onboarding never completed and no provider configured)
let settings = ironclaw::settings::Settings::load();
let session_path = ironclaw::llm::session::default_session_path();
let has_provider = std::env::var("LLM_BACKEND").is_ok()
|| settings.llm_backend.is_some()
|| session_path.exists();
if !bootstrap.onboard_completed && !has_provider {
return Some("First run");
}
None
}
+39 -2
View File
@@ -695,12 +695,21 @@ pub mod testing {
}
async fn get(&self, user_id: &str, name: &str) -> Result<Secret, SecretError> {
self.secrets
let secret = self
.secrets
.read()
.await
.get(&(user_id.to_string(), name.to_string()))
.cloned()
.ok_or_else(|| SecretError::NotFound(name.to_string()))
.ok_or_else(|| SecretError::NotFound(name.to_string()))?;
if let Some(expires_at) = secret.expires_at
&& expires_at < Utc::now()
{
return Err(SecretError::Expired);
}
Ok(secret)
}
async fn get_decrypted(
@@ -889,6 +898,34 @@ mod tests {
);
}
#[tokio::test]
async fn test_expired_secret_returns_error() {
let store = test_store();
let expires_at = chrono::Utc::now() - chrono::Duration::hours(1);
let params = CreateSecretParams::new("expired_key", "value").with_expiry(expires_at);
store.create("user1", params).await.unwrap();
let result = store.get("user1", "expired_key").await;
assert!(result.is_err());
assert!(matches!(
result.unwrap_err(),
crate::secrets::SecretError::Expired
));
}
#[tokio::test]
async fn test_non_expired_secret_succeeds() {
let store = test_store();
let expires_at = chrono::Utc::now() + chrono::Duration::hours(1);
let params = CreateSecretParams::new("fresh_key", "value").with_expiry(expires_at);
store.create("user1", params).await.unwrap();
let result = store.get("user1", "fresh_key").await;
assert!(result.is_ok());
}
#[tokio::test]
async fn test_user_isolation() {
let store = test_store();
+21 -74
View File
@@ -509,14 +509,6 @@ impl Default for BuilderSettings {
}
impl Settings {
/// Get the default settings file path (~/.ironclaw/settings.json).
pub fn default_path() -> PathBuf {
dirs::home_dir()
.unwrap_or_else(|| PathBuf::from("."))
.join(".ironclaw")
.join("settings.json")
}
/// Reconstruct Settings from a flat key-value map (as stored in the DB).
///
/// Each key is a dotted path (e.g., "agent.name"), value is a JSONB value.
@@ -572,50 +564,27 @@ impl Settings {
map
}
/// Get the default settings file path (~/.ironclaw/settings.json).
pub fn default_path() -> std::path::PathBuf {
dirs::home_dir()
.unwrap_or_else(|| std::path::PathBuf::from("."))
.join(".ironclaw")
.join("settings.json")
}
/// Load settings from disk, returning default if not found.
pub fn load() -> Self {
Self::load_from(&Self::default_path())
}
/// Load settings from a specific path.
pub fn load_from(path: &PathBuf) -> Self {
/// Load settings from a specific path (used by bootstrap legacy migration).
pub fn load_from(path: &std::path::Path) -> Self {
match std::fs::read_to_string(path) {
Ok(data) => serde_json::from_str(&data).unwrap_or_default(),
Err(_) => Self::default(),
}
}
/// Save settings to disk.
pub fn save(&self) -> std::io::Result<()> {
self.save_to(&Self::default_path())
}
/// Save settings to a specific path.
pub fn save_to(&self, path: &PathBuf) -> std::io::Result<()> {
// Ensure parent directory exists
if let Some(parent) = path.parent() {
std::fs::create_dir_all(parent)?;
}
let json = serde_json::to_string_pretty(self)
.map_err(|e| std::io::Error::new(std::io::ErrorKind::InvalidData, e.to_string()))?;
std::fs::write(path, json)
}
/// Get the selected model, falling back to the provided default.
pub fn model_or(&self, default: &str) -> String {
self.selected_model
.clone()
.unwrap_or_else(|| default.to_string())
}
/// Set the selected model and save.
pub fn set_model(&mut self, model: &str) -> std::io::Result<()> {
self.selected_model = Some(model.to_string());
self.save()
}
/// Get a setting value by dotted path (e.g., "agent.max_parallel_jobs").
pub fn get(&self, path: &str) -> Option<String> {
let json = serde_json::to_value(self).ok()?;
@@ -800,42 +769,22 @@ fn collect_settings(
#[cfg(test)]
mod tests {
use super::*;
use tempfile::tempdir;
#[test]
fn test_settings_save_load() {
let dir = tempdir().unwrap();
let path = dir.path().join("settings.json");
fn test_db_map_round_trip() {
let settings = Settings {
selected_model: Some("claude-3-5-sonnet-20241022".to_string()),
..Default::default()
};
settings.save_to(&path).unwrap();
let loaded = Settings::load_from(&path);
let map = settings.to_db_map();
let restored = Settings::from_db_map(&map);
assert_eq!(
loaded.selected_model,
restored.selected_model,
Some("claude-3-5-sonnet-20241022".to_string())
);
}
#[test]
fn test_model_or_default() {
let settings = Settings::default();
assert_eq!(
settings.model_or("default-model"),
"default-model".to_string()
);
let settings = Settings {
selected_model: Some("my-model".to_string()),
..Default::default()
};
assert_eq!(settings.model_or("default-model"), "my-model".to_string());
}
#[test]
fn test_get_setting() {
let settings = Settings::default();
@@ -906,16 +855,13 @@ mod tests {
}
#[test]
fn test_telegram_owner_id_round_trip() {
let dir = tempdir().unwrap();
let path = dir.path().join("settings.json");
fn test_telegram_owner_id_db_round_trip() {
let mut settings = Settings::default();
settings.channels.telegram_owner_id = Some(123456789);
settings.save_to(&path).unwrap();
let loaded = Settings::load_from(&path);
assert_eq!(loaded.channels.telegram_owner_id, Some(123456789));
let map = settings.to_db_map();
let restored = Settings::from_db_map(&map);
assert_eq!(restored.channels.telegram_owner_id, Some(123456789));
}
#[test]
@@ -935,7 +881,7 @@ mod tests {
#[test]
fn test_llm_backend_round_trip() {
let dir = tempdir().unwrap();
let dir = tempfile::tempdir().unwrap();
let path = dir.path().join("settings.json");
let settings = Settings {
@@ -944,7 +890,8 @@ mod tests {
openai_compatible_base_url: Some("http://my-vllm:8000/v1".to_string()),
..Default::default()
};
settings.save_to(&path).unwrap();
let json = serde_json::to_string_pretty(&settings).unwrap();
std::fs::write(&path, json).unwrap();
let loaded = Settings::load_from(&path);
assert_eq!(loaded.llm_backend, Some("anthropic".to_string()));
+11 -23
View File
@@ -15,7 +15,7 @@ use serde::Deserialize;
#[cfg(feature = "postgres")]
use crate::secrets::SecretsCrypto;
use crate::secrets::{CreateSecretParams, SecretsStore};
use crate::settings::Settings;
use crate::settings::{Settings, TunnelSettings};
use crate::setup::prompts::{
confirm, input, optional_input, print_error, print_info, print_success, secret_input,
};
@@ -155,6 +155,7 @@ struct TelegramUpdateUser {
/// 4. Saving the token to the database
pub async fn setup_telegram(
secrets: &SecretsContext,
settings: &Settings,
) -> Result<TelegramSetupResult, ChannelSetupError> {
println!("Telegram Setup:");
println!();
@@ -169,8 +170,8 @@ pub async fn setup_telegram(
print_info("Existing Telegram token found in database.");
if !confirm("Replace existing token?", false)? {
// Still offer to configure webhook secret and owner binding
let webhook_secret = setup_telegram_webhook_secret(secrets).await?;
let owner_id = bind_telegram_owner_flow(secrets).await?;
let webhook_secret = setup_telegram_webhook_secret(secrets, &settings.tunnel).await?;
let owner_id = bind_telegram_owner_flow(secrets, settings).await?;
return Ok(TelegramSetupResult {
enabled: true,
bot_username: None,
@@ -201,7 +202,8 @@ pub async fn setup_telegram(
let owner_id = bind_telegram_owner(&token).await?;
// Offer webhook secret configuration
let webhook_secret = setup_telegram_webhook_secret(secrets).await?;
let webhook_secret =
setup_telegram_webhook_secret(secrets, &settings.tunnel).await?;
return Ok(TelegramSetupResult {
enabled: true,
@@ -336,9 +338,8 @@ async fn bind_telegram_owner(token: &SecretString) -> Result<Option<i64>, Channe
/// Retrieves the saved bot token and delegates to `bind_telegram_owner`.
async fn bind_telegram_owner_flow(
secrets: &SecretsContext,
settings: &Settings,
) -> Result<Option<i64>, ChannelSetupError> {
// Check current settings first
let settings = Settings::load();
if settings.channels.telegram_owner_id.is_some() {
print_info("Bot is already bound to a Telegram account.");
if !confirm("Re-bind to a different account?", false)? {
@@ -356,9 +357,7 @@ async fn bind_telegram_owner_flow(
///
/// This is shared across all channels that need webhook endpoints.
/// Returns the tunnel URL if configured.
pub fn setup_tunnel() -> Result<Option<String>, ChannelSetupError> {
// Check if already configured
let settings = Settings::load();
pub fn setup_tunnel(settings: &Settings) -> Result<Option<String>, ChannelSetupError> {
if let Some(ref url) = settings.tunnel.public_url {
print_info(&format!("Existing tunnel configured: {}", url));
if !confirm("Change tunnel configuration?", false)? {
@@ -398,17 +397,7 @@ pub fn setup_tunnel() -> Result<Option<String>, ChannelSetupError> {
// Remove trailing slash if present
let tunnel_url = tunnel_url.trim_end_matches('/').to_string();
// Save to settings
let mut settings = Settings::load();
settings.tunnel.public_url = Some(tunnel_url.clone());
settings.save().map_err(|e| {
ChannelSetupError::Io(std::io::Error::other(format!(
"Failed to save settings: {}",
e
)))
})?;
print_success(&format!("Tunnel URL saved: {}", tunnel_url));
print_success(&format!("Tunnel URL configured: {}", tunnel_url));
print_info("");
print_info("Make sure your tunnel is running before starting the agent.");
print_info("You can also set TUNNEL_URL environment variable to override.");
@@ -421,10 +410,9 @@ pub fn setup_tunnel() -> Result<Option<String>, ChannelSetupError> {
/// Returns the webhook secret if configured.
async fn setup_telegram_webhook_secret(
secrets: &SecretsContext,
tunnel: &TunnelSettings,
) -> Result<Option<String>, ChannelSetupError> {
// Check if tunnel is configured
let settings = Settings::load();
if settings.tunnel.public_url.is_none() {
if tunnel.public_url.is_none() {
print_info("");
print_info("No tunnel configured. Telegram will use polling mode (30s+ delay).");
print_info("Run setup again to configure a tunnel for instant delivery.");
+74 -11
View File
@@ -89,7 +89,7 @@ impl SetupWizard {
pub fn new() -> Self {
Self {
config: SetupConfig::default(),
settings: Settings::load(),
settings: Settings::default(),
session_manager: None,
#[cfg(feature = "postgres")]
db_pool: None,
@@ -103,7 +103,7 @@ impl SetupWizard {
pub fn with_config(config: SetupConfig) -> Self {
Self {
config,
settings: Settings::load(),
settings: Settings::default(),
session_manager: None,
#[cfg(feature = "postgres")]
db_pool: None,
@@ -164,7 +164,7 @@ impl SetupWizard {
}
// Save settings and print summary
self.save_and_summarize()?;
self.save_and_summarize().await?;
Ok(())
}
@@ -999,6 +999,8 @@ impl SetupWizard {
session_path: crate::llm::session::default_session_path(),
api_mode: crate::config::NearAiApiMode::Responses,
api_key: None,
fallback_model: None,
max_retries: 3,
},
openai: None,
anthropic: None,
@@ -1198,7 +1200,7 @@ impl SetupWizard {
/// Step 6: Channel configuration.
async fn step_channels(&mut self) -> Result<(), SetupError> {
// First, configure tunnel (shared across all channels that need webhooks)
match setup_tunnel() {
match setup_tunnel(&self.settings) {
Ok(Some(url)) => {
self.settings.tunnel.public_url = Some(url);
}
@@ -1312,7 +1314,7 @@ impl SetupWizard {
if !cap_file.setup.required_secrets.is_empty() {
setup_wasm_channel(ctx, &channel_name, &cap_file.setup).await?
} else if channel_name == "telegram" {
let telegram_result = setup_telegram(ctx).await?;
let telegram_result = setup_telegram(ctx, &self.settings).await?;
if let Some(owner_id) = telegram_result.owner_id {
self.settings.channels.telegram_owner_id = Some(owner_id);
}
@@ -1395,16 +1397,77 @@ impl SetupWizard {
Ok(())
}
/// Save settings and print summary.
fn save_and_summarize(&mut self) -> Result<(), SetupError> {
/// Save settings to the database and `~/.ironclaw/.env`, then print summary.
async fn save_and_summarize(&mut self) -> Result<(), SetupError> {
self.settings.onboard_completed = true;
self.settings
.save()
.map_err(|e| std::io::Error::other(format!("Failed to save settings: {}", e)))?;
// Write all settings to the database (whichever backend is active).
{
let db_map = self.settings.to_db_map();
let saved = false;
#[cfg(feature = "postgres")]
let saved = if !saved {
if let Some(ref pool) = self.db_pool {
let store = crate::history::Store::from_pool(pool.clone());
store
.set_all_settings("default", &db_map)
.await
.map_err(|e| {
SetupError::Database(format!(
"Failed to save settings to database: {}",
e
))
})?;
true
} else {
false
}
} else {
saved
};
#[cfg(feature = "libsql")]
let saved = if !saved {
if let Some(ref backend) = self.db_backend {
use crate::db::Database as _;
backend
.set_all_settings("default", &db_map)
.await
.map_err(|e| {
SetupError::Database(format!(
"Failed to save settings to database: {}",
e
))
})?;
true
} else {
false
}
} else {
saved
};
if !saved {
return Err(SetupError::Database(
"No database connection, cannot save settings".to_string(),
));
}
}
// Save DATABASE_URL to ~/.ironclaw/.env (the only field that needs
// disk persistence before the DB is available).
if let Some(ref url) = self.settings.database_url {
crate::bootstrap::save_database_url(url).map_err(|e| {
SetupError::Io(std::io::Error::other(format!(
"Failed to save DATABASE_URL to .env: {}",
e
)))
})?;
}
println!();
print_success("Configuration saved to ~/.ironclaw/");
print_success("Configuration saved to database");
println!();
// Print summary
+42
View File
@@ -527,6 +527,48 @@ mod tests {
));
}
/// Replicate the extraction logic from agent_loop.rs to prove it works
/// when `arguments` is a `serde_json::Value::Object` (the common case
/// that was previously broken because `Value::Object.as_str()` returns None).
#[test]
fn test_destructive_command_extraction_from_object_args() {
let arguments = serde_json::json!({"command": "rm -rf /tmp/stuff"});
let cmd = arguments
.get("command")
.and_then(|c| c.as_str().map(String::from))
.or_else(|| {
arguments
.as_str()
.and_then(|s| serde_json::from_str::<serde_json::Value>(s).ok())
.and_then(|v| v.get("command").and_then(|c| c.as_str().map(String::from)))
});
assert_eq!(cmd.as_deref(), Some("rm -rf /tmp/stuff"));
assert!(requires_explicit_approval(cmd.as_deref().unwrap()));
}
/// Verify extraction still works when `arguments` is a JSON string
/// (rare, but possible if the LLM provider returns string-encoded JSON).
#[test]
fn test_destructive_command_extraction_from_string_args() {
let arguments =
serde_json::Value::String(r#"{"command": "git push --force origin main"}"#.to_string());
let cmd = arguments
.get("command")
.and_then(|c| c.as_str().map(String::from))
.or_else(|| {
arguments
.as_str()
.and_then(|s| serde_json::from_str::<serde_json::Value>(s).ok())
.and_then(|v| v.get("command").and_then(|c| c.as_str().map(String::from)))
});
assert_eq!(cmd.as_deref(), Some("git push --force origin main"));
assert!(requires_explicit_approval(cmd.as_deref().unwrap()));
}
#[test]
fn test_sandbox_policy_builder() {
let tool = ShellTool::new()
+16 -71
View File
@@ -11,9 +11,9 @@ use base64::{Engine, engine::general_purpose::URL_SAFE_NO_PAD};
use rand::RngCore;
use serde::{Deserialize, Serialize};
use sha2::{Digest, Sha256};
use tokio::io::{AsyncBufReadExt, AsyncWriteExt, BufReader};
use tokio::net::TcpListener;
use crate::cli::oauth_defaults::{self, OAUTH_CALLBACK_PORT};
use crate::secrets::{CreateSecretParams, SecretsStore};
use crate::tools::mcp::config::McpServerConfig;
@@ -466,14 +466,12 @@ pub async fn authorize_mcp_server(
Ok(token)
}
/// Find an available port for the OAuth callback.
/// Bind the OAuth callback listener on the shared fixed port.
pub async fn find_available_port() -> Result<(TcpListener, u16), AuthError> {
for port in 9876..=9886 {
if let Ok(listener) = TcpListener::bind(format!("127.0.0.1:{}", port)).await {
return Ok((listener, port));
}
}
Err(AuthError::PortUnavailable)
let listener = oauth_defaults::bind_callback_listener()
.await
.map_err(|_| AuthError::PortUnavailable)?;
Ok((listener, OAUTH_CALLBACK_PORT))
}
/// Build the authorization URL with all required parameters.
@@ -522,69 +520,16 @@ pub async fn wait_for_authorization_callback(
listener: TcpListener,
server_name: &str,
) -> Result<String, AuthError> {
let timeout = Duration::from_secs(300);
tokio::time::timeout(timeout, async {
loop {
let (mut socket, _) = listener
.accept()
.await
.map_err(|e| AuthError::Http(e.to_string()))?;
let mut reader = BufReader::new(&mut socket);
let mut request_line = String::new();
reader
.read_line(&mut request_line)
.await
.map_err(|e| AuthError::Http(e.to_string()))?;
// Parse GET /callback?code=xxx HTTP/1.1
if let Some(path) = request_line.split_whitespace().nth(1)
&& path.starts_with("/callback")
&& let Some(query) = path.split('?').nth(1) {
// Check for error first
if query.contains("error=") {
let response = "HTTP/1.1 400 Bad Request\r\n\r\nAuthorization denied";
let _ = socket.write_all(response.as_bytes()).await;
return Err(AuthError::AuthorizationDenied);
}
// Look for code
for param in query.split('&') {
let parts: Vec<&str> = param.splitn(2, '=').collect();
if parts.len() == 2 && parts[0] == "code" {
let code = urlencoding::decode(parts[1])
.unwrap_or_else(|_| parts[1].into())
.into_owned();
// Send success response
let response = format!(
"HTTP/1.1 200 OK\r\n\
Content-Type: text/html\r\n\
\r\n\
<!DOCTYPE html><html><body style=\"font-family: sans-serif; \
display: flex; justify-content: center; align-items: center; \
height: 100vh; margin: 0; background: #191919; color: white;\">\
<div style=\"text-align: center;\">\
<h1> {} Connected!</h1>\
<p>You can close this window.</p>\
</div></body></html>",
server_name
);
let _ = socket.write_all(response.as_bytes()).await;
let _ = socket.shutdown().await;
return Ok(code);
}
}
}
let response = "HTTP/1.1 404 Not Found\r\n\r\n";
let _ = socket.write_all(response.as_bytes()).await;
}
})
.await
.map_err(|_| AuthError::Timeout)?
oauth_defaults::wait_for_callback(listener, "/callback", "code", server_name)
.await
.map_err(|e| match e {
oauth_defaults::OAuthCallbackError::Denied => AuthError::AuthorizationDenied,
oauth_defaults::OAuthCallbackError::Timeout => AuthError::Timeout,
oauth_defaults::OAuthCallbackError::PortInUse(_, msg) => {
AuthError::Http(format!("Port error: {}", msg))
}
oauth_defaults::OAuthCallbackError::Io(msg) => AuthError::Http(msg),
})
}
/// Exchange the authorization code for an access token.
+9 -4
View File
@@ -190,10 +190,15 @@ impl McpClient {
req_builder = req_builder.header("Mcp-Session-Id", session_id);
}
let response = req_builder
.send()
.await
.map_err(|e| ToolError::ExternalService(format!("MCP request failed: {}", e)))?;
let response = req_builder.send().await.map_err(|e| {
let mut chain = format!("MCP request failed: {}", e);
let mut source = std::error::Error::source(&e);
while let Some(cause) = source {
chain.push_str(&format!(" -> {}", cause));
source = cause.source();
}
ToolError::ExternalService(chain)
})?;
// Check for 401 Unauthorized - try to refresh token on first attempt
if response.status() == reqwest::StatusCode::UNAUTHORIZED {
+82 -1
View File
@@ -88,8 +88,18 @@ impl McpServerConfig {
}
/// Check if this server requires authentication.
///
/// Returns true if OAuth is pre-configured OR if this is a remote HTTPS server
/// (which likely supports Dynamic Client Registration even without pre-configured OAuth).
pub fn requires_auth(&self) -> bool {
self.oauth.is_some()
if self.oauth.is_some() {
return true;
}
// Remote HTTPS servers need auth handling (DCR, token refresh, 401 detection).
// Localhost/127.0.0.1 servers are assumed to be dev servers without auth.
let url_lower = self.url.to_lowercase();
let is_localhost = is_localhost_url(&url_lower);
url_lower.starts_with("https://") && !is_localhost
}
/// Get the secret name used to store the access token.
@@ -402,11 +412,43 @@ pub async fn remove_mcp_server_db(
Ok(())
}
/// Check if a URL points to a loopback address (localhost, 127.0.0.1, [::1]).
///
/// Uses `url::Url` for proper parsing so edge cases (IPv6, userinfo, ports)
/// are handled correctly without manual string splitting.
fn is_localhost_url(url: &str) -> bool {
let Ok(parsed) = url::Url::parse(url) else {
return false;
};
match parsed.host() {
Some(url::Host::Domain(d)) => d.eq_ignore_ascii_case("localhost"),
Some(url::Host::Ipv4(ip)) => ip.is_loopback(),
Some(url::Host::Ipv6(ip)) => ip.is_loopback(),
None => false,
}
}
#[cfg(test)]
mod tests {
use super::*;
use tempfile::tempdir;
#[test]
fn test_is_localhost_url() {
assert!(is_localhost_url("http://localhost:3000/path"));
assert!(is_localhost_url("https://localhost/path"));
assert!(is_localhost_url("http://127.0.0.1:8080"));
assert!(is_localhost_url("http://127.0.0.1"));
assert!(!is_localhost_url("https://notlocalhost.com/path"));
assert!(!is_localhost_url("https://example-localhost.io"));
assert!(!is_localhost_url("https://mcp.notion.com"));
assert!(is_localhost_url("http://user:pass@localhost:3000/path"));
// IPv6 loopback
assert!(is_localhost_url("http://[::1]:8080/path"));
assert!(is_localhost_url("http://[::1]/path"));
assert!(!is_localhost_url("http://[::2]:8080/path"));
}
#[test]
fn test_server_config_validation() {
// Valid HTTPS server
@@ -514,4 +556,43 @@ mod tests {
"mcp_notion_refresh_token"
);
}
#[test]
fn test_requires_auth_with_oauth() {
let config = McpServerConfig::new("notion", "https://mcp.notion.com")
.with_oauth(OAuthConfig::new("client-123"));
assert!(config.requires_auth());
}
#[test]
fn test_requires_auth_remote_https_without_oauth() {
// Remote HTTPS servers need auth even without pre-configured OAuth (DCR)
let config = McpServerConfig::new("github-copilot", "https://api.githubcopilot.com/mcp/");
assert!(config.requires_auth());
let config = McpServerConfig::new("notion", "https://mcp.notion.com");
assert!(config.requires_auth());
}
#[test]
fn test_requires_auth_localhost_no_auth() {
// Localhost servers are dev servers, no auth needed
let config = McpServerConfig::new("local", "http://localhost:8080");
assert!(!config.requires_auth());
let config = McpServerConfig::new("local", "http://127.0.0.1:3000/mcp");
assert!(!config.requires_auth());
// Even HTTPS localhost doesn't require auth
let config = McpServerConfig::new("local", "https://localhost:8443");
assert!(!config.requires_auth());
}
#[test]
fn test_requires_auth_http_remote_no_auth() {
// HTTP remote servers won't pass validation, but if they existed
// they wouldn't trigger HTTPS auth detection
let config = McpServerConfig::new("bad", "http://mcp.example.com");
assert!(!config.requires_auth());
}
}
+15 -2
View File
@@ -11,6 +11,7 @@ use crate::extensions::ExtensionManager;
use crate::llm::{LlmProvider, ToolDefinition};
use crate::orchestrator::job_manager::ContainerJobManager;
use crate::safety::SafetyLayer;
use crate::secrets::SecretsStore;
use crate::tools::builder::{BuildSoftwareTool, BuilderConfig, LlmSoftwareBuilder};
use crate::tools::builtin::{
ApplyPatchTool, CancelJobTool, CreateJobTool, EchoTool, HttpTool, JobStatusTool, JsonTool,
@@ -20,8 +21,8 @@ use crate::tools::builtin::{
};
use crate::tools::tool::{Tool, ToolDomain};
use crate::tools::wasm::{
Capabilities, ResourceLimits, WasmError, WasmStorageError, WasmToolRuntime, WasmToolStore,
WasmToolWrapper,
Capabilities, OAuthRefreshConfig, ResourceLimits, WasmError, WasmStorageError, WasmToolRuntime,
WasmToolStore, WasmToolWrapper,
};
use crate::workspace::Workspace;
@@ -366,6 +367,12 @@ impl ToolRegistry {
if let Some(s) = reg.schema {
wrapper = wrapper.with_schema(s);
}
if let Some(store) = reg.secrets_store {
wrapper = wrapper.with_secrets_store(store);
}
if let Some(oauth) = reg.oauth_refresh {
wrapper = wrapper.with_oauth_refresh(oauth);
}
// Register the tool
self.register(Arc::new(wrapper)).await;
@@ -421,6 +428,8 @@ impl ToolRegistry {
limits: None,
description: Some(&tool_with_binary.tool.description),
schema: Some(tool_with_binary.tool.parameters_schema.clone()),
secrets_store: None,
oauth_refresh: None,
})
.await
.map_err(WasmRegistrationError::Wasm)?;
@@ -462,6 +471,10 @@ pub struct WasmToolRegistration<'a> {
pub description: Option<&'a str>,
/// Optional parameter schema override.
pub schema: Option<serde_json::Value>,
/// Secrets store for credential injection at request time.
pub secrets_store: Option<Arc<dyn SecretsStore + Send + Sync>>,
/// OAuth refresh configuration for auto-refreshing expired tokens.
pub oauth_refresh: Option<OAuthRefreshConfig>,
}
impl Default for ToolRegistry {
+2 -2
View File
@@ -169,7 +169,7 @@ impl CredentialInjector {
}
/// Inject a single credential into the result.
fn inject_credential(
pub(crate) fn inject_credential(
result: &mut InjectedCredentials,
location: &CredentialLocation,
secret: &DecryptedSecret,
@@ -208,7 +208,7 @@ fn inject_credential(
}
/// Check if a host matches a pattern (supports wildcards).
fn host_matches_pattern(host: &str, pattern: &str) -> bool {
pub(crate) fn host_matches_pattern(host: &str, pattern: &str) -> bool {
if pattern == host {
return true;
}
+179 -7
View File
@@ -39,10 +39,11 @@ use std::sync::Arc;
use tokio::fs;
use crate::secrets::SecretsStore;
use crate::tools::registry::{ToolRegistry, WasmRegistrationError, WasmToolRegistration};
use crate::tools::wasm::capabilities_schema::CapabilitiesFile;
use crate::tools::wasm::{
Capabilities, WasmError, WasmStorageError, WasmToolRuntime, WasmToolStore,
Capabilities, OAuthRefreshConfig, WasmError, WasmStorageError, WasmToolRuntime, WasmToolStore,
};
/// Error during WASM tool loading.
@@ -77,12 +78,23 @@ pub enum WasmLoadError {
pub struct WasmToolLoader {
runtime: Arc<WasmToolRuntime>,
registry: Arc<ToolRegistry>,
secrets_store: Option<Arc<dyn SecretsStore + Send + Sync>>,
}
impl WasmToolLoader {
/// Create a new loader with the given runtime and registry.
pub fn new(runtime: Arc<WasmToolRuntime>, registry: Arc<ToolRegistry>) -> Self {
Self { runtime, registry }
Self {
runtime,
registry,
secrets_store: None,
}
}
/// Set the secrets store for credential injection in WASM tools.
pub fn with_secrets_store(mut self, store: Arc<dyn SecretsStore + Send + Sync>) -> Self {
self.secrets_store = Some(store);
self
}
/// Load a single WASM tool from a file pair.
@@ -108,22 +120,24 @@ impl WasmToolLoader {
}
let wasm_bytes = fs::read(wasm_path).await?;
// Read capabilities (optional)
let capabilities = if let Some(cap_path) = capabilities_path {
// Read capabilities (optional) and extract OAuth refresh config
let (capabilities, oauth_refresh) = if let Some(cap_path) = capabilities_path {
if cap_path.exists() {
let cap_bytes = fs::read(cap_path).await?;
let cap_file = CapabilitiesFile::from_bytes(&cap_bytes)
.map_err(|e| WasmLoadError::InvalidCapabilities(e.to_string()))?;
cap_file.to_capabilities()
let caps = cap_file.to_capabilities();
let oauth = resolve_oauth_refresh_config(&cap_file);
(caps, oauth)
} else {
tracing::warn!(
path = %cap_path.display(),
"Capabilities file not found, using default (no permissions)"
);
Capabilities::default()
(Capabilities::default(), None)
}
} else {
Capabilities::default()
(Capabilities::default(), None)
};
// Register the tool
@@ -136,6 +150,8 @@ impl WasmToolLoader {
limits: None,
description: None,
schema: None,
secrets_store: self.secrets_store.clone(),
oauth_refresh,
})
.await?;
@@ -293,6 +309,50 @@ impl WasmToolLoader {
}
}
/// Extract OAuth refresh configuration from a parsed capabilities file.
///
/// Returns `None` if there's no `auth.oauth` section or if the client_id
/// can't be resolved from any source (inline, env var, or built-in defaults).
///
/// Fallback chain for client_id:
/// `oauth.client_id` > env var (`oauth.client_id_env`) > `builtin_credentials()`
fn resolve_oauth_refresh_config(cap_file: &CapabilitiesFile) -> Option<OAuthRefreshConfig> {
let auth = cap_file.auth.as_ref()?;
let oauth = auth.oauth.as_ref()?;
let builtin = crate::cli::oauth_defaults::builtin_credentials(&auth.secret_name);
let client_id = oauth
.client_id
.clone()
.or_else(|| {
oauth
.client_id_env
.as_ref()
.and_then(|env| std::env::var(env).ok())
})
.or_else(|| builtin.as_ref().map(|c| c.client_id.to_string()))?;
let client_secret = oauth
.client_secret
.clone()
.or_else(|| {
oauth
.client_secret_env
.as_ref()
.and_then(|env| std::env::var(env).ok())
})
.or_else(|| builtin.as_ref().map(|c| c.client_secret.to_string()));
Some(OAuthRefreshConfig {
token_url: oauth.token_url.clone(),
client_id,
client_secret,
secret_name: auth.secret_name.clone(),
provider: auth.provider.clone(),
})
}
/// Results from loading multiple tools.
#[derive(Debug, Default)]
pub struct LoadResults {
@@ -618,4 +678,116 @@ mod tests {
);
}
}
#[test]
fn test_resolve_oauth_refresh_config_with_oauth() {
use crate::tools::wasm::capabilities_schema::{
AuthCapabilitySchema, CapabilitiesFile, OAuthConfigSchema,
};
let caps = CapabilitiesFile {
auth: Some(AuthCapabilitySchema {
secret_name: "google_oauth_token".to_string(),
provider: Some("google".to_string()),
oauth: Some(OAuthConfigSchema {
authorization_url: "https://accounts.google.com/o/oauth2/v2/auth".to_string(),
token_url: "https://oauth2.googleapis.com/token".to_string(),
client_id: Some("test-client-id".to_string()),
client_secret: Some("test-client-secret".to_string()),
..Default::default()
}),
..Default::default()
}),
..Default::default()
};
let config = super::resolve_oauth_refresh_config(&caps);
assert!(config.is_some());
let config = config.unwrap();
assert_eq!(config.token_url, "https://oauth2.googleapis.com/token");
assert_eq!(config.client_id, "test-client-id");
assert_eq!(config.client_secret, Some("test-client-secret".to_string()));
assert_eq!(config.secret_name, "google_oauth_token");
assert_eq!(config.provider, Some("google".to_string()));
}
#[test]
fn test_resolve_oauth_refresh_config_no_auth() {
use crate::tools::wasm::capabilities_schema::CapabilitiesFile;
let caps = CapabilitiesFile::default();
let config = super::resolve_oauth_refresh_config(&caps);
assert!(config.is_none());
}
#[test]
fn test_resolve_oauth_refresh_config_no_oauth() {
use crate::tools::wasm::capabilities_schema::{AuthCapabilitySchema, CapabilitiesFile};
let caps = CapabilitiesFile {
auth: Some(AuthCapabilitySchema {
secret_name: "manual_token".to_string(),
..Default::default()
}),
..Default::default()
};
let config = super::resolve_oauth_refresh_config(&caps);
assert!(config.is_none());
}
#[test]
fn test_resolve_oauth_refresh_config_no_client_id() {
use crate::tools::wasm::capabilities_schema::{
AuthCapabilitySchema, CapabilitiesFile, OAuthConfigSchema,
};
// A non-Google provider with no client_id anywhere should return None
let caps = CapabilitiesFile {
auth: Some(AuthCapabilitySchema {
secret_name: "unknown_provider_token".to_string(),
oauth: Some(OAuthConfigSchema {
authorization_url: "https://example.com/auth".to_string(),
token_url: "https://example.com/token".to_string(),
// No client_id, no client_id_env, no builtin
..Default::default()
}),
..Default::default()
}),
..Default::default()
};
let config = super::resolve_oauth_refresh_config(&caps);
assert!(config.is_none());
}
#[test]
fn test_resolve_oauth_refresh_config_builtin_google() {
use crate::tools::wasm::capabilities_schema::{
AuthCapabilitySchema, CapabilitiesFile, OAuthConfigSchema,
};
// google_oauth_token should fall back to built-in credentials
let caps = CapabilitiesFile {
auth: Some(AuthCapabilitySchema {
secret_name: "google_oauth_token".to_string(),
provider: Some("google".to_string()),
oauth: Some(OAuthConfigSchema {
authorization_url: "https://accounts.google.com/o/oauth2/v2/auth".to_string(),
token_url: "https://oauth2.googleapis.com/token".to_string(),
// No inline client_id, should fall back to builtin
..Default::default()
}),
..Default::default()
}),
..Default::default()
};
let config = super::resolve_oauth_refresh_config(&caps);
assert!(config.is_some());
let config = config.unwrap();
assert!(!config.client_id.is_empty());
assert!(config.client_secret.is_some());
}
}
+1 -1
View File
@@ -94,7 +94,7 @@ pub use limits::{
WasmResourceLimiter,
};
pub use runtime::{PreparedModule, WasmRuntimeConfig, WasmToolRuntime};
pub use wrapper::WasmToolWrapper;
pub use wrapper::{OAuthRefreshConfig, WasmToolWrapper};
// Capabilities (V2)
pub use capabilities::{
+877 -14
View File
File diff suppressed because it is too large Load Diff
+3 -2
View File
@@ -313,6 +313,7 @@ Work independently to complete this job. Report when done."#,
parameters: tc.arguments.clone(),
reasoning: String::new(),
alternatives: vec![],
tool_call_id: tc.id.clone(),
};
self.process_result(reason_ctx, &selection, result);
}
@@ -422,7 +423,7 @@ Work independently to complete this job. Report when done."#,
);
reason_ctx.messages.push(ChatMessage::tool_result(
"tool_call_id",
&selection.tool_call_id,
&selection.tool_name,
wrapped,
));
@@ -436,7 +437,7 @@ Work independently to complete this job. Report when done."#,
Err(e) => {
tracing::warn!("Tool {} failed: {}", selection.tool_name, e);
reason_ctx.messages.push(ChatMessage::tool_result(
"tool_call_id",
&selection.tool_call_id,
&selection.tool_name,
format!("Error: {}", e),
));
+44 -107
View File
@@ -53,119 +53,56 @@ impl exports::near::agent::tool::Guest for GmailTool {
r#"{
"type": "object",
"required": ["action"],
"oneOf": [
{
"properties": {
"action": { "const": "list_messages" },
"query": {
"type": "string",
"description": "Gmail search query (same syntax as Gmail search box). Examples: 'is:unread', 'from:[email protected]', 'subject:meeting after:2025/01/01'"
},
"max_results": {
"type": "integer",
"description": "Maximum number of messages to return (default: 20)",
"default": 20
},
"label_ids": {
"type": "array",
"items": { "type": "string" },
"description": "Label IDs to filter by (e.g., 'INBOX', 'SENT', 'DRAFT')"
}
},
"required": ["action"]
"properties": {
"action": {
"type": "string",
"enum": ["list_messages", "get_message", "send_message", "create_draft", "reply_to_message", "trash_message"],
"description": "The Gmail operation to perform"
},
{
"properties": {
"action": { "const": "get_message" },
"message_id": {
"type": "string",
"description": "The message ID to retrieve"
}
},
"required": ["action", "message_id"]
"query": {
"type": "string",
"description": "Gmail search query (same syntax as Gmail search box, e.g., 'is:unread', 'from:[email protected]'). Used by: list_messages"
},
{
"properties": {
"action": { "const": "send_message" },
"to": {
"type": "string",
"description": "Recipient email address(es), comma-separated"
},
"subject": {
"type": "string",
"description": "Email subject"
},
"body": {
"type": "string",
"description": "Email body (plain text)"
},
"cc": {
"type": "string",
"description": "CC recipients, comma-separated"
},
"bcc": {
"type": "string",
"description": "BCC recipients, comma-separated"
}
},
"required": ["action", "to", "subject", "body"]
"max_results": {
"type": "integer",
"description": "Maximum number of messages to return (default: 20). Used by: list_messages",
"default": 20
},
{
"properties": {
"action": { "const": "create_draft" },
"to": {
"type": "string",
"description": "Recipient email address(es), comma-separated"
},
"subject": {
"type": "string",
"description": "Email subject"
},
"body": {
"type": "string",
"description": "Email body (plain text)"
},
"cc": {
"type": "string",
"description": "CC recipients, comma-separated"
},
"bcc": {
"type": "string",
"description": "BCC recipients, comma-separated"
}
},
"required": ["action", "to", "subject", "body"]
"label_ids": {
"type": "array",
"items": { "type": "string" },
"description": "Label IDs to filter by (e.g., 'INBOX', 'SENT', 'DRAFT'). Used by: list_messages"
},
{
"properties": {
"action": { "const": "reply_to_message" },
"message_id": {
"type": "string",
"description": "The message ID to reply to"
},
"body": {
"type": "string",
"description": "Reply body (plain text)"
},
"reply_all": {
"type": "boolean",
"description": "If true, reply to all recipients (default: false)",
"default": false
}
},
"required": ["action", "message_id", "body"]
"message_id": {
"type": "string",
"description": "Message ID. Required for: get_message, reply_to_message, trash_message"
},
{
"properties": {
"action": { "const": "trash_message" },
"message_id": {
"type": "string",
"description": "The message ID to move to trash"
}
},
"required": ["action", "message_id"]
"to": {
"type": "string",
"description": "Recipient email address(es), comma-separated. Required for: send_message, create_draft"
},
"subject": {
"type": "string",
"description": "Email subject. Required for: send_message, create_draft"
},
"body": {
"type": "string",
"description": "Email body (plain text). Required for: send_message, create_draft, reply_to_message"
},
"cc": {
"type": "string",
"description": "CC recipients, comma-separated. Used by: send_message, create_draft"
},
"bcc": {
"type": "string",
"description": "BCC recipients, comma-separated. Used by: send_message, create_draft"
},
"reply_all": {
"type": "boolean",
"description": "If true, reply to all recipients (default: false). Used by: reply_to_message",
"default": false
}
]
}
}"#
.to_string()
}
+65 -155
View File
@@ -52,166 +52,76 @@ impl exports::near::agent::tool::Guest for GoogleCalendarTool {
r#"{
"type": "object",
"required": ["action"],
"oneOf": [
{
"properties": {
"action": { "const": "list_events" },
"calendar_id": {
"type": "string",
"description": "Calendar ID (default: 'primary')",
"default": "primary"
},
"time_min": {
"type": "string",
"description": "Lower bound for event start time (RFC3339, e.g., '2025-01-15T00:00:00Z')"
},
"time_max": {
"type": "string",
"description": "Upper bound for event end time (RFC3339)"
},
"max_results": {
"type": "integer",
"description": "Maximum number of events to return (default: 25)",
"default": 25
},
"query": {
"type": "string",
"description": "Free text search terms to filter events"
}
},
"required": ["action"]
"properties": {
"action": {
"type": "string",
"enum": ["list_events", "get_event", "create_event", "update_event", "delete_event"],
"description": "The calendar operation to perform"
},
{
"properties": {
"action": { "const": "get_event" },
"calendar_id": {
"type": "string",
"description": "Calendar ID (default: 'primary')",
"default": "primary"
},
"event_id": {
"type": "string",
"description": "The event ID to retrieve"
}
},
"required": ["action", "event_id"]
"calendar_id": {
"type": "string",
"description": "Calendar ID (default: 'primary')",
"default": "primary"
},
{
"properties": {
"action": { "const": "create_event" },
"calendar_id": {
"type": "string",
"description": "Calendar ID (default: 'primary')",
"default": "primary"
},
"summary": {
"type": "string",
"description": "Event title"
},
"description": {
"type": "string",
"description": "Event description"
},
"location": {
"type": "string",
"description": "Event location"
},
"start_datetime": {
"type": "string",
"description": "Start time as RFC3339 (e.g., '2025-01-15T09:00:00-05:00'). Use start_date for all-day events."
},
"end_datetime": {
"type": "string",
"description": "End time as RFC3339. Use end_date for all-day events."
},
"start_date": {
"type": "string",
"description": "Start date for all-day events (e.g., '2025-01-15')"
},
"end_date": {
"type": "string",
"description": "End date for all-day events (exclusive, e.g., '2025-01-16' for a single day)"
},
"timezone": {
"type": "string",
"description": "Timezone (e.g., 'America/New_York')"
},
"attendees": {
"type": "array",
"items": { "type": "string" },
"description": "Attendee email addresses"
}
},
"required": ["action", "summary"]
"event_id": {
"type": "string",
"description": "Event ID. Required for: get_event, update_event, delete_event"
},
{
"properties": {
"action": { "const": "update_event" },
"calendar_id": {
"type": "string",
"description": "Calendar ID (default: 'primary')",
"default": "primary"
},
"event_id": {
"type": "string",
"description": "The event ID to update"
},
"summary": {
"type": "string",
"description": "New event title"
},
"description": {
"type": "string",
"description": "New event description"
},
"location": {
"type": "string",
"description": "New event location"
},
"start_datetime": {
"type": "string",
"description": "New start time (RFC3339)"
},
"end_datetime": {
"type": "string",
"description": "New end time (RFC3339)"
},
"start_date": {
"type": "string",
"description": "New start date for all-day events"
},
"end_date": {
"type": "string",
"description": "New end date for all-day events"
},
"timezone": {
"type": "string",
"description": "Timezone for datetime fields"
},
"attendees": {
"type": "array",
"items": { "type": "string" },
"description": "Replace attendees with these email addresses"
}
},
"required": ["action", "event_id"]
"time_min": {
"type": "string",
"description": "Lower bound for event start time (RFC3339, e.g., '2025-01-15T00:00:00Z'). Used by: list_events"
},
{
"properties": {
"action": { "const": "delete_event" },
"calendar_id": {
"type": "string",
"description": "Calendar ID (default: 'primary')",
"default": "primary"
},
"event_id": {
"type": "string",
"description": "The event ID to delete"
}
},
"required": ["action", "event_id"]
"time_max": {
"type": "string",
"description": "Upper bound for event end time (RFC3339). Used by: list_events"
},
"max_results": {
"type": "integer",
"description": "Maximum number of events to return (default: 25). Used by: list_events",
"default": 25
},
"query": {
"type": "string",
"description": "Free text search terms to filter events. Used by: list_events"
},
"summary": {
"type": "string",
"description": "Event title. Required for: create_event. Optional for: update_event"
},
"description": {
"type": "string",
"description": "Event description. Used by: create_event, update_event"
},
"location": {
"type": "string",
"description": "Event location. Used by: create_event, update_event"
},
"start_datetime": {
"type": "string",
"description": "Start time (RFC3339, e.g., '2025-01-15T09:00:00-05:00'). For all-day events use start_date. Used by: create_event, update_event"
},
"end_datetime": {
"type": "string",
"description": "End time (RFC3339). For all-day events use end_date. Used by: create_event, update_event"
},
"start_date": {
"type": "string",
"description": "Start date for all-day events (e.g., '2025-01-15'). Used by: create_event, update_event"
},
"end_date": {
"type": "string",
"description": "End date for all-day events (exclusive, e.g., '2025-01-16'). Used by: create_event, update_event"
},
"timezone": {
"type": "string",
"description": "Timezone (e.g., 'America/New_York'). Used by: create_event, update_event"
},
"attendees": {
"type": "array",
"items": { "type": "string" },
"description": "Attendee email addresses. Used by: create_event, update_event"
}
]
}
}"#
.to_string()
}
+103 -234
View File
@@ -74,251 +74,120 @@ impl exports::near::agent::tool::Guest for GoogleDocsTool {
r#"{
"type": "object",
"required": ["action"],
"oneOf": [
{
"properties": {
"action": { "const": "create_document" },
"title": {
"type": "string",
"description": "Document title"
}
},
"required": ["action", "title"]
"properties": {
"action": {
"type": "string",
"enum": ["create_document", "get_document", "read_content", "insert_text", "delete_content", "replace_text", "format_text", "format_paragraph", "insert_table", "create_list", "batch_update"],
"description": "The Google Docs operation to perform"
},
{
"properties": {
"action": { "const": "get_document" },
"document_id": {
"type": "string",
"description": "The document ID (same as Google Drive file ID)"
}
},
"required": ["action", "document_id"]
"title": {
"type": "string",
"description": "Document title. Required for: create_document"
},
{
"properties": {
"action": { "const": "read_content" },
"document_id": {
"type": "string",
"description": "The document ID"
}
},
"required": ["action", "document_id"]
"document_id": {
"type": "string",
"description": "The document ID (same as Google Drive file ID). Required for all actions except create_document"
},
{
"properties": {
"action": { "const": "insert_text" },
"document_id": {
"type": "string",
"description": "The document ID"
},
"text": {
"type": "string",
"description": "Text to insert"
},
"index": {
"type": "integer",
"description": "Character index to insert at (1 for start of body). Use -1 to append at end.",
"default": -1
},
"segment_id": {
"type": "string",
"description": "Segment ID (empty string for body, or a header/footer ID)",
"default": ""
}
},
"required": ["action", "document_id", "text"]
"text": {
"type": "string",
"description": "Text to insert. Required for: insert_text"
},
{
"properties": {
"action": { "const": "delete_content" },
"document_id": {
"type": "string",
"description": "The document ID"
},
"start_index": {
"type": "integer",
"description": "Start index (inclusive)"
},
"end_index": {
"type": "integer",
"description": "End index (exclusive)"
},
"segment_id": {
"type": "string",
"description": "Segment ID (empty for body)",
"default": ""
}
},
"required": ["action", "document_id", "start_index", "end_index"]
"index": {
"type": "integer",
"description": "Character index (1 for start of body, -1 to append at end). Required for: insert_table. Used by: insert_text (default: -1)"
},
{
"properties": {
"action": { "const": "replace_text" },
"document_id": {
"type": "string",
"description": "The document ID"
},
"find": {
"type": "string",
"description": "Text to search for"
},
"replace": {
"type": "string",
"description": "Replacement text"
},
"match_case": {
"type": "boolean",
"description": "Case-sensitive match (default: true)",
"default": true
}
},
"required": ["action", "document_id", "find", "replace"]
"segment_id": {
"type": "string",
"description": "Segment ID (empty for body, or a header/footer ID). Used by: insert_text, delete_content",
"default": ""
},
{
"properties": {
"action": { "const": "format_text" },
"document_id": {
"type": "string",
"description": "The document ID"
},
"start_index": {
"type": "integer",
"description": "Start index (inclusive)"
},
"end_index": {
"type": "integer",
"description": "End index (exclusive)"
},
"bold": {
"type": "boolean",
"description": "Make text bold"
},
"italic": {
"type": "boolean",
"description": "Make text italic"
},
"underline": {
"type": "boolean",
"description": "Underline text"
},
"strikethrough": {
"type": "boolean",
"description": "Strikethrough text"
},
"font_size": {
"type": "number",
"description": "Font size in points (e.g., 12, 14, 18)"
},
"font_family": {
"type": "string",
"description": "Font family (e.g., 'Arial', 'Times New Roman', 'Courier New')"
},
"foreground_color": {
"type": "string",
"description": "Text color as hex (e.g., '#FF0000' for red)"
},
"background_color": {
"type": "string",
"description": "Text background/highlight color as hex"
}
},
"required": ["action", "document_id", "start_index", "end_index"]
"start_index": {
"type": "integer",
"description": "Start index (inclusive). Required for: delete_content, format_text, format_paragraph, create_list"
},
{
"properties": {
"action": { "const": "format_paragraph" },
"document_id": {
"type": "string",
"description": "The document ID"
},
"start_index": {
"type": "integer",
"description": "Start index (inclusive)"
},
"end_index": {
"type": "integer",
"description": "End index (exclusive)"
},
"named_style": {
"type": "string",
"enum": ["NORMAL_TEXT", "TITLE", "SUBTITLE", "HEADING_1", "HEADING_2", "HEADING_3", "HEADING_4", "HEADING_5", "HEADING_6"],
"description": "Paragraph style (heading level)"
},
"alignment": {
"type": "string",
"enum": ["START", "CENTER", "END", "JUSTIFIED"],
"description": "Text alignment"
},
"line_spacing": {
"type": "number",
"description": "Line spacing as percentage (e.g., 100 for single, 150 for 1.5x, 200 for double)"
}
},
"required": ["action", "document_id", "start_index", "end_index"]
"end_index": {
"type": "integer",
"description": "End index (exclusive). Required for: delete_content, format_text, format_paragraph, create_list"
},
{
"properties": {
"action": { "const": "insert_table" },
"document_id": {
"type": "string",
"description": "The document ID"
},
"rows": {
"type": "integer",
"description": "Number of rows"
},
"columns": {
"type": "integer",
"description": "Number of columns"
},
"index": {
"type": "integer",
"description": "Character index to insert the table at"
}
},
"required": ["action", "document_id", "rows", "columns", "index"]
"find": {
"type": "string",
"description": "Text to search for. Required for: replace_text"
},
{
"properties": {
"action": { "const": "create_list" },
"document_id": {
"type": "string",
"description": "The document ID"
},
"start_index": {
"type": "integer",
"description": "Start index (inclusive)"
},
"end_index": {
"type": "integer",
"description": "End index (exclusive)"
},
"bullet_preset": {
"type": "string",
"enum": ["BULLET_DISC_CIRCLE_SQUARE", "BULLET_CHECKBOX", "BULLET_ARROW_DIAMOND_DISC", "NUMBERED_DECIMAL_ALPHA_ROMAN", "NUMBERED_DECIMAL_NESTED", "NUMBERED_UPPERALPHA_ALPHA_ROMAN"],
"description": "Bullet style preset (default: BULLET_DISC_CIRCLE_SQUARE)",
"default": "BULLET_DISC_CIRCLE_SQUARE"
}
},
"required": ["action", "document_id", "start_index", "end_index"]
"replace": {
"type": "string",
"description": "Replacement text. Required for: replace_text"
},
{
"properties": {
"action": { "const": "batch_update" },
"document_id": {
"type": "string",
"description": "The document ID"
},
"requests": {
"type": "array",
"items": { "type": "object" },
"description": "Array of raw Docs API batchUpdate request objects"
}
},
"required": ["action", "document_id", "requests"]
"match_case": {
"type": "boolean",
"description": "Case-sensitive match (default: true). Used by: replace_text",
"default": true
},
"bold": {
"type": "boolean",
"description": "Make text bold. Used by: format_text"
},
"italic": {
"type": "boolean",
"description": "Make text italic. Used by: format_text"
},
"underline": {
"type": "boolean",
"description": "Underline text. Used by: format_text"
},
"strikethrough": {
"type": "boolean",
"description": "Strikethrough text. Used by: format_text"
},
"font_size": {
"type": "number",
"description": "Font size in points (e.g., 12, 14, 18). Used by: format_text"
},
"font_family": {
"type": "string",
"description": "Font family (e.g., 'Arial', 'Times New Roman'). Used by: format_text"
},
"foreground_color": {
"type": "string",
"description": "Text color as hex (e.g., '#FF0000'). Used by: format_text"
},
"background_color": {
"type": "string",
"description": "Text background/highlight color as hex. Used by: format_text"
},
"named_style": {
"type": "string",
"enum": ["NORMAL_TEXT", "TITLE", "SUBTITLE", "HEADING_1", "HEADING_2", "HEADING_3", "HEADING_4", "HEADING_5", "HEADING_6"],
"description": "Paragraph style (heading level). Used by: format_paragraph"
},
"alignment": {
"type": "string",
"enum": ["START", "CENTER", "END", "JUSTIFIED"],
"description": "Text alignment. Used by: format_paragraph"
},
"line_spacing": {
"type": "number",
"description": "Line spacing as percentage (100=single, 150=1.5x, 200=double). Used by: format_paragraph"
},
"rows": {
"type": "integer",
"description": "Number of rows. Required for: insert_table"
},
"columns": {
"type": "integer",
"description": "Number of columns. Required for: insert_table"
},
"bullet_preset": {
"type": "string",
"enum": ["BULLET_DISC_CIRCLE_SQUARE", "BULLET_CHECKBOX", "BULLET_ARROW_DIAMOND_DISC", "NUMBERED_DECIMAL_ALPHA_ROMAN", "NUMBERED_DECIMAL_NESTED", "NUMBERED_UPPERALPHA_ALPHA_ROMAN"],
"description": "Bullet style preset (default: BULLET_DISC_CIRCLE_SQUARE). Used by: create_list",
"default": "BULLET_DISC_CIRCLE_SQUARE"
},
"requests": {
"type": "array",
"items": { "type": "object" },
"description": "Array of raw Docs API batchUpdate request objects. Required for: batch_update"
}
]
}
}"#
.to_string()
}
+77 -197
View File
@@ -62,215 +62,95 @@ impl exports::near::agent::tool::Guest for GoogleDriveTool {
r#"{
"type": "object",
"required": ["action"],
"oneOf": [
{
"properties": {
"action": { "const": "list_files" },
"query": {
"type": "string",
"description": "Drive search query. Examples: \"name contains 'report'\", \"mimeType = 'application/pdf'\", \"'folderId' in parents\", \"sharedWithMe = true\""
},
"page_size": {
"type": "integer",
"description": "Max results (default: 25, max: 1000)",
"default": 25
},
"order_by": {
"type": "string",
"description": "Sort order (e.g., 'modifiedTime desc', 'name')"
},
"corpora": {
"type": "string",
"enum": ["user", "drive", "domain", "allDrives"],
"description": "Search scope: 'user' (personal, default), 'drive' (specific shared drive), 'domain' (org-wide), 'allDrives' (everything)",
"default": "user"
},
"drive_id": {
"type": "string",
"description": "Shared drive ID (required when corpora is 'drive')"
},
"page_token": {
"type": "string",
"description": "Token for next page of results"
}
},
"required": ["action"]
"properties": {
"action": {
"type": "string",
"enum": ["list_files", "get_file", "download_file", "upload_file", "update_file", "create_folder", "delete_file", "trash_file", "share_file", "list_permissions", "remove_permission", "list_shared_drives"],
"description": "The Google Drive operation to perform"
},
{
"properties": {
"action": { "const": "get_file" },
"file_id": {
"type": "string",
"description": "The file ID"
}
},
"required": ["action", "file_id"]
"file_id": {
"type": "string",
"description": "File ID. Required for: get_file, download_file, update_file, delete_file, trash_file, share_file, list_permissions, remove_permission"
},
{
"properties": {
"action": { "const": "download_file" },
"file_id": {
"type": "string",
"description": "The file ID to download"
},
"export_mime_type": {
"type": "string",
"description": "Export format for Google Workspace files (e.g., 'text/plain', 'text/csv', 'application/pdf')"
}
},
"required": ["action", "file_id"]
"query": {
"type": "string",
"description": "Drive search query (e.g., \"name contains 'report'\", \"mimeType = 'application/pdf'\"). Used by: list_files"
},
{
"properties": {
"action": { "const": "upload_file" },
"name": {
"type": "string",
"description": "File name"
},
"content": {
"type": "string",
"description": "File content (text)"
},
"mime_type": {
"type": "string",
"description": "MIME type (default: 'text/plain')",
"default": "text/plain"
},
"parent_id": {
"type": "string",
"description": "Parent folder ID (omit for root)"
},
"description": {
"type": "string",
"description": "File description"
}
},
"required": ["action", "name", "content"]
"page_size": {
"type": "integer",
"description": "Max results (default: 25, max: 1000). Used by: list_files, list_shared_drives",
"default": 25
},
{
"properties": {
"action": { "const": "update_file" },
"file_id": {
"type": "string",
"description": "The file ID to update"
},
"name": {
"type": "string",
"description": "New file name"
},
"description": {
"type": "string",
"description": "New description"
},
"move_to_parent": {
"type": "string",
"description": "Move file to this folder ID"
},
"starred": {
"type": "boolean",
"description": "Star or unstar the file"
}
},
"required": ["action", "file_id"]
"order_by": {
"type": "string",
"description": "Sort order (e.g., 'modifiedTime desc', 'name'). Used by: list_files"
},
{
"properties": {
"action": { "const": "create_folder" },
"name": {
"type": "string",
"description": "Folder name"
},
"parent_id": {
"type": "string",
"description": "Parent folder ID (omit for root)"
},
"description": {
"type": "string",
"description": "Folder description"
}
},
"required": ["action", "name"]
"corpora": {
"type": "string",
"enum": ["user", "drive", "domain", "allDrives"],
"description": "Search scope: 'user' (default), 'drive' (shared drive), 'domain', 'allDrives'. Used by: list_files",
"default": "user"
},
{
"properties": {
"action": { "const": "delete_file" },
"file_id": {
"type": "string",
"description": "The file ID to permanently delete"
}
},
"required": ["action", "file_id"]
"drive_id": {
"type": "string",
"description": "Shared drive ID (required when corpora is 'drive'). Used by: list_files"
},
{
"properties": {
"action": { "const": "trash_file" },
"file_id": {
"type": "string",
"description": "The file ID to move to trash"
}
},
"required": ["action", "file_id"]
"page_token": {
"type": "string",
"description": "Token for next page of results. Used by: list_files"
},
{
"properties": {
"action": { "const": "share_file" },
"file_id": {
"type": "string",
"description": "The file ID to share"
},
"email": {
"type": "string",
"description": "Recipient email address"
},
"role": {
"type": "string",
"enum": ["reader", "commenter", "writer", "organizer"],
"description": "Permission level (default: 'reader')",
"default": "reader"
},
"message": {
"type": "string",
"description": "Optional message in sharing notification"
}
},
"required": ["action", "file_id", "email"]
"export_mime_type": {
"type": "string",
"description": "Export format for Google Workspace files (e.g., 'text/plain', 'text/csv'). Used by: download_file"
},
{
"properties": {
"action": { "const": "list_permissions" },
"file_id": {
"type": "string",
"description": "The file ID to check permissions for"
}
},
"required": ["action", "file_id"]
"name": {
"type": "string",
"description": "File/folder name. Required for: upload_file, create_folder. Optional for: update_file"
},
{
"properties": {
"action": { "const": "remove_permission" },
"file_id": {
"type": "string",
"description": "The file ID"
},
"permission_id": {
"type": "string",
"description": "The permission ID to remove (get from list_permissions)"
}
},
"required": ["action", "file_id", "permission_id"]
"content": {
"type": "string",
"description": "File content (text). Required for: upload_file"
},
{
"properties": {
"action": { "const": "list_shared_drives" },
"page_size": {
"type": "integer",
"description": "Max results (default: 25)",
"default": 25
}
},
"required": ["action"]
"mime_type": {
"type": "string",
"description": "MIME type (default: 'text/plain'). Used by: upload_file",
"default": "text/plain"
},
"parent_id": {
"type": "string",
"description": "Parent folder ID (omit for root). Used by: upload_file, create_folder"
},
"description": {
"type": "string",
"description": "File/folder description. Used by: upload_file, update_file, create_folder"
},
"move_to_parent": {
"type": "string",
"description": "Move file to this folder ID. Used by: update_file"
},
"starred": {
"type": "boolean",
"description": "Star or unstar the file. Used by: update_file"
},
"email": {
"type": "string",
"description": "Recipient email address. Required for: share_file"
},
"role": {
"type": "string",
"enum": ["reader", "commenter", "writer", "organizer"],
"description": "Permission level (default: 'reader'). Used by: share_file",
"default": "reader"
},
"message": {
"type": "string",
"description": "Optional message in sharing notification. Used by: share_file"
},
"permission_id": {
"type": "string",
"description": "Permission ID to remove (from list_permissions). Required for: remove_permission"
}
]
}
}"#
.to_string()
}
+83 -219
View File
@@ -70,236 +70,100 @@ impl exports::near::agent::tool::Guest for GoogleSheetsTool {
r#"{
"type": "object",
"required": ["action"],
"oneOf": [
{
"properties": {
"action": { "const": "create_spreadsheet" },
"title": {
"type": "string",
"description": "Spreadsheet title"
},
"sheet_names": {
"type": "array",
"items": { "type": "string" },
"description": "Names for sheets (tabs). Defaults to ['Sheet1'] if omitted."
}
},
"required": ["action", "title"]
"properties": {
"action": {
"type": "string",
"enum": ["create_spreadsheet", "get_spreadsheet", "read_values", "batch_read_values", "write_values", "append_values", "clear_values", "add_sheet", "delete_sheet", "rename_sheet", "format_cells"],
"description": "The Google Sheets operation to perform"
},
{
"properties": {
"action": { "const": "get_spreadsheet" },
"spreadsheet_id": {
"type": "string",
"description": "The spreadsheet ID (same as Google Drive file ID)"
}
},
"required": ["action", "spreadsheet_id"]
"spreadsheet_id": {
"type": "string",
"description": "Spreadsheet ID (same as Google Drive file ID). Required for all actions except create_spreadsheet"
},
{
"properties": {
"action": { "const": "read_values" },
"spreadsheet_id": {
"type": "string",
"description": "The spreadsheet ID"
},
"range": {
"type": "string",
"description": "A1 notation range (e.g., 'Sheet1!A1:D10', 'A1:B5')"
}
},
"required": ["action", "spreadsheet_id", "range"]
"title": {
"type": "string",
"description": "Title/name. Required for: create_spreadsheet, add_sheet, rename_sheet"
},
{
"properties": {
"action": { "const": "batch_read_values" },
"spreadsheet_id": {
"type": "string",
"description": "The spreadsheet ID"
},
"ranges": {
"type": "array",
"items": { "type": "string" },
"description": "List of A1 notation ranges to read"
}
},
"required": ["action", "spreadsheet_id", "ranges"]
"sheet_names": {
"type": "array",
"items": { "type": "string" },
"description": "Names for sheets (tabs, defaults to ['Sheet1']). Used by: create_spreadsheet"
},
{
"properties": {
"action": { "const": "write_values" },
"spreadsheet_id": {
"type": "string",
"description": "The spreadsheet ID"
},
"range": {
"type": "string",
"description": "A1 notation range (e.g., 'Sheet1!A1')"
},
"values": {
"type": "array",
"items": { "type": "array" },
"description": "2D array of values (rows of columns)"
},
"value_input_option": {
"type": "string",
"enum": ["RAW", "USER_ENTERED"],
"description": "How to interpret input. USER_ENTERED (default) parses like typing in the UI. RAW stores as-is.",
"default": "USER_ENTERED"
}
},
"required": ["action", "spreadsheet_id", "range", "values"]
"range": {
"type": "string",
"description": "A1 notation range (e.g., 'Sheet1!A1:D10'). Required for: read_values, write_values, append_values, clear_values"
},
{
"properties": {
"action": { "const": "append_values" },
"spreadsheet_id": {
"type": "string",
"description": "The spreadsheet ID"
},
"range": {
"type": "string",
"description": "A1 notation range to find the table (e.g., 'Sheet1!A:E')"
},
"values": {
"type": "array",
"items": { "type": "array" },
"description": "Rows to append (2D array)"
},
"value_input_option": {
"type": "string",
"enum": ["RAW", "USER_ENTERED"],
"description": "How to interpret input (default: USER_ENTERED)",
"default": "USER_ENTERED"
}
},
"required": ["action", "spreadsheet_id", "range", "values"]
"ranges": {
"type": "array",
"items": { "type": "string" },
"description": "List of A1 notation ranges. Required for: batch_read_values"
},
{
"properties": {
"action": { "const": "clear_values" },
"spreadsheet_id": {
"type": "string",
"description": "The spreadsheet ID"
},
"range": {
"type": "string",
"description": "A1 notation range to clear"
}
},
"required": ["action", "spreadsheet_id", "range"]
"values": {
"type": "array",
"items": { "type": "array" },
"description": "2D array of values (rows of columns). Required for: write_values, append_values"
},
{
"properties": {
"action": { "const": "add_sheet" },
"spreadsheet_id": {
"type": "string",
"description": "The spreadsheet ID"
},
"title": {
"type": "string",
"description": "Name for the new sheet (tab)"
}
},
"required": ["action", "spreadsheet_id", "title"]
"value_input_option": {
"type": "string",
"enum": ["RAW", "USER_ENTERED"],
"description": "How to interpret input (USER_ENTERED parses like the UI, RAW stores as-is, default: USER_ENTERED). Used by: write_values, append_values",
"default": "USER_ENTERED"
},
{
"properties": {
"action": { "const": "delete_sheet" },
"spreadsheet_id": {
"type": "string",
"description": "The spreadsheet ID"
},
"sheet_id": {
"type": "integer",
"description": "Numeric sheet ID (get from get_spreadsheet, NOT the sheet name)"
}
},
"required": ["action", "spreadsheet_id", "sheet_id"]
"sheet_id": {
"type": "integer",
"description": "Numeric sheet ID (from get_spreadsheet, NOT the sheet name). Required for: delete_sheet, rename_sheet, format_cells"
},
{
"properties": {
"action": { "const": "rename_sheet" },
"spreadsheet_id": {
"type": "string",
"description": "The spreadsheet ID"
},
"sheet_id": {
"type": "integer",
"description": "Numeric sheet ID"
},
"title": {
"type": "string",
"description": "New name for the sheet"
}
},
"required": ["action", "spreadsheet_id", "sheet_id", "title"]
"start_row": {
"type": "integer",
"description": "Start row (0-indexed, inclusive). Required for: format_cells"
},
{
"properties": {
"action": { "const": "format_cells" },
"spreadsheet_id": {
"type": "string",
"description": "The spreadsheet ID"
},
"sheet_id": {
"type": "integer",
"description": "Numeric sheet ID"
},
"start_row": {
"type": "integer",
"description": "Start row (0-indexed, inclusive)"
},
"end_row": {
"type": "integer",
"description": "End row (0-indexed, exclusive)"
},
"start_column": {
"type": "integer",
"description": "Start column (0-indexed, inclusive)"
},
"end_column": {
"type": "integer",
"description": "End column (0-indexed, exclusive)"
},
"bold": {
"type": "boolean",
"description": "Make text bold"
},
"italic": {
"type": "boolean",
"description": "Make text italic"
},
"font_size": {
"type": "integer",
"description": "Font size in points"
},
"text_color": {
"type": "string",
"description": "Text color as hex (e.g., '#FF0000' for red)"
},
"background_color": {
"type": "string",
"description": "Cell background color as hex (e.g., '#FFFF00' for yellow)"
},
"horizontal_alignment": {
"type": "string",
"enum": ["LEFT", "CENTER", "RIGHT"],
"description": "Horizontal text alignment"
},
"number_format": {
"type": "string",
"description": "Number format pattern (e.g., '#,##0.00', 'yyyy-mm-dd', '$#,##0')"
},
"number_format_type": {
"type": "string",
"enum": ["NUMBER", "CURRENCY", "PERCENT", "DATE", "TIME", "TEXT"],
"description": "Type of number format (default: NUMBER)"
}
},
"required": ["action", "spreadsheet_id", "sheet_id", "start_row", "end_row", "start_column", "end_column"]
"end_row": {
"type": "integer",
"description": "End row (0-indexed, exclusive). Required for: format_cells"
},
"start_column": {
"type": "integer",
"description": "Start column (0-indexed, inclusive). Required for: format_cells"
},
"end_column": {
"type": "integer",
"description": "End column (0-indexed, exclusive). Required for: format_cells"
},
"bold": {
"type": "boolean",
"description": "Make text bold. Used by: format_cells"
},
"italic": {
"type": "boolean",
"description": "Make text italic. Used by: format_cells"
},
"font_size": {
"type": "integer",
"description": "Font size in points. Used by: format_cells"
},
"text_color": {
"type": "string",
"description": "Text color as hex (e.g., '#FF0000'). Used by: format_cells"
},
"background_color": {
"type": "string",
"description": "Cell background color as hex (e.g., '#FFFF00'). Used by: format_cells"
},
"horizontal_alignment": {
"type": "string",
"enum": ["LEFT", "CENTER", "RIGHT"],
"description": "Horizontal text alignment. Used by: format_cells"
},
"number_format": {
"type": "string",
"description": "Number format pattern (e.g., '#,##0.00', 'yyyy-mm-dd'). Used by: format_cells"
},
"number_format_type": {
"type": "string",
"enum": ["NUMBER", "CURRENCY", "PERCENT", "DATE", "TIME", "TEXT"],
"description": "Type of number format (default: NUMBER). Used by: format_cells"
}
]
}
}"#
.to_string()
}
+104 -306
View File
@@ -79,326 +79,124 @@ impl exports::near::agent::tool::Guest for GoogleSlidesTool {
r#"{
"type": "object",
"required": ["action"],
"oneOf": [
{
"properties": {
"action": { "const": "create_presentation" },
"title": {
"type": "string",
"description": "Presentation title"
}
},
"required": ["action", "title"]
"properties": {
"action": {
"type": "string",
"enum": ["create_presentation", "get_presentation", "get_thumbnail", "create_slide", "delete_object", "insert_text", "delete_text", "replace_all_text", "create_shape", "insert_image", "format_text", "format_paragraph", "replace_shapes_with_image", "batch_update"],
"description": "The Google Slides operation to perform"
},
{
"properties": {
"action": { "const": "get_presentation" },
"presentation_id": {
"type": "string",
"description": "The presentation ID (same as Google Drive file ID)"
}
},
"required": ["action", "presentation_id"]
"title": {
"type": "string",
"description": "Presentation title. Required for: create_presentation"
},
{
"properties": {
"action": { "const": "get_thumbnail" },
"presentation_id": {
"type": "string",
"description": "The presentation ID"
},
"slide_object_id": {
"type": "string",
"description": "The slide's object ID"
}
},
"required": ["action", "presentation_id", "slide_object_id"]
"presentation_id": {
"type": "string",
"description": "Presentation ID (same as Google Drive file ID). Required for all actions except create_presentation"
},
{
"properties": {
"action": { "const": "create_slide" },
"presentation_id": {
"type": "string",
"description": "The presentation ID"
},
"insertion_index": {
"type": "integer",
"description": "Position to insert (0-based). Omit to append at end."
},
"layout": {
"type": "string",
"enum": ["BLANK", "TITLE", "TITLE_AND_BODY", "TITLE_AND_TWO_COLUMNS", "TITLE_ONLY", "SECTION_HEADER", "CAPTION_ONLY", "BIG_NUMBER", "ONE_COLUMN_TEXT", "MAIN_POINT"],
"description": "Predefined layout (default: BLANK)",
"default": "BLANK"
}
},
"required": ["action", "presentation_id"]
"slide_object_id": {
"type": "string",
"description": "Slide object ID. Required for: get_thumbnail, create_shape, insert_image"
},
{
"properties": {
"action": { "const": "delete_object" },
"presentation_id": {
"type": "string",
"description": "The presentation ID"
},
"object_id": {
"type": "string",
"description": "Object ID of the slide or element to delete"
}
},
"required": ["action", "presentation_id", "object_id"]
"object_id": {
"type": "string",
"description": "Object ID of a slide element. Required for: delete_object, insert_text, delete_text, format_text, format_paragraph"
},
{
"properties": {
"action": { "const": "insert_text" },
"presentation_id": {
"type": "string",
"description": "The presentation ID"
},
"object_id": {
"type": "string",
"description": "Object ID of the shape or text box"
},
"text": {
"type": "string",
"description": "Text to insert"
},
"insertion_index": {
"type": "integer",
"description": "Character index to insert at (0-based). Default: 0.",
"default": 0
}
},
"required": ["action", "presentation_id", "object_id", "text"]
"text": {
"type": "string",
"description": "Text to insert. Required for: insert_text"
},
{
"properties": {
"action": { "const": "delete_text" },
"presentation_id": {
"type": "string",
"description": "The presentation ID"
},
"object_id": {
"type": "string",
"description": "Object ID of the shape"
},
"start_index": {
"type": "integer",
"description": "Start index (inclusive, 0-based)",
"default": 0
},
"end_index": {
"type": "integer",
"description": "End index (exclusive). Omit to delete from start_index to end."
}
},
"required": ["action", "presentation_id", "object_id"]
"insertion_index": {
"type": "integer",
"description": "Position to insert at (0-based). Used by: create_slide (omit to append at end), insert_text (default: 0)"
},
{
"properties": {
"action": { "const": "replace_all_text" },
"presentation_id": {
"type": "string",
"description": "The presentation ID"
},
"find": {
"type": "string",
"description": "Text to search for"
},
"replace": {
"type": "string",
"description": "Replacement text"
},
"match_case": {
"type": "boolean",
"description": "Case-sensitive match (default: true)",
"default": true
}
},
"required": ["action", "presentation_id", "find", "replace"]
"layout": {
"type": "string",
"enum": ["BLANK", "TITLE", "TITLE_AND_BODY", "TITLE_AND_TWO_COLUMNS", "TITLE_ONLY", "SECTION_HEADER", "CAPTION_ONLY", "BIG_NUMBER", "ONE_COLUMN_TEXT", "MAIN_POINT"],
"description": "Predefined slide layout (default: BLANK). Used by: create_slide",
"default": "BLANK"
},
{
"properties": {
"action": { "const": "create_shape" },
"presentation_id": {
"type": "string",
"description": "The presentation ID"
},
"slide_object_id": {
"type": "string",
"description": "Slide object ID to place the shape on"
},
"shape_type": {
"type": "string",
"enum": ["TEXT_BOX", "RECTANGLE", "ROUND_RECTANGLE", "ELLIPSE"],
"description": "Shape type (default: TEXT_BOX)",
"default": "TEXT_BOX"
},
"x": {
"type": "number",
"description": "X position in points from left edge"
},
"y": {
"type": "number",
"description": "Y position in points from top edge"
},
"width": {
"type": "number",
"description": "Width in points"
},
"height": {
"type": "number",
"description": "Height in points"
}
},
"required": ["action", "presentation_id", "slide_object_id", "x", "y", "width", "height"]
"start_index": {
"type": "integer",
"description": "Start index (inclusive, 0-based). Used by: delete_text, format_text, format_paragraph"
},
{
"properties": {
"action": { "const": "insert_image" },
"presentation_id": {
"type": "string",
"description": "The presentation ID"
},
"slide_object_id": {
"type": "string",
"description": "Slide object ID to place the image on"
},
"image_url": {
"type": "string",
"description": "Publicly accessible image URL"
},
"x": {
"type": "number",
"description": "X position in points"
},
"y": {
"type": "number",
"description": "Y position in points"
},
"width": {
"type": "number",
"description": "Width in points"
},
"height": {
"type": "number",
"description": "Height in points"
}
},
"required": ["action", "presentation_id", "slide_object_id", "image_url", "x", "y", "width", "height"]
"end_index": {
"type": "integer",
"description": "End index (exclusive). Used by: delete_text, format_text, format_paragraph"
},
{
"properties": {
"action": { "const": "format_text" },
"presentation_id": {
"type": "string",
"description": "The presentation ID"
},
"object_id": {
"type": "string",
"description": "Object ID of the shape"
},
"start_index": {
"type": "integer",
"description": "Start index (inclusive). Omit to format all text."
},
"end_index": {
"type": "integer",
"description": "End index (exclusive). Omit to format to end."
},
"bold": {
"type": "boolean",
"description": "Make text bold"
},
"italic": {
"type": "boolean",
"description": "Make text italic"
},
"underline": {
"type": "boolean",
"description": "Underline text"
},
"font_size": {
"type": "number",
"description": "Font size in points (e.g., 12, 18, 24)"
},
"font_family": {
"type": "string",
"description": "Font family (e.g., 'Arial', 'Roboto', 'Times New Roman')"
},
"foreground_color": {
"type": "string",
"description": "Text color as hex (e.g., '#FF0000' for red)"
}
},
"required": ["action", "presentation_id", "object_id"]
"find": {
"type": "string",
"description": "Text to search for. Required for: replace_all_text, replace_shapes_with_image"
},
{
"properties": {
"action": { "const": "format_paragraph" },
"presentation_id": {
"type": "string",
"description": "The presentation ID"
},
"object_id": {
"type": "string",
"description": "Object ID of the shape"
},
"alignment": {
"type": "string",
"enum": ["START", "CENTER", "END", "JUSTIFIED"],
"description": "Paragraph alignment"
},
"start_index": {
"type": "integer",
"description": "Start index (inclusive). Omit to format all."
},
"end_index": {
"type": "integer",
"description": "End index (exclusive). Omit to format to end."
}
},
"required": ["action", "presentation_id", "object_id", "alignment"]
"replace": {
"type": "string",
"description": "Replacement text. Required for: replace_all_text"
},
{
"properties": {
"action": { "const": "replace_shapes_with_image" },
"presentation_id": {
"type": "string",
"description": "The presentation ID"
},
"find": {
"type": "string",
"description": "Text to match in shapes"
},
"image_url": {
"type": "string",
"description": "Image URL to replace matched shapes with"
},
"match_case": {
"type": "boolean",
"description": "Case-sensitive match (default: true)",
"default": true
}
},
"required": ["action", "presentation_id", "find", "image_url"]
"match_case": {
"type": "boolean",
"description": "Case-sensitive match (default: true). Used by: replace_all_text, replace_shapes_with_image",
"default": true
},
{
"properties": {
"action": { "const": "batch_update" },
"presentation_id": {
"type": "string",
"description": "The presentation ID"
},
"requests": {
"type": "array",
"items": { "type": "object" },
"description": "Array of raw Slides API batchUpdate request objects"
}
},
"required": ["action", "presentation_id", "requests"]
"shape_type": {
"type": "string",
"enum": ["TEXT_BOX", "RECTANGLE", "ROUND_RECTANGLE", "ELLIPSE"],
"description": "Shape type (default: TEXT_BOX). Used by: create_shape",
"default": "TEXT_BOX"
},
"x": {
"type": "number",
"description": "X position in points from left edge. Required for: create_shape, insert_image"
},
"y": {
"type": "number",
"description": "Y position in points from top edge. Required for: create_shape, insert_image"
},
"width": {
"type": "number",
"description": "Width in points. Required for: create_shape, insert_image"
},
"height": {
"type": "number",
"description": "Height in points. Required for: create_shape, insert_image"
},
"image_url": {
"type": "string",
"description": "Publicly accessible image URL. Required for: insert_image, replace_shapes_with_image"
},
"bold": {
"type": "boolean",
"description": "Make text bold. Used by: format_text"
},
"italic": {
"type": "boolean",
"description": "Make text italic. Used by: format_text"
},
"underline": {
"type": "boolean",
"description": "Underline text. Used by: format_text"
},
"font_size": {
"type": "number",
"description": "Font size in points (e.g., 12, 18, 24). Used by: format_text"
},
"font_family": {
"type": "string",
"description": "Font family (e.g., 'Arial', 'Roboto'). Used by: format_text"
},
"foreground_color": {
"type": "string",
"description": "Text color as hex (e.g., '#FF0000'). Used by: format_text"
},
"alignment": {
"type": "string",
"enum": ["START", "CENTER", "END", "JUSTIFIED"],
"description": "Paragraph alignment. Required for: format_paragraph"
},
"requests": {
"type": "array",
"items": { "type": "object" },
"description": "Array of raw Slides API batchUpdate request objects. Required for: batch_update"
}
]
}
}"#
.to_string()
}
+15 -46
View File
@@ -54,56 +54,25 @@ impl exports::near::agent::tool::Guest for OktaTool {
r#"{
"type": "object",
"required": ["action"],
"oneOf": [
{
"properties": {
"action": { "const": "get_profile" }
},
"required": ["action"]
"properties": {
"action": {
"type": "string",
"enum": ["get_profile", "update_profile", "list_apps", "search_apps", "get_app_sso_link", "get_org_info"],
"description": "The Okta operation to perform"
},
{
"properties": {
"action": { "const": "update_profile" },
"fields": {
"type": "object",
"description": "Profile fields to update. Common: firstName, lastName, email, mobilePhone, displayName, nickName, title, department, organization"
}
},
"required": ["action", "fields"]
"fields": {
"type": "object",
"description": "Profile fields to update (e.g., firstName, lastName, email, mobilePhone, displayName, title, department). Required for: update_profile"
},
{
"properties": {
"action": { "const": "list_apps" }
},
"required": ["action"]
"query": {
"type": "string",
"description": "Case-insensitive search query to match against app labels and names. Required for: search_apps"
},
{
"properties": {
"action": { "const": "search_apps" },
"query": {
"type": "string",
"description": "Case-insensitive search query to match against app labels and names"
}
},
"required": ["action", "query"]
},
{
"properties": {
"action": { "const": "get_app_sso_link" },
"app": {
"type": "string",
"description": "App instance ID (e.g., '0oa1xxx') or app label (e.g., 'Google Workspace')"
}
},
"required": ["action", "app"]
},
{
"properties": {
"action": { "const": "get_org_info" }
},
"required": ["action"]
"app": {
"type": "string",
"description": "App instance ID (e.g., '0oa1xxx') or app label (e.g., 'Google Workspace'). Required for: get_app_sso_link"
}
]
}
}"#
.to_string()
}
+30 -70
View File
@@ -53,84 +53,44 @@ impl exports::near::agent::tool::Guest for SlackTool {
}
fn schema() -> String {
// JSON Schema for the tool's parameters
r#"{
"type": "object",
"required": ["action"],
"oneOf": [
{
"properties": {
"action": { "const": "send_message" },
"channel": {
"type": "string",
"description": "Channel ID or name (e.g., '#general' or 'C1234567890')"
},
"text": {
"type": "string",
"description": "Message text (supports Slack mrkdwn formatting)"
},
"thread_ts": {
"type": "string",
"description": "Optional thread timestamp to reply in a thread"
}
},
"required": ["action", "channel", "text"]
"properties": {
"action": {
"type": "string",
"enum": ["send_message", "list_channels", "get_channel_history", "post_reaction", "get_user_info"],
"description": "The Slack operation to perform"
},
{
"properties": {
"action": { "const": "list_channels" },
"limit": {
"type": "integer",
"description": "Maximum number of channels to return (default: 100)",
"default": 100
}
},
"required": ["action"]
"channel": {
"type": "string",
"description": "Channel ID or name (e.g., '#general' or 'C1234567890'). Required for: send_message, get_channel_history, post_reaction"
},
{
"properties": {
"action": { "const": "get_channel_history" },
"channel": {
"type": "string",
"description": "Channel ID (e.g., 'C1234567890')"
},
"limit": {
"type": "integer",
"description": "Maximum number of messages to return (default: 20)",
"default": 20
}
},
"required": ["action", "channel"]
"text": {
"type": "string",
"description": "Message text (supports Slack mrkdwn formatting). Required for: send_message"
},
{
"properties": {
"action": { "const": "post_reaction" },
"channel": {
"type": "string",
"description": "Channel ID containing the message"
},
"timestamp": {
"type": "string",
"description": "Timestamp of the message to react to"
},
"emoji": {
"type": "string",
"description": "Emoji name without colons (e.g., 'thumbsup')"
}
},
"required": ["action", "channel", "timestamp", "emoji"]
"thread_ts": {
"type": "string",
"description": "Thread timestamp to reply in a thread. Used by: send_message"
},
{
"properties": {
"action": { "const": "get_user_info" },
"user_id": {
"type": "string",
"description": "User ID (e.g., 'U1234567890')"
}
},
"required": ["action", "user_id"]
"limit": {
"type": "integer",
"description": "Maximum number of results to return. Used by: list_channels, get_channel_history"
},
"timestamp": {
"type": "string",
"description": "Timestamp of the message to react to. Required for: post_reaction"
},
"emoji": {
"type": "string",
"description": "Emoji name without colons (e.g., 'thumbsup'). Required for: post_reaction"
},
"user_id": {
"type": "string",
"description": "User ID (e.g., 'U1234567890'). Required for: get_user_info"
}
]
}
}"#
.to_string()
}
+46 -136
View File
@@ -248,154 +248,64 @@ fn get_api_hash() -> Result<String, String> {
const SCHEMA: &str = r#"{
"type": "object",
"required": ["action"],
"oneOf": [
{
"properties": {
"action": { "const": "login" },
"phone_number": {
"type": "string",
"description": "Phone number in international format (e.g., '+1234567890')"
}
},
"required": ["action", "phone_number"]
"properties": {
"action": {
"type": "string",
"enum": ["login", "submit_auth_code", "submit_2fa_password", "get_me", "get_contacts", "get_chats", "get_messages", "send_message", "forward_message", "delete_message", "search_messages", "get_updates"],
"description": "The Telegram operation to perform"
},
{
"properties": {
"action": { "const": "submit_auth_code" },
"code": {
"type": "string",
"description": "Verification code received via SMS or Telegram"
}
},
"required": ["action", "code"]
"phone_number": {
"type": "string",
"description": "Phone number in international format (e.g., '+1234567890'). Required for: login"
},
{
"properties": {
"action": { "const": "submit_2fa_password" },
"password": {
"type": "string",
"description": "Two-factor authentication password"
}
},
"required": ["action", "password"]
"code": {
"type": "string",
"description": "Verification code received via SMS or Telegram. Required for: submit_auth_code"
},
{
"properties": {
"action": { "const": "get_me" }
},
"required": ["action"]
"password": {
"type": "string",
"description": "Two-factor authentication password. Required for: submit_2fa_password"
},
{
"properties": {
"action": { "const": "get_contacts" }
},
"required": ["action"]
"chat_id": {
"type": "integer",
"description": "Chat ID (negative for groups/channels). Required for: get_messages, send_message. Optional for: search_messages"
},
{
"properties": {
"action": { "const": "get_chats" },
"limit": {
"type": "integer",
"description": "Maximum number of chats to return (default: 20)",
"default": 20
}
},
"required": ["action"]
"limit": {
"type": "integer",
"description": "Maximum number of results (default: 20). Used by: get_chats, get_messages, search_messages",
"default": 20
},
{
"properties": {
"action": { "const": "get_messages" },
"chat_id": {
"type": "integer",
"description": "Chat ID (negative for groups/channels)"
},
"limit": {
"type": "integer",
"description": "Maximum number of messages (default: 20)",
"default": 20
},
"from_message_id": {
"type": "integer",
"description": "Start from this message ID for pagination"
}
},
"required": ["action", "chat_id"]
"from_message_id": {
"type": "integer",
"description": "Start from this message ID for pagination. Used by: get_messages"
},
{
"properties": {
"action": { "const": "send_message" },
"chat_id": {
"type": "integer",
"description": "Chat ID to send the message to"
},
"text": {
"type": "string",
"description": "Message text"
}
},
"required": ["action", "chat_id", "text"]
"text": {
"type": "string",
"description": "Message text. Required for: send_message"
},
{
"properties": {
"action": { "const": "forward_message" },
"from_chat_id": {
"type": "integer",
"description": "Source chat ID"
},
"to_chat_id": {
"type": "integer",
"description": "Destination chat ID"
},
"message_ids": {
"type": "array",
"items": { "type": "integer" },
"description": "Message IDs to forward"
}
},
"required": ["action", "from_chat_id", "to_chat_id", "message_ids"]
"from_chat_id": {
"type": "integer",
"description": "Source chat ID. Required for: forward_message"
},
{
"properties": {
"action": { "const": "delete_message" },
"message_ids": {
"type": "array",
"items": { "type": "integer" },
"description": "Message IDs to delete"
},
"revoke": {
"type": "boolean",
"description": "Also delete for other participants (default: false)",
"default": false
}
},
"required": ["action", "message_ids"]
"to_chat_id": {
"type": "integer",
"description": "Destination chat ID. Required for: forward_message"
},
{
"properties": {
"action": { "const": "search_messages" },
"query": {
"type": "string",
"description": "Search query"
},
"chat_id": {
"type": "integer",
"description": "Chat ID to search within (omit for global search)"
},
"limit": {
"type": "integer",
"description": "Maximum number of results (default: 20)",
"default": 20
}
},
"required": ["action", "query"]
"message_ids": {
"type": "array",
"items": { "type": "integer" },
"description": "Message IDs. Required for: forward_message, delete_message"
},
{
"properties": {
"action": { "const": "get_updates" }
},
"required": ["action"]
"revoke": {
"type": "boolean",
"description": "Also delete for other participants (default: false). Used by: delete_message",
"default": false
},
"query": {
"type": "string",
"description": "Search query. Required for: search_messages"
}
]
}
}"#;
export!(TelegramTool);