From f0a0642e7d7c6e686ef1bb5e1ffe6227de6e83f0 Mon Sep 17 00:00:00 2001 From: Illia Polosukhin Date: Sun, 15 Feb 2026 00:24:51 -0800 Subject: [PATCH] feat: multi-provider inference + libSQL onboarding selection (#92) * feat: add interactive database backend selection during onboarding Previously the onboarding wizard silently defaulted to PostgreSQL because libsql wasn't in the default feature set. Now both backends ship by default and the wizard presents a selection prompt when both are available. DATABASE_BACKEND env var still bypasses the prompt for headless/CI use. Co-Authored-By: Claude Opus 4.6 * fix: resolve libSQL onboarding crash, keychain double-prompt, and setup audit findings Three bugs fixed: 1. libSQL onboarding crash ("Missing required setting 'database_url'"): DatabaseConfig::resolve() only checked DATABASE_BACKEND env var, falling back to Postgres default. Now reads settings.database_backend, plus settings.libsql_path and settings.libsql_url as fallbacks. 2. OS keychain prompts twice during startup: Config::from_env() and Config::from_db() both called get_master_key(). Now caches the key in SECRETS_MASTER_KEY env var after first read so from_db() skips keychain. 3. "Path not found: nearai.session" warning: from_db_map() tried to apply app-specific DB keys (nearai.session_token) to the Settings struct. Now skips keys that don't map to known Settings fields. Also fixed bootstrap migration key mismatch (nearai.session -> nearai.session_token). Setup module audit fixes (14 findings): - Replace unreachable!() with proper error in provider match - Extract setup_api_key_provider() to deduplicate setup_anthropic/setup_openai - Add SAFETY comments to all unsafe std::env::set_var blocks - Fix .unwrap() calls with proper error handling - Remove incorrect #[allow(dead_code)] on used TelegramUpdate::update_id - Log warnings instead of silently discarding HTTP errors in Telegram binding - Guard select_many against empty options, fix mask_api_key for non-ASCII - Update stale doc comment in mod.rs, rename misleading variable - Add 7 new tests (model fetcher fallbacks, channel discovery, secret gen) Co-Authored-By: Claude Opus 4.6 * fix: address PR review feedback (set_var safety, parse warnings, db_map efficiency) 1. Replace unsafe set_var keychain caching with OnceLock in SecretsConfig::resolve(). Eliminates the env var write from main.rs entirely, using a process-wide OnceLock cache instead. 2. Log tracing::warn when database_backend or llm_backend settings fail to parse, instead of silently falling back to defaults. 3. Remove O(K*S) get() pre-check in from_db_map(). Instead, let set() run and match on "Path not found" errors to skip unknown keys, avoiding full Settings serialization per key. Co-Authored-By: Claude Opus 4.6 * fix: address critical/high audit findings across WASM sub-crates - Telegram: remove .unwrap() panic on workspace_read (owner_id check) - WhatsApp: use configured api_version instead of hardcoded v18.0 - WhatsApp: log config parse errors before falling back to defaults - Slack: log serialization errors in emit_message and json_response - Google Docs: safe array access for batch update replies - Google Sheets: safe array access for add_sheet replies - Google Calendar: fix doc comment secret name mismatch - Gmail: avoid unnecessary String allocation in UNREAD check Co-Authored-By: Claude Opus 4.6 * fix: address second-round PR review feedback - Validate custom model ID is non-empty (loop until valid input) - Warn on unknown DATABASE_BACKEND env var before defaulting to Postgres - Force re-selection when llm_backend contains unknown provider value - Use ok_or_else for proper String error type in google-sheets Co-Authored-By: Claude Opus 4.6 * fix: harden setup module error handling and secret safety - Introduce ChannelSetupError typed enum replacing raw String errors across all channel setup functions (setup_telegram, setup_http, setup_tunnel, setup_wasm_channel, validate_telegram_token) - Add From for SetupError to simplify call sites - Convert setup_telegram retry from recursion to loop (unbounded stack) - Stop printing HTTP webhook secret plaintext to terminal - Use secret_input() for Turso auth token (was visible input()) - Replace dirs::home_dir().unwrap_or_default() with proper error - Fix UTF-8 panic in model name truncation (byte-index to chars-based) - Log warning in secret_exists() instead of silently swallowing errors - Deduplicate generate_webhook_secret() to delegate to shared helper Co-Authored-By: Claude Opus 4.6 * fix: replace unreachable!() with error return in setup wizard The provider match in step_inference_provider was guarded by is_known but used unreachable!() as the catch-all. If a new provider is added to the is_known check without a corresponding match arm, this would panic at runtime. Return a typed error instead. Co-Authored-By: Claude Opus 4.6 * fix: remove unsafe set_var, use thread-safe overlay for injected secrets Address PR #92 review comments: - Replace all 5 unsafe `std::env::set_var()` calls with safe alternatives - Add INJECTED_VARS OnceLock overlay in config.rs, checked by optional_env() before falling back to std::env::var() - Cache wizard API key in SetupWizard.llm_api_key field instead of env - Pass explicit key param to fetch_anthropic_models/fetch_openai_models - Persist env-provided API keys to secrets store during onboarding Co-Authored-By: Claude Opus 4.6 * fix: address remaining PR review comments (clippy, TODO, secrets backend ordering) - Fix empty line after doc comment (clippy: empty_line_after_doc_comments) - Collapse nested if in optional_env overlay check (clippy: collapsible_if) - Remove dangling TODO(#XX) placeholder issue ref in channels.rs - Fix init_secrets_context to respect selected database_backend when both postgres and libsql features are compiled, preventing wrong-backend secrets storage when DATABASE_URL is set but libsql was chosen Co-Authored-By: Claude Opus 4.6 * fix: address latest PR review comments (SecretString, empty env, docs, embeddings) - Change wizard llm_api_key from String to SecretString to prevent accidental logging of API keys - Fix inject_llm_keys_from_secrets skipping when env var is set but empty, matching optional_env's treatment of empty as unset - Fix inverted doc comment on INJECTED_VARS (env checked first, overlay is the fallback, not the other way around) - Update stale "env vars" comments in main.rs to reflect overlay pattern - Fix step_embeddings not seeing cached OpenAI key from wizard session Co-Authored-By: Claude Opus 4.6 * fix: OAuth callback listener binds IPv4 first to match redirect URLs The listener was binding to [::1] (IPv6) first, but NEAR AI and other OAuth flows redirect to http://127.0.0.1:9876/... (IPv4 explicit). On macOS and most systems, [::1] and 127.0.0.1 are separate addresses, so the browser's connection to 127.0.0.1 was refused when the listener was on [::1]. Reversed the bind order: try 127.0.0.1 first, fall back to [::1] if IPv4 is unavailable. Co-Authored-By: Claude Opus 4.6 * fix: cache keychain key eagerly to avoid redundant macOS password dialogs Replace has_master_key() with get_master_key() in step_security() and immediately build SecretsCrypto from the result. This eliminates redundant keychain accesses later in init_secrets_context(), each of which triggers macOS system dialogs (keychain unlock + app authorization). Co-Authored-By: Claude Opus 4.6 * fix: persist DATABASE_BACKEND to ~/.ironclaw/.env for libSQL startup The wizard saved database_backend only to the database, but Config::from_env() needs it BEFORE connecting to any database (to decide which backend to use). Without it, the backend defaults to Postgres and then fails with "Missing required setting database_url". Now save all database bootstrap vars (DATABASE_BACKEND, DATABASE_URL, LIBSQL_PATH, LIBSQL_URL) to ~/.ironclaw/.env via save_bootstrap_env(). Co-Authored-By: Claude Opus 4.6 * fix: status command shows libSQL backend and skips keychain probe The status command only checked DATABASE_URL (postgres), showing "not configured" for libSQL users. Now detects the DATABASE_BACKEND env var and reports libSQL path and Turso sync status. Also remove the keychain probe from status. get_generic_password() triggers macOS unlock+authorization dialogs which is terrible UX for a read-only diagnostic command. Co-Authored-By: Claude Opus 4.6 * style: fix rustfmt formatting in bootstrap test Co-Authored-By: Claude Opus 4.6 --------- Co-authored-by: Claude Opus 4.6 --- .gitignore | 1 + Cargo.toml | 2 +- channels-src/slack/src/lib.rs | 16 +- channels-src/telegram/src/lib.rs | 28 +- channels-src/whatsapp/src/lib.rs | 29 +- src/bootstrap.rs | 86 ++- src/cli/oauth_defaults.rs | 17 +- src/cli/status.rs | 50 +- src/config.rs | 93 ++- src/main.rs | 110 +-- src/settings.rs | 68 +- src/setup/channels.rs | 255 ++++--- src/setup/mod.rs | 5 +- src/setup/prompts.rs | 5 + src/setup/wizard.rs | 975 ++++++++++++++++++++++++--- tools-src/gmail/src/api.rs | 4 +- tools-src/google-calendar/src/lib.rs | 2 +- tools-src/google-docs/src/api.rs | 9 +- tools-src/google-sheets/src/api.rs | 8 +- 19 files changed, 1417 insertions(+), 346 deletions(-) diff --git a/.gitignore b/.gitignore index 0f80f04c..e1863e31 100644 --- a/.gitignore +++ b/.gitignore @@ -1,6 +1,7 @@ .env .env.local +.env.* target/ diff --git a/Cargo.toml b/Cargo.toml index e00dd628..ca33c35e 100644 --- a/Cargo.toml +++ b/Cargo.toml @@ -139,7 +139,7 @@ pretty_assertions = "1" tempfile = "3" [features] -default = ["postgres"] +default = ["postgres", "libsql"] postgres = [ "dep:deadpool-postgres", "dep:tokio-postgres", diff --git a/channels-src/slack/src/lib.rs b/channels-src/slack/src/lib.rs index 5cf7b10d..e4f47692 100644 --- a/channels-src/slack/src/lib.rs +++ b/channels-src/slack/src/lib.rs @@ -338,7 +338,13 @@ fn emit_message( team_id, }; - let metadata_json = serde_json::to_string(&metadata).unwrap_or_else(|_| "{}".to_string()); + let metadata_json = serde_json::to_string(&metadata).unwrap_or_else(|e| { + channel_host::log( + channel_host::LogLevel::Error, + &format!("Failed to serialize Slack metadata: {}", e), + ); + "{}".to_string() + }); // Strip @ mentions of the bot from the text for cleaner messages let cleaned_text = strip_bot_mention(&text); @@ -366,7 +372,13 @@ fn strip_bot_mention(text: &str) -> String { /// Create a JSON HTTP response. fn json_response(status: u16, value: serde_json::Value) -> OutgoingHttpResponse { - let body = serde_json::to_vec(&value).unwrap_or_default(); + let body = serde_json::to_vec(&value).unwrap_or_else(|e| { + channel_host::log( + channel_host::LogLevel::Error, + &format!("Failed to serialize JSON response: {}", e), + ); + Vec::new() + }); let headers = serde_json::json!({"Content-Type": "application/json"}); OutgoingHttpResponse { diff --git a/channels-src/telegram/src/lib.rs b/channels-src/telegram/src/lib.rs index 08e82804..a7f7f5cb 100644 --- a/channels-src/telegram/src/lib.rs +++ b/channels-src/telegram/src/lib.rs @@ -285,11 +285,7 @@ impl Guest for TelegramChannel { } // Persist dm_policy and allow_from for DM pairing in handle_message - let dm_policy = config - .dm_policy - .as_deref() - .unwrap_or("pairing") - .to_string(); + let dm_policy = config.dm_policy.as_deref().unwrap_or("pairing").to_string(); let _ = channel_host::workspace_write(DM_POLICY_PATH, &dm_policy); let allow_from_json = serde_json::to_string(&config.allow_from.unwrap_or_default()) @@ -844,8 +840,8 @@ fn send_pairing_reply(chat_id: i64, code: &str) -> Result<(), String> { "parse_mode": "Markdown", }); - let payload_bytes = serde_json::to_vec(&payload) - .map_err(|e| format!("Failed to serialize payload: {}", e))?; + let payload_bytes = + serde_json::to_vec(&payload).map_err(|e| format!("Failed to serialize payload: {}", e))?; let headers = serde_json::json!({ "Content-Type": "application/json" @@ -915,15 +911,10 @@ fn handle_message(message: TelegramMessage) { let is_private = message.chat.chat_type == "private"; // Owner validation: when owner_id is set, only that user can message - let owner_configured = channel_host::workspace_read(OWNER_ID_PATH) - .map(|s| !s.is_empty()) - .unwrap_or(false); + let owner_id_str = channel_host::workspace_read(OWNER_ID_PATH).filter(|s| !s.is_empty()); - if owner_configured { - if let Ok(owner_id) = channel_host::workspace_read(OWNER_ID_PATH) - .unwrap() - .parse::() - { + if let Some(ref id_str) = owner_id_str { + if let Ok(owner_id) = id_str.parse::() { if from.id != owner_id { channel_host::log( channel_host::LogLevel::Debug, @@ -937,8 +928,8 @@ fn handle_message(message: TelegramMessage) { } } else if is_private { // No owner_id: apply dm_policy for private chats - let dm_policy = channel_host::workspace_read(DM_POLICY_PATH) - .unwrap_or_else(|| "pairing".to_string()); + let dm_policy = + channel_host::workspace_read(DM_POLICY_PATH).unwrap_or_else(|| "pairing".to_string()); if dm_policy != "open" { // Build effective allow list: config allow_from + pairing store @@ -1001,8 +992,7 @@ fn handle_message(message: TelegramMessage) { if !respond_to_all { let has_command = content.starts_with('/'); - let bot_username = channel_host::workspace_read(BOT_USERNAME_PATH) - .unwrap_or_default(); + let bot_username = channel_host::workspace_read(BOT_USERNAME_PATH).unwrap_or_default(); let has_bot_mention = if bot_username.is_empty() { content.contains('@') } else { diff --git a/channels-src/whatsapp/src/lib.rs b/channels-src/whatsapp/src/lib.rs index 27d79e2c..7913fcd4 100644 --- a/channels-src/whatsapp/src/lib.rs +++ b/channels-src/whatsapp/src/lib.rs @@ -254,10 +254,19 @@ struct WhatsAppChannel; impl Guest for WhatsAppChannel { fn on_start(config_json: String) -> Result { - let config: WhatsAppConfig = serde_json::from_str(&config_json).unwrap_or(WhatsAppConfig { - api_version: default_api_version(), - reply_to_message: default_reply_to_message(), - }); + let config: WhatsAppConfig = match serde_json::from_str(&config_json) { + Ok(c) => c, + Err(e) => { + channel_host::log( + channel_host::LogLevel::Warn, + &format!("Failed to parse WhatsApp config, using defaults: {}", e), + ); + WhatsAppConfig { + api_version: default_api_version(), + reply_to_message: default_reply_to_message(), + } + } + }; channel_host::log( channel_host::LogLevel::Info, @@ -267,6 +276,9 @@ impl Guest for WhatsAppChannel { ), ); + // Persist api_version in workspace so on_respond() can read it + let _ = channel_host::workspace_write("channels/whatsapp/api_version", &config.api_version); + // WhatsApp Cloud API is webhook-only, no polling available Ok(ChannelConfig { display_name: "WhatsApp".to_string(), @@ -327,11 +339,16 @@ impl Guest for WhatsAppChannel { let metadata: WhatsAppMessageMetadata = serde_json::from_str(&response.metadata_json) .map_err(|e| format!("Failed to parse metadata: {}", e))?; + // Read api_version from workspace (set during on_start), fallback to default + let api_version = channel_host::workspace_read("channels/whatsapp/api_version") + .filter(|s| !s.is_empty()) + .unwrap_or_else(|| "v18.0".to_string()); + // Build WhatsApp API URL with token placeholder // Host will replace {WHATSAPP_ACCESS_TOKEN} with actual token in Authorization header let api_url = format!( - "https://graph.facebook.com/v18.0/{}/messages", - metadata.phone_number_id + "https://graph.facebook.com/{}/{}/messages", + api_version, metadata.phone_number_id ); // Build sendMessage payload diff --git a/src/bootstrap.rs b/src/bootstrap.rs index 6c14efdf..bf7d0223 100644 --- a/src/bootstrap.rs +++ b/src/bootstrap.rs @@ -81,17 +81,34 @@ fn migrate_bootstrap_json_to_env(env_path: &std::path::Path) { } } -/// Write `DATABASE_URL` to `~/.ironclaw/.env`. +/// Write database bootstrap vars to `~/.ironclaw/.env`. +/// +/// These settings form the chicken-and-egg layer: they must be available +/// from the filesystem (env vars) BEFORE any database connection, because +/// they determine which database to connect to. Everything else is stored +/// in the database itself. /// /// Creates the parent directory if it doesn't exist. -/// The value is double-quoted so that `#` (common in URL-encoded passwords) +/// Values are double-quoted so that `#` (common in URL-encoded passwords) /// and other shell-special characters are preserved by dotenvy. -pub fn save_database_url(url: &str) -> std::io::Result<()> { +pub fn save_bootstrap_env(vars: &[(&str, &str)]) -> std::io::Result<()> { let path = ironclaw_env_path(); if let Some(parent) = path.parent() { std::fs::create_dir_all(parent)?; } - std::fs::write(&path, format!("DATABASE_URL=\"{}\"\n", url)) + let mut content = String::new(); + for (key, value) in vars { + content.push_str(&format!("{}=\"{}\"\n", key, value)); + } + std::fs::write(&path, content) +} + +/// Write `DATABASE_URL` to `~/.ironclaw/.env`. +/// +/// Convenience wrapper around `save_bootstrap_env` for single-value migration +/// paths. Prefer `save_bootstrap_env` for new code. +pub fn save_database_url(url: &str) -> std::io::Result<()> { + save_bootstrap_env(&[("DATABASE_URL", url)]) } /// One-time migration of legacy `~/.ironclaw/settings.json` into the database. @@ -184,7 +201,7 @@ pub async fn migrate_disk_to_db( Ok(content) => match serde_json::from_str::(&content) { Ok(value) => { store - .set_setting(user_id, "nearai.session", &value) + .set_setting(user_id, "nearai.session_token", &value) .await .map_err(|e| { MigrationError::Database(format!( @@ -385,4 +402,63 @@ mod tests { // Nothing should happen assert!(!env_path.exists()); } + + #[test] + fn test_save_bootstrap_env_multiple_vars() { + let dir = tempdir().unwrap(); + let env_path = dir.path().join("nested").join(".env"); + + std::fs::create_dir_all(env_path.parent().unwrap()).unwrap(); + + let vars = [ + ("DATABASE_BACKEND", "libsql"), + ("LIBSQL_PATH", "/home/user/.ironclaw/ironclaw.db"), + ]; + + // Write manually to the temp path (save_bootstrap_env uses the global path) + let mut content = String::new(); + for (key, value) in &vars { + content.push_str(&format!("{}=\"{}\"\n", key, value)); + } + std::fs::write(&env_path, &content).unwrap(); + + // Verify dotenvy can parse all entries + let parsed: Vec<(String, String)> = dotenvy::from_path_iter(&env_path) + .unwrap() + .filter_map(|r| r.ok()) + .collect(); + assert_eq!(parsed.len(), 2); + assert_eq!( + parsed[0], + ("DATABASE_BACKEND".to_string(), "libsql".to_string()) + ); + assert_eq!( + parsed[1], + ( + "LIBSQL_PATH".to_string(), + "/home/user/.ironclaw/ironclaw.db".to_string() + ) + ); + } + + #[test] + fn test_save_bootstrap_env_overwrites_previous() { + let dir = tempdir().unwrap(); + let env_path = dir.path().join(".env"); + + // Write initial content + std::fs::write(&env_path, "DATABASE_URL=\"postgres://old\"\n").unwrap(); + + // Overwrite with new vars (simulating save_bootstrap_env behavior) + let content = "DATABASE_BACKEND=\"libsql\"\nLIBSQL_PATH=\"/new/path.db\"\n"; + std::fs::write(&env_path, content).unwrap(); + + let parsed: Vec<(String, String)> = dotenvy::from_path_iter(&env_path) + .unwrap() + .filter_map(|r| r.ok()) + .collect(); + // Old DATABASE_URL should be gone + assert_eq!(parsed.len(), 2); + assert!(parsed.iter().all(|(k, _)| k != "DATABASE_URL")); + } } diff --git a/src/cli/oauth_defaults.rs b/src/cli/oauth_defaults.rs index eea91a71..8ea89c3c 100644 --- a/src/cli/oauth_defaults.rs +++ b/src/cli/oauth_defaults.rs @@ -80,14 +80,13 @@ pub enum OAuthCallbackError { /// Bind the OAuth callback listener on the fixed port. /// -/// Tries IPv6 loopback (`[::1]`) first so that `http://localhost:…` redirects -/// work on systems where `localhost` resolves to `::1`. Falls back to IPv4 -/// (`127.0.0.1`) only if IPv6 fails for a reason other than `AddrInUse` -/// (e.g., IPv6 not supported on the host). If the port is already occupied -/// on IPv6, the port is occupied period, so we fail immediately. +/// Binds to IPv4 `127.0.0.1` first because callback URLs use `127.0.0.1` +/// explicitly (e.g., NEAR AI redirects to `http://127.0.0.1:9876/auth/callback`). +/// Falls back to IPv6 `[::1]` only if IPv4 binding fails for a reason other +/// than `AddrInUse`. If the port is already occupied, fails immediately. pub async fn bind_callback_listener() -> Result { - let ipv6_addr = format!("[::1]:{}", OAUTH_CALLBACK_PORT); - match TcpListener::bind(&ipv6_addr).await { + let ipv4_addr = format!("127.0.0.1:{}", OAUTH_CALLBACK_PORT); + match TcpListener::bind(&ipv4_addr).await { Ok(listener) => return Ok(listener), Err(e) if e.kind() == std::io::ErrorKind::AddrInUse => { return Err(OAuthCallbackError::PortInUse( @@ -96,10 +95,10 @@ pub async fn bind_callback_listener() -> Result )); } Err(_) => { - // IPv6 not available on this host, fall back to IPv4 + // IPv4 not available, fall back to IPv6 } } - TcpListener::bind(format!("127.0.0.1:{}", OAUTH_CALLBACK_PORT)) + TcpListener::bind(format!("[::1]:{}", OAUTH_CALLBACK_PORT)) .await .map_err(|e| { if e.kind() == std::io::ErrorKind::AddrInUse { diff --git a/src/cli/status.rs b/src/cli/status.rs index 14947f6f..2f9bf28d 100644 --- a/src/cli/status.rs +++ b/src/cli/status.rs @@ -22,15 +22,36 @@ pub async fn run_status_command() -> anyhow::Result<()> { ); // Database - let db_url_set = std::env::var("DATABASE_URL").is_ok(); print!(" Database: "); - if db_url_set { - match check_database().await { - Ok(()) => println!("connected"), - Err(e) => println!("error ({})", e), + let db_backend = std::env::var("DATABASE_BACKEND") + .ok() + .unwrap_or_else(|| "postgres".to_string()); + match db_backend.as_str() { + "libsql" | "turso" | "sqlite" => { + let path = std::env::var("LIBSQL_PATH") + .map(std::path::PathBuf::from) + .unwrap_or_else(|_| crate::config::default_libsql_path()); + if path.exists() { + let turso = if std::env::var("LIBSQL_URL").is_ok() { + " + Turso sync" + } else { + "" + }; + println!("libSQL ({}{})", path.display(), turso); + } else { + println!("libSQL (file missing: {})", path.display()); + } + } + _ => { + if std::env::var("DATABASE_URL").is_ok() { + match check_database().await { + Ok(()) => println!("connected (PostgreSQL)"), + Err(e) => println!("error ({})", e), + } + } else { + println!("not configured"); + } } - } else { - println!("not configured"); } // Session / Auth @@ -42,16 +63,17 @@ pub async fn run_status_command() -> anyhow::Result<()> { println!("not found (run `ironclaw onboard`)"); } - // Secrets (auto-detect: env var or keychain) + // Secrets (auto-detect from env only; skip keychain probe to avoid + // triggering macOS system password dialogs on a simple status check) print!(" Secrets: "); - let has_env_key = std::env::var("SECRETS_MASTER_KEY").is_ok(); - let has_keychain = crate::secrets::keychain::has_master_key().await; - if has_env_key { + if std::env::var("SECRETS_MASTER_KEY").is_ok() { println!("configured (env)"); - } else if has_keychain { - println!("configured (keychain)"); } else { - println!("not configured"); + // We don't probe the keychain here because get_generic_password() + // triggers macOS unlock+authorization dialogs, which is bad UX for + // a read-only status command. If onboarding completed with keychain + // storage, the key is there; we just can't cheaply verify it. + println!("env not set (keychain may be configured)"); } // Embeddings diff --git a/src/config.rs b/src/config.rs index 9199f887..2b40f71d 100644 --- a/src/config.rs +++ b/src/config.rs @@ -5,7 +5,9 @@ //! in startup). Everything else comes from env vars, the DB settings //! table, or auto-detection. +use std::collections::HashMap; use std::path::PathBuf; +use std::sync::OnceLock; use std::time::Duration; use secrecy::{ExposeSecret, SecretString}; @@ -13,6 +15,13 @@ use secrecy::{ExposeSecret, SecretString}; use crate::error::ConfigError; use crate::settings::Settings; +/// Thread-safe overlay for injected env vars (secrets loaded from DB). +/// +/// Used by `inject_llm_keys_from_secrets()` to make API keys available to +/// `optional_env()` without unsafe `set_var` calls. `optional_env()` checks +/// real env vars first, then falls back to this overlay. +static INJECTED_VARS: OnceLock> = OnceLock::new(); + /// Main configuration for the agent. #[derive(Debug, Clone)] pub struct Config { @@ -402,12 +411,24 @@ pub struct NearAiConfig { impl LlmConfig { fn resolve(settings: &Settings) -> Result { - // Determine backend (default: NearAi) + // Determine backend: env var > settings > default (NearAi) let backend: LlmBackend = if let Some(b) = optional_env("LLM_BACKEND")? { b.parse().map_err(|e| ConfigError::InvalidValue { key: "LLM_BACKEND".to_string(), message: e, })? + } else if let Some(ref b) = settings.llm_backend { + match b.parse() { + Ok(backend) => backend, + Err(e) => { + tracing::warn!( + "Invalid llm_backend '{}' in settings: {}. Using default NearAi.", + b, + e + ); + LlmBackend::NearAi + } + } } else { LlmBackend::NearAi }; @@ -476,6 +497,7 @@ impl LlmConfig { let ollama = if backend == LlmBackend::Ollama { let base_url = optional_env("OLLAMA_BASE_URL")? + .or_else(|| settings.ollama_base_url.clone()) .unwrap_or_else(|| "http://localhost:11434".to_string()); let model = optional_env("OLLAMA_MODEL")?.unwrap_or_else(|| "llama3".to_string()); Some(OllamaConfig { base_url, model }) @@ -484,8 +506,9 @@ impl LlmConfig { }; let openai_compatible = if backend == LlmBackend::OpenAiCompatible { - let base_url = - optional_env("LLM_BASE_URL")?.ok_or_else(|| ConfigError::MissingRequired { + let base_url = optional_env("LLM_BASE_URL")? + .or_else(|| settings.openai_compatible_base_url.clone()) + .ok_or_else(|| ConfigError::MissingRequired { key: "LLM_BASE_URL".to_string(), hint: "Set LLM_BASE_URL when LLM_BACKEND=openai_compatible".to_string(), })?; @@ -855,6 +878,11 @@ impl std::fmt::Debug for SecretsConfig { } } +/// Process-wide cache for the keychain master key. +/// +/// Avoids re-prompting the OS keychain on every `SecretsConfig::resolve()` call +/// (e.g. `Config::from_env()` then `Config::from_db()`). Thread-safe alternative +/// to caching in a process env var. impl SecretsConfig { /// Auto-detect secrets master key from env var, then OS keychain. /// @@ -1338,17 +1366,64 @@ impl ClaudeCodeConfig { } } +/// Load API keys from the encrypted secrets store into a thread-safe overlay. +/// +/// This bridges the gap between secrets stored during onboarding and the +/// env-var-first resolution in `LlmConfig::resolve()`. Keys in the overlay +/// are read by `optional_env()` before falling back to `std::env::var()`, +/// so explicit env vars always win. +pub async fn inject_llm_keys_from_secrets( + secrets: &dyn crate::secrets::SecretsStore, + user_id: &str, +) { + let mappings = [ + ("llm_openai_api_key", "OPENAI_API_KEY"), + ("llm_anthropic_api_key", "ANTHROPIC_API_KEY"), + ("llm_compatible_api_key", "LLM_API_KEY"), + ]; + + let mut injected = HashMap::new(); + + for (secret_name, env_var) in mappings { + match std::env::var(env_var) { + Ok(val) if !val.is_empty() => continue, + _ => {} + } + match secrets.get_decrypted(user_id, secret_name).await { + Ok(decrypted) => { + injected.insert(env_var.to_string(), decrypted.expose().to_string()); + tracing::debug!("Loaded secret '{}' for env var '{}'", secret_name, env_var); + } + Err(_) => { + // Secret doesn't exist, that's fine + } + } + } + + let _ = INJECTED_VARS.set(injected); +} + // Helper functions fn optional_env(key: &str) -> Result, ConfigError> { + // Check real env vars first (always win over injected secrets) match std::env::var(key) { - Ok(val) if val.is_empty() => Ok(None), - Ok(val) => Ok(Some(val)), - Err(std::env::VarError::NotPresent) => Ok(None), - Err(e) => Err(ConfigError::ParseError(format!( - "failed to read {key}: {e}" - ))), + Ok(val) if val.is_empty() => {} + Ok(val) => return Ok(Some(val)), + Err(std::env::VarError::NotPresent) => {} + Err(e) => { + return Err(ConfigError::ParseError(format!( + "failed to read {key}: {e}" + ))); + } } + + // Fall back to thread-safe overlay (secrets injected from DB) + if let Some(val) = INJECTED_VARS.get().and_then(|map| map.get(key)) { + return Ok(Some(val.clone())); + } + + Ok(None) } fn parse_optional_env(key: &str, default: T) -> Result diff --git a/src/main.rs b/src/main.rs index 89407608..30d8a5e8 100644 --- a/src/main.rs +++ b/src/main.rs @@ -48,7 +48,6 @@ use ironclaw::secrets::PostgresSecretsStore; use ironclaw::secrets::SecretsCrypto; #[cfg(any(feature = "postgres", feature = "libsql"))] use ironclaw::setup::{SetupConfig, SetupWizard}; - #[tokio::main] async fn main() -> anyhow::Result<()> { let cli = Cli::parse(); @@ -444,6 +443,72 @@ async fn main() -> anyhow::Result<()> { tracing::warn!("Failed to cleanup stale sandbox jobs: {}", e); } } + + // Create secrets store early: needed for injecting LLM API keys from encrypted + // storage before creating the LLM provider, and later for MCP auth + WASM channels. + // + // When both `postgres` and `libsql` features are compiled, the runtime-selected + // backend determines which store is created: whichever DB init branch ran will + // have set its handle (pg_pool or libsql_db), and the or_else chain picks it up. + let secrets_store: Option> = + if let Some(master_key) = config.secrets.master_key() { + match SecretsCrypto::new(master_key.clone()) { + Ok(crypto) => { + let crypto = Arc::new(crypto); + let store: Option> = None; + + #[cfg(feature = "libsql")] + let store = store.or_else(|| { + libsql_db.take().map(|db| { + Arc::new(LibSqlSecretsStore::new(db, Arc::clone(&crypto))) + as Arc + }) + }); + + #[cfg(feature = "postgres")] + let store = store.or_else(|| { + pg_pool.as_ref().map(|pool| { + Arc::new(PostgresSecretsStore::new(pool.clone(), Arc::clone(&crypto))) + as Arc + }) + }); + + store + } + Err(e) => { + tracing::warn!("Failed to initialize secrets crypto: {}", e); + #[cfg(feature = "libsql")] + let _ = libsql_db.take(); + None + } + } + } else { + #[cfg(feature = "libsql")] + let _ = libsql_db.take(); + None + }; + + // Inject LLM API keys from the encrypted secrets store into a thread-safe + // overlay so that optional_env() (used by LlmConfig::resolve()) picks them + // up. Then re-resolve LlmConfig with the newly available keys (backend may + // have been set during onboarding but the API key is in the secrets store). + if let Some(ref secrets) = secrets_store { + ironclaw::config::inject_llm_keys_from_secrets(secrets.as_ref(), "default").await; + + // Re-resolve LlmConfig now that secrets overlay has been populated + if let Some(ref db_ref) = db { + match Config::from_db(db_ref.as_ref(), "default").await { + Ok(refreshed) => { + config = refreshed; + tracing::debug!("LlmConfig re-resolved after secret injection"); + } + Err(e) => { + tracing::warn!("Failed to re-resolve config after secret injection: {}", e); + } + } + } + } + // Initialize LLM provider (clone session so we can reuse it for embeddings) let llm = create_llm_provider(&config.llm, session.clone())?; tracing::info!("LLM provider initialized: {}", llm.model_name()); @@ -542,49 +607,6 @@ async fn main() -> anyhow::Result<()> { tracing::info!("Builder mode enabled"); } - // Create secrets store if master key is configured (needed for MCP auth and WASM channels). - // - // When both `postgres` and `libsql` features are compiled, the runtime-selected - // backend determines which store is created: whichever DB init branch ran will - // have set its handle (pg_pool or libsql_db), and the or_else chain picks it up. - let secrets_store: Option> = - if let Some(master_key) = config.secrets.master_key() { - match SecretsCrypto::new(master_key.clone()) { - Ok(crypto) => { - let crypto = Arc::new(crypto); - let store: Option> = None; - - #[cfg(feature = "libsql")] - let store = store.or_else(|| { - libsql_db.take().map(|db| { - Arc::new(LibSqlSecretsStore::new(db, Arc::clone(&crypto))) - as Arc - }) - }); - - #[cfg(feature = "postgres")] - let store = store.or_else(|| { - pg_pool.as_ref().map(|pool| { - Arc::new(PostgresSecretsStore::new(pool.clone(), Arc::clone(&crypto))) - as Arc - }) - }); - - store - } - Err(e) => { - tracing::warn!("Failed to initialize secrets crypto: {}", e); - #[cfg(feature = "libsql")] - let _ = libsql_db.take(); - None - } - } - } else { - #[cfg(feature = "libsql")] - let _ = libsql_db.take(); - None - }; - let mcp_session_manager = Arc::new(McpSessionManager::new()); // Create WASM tool runtime (sync, just builds the wasmtime engine) diff --git a/src/settings.rs b/src/settings.rs index 24e710cf..60b14cd2 100644 --- a/src/settings.rs +++ b/src/settings.rs @@ -40,8 +40,18 @@ pub struct Settings { #[serde(default)] pub secrets_master_key_source: KeySource, - // === Step 3: NEAR AI Auth === - // Session stored separately in session.json + // === Step 3: Inference Provider === + /// LLM backend: "nearai", "anthropic", "openai", "ollama", "openai_compatible". + #[serde(default)] + pub llm_backend: Option, + + /// Ollama base URL (when llm_backend = "ollama"). + #[serde(default)] + pub ollama_base_url: Option, + + /// OpenAI-compatible endpoint base URL (when llm_backend = "openai_compatible"). + #[serde(default)] + pub openai_compatible_base_url: Option, // === Step 4: Model Selection === /// Currently selected model. @@ -504,7 +514,11 @@ impl Settings { /// Each key is a dotted path (e.g., "agent.name"), value is a JSONB value. /// Missing keys get their default value. pub fn from_db_map(map: &std::collections::HashMap) -> Self { - // Start with defaults, then overlay each DB setting + // Start with defaults, then overlay each DB setting. + // + // The settings table stores both Settings struct fields and app-specific + // data (e.g. nearai.session_token). Skip keys that don't correspond to + // a known Settings path. let mut settings = Self::default(); for (key, value) in map { @@ -513,17 +527,23 @@ impl Settings { serde_json::Value::String(s) => s.clone(), serde_json::Value::Bool(b) => b.to_string(), serde_json::Value::Number(n) => n.to_string(), - serde_json::Value::Null => "null".to_string(), + serde_json::Value::Null => continue, // null means default, skip other => other.to_string(), }; - if let Err(e) = settings.set(key, &value_str) { - tracing::warn!( - "Failed to apply DB setting '{}' = '{}': {}", - key, - value_str, - e - ); + match settings.set(key, &value_str) { + Ok(()) => {} + // The settings table stores both Settings fields and app-specific + // data (e.g. nearai.session_token). Silently skip unknown paths. + Err(e) if e.starts_with("Path not found") => {} + Err(e) => { + tracing::warn!( + "Failed to apply DB setting '{}' = '{}': {}", + key, + value_str, + e + ); + } } } @@ -858,4 +878,30 @@ mod tests { .unwrap(); assert_eq!(settings.channels.telegram_owner_id, Some(987654321)); } + + #[test] + fn test_llm_backend_round_trip() { + let dir = tempfile::tempdir().unwrap(); + let path = dir.path().join("settings.json"); + + let settings = Settings { + llm_backend: Some("anthropic".to_string()), + ollama_base_url: Some("http://localhost:11434".to_string()), + openai_compatible_base_url: Some("http://my-vllm:8000/v1".to_string()), + ..Default::default() + }; + let json = serde_json::to_string_pretty(&settings).unwrap(); + std::fs::write(&path, json).unwrap(); + + let loaded = Settings::load_from(&path); + assert_eq!(loaded.llm_backend, Some("anthropic".to_string())); + assert_eq!( + loaded.ollama_base_url, + Some("http://localhost:11434".to_string()) + ); + assert_eq!( + loaded.openai_compatible_base_url, + Some("http://my-vllm:8000/v1".to_string()) + ); + } } diff --git a/src/setup/channels.rs b/src/setup/channels.rs index fae447dd..34811358 100644 --- a/src/setup/channels.rs +++ b/src/setup/channels.rs @@ -20,6 +20,22 @@ use crate::setup::prompts::{ confirm, input, optional_input, print_error, print_info, print_success, secret_input, }; +/// Typed errors for channel setup flows. +#[derive(Debug, thiserror::Error)] +pub enum ChannelSetupError { + #[error("I/O error: {0}")] + Io(#[from] std::io::Error), + + #[error("{0}")] + Network(String), + + #[error("{0}")] + Secrets(String), + + #[error("{0}")] + Validation(String), +} + /// Context for saving secrets during setup. pub struct SecretsContext { store: Arc, @@ -45,32 +61,39 @@ impl SecretsContext { } /// Save a secret to the database. - pub async fn save_secret(&self, name: &str, value: &SecretString) -> Result<(), String> { + pub async fn save_secret( + &self, + name: &str, + value: &SecretString, + ) -> Result<(), ChannelSetupError> { let params = CreateSecretParams::new(name, value.expose_secret()); self.store .create(&self.user_id, params) .await - .map_err(|e| format!("Failed to save secret: {}", e))?; + .map_err(|e| ChannelSetupError::Secrets(format!("Failed to save secret: {}", e)))?; Ok(()) } /// Check if a secret exists. pub async fn secret_exists(&self, name: &str) -> bool { - self.store - .exists(&self.user_id, name) - .await - .unwrap_or(false) + match self.store.exists(&self.user_id, name).await { + Ok(exists) => exists, + Err(e) => { + tracing::warn!(secret = name, error = %e, "Failed to check if secret exists, assuming absent"); + false + } + } } /// Read a secret from the database (decrypted). - pub async fn get_secret(&self, name: &str) -> Result { + pub async fn get_secret(&self, name: &str) -> Result { let decrypted = self .store .get_decrypted(&self.user_id, name) .await - .map_err(|e| format!("Failed to read secret: {}", e))?; + .map_err(|e| ChannelSetupError::Secrets(format!("Failed to read secret: {}", e)))?; Ok(SecretString::from(decrypted.expose().to_string())) } } @@ -107,7 +130,6 @@ struct TelegramGetUpdatesResponse { #[derive(Debug, Deserialize)] struct TelegramUpdate { - #[allow(dead_code)] update_id: i64, message: Option, } @@ -134,7 +156,7 @@ struct TelegramUpdateUser { pub async fn setup_telegram( secrets: &SecretsContext, settings: &Settings, -) -> Result { +) -> Result { println!("Telegram Setup:"); println!(); print_info("To create a Telegram bot:"); @@ -146,7 +168,7 @@ pub async fn setup_telegram( // Check if token already exists if secrets.secret_exists("telegram_bot_token").await { print_info("Existing Telegram token found in database."); - if !confirm("Replace existing token?", false).map_err(|e| e.to_string())? { + if !confirm("Replace existing token?", false)? { // Still offer to configure webhook secret and owner binding let webhook_secret = setup_telegram_webhook_secret(secrets, &settings.tunnel).await?; let owner_id = bind_telegram_owner_flow(secrets, settings).await?; @@ -159,47 +181,48 @@ pub async fn setup_telegram( } } - let token = secret_input("Bot token (from @BotFather)").map_err(|e| e.to_string())?; + loop { + let token = secret_input("Bot token (from @BotFather)")?; - // Validate the token - print_info("Validating bot token..."); + // Validate the token + print_info("Validating bot token..."); - match validate_telegram_token(&token).await { - Ok(username) => { - print_success(&format!( - "Bot validated: @{}", - username.as_deref().unwrap_or("unknown") - )); + match validate_telegram_token(&token).await { + Ok(username) => { + print_success(&format!( + "Bot validated: @{}", + username.as_deref().unwrap_or("unknown") + )); - // Save to database - secrets.save_secret("telegram_bot_token", &token).await?; - print_success("Token saved to database"); + // Save to database + secrets.save_secret("telegram_bot_token", &token).await?; + print_success("Token saved to database"); - // Bind bot to owner's Telegram account - let owner_id = bind_telegram_owner(&token).await?; + // Bind bot to owner's Telegram account + let owner_id = bind_telegram_owner(&token).await?; - // Offer webhook secret configuration - let webhook_secret = setup_telegram_webhook_secret(secrets, &settings.tunnel).await?; + // Offer webhook secret configuration + let webhook_secret = + setup_telegram_webhook_secret(secrets, &settings.tunnel).await?; - Ok(TelegramSetupResult { - enabled: true, - bot_username: username, - webhook_secret, - owner_id, - }) - } - Err(e) => { - print_error(&format!("Token validation failed: {}", e)); + return Ok(TelegramSetupResult { + enabled: true, + bot_username: username, + webhook_secret, + owner_id, + }); + } + Err(e) => { + print_error(&format!("Token validation failed: {}", e)); - if confirm("Try again?", true).map_err(|e| e.to_string())? { - Box::pin(setup_telegram(secrets, settings)).await - } else { - Ok(TelegramSetupResult { - enabled: false, - bot_username: None, - webhook_secret: None, - owner_id: None, - }) + if !confirm("Try again?", true)? { + return Ok(TelegramSetupResult { + enabled: false, + bot_username: None, + webhook_secret: None, + owner_id: None, + }); + } } } } @@ -209,14 +232,14 @@ pub async fn setup_telegram( /// /// Polls `getUpdates` until a message arrives, then captures the sender's user ID. /// Returns `None` if the user declines or the flow times out. -async fn bind_telegram_owner(token: &SecretString) -> Result, String> { +async fn bind_telegram_owner(token: &SecretString) -> Result, ChannelSetupError> { println!(); print_info("Account Binding (recommended):"); print_info("Binding restricts the bot so only YOU can use it."); print_info("Without this, anyone who finds your bot can send it messages."); println!(); - if !confirm("Bind bot to your Telegram account?", true).map_err(|e| e.to_string())? { + if !confirm("Bind bot to your Telegram account?", true)? { print_info("Skipping account binding. Bot will accept messages from all users."); return Ok(None); } @@ -227,14 +250,16 @@ async fn bind_telegram_owner(token: &SecretString) -> Result, String let client = Client::builder() .timeout(std::time::Duration::from_secs(35)) .build() - .map_err(|e| format!("Failed to create HTTP client: {}", e))?; + .map_err(|e| ChannelSetupError::Network(format!("Failed to create HTTP client: {}", e)))?; // Clear any existing webhook so getUpdates works let delete_url = format!( "https://api.telegram.org/bot{}/deleteWebhook", token.expose_secret() ); - let _ = client.post(&delete_url).send().await; + if let Err(e) = client.post(&delete_url).send().await { + tracing::warn!("Failed to delete webhook (getUpdates may not work): {e}"); + } let updates_url = format!( "https://api.telegram.org/bot{}/getUpdates", @@ -249,19 +274,23 @@ async fn bind_telegram_owner(token: &SecretString) -> Result, String .query(&[("timeout", "30"), ("allowed_updates", "[\"message\"]")]) .send() .await - .map_err(|e| format!("getUpdates request failed: {}", e))?; + .map_err(|e| ChannelSetupError::Network(format!("getUpdates request failed: {}", e)))?; if !response.status().is_success() { - return Err(format!("getUpdates returned status {}", response.status())); + return Err(ChannelSetupError::Network(format!( + "getUpdates returned status {}", + response.status() + ))); } - let body: TelegramGetUpdatesResponse = response - .json() - .await - .map_err(|e| format!("Failed to parse getUpdates response: {}", e))?; + let body: TelegramGetUpdatesResponse = response.json().await.map_err(|e| { + ChannelSetupError::Network(format!("Failed to parse getUpdates response: {}", e)) + })?; if !body.ok { - return Err("Telegram API returned error for getUpdates".to_string()); + return Err(ChannelSetupError::Network( + "Telegram API returned error for getUpdates".to_string(), + )); } // Find the first message with a sender @@ -285,11 +314,14 @@ async fn bind_telegram_owner(token: &SecretString) -> Result, String "https://api.telegram.org/bot{}/getUpdates", token.expose_secret() ); - let _ = client + if let Err(e) = client .get(&ack_url) .query(&[("offset", &(update.update_id + 1).to_string())]) .send() - .await; + .await + { + tracing::warn!("Failed to acknowledge Telegram update: {e}"); + } return Ok(Some(from.id)); } @@ -307,10 +339,10 @@ async fn bind_telegram_owner(token: &SecretString) -> Result, String async fn bind_telegram_owner_flow( secrets: &SecretsContext, settings: &Settings, -) -> Result, String> { +) -> Result, ChannelSetupError> { if settings.channels.telegram_owner_id.is_some() { print_info("Bot is already bound to a Telegram account."); - if !confirm("Re-bind to a different account?", false).map_err(|e| e.to_string())? { + if !confirm("Re-bind to a different account?", false)? { return Ok(settings.channels.telegram_owner_id); } } @@ -325,10 +357,10 @@ async fn bind_telegram_owner_flow( /// /// This is shared across all channels that need webhook endpoints. /// Returns the tunnel URL if configured. -pub fn setup_tunnel(settings: &Settings) -> Result, String> { +pub fn setup_tunnel(settings: &Settings) -> Result, ChannelSetupError> { if let Some(ref url) = settings.tunnel.public_url { print_info(&format!("Existing tunnel configured: {}", url)); - if !confirm("Change tunnel configuration?", false).map_err(|e| e.to_string())? { + if !confirm("Change tunnel configuration?", false)? { return Ok(Some(url.clone())); } } @@ -348,17 +380,18 @@ pub fn setup_tunnel(settings: &Settings) -> Result, String> { print_info("Security comes from provider-specific secrets (e.g., Telegram webhook secret)."); println!(); - if !confirm("Configure a tunnel?", false).map_err(|e| e.to_string())? { + if !confirm("Configure a tunnel?", false)? { return Ok(None); } - let tunnel_url = - input("Tunnel URL (e.g., https://abc123.ngrok.io)").map_err(|e| e.to_string())?; + let tunnel_url = input("Tunnel URL (e.g., https://abc123.ngrok.io)")?; // Validate URL format if !tunnel_url.starts_with("https://") { print_error("URL must start with https:// (webhooks require HTTPS)"); - return Err("Invalid tunnel URL: must use HTTPS".to_string()); + return Err(ChannelSetupError::Validation( + "Invalid tunnel URL: must use HTTPS".to_string(), + )); } // Remove trailing slash if present @@ -378,7 +411,7 @@ pub fn setup_tunnel(settings: &Settings) -> Result, String> { async fn setup_telegram_webhook_secret( secrets: &SecretsContext, tunnel: &TunnelSettings, -) -> Result, String> { +) -> Result, ChannelSetupError> { if tunnel.public_url.is_none() { print_info(""); print_info("No tunnel configured. Telegram will use polling mode (30s+ delay)."); @@ -391,7 +424,7 @@ async fn setup_telegram_webhook_secret( print_info("A webhook secret adds an extra layer of security by validating"); print_info("that requests actually come from Telegram's servers."); - if !confirm("Generate a webhook secret?", true).map_err(|e| e.to_string())? { + if !confirm("Generate a webhook secret?", true)? { return Ok(None); } @@ -410,11 +443,13 @@ async fn setup_telegram_webhook_secret( /// Validate a Telegram bot token by calling the getMe API. /// /// Returns the bot's username if valid. -pub async fn validate_telegram_token(token: &SecretString) -> Result, String> { +pub async fn validate_telegram_token( + token: &SecretString, +) -> Result, ChannelSetupError> { let client = Client::builder() .timeout(std::time::Duration::from_secs(10)) .build() - .map_err(|e| format!("Failed to create HTTP client: {}", e))?; + .map_err(|e| ChannelSetupError::Network(format!("Failed to create HTTP client: {}", e)))?; let url = format!( "https://api.telegram.org/bot{}/getMe", @@ -425,21 +460,26 @@ pub async fn validate_telegram_token(token: &SecretString) -> Result Result { +pub async fn setup_http(secrets: &SecretsContext) -> Result { println!("HTTP Webhook Setup:"); println!(); print_info("The HTTP webhook allows external services to send messages to the agent."); println!(); - let port_str = optional_input("Port", Some("default: 8080")).map_err(|e| e.to_string())?; + let port_str = optional_input("Port", Some("default: 8080"))?; let port: u16 = port_str .as_deref() .unwrap_or("8080") .parse() - .map_err(|e| format!("Invalid port: {}", e))?; + .map_err(|e| ChannelSetupError::Validation(format!("Invalid port: {}", e)))?; if port < 1024 { print_info("Note: Ports below 1024 may require root privileges"); } - let host = optional_input("Host", Some("default: 0.0.0.0")) - .map_err(|e| e.to_string())? - .unwrap_or_else(|| "0.0.0.0".to_string()); + let host = + optional_input("Host", Some("default: 0.0.0.0"))?.unwrap_or_else(|| "0.0.0.0".to_string()); // Generate a webhook secret - if confirm("Generate a webhook secret for authentication?", true).map_err(|e| e.to_string())? { + if confirm("Generate a webhook secret for authentication?", true)? { let secret = generate_webhook_secret(); secrets - .save_secret("http_webhook_secret", &SecretString::from(secret.clone())) + .save_secret("http_webhook_secret", &SecretString::from(secret)) .await?; print_success("Webhook secret generated and saved to database"); - print_info(&format!( - "Secret: {} (store this for your webhook clients)", - secret - )); + print_info("Retrieve it later with: ironclaw secret get http_webhook_secret"); } print_success(&format!("HTTP webhook will listen on {}:{}", host, port)); @@ -497,11 +533,7 @@ pub async fn setup_http(secrets: &SecretsContext) -> Result String { - use rand::RngCore; - let mut rng = rand::thread_rng(); - let mut bytes = [0u8; 32]; - rng.fill_bytes(&mut bytes); - bytes.iter().map(|b| format!("{:02x}", b)).collect() + generate_secret_with_length(32) } /// Result of WASM channel setup. @@ -519,7 +551,7 @@ pub async fn setup_wasm_channel( secrets: &SecretsContext, channel_name: &str, setup: &crate::channels::wasm::SetupSchema, -) -> Result { +) -> Result { println!("{} Setup:", channel_name); println!(); @@ -530,7 +562,7 @@ pub async fn setup_wasm_channel( "Existing {} found in database.", secret_config.name )); - if !confirm("Replace existing value?", false).map_err(|e| e.to_string())? { + if !confirm("Replace existing value?", false)? { continue; } } @@ -538,8 +570,7 @@ pub async fn setup_wasm_channel( // Get the value from user or auto-generate let value = if secret_config.optional { let input_value = - optional_input(&secret_config.prompt, Some("leave empty to auto-generate")) - .map_err(|e| e.to_string())?; + optional_input(&secret_config.prompt, Some("leave empty to auto-generate"))?; if let Some(v) = input_value { if !v.is_empty() { @@ -566,18 +597,21 @@ pub async fn setup_wasm_channel( } } else { // Required secret - let input_value = secret_input(&secret_config.prompt).map_err(|e| e.to_string())?; + let input_value = secret_input(&secret_config.prompt)?; // Validate if pattern is provided if let Some(ref pattern) = secret_config.validation { - let re = regex::Regex::new(pattern) - .map_err(|e| format!("Invalid validation pattern: {}", e))?; + let re = regex::Regex::new(pattern).map_err(|e| { + ChannelSetupError::Validation(format!("Invalid validation pattern: {}", e)) + })?; if !re.is_match(input_value.expose_secret()) { print_error(&format!( "Value does not match expected format: {}", pattern )); - return Err("Validation failed".to_string()); + return Err(ChannelSetupError::Validation( + "Validation failed".to_string(), + )); } } @@ -589,14 +623,11 @@ pub async fn setup_wasm_channel( print_success(&format!("{} saved to database", secret_config.name)); } - // Optionally validate the configuration + // TODO: Substitute secrets into the validation URL and make a + // GET request to verify the configured credentials actually work. if let Some(ref validation_endpoint) = setup.validation_endpoint { - print_info("Validating configuration..."); - // The validation endpoint may contain placeholders like {telegram_bot_token} - // For now, we skip validation since we'd need to substitute secrets - // A full implementation would fetch secrets and substitute them print_info(&format!( - "Validation endpoint configured: {} (validation skipped)", + "Validation endpoint configured: {} (validation not yet implemented)", validation_endpoint )); } @@ -620,11 +651,23 @@ fn generate_secret_with_length(length: usize) -> String { #[cfg(test)] mod tests { - use super::*; + use crate::setup::channels::generate_webhook_secret; #[test] fn test_generate_webhook_secret() { let secret = generate_webhook_secret(); assert_eq!(secret.len(), 64); // 32 bytes = 64 hex chars } + + #[test] + fn test_generate_secret_with_length() { + use super::generate_secret_with_length; + + let s = generate_secret_with_length(16); + assert_eq!(s.len(), 32); // 16 bytes = 32 hex chars + assert!(s.chars().all(|c| c.is_ascii_hexdigit())); + + let s2 = generate_secret_with_length(1); + assert_eq!(s2.len(), 2); + } } diff --git a/src/setup/mod.rs b/src/setup/mod.rs index ca4d4c56..f2501a57 100644 --- a/src/setup/mod.rs +++ b/src/setup/mod.rs @@ -3,7 +3,7 @@ //! Provides a guided setup experience for: //! 1. Database connection //! 2. Security (secrets master key) -//! 3. NEAR AI authentication +//! 3. Inference provider selection //! 4. Model selection //! 5. Embeddings //! 6. Channel configuration (HTTP, Telegram, etc.) @@ -24,7 +24,8 @@ mod prompts; mod wizard; pub use channels::{ - SecretsContext, setup_http, setup_telegram, setup_tunnel, validate_telegram_token, + ChannelSetupError, SecretsContext, setup_http, setup_telegram, setup_tunnel, + validate_telegram_token, }; pub use prompts::{ confirm, input, optional_input, print_error, print_header, print_info, print_step, diff --git a/src/setup/prompts.rs b/src/setup/prompts.rs index 7d77c245..ce075572 100644 --- a/src/setup/prompts.rs +++ b/src/setup/prompts.rs @@ -21,6 +21,7 @@ use secrecy::SecretString; /// Display a numbered menu and get user selection. /// /// Returns the index (0-based) of the selected option. +/// Pressing Enter without input selects the first option (index 0). /// /// # Example /// @@ -84,6 +85,10 @@ pub fn select_one(prompt: &str, options: &[&str]) -> io::Result { /// ])?; /// ``` pub fn select_many(prompt: &str, options: &[(&str, bool)]) -> io::Result> { + if options.is_empty() { + return Ok(vec![]); + } + let mut stdout = io::stdout(); let mut selected: Vec = options.iter().map(|(_, s)| *s).collect(); let mut cursor_pos = 0; diff --git a/src/setup/wizard.rs b/src/setup/wizard.rs index 2e495028..63f955be 100644 --- a/src/setup/wizard.rs +++ b/src/setup/wizard.rs @@ -3,7 +3,7 @@ //! The wizard guides users through: //! 1. Database connection //! 2. Security (secrets master key) -//! 3. NEAR AI authentication +//! 3. Inference provider (NEAR AI, Anthropic, OpenAI, Ollama, OpenAI-compatible) //! 4. Model selection //! 5. Embeddings //! 6. Channel configuration @@ -14,7 +14,7 @@ use std::sync::Arc; #[cfg(feature = "postgres")] use deadpool_postgres::{Config as PoolConfig, Runtime}; -use secrecy::SecretString; +use secrecy::{ExposeSecret, SecretString}; #[cfg(feature = "postgres")] use tokio_postgres::NoTls; @@ -29,7 +29,7 @@ use crate::setup::channels::{ }; use crate::setup::prompts::{ confirm, input, optional_input, print_error, print_header, print_info, print_step, - print_success, select_many, select_one, + print_success, secret_input, select_many, select_one, }; /// Setup wizard error. @@ -54,6 +54,12 @@ pub enum SetupError { Cancelled, } +impl From for SetupError { + fn from(e: crate::setup::channels::ChannelSetupError) -> Self { + SetupError::Channel(e.to_string()) + } +} + /// Setup wizard configuration. #[derive(Debug, Clone, Default)] pub struct SetupConfig { @@ -76,6 +82,8 @@ pub struct SetupWizard { db_backend: Option, /// Secrets crypto (created during setup). secrets_crypto: Option>, + /// Cached API key from provider setup (used by model fetcher without env mutation). + llm_api_key: Option, } impl SetupWizard { @@ -90,6 +98,7 @@ impl SetupWizard { #[cfg(feature = "libsql")] db_backend: None, secrets_crypto: None, + llm_api_key: None, } } @@ -104,6 +113,7 @@ impl SetupWizard { #[cfg(feature = "libsql")] db_backend: None, secrets_crypto: None, + llm_api_key: None, } } @@ -132,12 +142,12 @@ impl SetupWizard { print_step(2, total_steps, "Security"); self.step_security().await?; - // Step 3: Authentication (unless skipped) + // Step 3: Inference provider selection (unless skipped) if !self.config.skip_auth { - print_step(3, total_steps, "NEAR AI Authentication"); - self.step_authentication().await?; + print_step(3, total_steps, "Inference Provider"); + self.step_inference_provider().await?; } else { - print_info("Skipping authentication (using existing session)"); + print_info("Skipping inference provider setup (using existing config)"); } // Step 4: Model selection @@ -165,20 +175,56 @@ impl SetupWizard { /// Step 1: Database connection. async fn step_database(&mut self) -> Result<(), SetupError> { - // Determine which backend to use based on compile-time features. - // When both features are enabled, prefer the currently configured backend - // or default to postgres. + // When both features are compiled, let the user choose. + // If DATABASE_BACKEND is already set in the environment, respect it. #[cfg(all(feature = "postgres", feature = "libsql"))] { - let backend = std::env::var("DATABASE_BACKEND") - .ok() - .or_else(|| self.settings.database_backend.clone()) - .unwrap_or_else(|| "postgres".to_string()); + // Check if a backend is already pinned via env var + let env_backend = std::env::var("DATABASE_BACKEND").ok(); - if backend == "libsql" || backend == "turso" || backend == "sqlite" { - return self.step_database_libsql().await; + if let Some(ref backend) = env_backend { + if backend == "libsql" || backend == "turso" || backend == "sqlite" { + return self.step_database_libsql().await; + } + if backend != "postgres" && backend != "postgresql" { + print_info(&format!( + "Unknown DATABASE_BACKEND '{}', defaulting to PostgreSQL", + backend + )); + } + return self.step_database_postgres().await; + } + + // Interactive selection + let pre_selected = self.settings.database_backend.as_deref().map(|b| match b { + "libsql" | "turso" | "sqlite" => 1, + _ => 0, + }); + + print_info("Which database backend would you like to use?"); + println!(); + + let options = &[ + "PostgreSQL - production-grade, requires a running server", + "libSQL - embedded SQLite, zero dependencies, optional Turso cloud sync", + ]; + let choice = + select_one("Select a database backend:", options).map_err(SetupError::Io)?; + + // If the user picked something different from what was pre-selected, clear + // stale connection settings so the next step starts fresh. + if let Some(prev) = pre_selected + && prev != choice + { + self.settings.database_url = None; + self.settings.libsql_path = None; + self.settings.libsql_url = None; + } + + match choice { + 1 => return self.step_database_libsql().await, + _ => return self.step_database_postgres().await, } - return self.step_database_postgres().await; } #[cfg(all(feature = "postgres", not(feature = "libsql")))] @@ -327,7 +373,8 @@ impl SetupWizard { print_error("Turso URL is required for cloud sync."); (None, None) } else { - let token = input("Auth token").map_err(SetupError::Io)?; + let token_secret = secret_input("Auth token").map_err(SetupError::Io)?; + let token = token_secret.expose_secret().to_string(); if token.is_empty() { print_error("Auth token is required for cloud sync."); (None, None) @@ -456,7 +503,6 @@ impl SetupWizard { async fn step_security(&mut self) -> Result<(), SetupError> { // Check current configuration let env_key_exists = std::env::var("SECRETS_MASTER_KEY").is_ok(); - let keychain_key_exists = crate::secrets::keychain::has_master_key().await; if env_key_exists { print_info("Secrets master key found in SECRETS_MASTER_KEY environment variable."); @@ -465,13 +511,30 @@ impl SetupWizard { return Ok(()); } - if keychain_key_exists { + // Try to retrieve existing key from keychain. We use get_master_key() + // instead of has_master_key() so we can cache the key bytes and build + // SecretsCrypto eagerly, avoiding redundant keychain accesses later + // (each access triggers macOS system dialogs). + print_info("Checking OS keychain for existing master key..."); + if let Ok(keychain_key_bytes) = crate::secrets::keychain::get_master_key().await { + let key_hex: String = keychain_key_bytes + .iter() + .map(|b| format!("{:02x}", b)) + .collect(); + self.secrets_crypto = Some(Arc::new( + SecretsCrypto::new(SecretString::from(key_hex)) + .map_err(|e| SetupError::Config(e.to_string()))?, + )); + print_info("Existing master key found in OS keychain."); if confirm("Use existing keychain key?", true).map_err(SetupError::Io)? { self.settings.secrets_master_key_source = KeySource::Keychain; print_success("Security configured (keychain)"); return Ok(()); } + // User declined the existing key; clear the cached crypto so a fresh + // key can be generated below. + self.secrets_crypto = None; } // Offer options @@ -531,8 +594,83 @@ impl SetupWizard { Ok(()) } - /// Step 3: NEAR AI authentication. - async fn step_authentication(&mut self) -> Result<(), SetupError> { + /// Step 3: Inference provider selection. + /// + /// Lets the user pick from all supported LLM backends, then runs the + /// provider-specific auth sub-flow (API key entry, NEAR AI login, etc.). + async fn step_inference_provider(&mut self) -> Result<(), SetupError> { + // Show current provider if already configured + if let Some(ref current) = self.settings.llm_backend { + let display = match current.as_str() { + "nearai" => "NEAR AI", + "anthropic" => "Anthropic (Claude)", + "openai" => "OpenAI", + "ollama" => "Ollama (local)", + "openai_compatible" => "OpenAI-compatible endpoint", + other => other, + }; + print_info(&format!("Current provider: {}", display)); + println!(); + + let is_known = matches!( + current.as_str(), + "nearai" | "anthropic" | "openai" | "ollama" | "openai_compatible" + ); + + if is_known && confirm("Keep current provider?", true).map_err(SetupError::Io)? { + // Still run the auth sub-flow in case they need to update keys + match current.as_str() { + "nearai" => return self.setup_nearai().await, + "anthropic" => return self.setup_anthropic().await, + "openai" => return self.setup_openai().await, + "ollama" => return self.setup_ollama(), + "openai_compatible" => return self.setup_openai_compatible().await, + _ => { + return Err(SetupError::Config(format!( + "Unhandled provider: {}", + current + ))); + } + } + } + + if !is_known { + print_info(&format!( + "Unknown provider '{}', please select a supported provider.", + current + )); + } + } + + print_info("Select your inference provider:"); + println!(); + + let options = &[ + "NEAR AI - multi-model access via NEAR account", + "Anthropic - Claude models (direct API key)", + "OpenAI - GPT models (direct API key)", + "Ollama - local models, no API key needed", + "OpenAI-compatible - custom endpoint (vLLM, LiteLLM, Together, etc.)", + ]; + + let choice = select_one("Provider:", options).map_err(SetupError::Io)?; + + match choice { + 0 => self.setup_nearai().await?, + 1 => self.setup_anthropic().await?, + 2 => self.setup_openai().await?, + 3 => self.setup_ollama()?, + 4 => self.setup_openai_compatible().await?, + _ => return Err(SetupError::Config("Invalid provider selection".to_string())), + } + + Ok(()) + } + + /// NEAR AI provider setup (extracted from the old step_authentication). + async fn setup_nearai(&mut self) -> Result<(), SetupError> { + self.settings.llm_backend = Some("nearai".to_string()); + // Check if we already have a session if let Some(ref session) = self.session_manager && session.has_token().await @@ -540,7 +678,7 @@ impl SetupWizard { print_info("Existing session found. Validating..."); match session.ensure_authenticated().await { Ok(()) => { - print_success("Session valid"); + print_success("NEAR AI session valid"); return Ok(()); } Err(e) => { @@ -564,10 +702,183 @@ impl SetupWizard { .map_err(|e| SetupError::Auth(e.to_string()))?; self.session_manager = Some(session); + print_success("NEAR AI configured"); + Ok(()) + } + + /// Anthropic provider setup: collect API key and store in secrets. + async fn setup_anthropic(&mut self) -> Result<(), SetupError> { + self.setup_api_key_provider( + "anthropic", + "ANTHROPIC_API_KEY", + "llm_anthropic_api_key", + "Anthropic API key", + "https://console.anthropic.com/settings/keys", + ) + .await + } + + /// OpenAI provider setup: collect API key and store in secrets. + async fn setup_openai(&mut self) -> Result<(), SetupError> { + self.setup_api_key_provider( + "openai", + "OPENAI_API_KEY", + "llm_openai_api_key", + "OpenAI API key", + "https://platform.openai.com/api-keys", + ) + .await + } + + /// Shared setup flow for API-key-based providers (Anthropic, OpenAI). + async fn setup_api_key_provider( + &mut self, + backend: &str, + env_var: &str, + secret_name: &str, + prompt_label: &str, + hint_url: &str, + ) -> Result<(), SetupError> { + let display_name = match backend { + "anthropic" => "Anthropic", + "openai" => "OpenAI", + other => other, + }; + + self.settings.llm_backend = Some(backend.to_string()); + if self.settings.selected_model.is_some() { + self.settings.selected_model = None; + } + + // Check env var first + if let Ok(existing) = std::env::var(env_var) { + print_info(&format!("{env_var} found: {}", mask_api_key(&existing))); + if confirm("Use this key?", true).map_err(SetupError::Io)? { + // Persist env-provided key to secrets store for future runs + if let Ok(ctx) = self.init_secrets_context().await { + let key = SecretString::from(existing.clone()); + if let Err(e) = ctx.save_secret(secret_name, &key).await { + tracing::warn!("Failed to persist env key to secrets: {}", e); + } + } + self.llm_api_key = Some(SecretString::from(existing)); + print_success(&format!("{display_name} configured (from env)")); + return Ok(()); + } + } + + println!(); + print_info(&format!("Get your API key from: {hint_url}")); + println!(); + + let key = secret_input(prompt_label).map_err(SetupError::Io)?; + let key_str = key.expose_secret(); + + if key_str.is_empty() { + return Err(SetupError::Config("API key cannot be empty".to_string())); + } + + // Store in secrets if available + if let Ok(ctx) = self.init_secrets_context().await { + ctx.save_secret(secret_name, &key) + .await + .map_err(|e| SetupError::Config(format!("Failed to save API key: {e}")))?; + print_success("API key encrypted and saved"); + } else { + print_info(&format!( + "Secrets not available. Set {env_var} in your environment." + )); + } + + // Cache key in memory for model fetching later in the wizard + self.llm_api_key = Some(SecretString::from(key_str.to_string())); + + print_success(&format!("{display_name} configured")); + Ok(()) + } + + /// Ollama provider setup: just needs a base URL, no API key. + fn setup_ollama(&mut self) -> Result<(), SetupError> { + self.settings.llm_backend = Some("ollama".to_string()); + if self.settings.selected_model.is_some() { + self.settings.selected_model = None; + } + + let default_url = self + .settings + .ollama_base_url + .as_deref() + .unwrap_or("http://localhost:11434"); + + let url_input = optional_input( + "Ollama base URL", + Some(&format!("default: {}", default_url)), + ) + .map_err(SetupError::Io)?; + + let url = url_input.unwrap_or_else(|| default_url.to_string()); + self.settings.ollama_base_url = Some(url.clone()); + + print_success(&format!("Ollama configured ({})", url)); + Ok(()) + } + + /// OpenAI-compatible provider setup: base URL + optional API key. + async fn setup_openai_compatible(&mut self) -> Result<(), SetupError> { + self.settings.llm_backend = Some("openai_compatible".to_string()); + if self.settings.selected_model.is_some() { + self.settings.selected_model = None; + } + + let existing_url = self + .settings + .openai_compatible_base_url + .clone() + .or_else(|| std::env::var("LLM_BASE_URL").ok()); + + let url = if let Some(ref u) = existing_url { + let url_input = optional_input("Base URL", Some(&format!("current: {}", u))) + .map_err(SetupError::Io)?; + url_input.unwrap_or_else(|| u.clone()) + } else { + input("Base URL (e.g., http://localhost:8000/v1)").map_err(SetupError::Io)? + }; + + if url.is_empty() { + return Err(SetupError::Config( + "Base URL is required for OpenAI-compatible provider".to_string(), + )); + } + + self.settings.openai_compatible_base_url = Some(url.clone()); + + // Optional API key + if confirm("Does this endpoint require an API key?", false).map_err(SetupError::Io)? { + let key = secret_input("API key").map_err(SetupError::Io)?; + let key_str = key.expose_secret(); + + if !key_str.is_empty() { + if let Ok(ctx) = self.init_secrets_context().await { + ctx.save_secret("llm_compatible_api_key", &key) + .await + .map_err(|e| { + SetupError::Config(format!("Failed to save API key: {}", e)) + })?; + print_success("API key encrypted and saved"); + } else { + print_info("Secrets not available. Set LLM_API_KEY in your environment."); + } + } + } + + print_success(&format!("OpenAI-compatible configured ({})", url)); Ok(()) } /// Step 4: Model selection. + /// + /// Branches on the selected LLM backend and fetches models from the + /// appropriate provider API, with static defaults as fallback. async fn step_model_selection(&mut self) -> Result<(), SetupError> { // Show current model if already configured if let Some(ref current) = self.settings.selected_model { @@ -584,58 +895,113 @@ impl SetupWizard { } } - // Try to fetch available models - let models = if let Some(ref session) = self.session_manager { - self.fetch_available_models(session).await - } else { - vec![] - }; + let backend = self.settings.llm_backend.as_deref().unwrap_or("nearai"); - // Default models if we couldn't fetch - let default_models = [ - ( - "fireworks::accounts/fireworks/models/llama4-maverick-instruct-basic", - "Llama 4 Maverick (default, fast)", - ), - ( - "anthropic::claude-sonnet-4-20250514", - "Claude Sonnet 4 (best quality)", - ), - ("openai::gpt-4o", "GPT-4o"), - ]; + match backend { + "anthropic" => { + let cached = self + .llm_api_key + .as_ref() + .map(|k| k.expose_secret().to_string()); + let models = fetch_anthropic_models(cached.as_deref()).await; + self.select_from_model_list(&models)?; + } + "openai" => { + let cached = self + .llm_api_key + .as_ref() + .map(|k| k.expose_secret().to_string()); + let models = fetch_openai_models(cached.as_deref()).await; + self.select_from_model_list(&models)?; + } + "ollama" => { + let base_url = self + .settings + .ollama_base_url + .as_deref() + .unwrap_or("http://localhost:11434"); + let models = fetch_ollama_models(base_url).await; + if models.is_empty() { + print_info("No models found. Pull one first: ollama pull llama3"); + } + self.select_from_model_list(&models)?; + } + "openai_compatible" => { + // No standard API for listing models on arbitrary endpoints + let model_id = input("Model name (e.g., meta-llama/Llama-3-8b-chat-hf)") + .map_err(SetupError::Io)?; + if model_id.is_empty() { + return Err(SetupError::Config("Model name is required".to_string())); + } + self.settings.selected_model = Some(model_id.clone()); + print_success(&format!("Selected {}", model_id)); + } + _ => { + // NEAR AI: use existing provider list_models() + let fetched = self.fetch_nearai_models().await; + let default_models: Vec<(String, String)> = vec![ + ( + "fireworks::accounts/fireworks/models/llama4-maverick-instruct-basic" + .into(), + "Llama 4 Maverick (default, fast)".into(), + ), + ( + "anthropic::claude-sonnet-4-20250514".into(), + "Claude Sonnet 4 (best quality)".into(), + ), + ("openai::gpt-4o".into(), "GPT-4o".into()), + ]; - println!("Available models:"); - println!(); - - let options: Vec<&str> = if models.is_empty() { - default_models.iter().map(|(_, desc)| *desc).collect() - } else { - models.iter().map(|m| m.as_str()).collect() - }; - - // Add custom option - let mut all_options = options.clone(); - all_options.push("Custom model ID"); - - let choice = select_one("Select a model:", &all_options).map_err(SetupError::Io)?; - - let selected_model = if choice == all_options.len() - 1 { - // Custom model - input("Enter model ID").map_err(SetupError::Io)? - } else if models.is_empty() { - default_models[choice].0.to_string() - } else { - models[choice].clone() - }; - - self.settings.selected_model = Some(selected_model.clone()); - print_success(&format!("Selected {}", selected_model)); + let models = if fetched.is_empty() { + default_models + } else { + fetched.iter().map(|m| (m.clone(), m.clone())).collect() + }; + self.select_from_model_list(&models)?; + } + } Ok(()) } - /// Fetch available models from the API. - async fn fetch_available_models(&self, session: &Arc) -> Vec { + /// Present a model list to the user, with a "Custom model ID" escape hatch. + /// + /// Each entry is `(model_id, display_label)`. + fn select_from_model_list(&mut self, models: &[(String, String)]) -> Result<(), SetupError> { + println!("Available models:"); + println!(); + + let mut options: Vec<&str> = models.iter().map(|(_, desc)| desc.as_str()).collect(); + options.push("Custom model ID"); + + let choice = select_one("Select a model:", &options).map_err(SetupError::Io)?; + + let selected = if choice == options.len() - 1 { + loop { + let raw = input("Enter model ID").map_err(SetupError::Io)?; + let trimmed = raw.trim().to_string(); + if trimmed.is_empty() { + println!("Model ID cannot be empty."); + continue; + } + break trimmed; + } + } else { + models[choice].0.clone() + }; + + self.settings.selected_model = Some(selected.clone()); + print_success(&format!("Selected {}", selected)); + Ok(()) + } + + /// Fetch available models from the NEAR AI API. + async fn fetch_nearai_models(&self) -> Vec { + let session = match self.session_manager { + Some(ref s) => Arc::clone(s), + None => return vec![], + }; + use crate::config::LlmConfig; use crate::llm::create_llm_provider; @@ -662,7 +1028,7 @@ impl SetupWizard { openai_compatible: None, }; - match create_llm_provider(&config, Arc::clone(session)) { + match create_llm_provider(&config, session) { Ok(provider) => match provider.list_models().await { Ok(models) => models, Err(e) => { @@ -691,23 +1057,52 @@ impl SetupWizard { return Ok(()); } - let options = [ - "NEAR AI (uses same auth, no extra cost)", - "OpenAI (requires API key)", - ]; + let backend = self.settings.llm_backend.as_deref().unwrap_or("nearai"); + let has_openai_key = std::env::var("OPENAI_API_KEY").is_ok() + || (backend == "openai" && self.llm_api_key.is_some()); + let has_nearai = backend == "nearai" || self.session_manager.is_some(); + + // If the LLM backend is OpenAI and we already have a key, default to OpenAI embeddings + if backend == "openai" && has_openai_key { + self.settings.embeddings.enabled = true; + self.settings.embeddings.provider = "openai".to_string(); + self.settings.embeddings.model = "text-embedding-3-small".to_string(); + print_success("Embeddings enabled via OpenAI (using existing API key)"); + return Ok(()); + } + + // If no NEAR AI session and no OpenAI key, only OpenAI is viable + if !has_nearai && !has_openai_key { + print_info("No NEAR AI session or OpenAI key found for embeddings."); + print_info("Set OPENAI_API_KEY in your environment to enable embeddings."); + self.settings.embeddings.enabled = false; + return Ok(()); + } + + let mut options = Vec::new(); + if has_nearai { + options.push("NEAR AI (uses same auth, no extra cost)"); + } + options.push("OpenAI (requires API key)"); let choice = select_one("Select embeddings provider:", &options).map_err(SetupError::Io)?; - match choice { - 0 => { + // Map choice back to provider name + let provider = if has_nearai && choice == 0 { + "nearai" + } else { + "openai" + }; + + match provider { + "nearai" => { self.settings.embeddings.enabled = true; self.settings.embeddings.provider = "nearai".to_string(); self.settings.embeddings.model = "text-embedding-3-small".to_string(); print_success("Embeddings enabled via NEAR AI"); } - 1 => { - // Check if API key is set - if std::env::var("OPENAI_API_KEY").is_err() { + _ => { + if !has_openai_key { print_info("OPENAI_API_KEY not set in environment."); print_info("Add it to your .env file or environment to enable embeddings."); } @@ -716,7 +1111,6 @@ impl SetupWizard { self.settings.embeddings.model = "text-embedding-3-small".to_string(); print_success("Embeddings configured for OpenAI"); } - _ => unreachable!(), } Ok(()) @@ -739,23 +1133,54 @@ impl SetupWizard { )); }; - let crypto = SecretsCrypto::new(SecretString::from(key)) - .map_err(|e| SetupError::Config(e.to_string()))?; - self.secrets_crypto = Some(Arc::new(crypto)); - Arc::clone(self.secrets_crypto.as_ref().unwrap()) + let crypto = Arc::new( + SecretsCrypto::new(SecretString::from(key)) + .map_err(|e| SetupError::Config(e.to_string()))?, + ); + self.secrets_crypto = Some(Arc::clone(&crypto)); + crypto }; - // Create backend-appropriate secrets store - #[cfg(feature = "postgres")] + // Create backend-appropriate secrets store. + // Respect the user's selected backend when both features are compiled, + // so we don't accidentally use a postgres pool from DATABASE_URL when + // libsql was chosen (or vice versa). + let selected_backend = self + .settings + .database_backend + .as_deref() + .unwrap_or("postgres"); + + #[cfg(all(feature = "libsql", feature = "postgres"))] { - // Try postgres path first when postgres feature is available + if selected_backend == "libsql" { + if let Some(store) = self.create_libsql_secrets_store(&crypto)? { + return Ok(SecretsContext::from_store(store, "default")); + } + if let Some(store) = self.create_postgres_secrets_store(&crypto).await? { + return Ok(SecretsContext::from_store(store, "default")); + } + } else { + if let Some(store) = self.create_postgres_secrets_store(&crypto).await? { + return Ok(SecretsContext::from_store(store, "default")); + } + if let Some(store) = self.create_libsql_secrets_store(&crypto)? { + return Ok(SecretsContext::from_store(store, "default")); + } + } + } + + #[cfg(all(feature = "postgres", not(feature = "libsql")))] + { + let _ = selected_backend; if let Some(store) = self.create_postgres_secrets_store(&crypto).await? { return Ok(SecretsContext::from_store(store, "default")); } } - #[cfg(feature = "libsql")] + #[cfg(all(feature = "libsql", not(feature = "postgres")))] { + let _ = selected_backend; if let Some(store) = self.create_libsql_secrets_store(&crypto)? { return Ok(SecretsContext::from_store(store, "default")); } @@ -785,7 +1210,14 @@ impl SetupWizard { if let Some(url) = url { self.test_database_connection_postgres(&url).await?; self.run_migrations_postgres().await?; - self.db_pool.clone().unwrap() + match self.db_pool.clone() { + Some(pool) => pool, + None => { + return Err(SetupError::Database( + "Database pool not initialized after connection test".to_string(), + )); + } + } } else { return Ok(None); } @@ -833,7 +1265,7 @@ impl SetupWizard { // Discover available WASM channels let channels_dir = dirs::home_dir() - .unwrap_or_default() + .ok_or_else(|| SetupError::Config("Could not determine home directory".into()))? .join(".ironclaw/channels"); let mut discovered_channels = discover_wasm_channels(&channels_dir).await; @@ -908,7 +1340,7 @@ impl SetupWizard { if selected.contains(&1) { println!(); if let Some(ref ctx) = secrets { - let result = setup_http(ctx).await.map_err(SetupError::Channel)?; + let result = setup_http(ctx).await?; self.settings.channels.http_enabled = result.enabled; self.settings.channels.http_port = Some(result.port); } else { @@ -930,13 +1362,9 @@ impl SetupWizard { if let Some(ref ctx) = secrets { let result = if let Some(cap_file) = discovered_by_name.get(&channel_name) { if !cap_file.setup.required_secrets.is_empty() { - setup_wasm_channel(ctx, &channel_name, &cap_file.setup) - .await - .map_err(SetupError::Channel)? + setup_wasm_channel(ctx, &channel_name, &cap_file.setup).await? } else if channel_name == "telegram" { - let telegram_result = setup_telegram(ctx, &self.settings) - .await - .map_err(SetupError::Channel)?; + 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); } @@ -1077,15 +1505,35 @@ impl SetupWizard { } } - // 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 - ))) - })?; + // Persist database bootstrap vars to ~/.ironclaw/.env. + // These are the chicken-and-egg settings: we need them to decide + // which database to connect to, so they can't live in the database. + { + let mut env_vars: Vec<(&str, String)> = Vec::new(); + + if let Some(ref backend) = self.settings.database_backend { + env_vars.push(("DATABASE_BACKEND", backend.clone())); + } + if let Some(ref url) = self.settings.database_url { + env_vars.push(("DATABASE_URL", url.clone())); + } + if let Some(ref path) = self.settings.libsql_path { + env_vars.push(("LIBSQL_PATH", path.clone())); + } + if let Some(ref url) = self.settings.libsql_url { + env_vars.push(("LIBSQL_URL", url.clone())); + } + + if !env_vars.is_empty() { + let pairs: Vec<(&str, &str)> = + env_vars.iter().map(|(k, v)| (*k, v.as_str())).collect(); + crate::bootstrap::save_bootstrap_env(&pairs).map_err(|e| { + SetupError::Io(std::io::Error::other(format!( + "Failed to save bootstrap env to .env: {}", + e + ))) + })?; + } } println!(); @@ -1125,10 +1573,23 @@ impl SetupWizard { KeySource::None => println!(" Security: disabled"), } + if let Some(ref provider) = self.settings.llm_backend { + let display = match provider.as_str() { + "nearai" => "NEAR AI", + "anthropic" => "Anthropic", + "openai" => "OpenAI", + "ollama" => "Ollama", + "openai_compatible" => "OpenAI-compatible", + other => other, + }; + println!(" Provider: {}", display); + } + if let Some(ref model) = self.settings.selected_model { - // Truncate long model names - let display = if model.len() > 40 { - format!("{}...", &model[..37]) + // Truncate long model names (char-based to avoid UTF-8 panic) + let display = if model.chars().count() > 40 { + let truncated: String = model.chars().take(37).collect(); + format!("{}...", truncated) } else { model.clone() }; @@ -1225,6 +1686,196 @@ fn mask_password_in_url(url: &str) -> String { format!("{}{}:****{}", scheme, username, after_at) } +/// Fetch models from the Anthropic API. +/// +/// Returns `(model_id, display_label)` pairs. Falls back to static defaults on error. +async fn fetch_anthropic_models(cached_key: Option<&str>) -> Vec<(String, String)> { + let static_defaults = vec![ + ("claude-sonnet-4-20250514".into(), "Claude Sonnet 4".into()), + ("claude-opus-4-20250514".into(), "Claude Opus 4".into()), + ( + "claude-3-5-haiku-20241022".into(), + "Claude 3.5 Haiku (fast)".into(), + ), + ]; + + let api_key = cached_key + .map(String::from) + .or_else(|| std::env::var("ANTHROPIC_API_KEY").ok()) + .filter(|k| !k.is_empty()); + + let api_key = match api_key { + Some(k) => k, + None => return static_defaults, + }; + + let client = reqwest::Client::new(); + let resp = match client + .get("https://api.anthropic.com/v1/models") + .header("x-api-key", &api_key) + .header("anthropic-version", "2023-06-01") + .timeout(std::time::Duration::from_secs(5)) + .send() + .await + { + Ok(r) if r.status().is_success() => r, + _ => return static_defaults, + }; + + #[derive(serde::Deserialize)] + struct ModelEntry { + id: String, + } + #[derive(serde::Deserialize)] + struct ModelsResponse { + data: Vec, + } + + match resp.json::().await { + Ok(body) => { + let mut models: Vec<(String, String)> = body + .data + .into_iter() + .filter(|m| !m.id.contains("embedding") && !m.id.contains("audio")) + .map(|m| { + let label = m.id.clone(); + (m.id, label) + }) + .collect(); + if models.is_empty() { + return static_defaults; + } + models.sort_by(|a, b| a.0.cmp(&b.0)); + models + } + Err(_) => static_defaults, + } +} + +/// Fetch models from the OpenAI API. +/// +/// Returns `(model_id, display_label)` pairs. Falls back to static defaults on error. +async fn fetch_openai_models(cached_key: Option<&str>) -> Vec<(String, String)> { + let static_defaults = vec![ + ("gpt-4o".into(), "GPT-4o".into()), + ("gpt-4o-mini".into(), "GPT-4o Mini (fast)".into()), + ("o3".into(), "o3 (reasoning)".into()), + ]; + + let api_key = cached_key + .map(String::from) + .or_else(|| std::env::var("OPENAI_API_KEY").ok()) + .filter(|k| !k.is_empty()); + + let api_key = match api_key { + Some(k) => k, + None => return static_defaults, + }; + + let client = reqwest::Client::new(); + let resp = match client + .get("https://api.openai.com/v1/models") + .bearer_auth(&api_key) + .timeout(std::time::Duration::from_secs(5)) + .send() + .await + { + Ok(r) if r.status().is_success() => r, + _ => return static_defaults, + }; + + #[derive(serde::Deserialize)] + struct ModelEntry { + id: String, + } + #[derive(serde::Deserialize)] + struct ModelsResponse { + data: Vec, + } + + // Prefixes that indicate chat-relevant models + let chat_prefixes = ["gpt-4", "gpt-3.5", "o1", "o3", "o4", "chatgpt"]; + + match resp.json::().await { + Ok(body) => { + let mut models: Vec<(String, String)> = body + .data + .into_iter() + .filter(|m| { + chat_prefixes.iter().any(|p| m.id.starts_with(p)) + && !m.id.contains("realtime") + && !m.id.contains("audio") + }) + .map(|m| { + let label = m.id.clone(); + (m.id, label) + }) + .collect(); + if models.is_empty() { + return static_defaults; + } + models.sort_by(|a, b| a.0.cmp(&b.0)); + models + } + Err(_) => static_defaults, + } +} + +/// Fetch installed models from a local Ollama instance. +/// +/// Returns `(model_name, display_label)` pairs. Falls back to static defaults on error. +async fn fetch_ollama_models(base_url: &str) -> Vec<(String, String)> { + let static_defaults = vec![ + ("llama3".into(), "llama3".into()), + ("mistral".into(), "mistral".into()), + ("codellama".into(), "codellama".into()), + ]; + + let url = format!("{}/api/tags", base_url.trim_end_matches('/')); + let client = reqwest::Client::new(); + + let resp = match client + .get(&url) + .timeout(std::time::Duration::from_secs(5)) + .send() + .await + { + Ok(r) if r.status().is_success() => r, + Ok(_) => return static_defaults, + Err(_) => { + print_info("Could not connect to Ollama. Is it running?"); + return static_defaults; + } + }; + + #[derive(serde::Deserialize)] + struct ModelEntry { + name: String, + } + #[derive(serde::Deserialize)] + struct TagsResponse { + models: Vec, + } + + match resp.json::().await { + Ok(body) => { + let models: Vec<(String, String)> = body + .models + .into_iter() + .map(|m| { + let label = m.name.clone(); + (m.name, label) + }) + .collect(); + if models.is_empty() { + return static_defaults; + } + models + } + Err(_) => static_defaults, + } +} + /// Discover WASM channels in a directory. /// /// Returns a list of (channel_name, capabilities_file) pairs. @@ -1244,14 +1895,14 @@ async fn discover_wasm_channels(dir: &std::path::Path) -> Vec<(String, ChannelCa let path = entry.path(); // Look for .capabilities.json files - let extension = path.file_name().and_then(|n| n.to_str()).unwrap_or(""); + let filename = path.file_name().and_then(|n| n.to_str()).unwrap_or(""); - if !extension.ends_with(".capabilities.json") { + if !filename.ends_with(".capabilities.json") { continue; } // Extract channel name - let name = extension.trim_end_matches(".capabilities.json").to_string(); + let name = filename.trim_end_matches(".capabilities.json").to_string(); if name.is_empty() { continue; } @@ -1291,6 +1942,20 @@ async fn discover_wasm_channels(dir: &std::path::Path) -> Vec<(String, ChannelCa channels } +/// Mask an API key for display: show first 6 + last 4 chars. +/// +/// Uses char-based indexing to avoid panicking on multi-byte UTF-8. +fn mask_api_key(key: &str) -> String { + let chars: Vec = key.chars().collect(); + if chars.len() < 12 { + let prefix: String = chars.iter().take(4).collect(); + return format!("{prefix}..."); + } + let prefix: String = chars[..6].iter().collect(); + let suffix: String = chars[chars.len() - 4..].iter().collect(); + format!("{prefix}...{suffix}") +} + /// Capitalize the first letter of a string. fn capitalize_first(s: &str) -> String { let mut chars = s.chars(); @@ -1408,6 +2073,20 @@ mod tests { assert_eq!(capitalize_first(""), ""); } + #[test] + fn test_mask_api_key() { + assert_eq!( + mask_api_key("sk-ant-api03-abcdef1234567890"), + "sk-ant...7890" + ); + assert_eq!(mask_api_key("short"), "shor..."); + assert_eq!(mask_api_key("exactly12ch"), "exac..."); + assert_eq!(mask_api_key("exactly12chr"), "exactl...2chr"); + assert_eq!(mask_api_key(""), "..."); + // Multi-byte chars should not panic + assert_eq!(mask_api_key("日本語キー"), "日本語キ..."); + } + #[tokio::test] async fn test_install_missing_bundled_channels_installs_telegram() { use crate::channels::wasm::available_channel_names; @@ -1456,4 +2135,76 @@ mod tests { "telegram should not be duplicated" ); } + + #[tokio::test] + async fn test_fetch_anthropic_models_static_fallback() { + // With no API key, should return static defaults + let _guard = EnvGuard::clear("ANTHROPIC_API_KEY"); + let models = fetch_anthropic_models(None).await; + assert!(!models.is_empty()); + assert!( + models.iter().any(|(id, _)| id.contains("claude")), + "static defaults should include a Claude model" + ); + } + + #[tokio::test] + async fn test_fetch_openai_models_static_fallback() { + let _guard = EnvGuard::clear("OPENAI_API_KEY"); + let models = fetch_openai_models(None).await; + assert!(!models.is_empty()); + assert!( + models.iter().any(|(id, _)| id.contains("gpt")), + "static defaults should include a GPT model" + ); + } + + #[tokio::test] + async fn test_fetch_ollama_models_unreachable_fallback() { + // Point at a port nothing listens on + let models = fetch_ollama_models("http://127.0.0.1:1").await; + assert!(!models.is_empty(), "should fall back to static defaults"); + } + + #[tokio::test] + async fn test_discover_wasm_channels_empty_dir() { + let dir = tempdir().unwrap(); + let channels = discover_wasm_channels(dir.path()).await; + assert!(channels.is_empty()); + } + + #[tokio::test] + async fn test_discover_wasm_channels_nonexistent_dir() { + let channels = + discover_wasm_channels(std::path::Path::new("/tmp/ironclaw_nonexistent_dir")).await; + assert!(channels.is_empty()); + } + + /// RAII guard that sets/clears an env var for the duration of a test. + struct EnvGuard { + key: &'static str, + original: Option, + } + + impl EnvGuard { + fn clear(key: &'static str) -> Self { + let original = std::env::var(key).ok(); + unsafe { + std::env::remove_var(key); + } + Self { key, original } + } + } + + impl Drop for EnvGuard { + fn drop(&mut self) { + unsafe { + if let Some(ref val) = self.original { + std::env::set_var(self.key, val); + } else { + std::env::remove_var(self.key); + } + } + } + } } diff --git a/tools-src/gmail/src/api.rs b/tools-src/gmail/src/api.rs index f1bc9764..7058e10f 100644 --- a/tools-src/gmail/src/api.rs +++ b/tools-src/gmail/src/api.rs @@ -136,7 +136,7 @@ fn parse_message(v: &serde_json::Value) -> Message { date: get_header(payload, "Date"), body: extract_body(payload), snippet: v["snippet"].as_str().unwrap_or("").to_string(), - is_unread: label_ids.contains(&"UNREAD".to_string()), + is_unread: label_ids.iter().any(|l| l == "UNREAD"), label_ids, } } @@ -198,7 +198,7 @@ pub fn list_messages( to: get_header(payload, "To"), date: get_header(payload, "Date"), snippet: msg["snippet"].as_str().unwrap_or("").to_string(), - is_unread: label_ids.contains(&"UNREAD".to_string()), + is_unread: label_ids.iter().any(|l| l == "UNREAD"), label_ids, }); } diff --git a/tools-src/google-calendar/src/lib.rs b/tools-src/google-calendar/src/lib.rs index a5e513d2..9cfd8ca3 100644 --- a/tools-src/google-calendar/src/lib.rs +++ b/tools-src/google-calendar/src/lib.rs @@ -6,7 +6,7 @@ //! # Capabilities Required //! //! - HTTP: `www.googleapis.com/calendar/v3/*` (GET, POST, PUT, PATCH, DELETE) -//! - Secrets: `google_calendar_token` (OAuth 2.0 token, injected automatically) +//! - Secrets: `google_oauth_token` (OAuth 2.0 token, injected automatically) //! //! # Supported Actions //! diff --git a/tools-src/google-docs/src/api.rs b/tools-src/google-docs/src/api.rs index 185ae979..f31ccf69 100644 --- a/tools-src/google-docs/src/api.rs +++ b/tools-src/google-docs/src/api.rs @@ -269,8 +269,13 @@ pub fn replace_text( let parsed = batch_update_raw(document_id, vec![request])?; - let occurrences = parsed["replies"][0]["replaceAllText"]["occurrencesChanged"] - .as_i64() + let first_reply = parsed["replies"].as_array().and_then(|arr| arr.first()); + let occurrences = first_reply + .map(|r| { + r["replaceAllText"]["occurrencesChanged"] + .as_i64() + .unwrap_or(0) + }) .unwrap_or(0); Ok(ReplaceResult { diff --git a/tools-src/google-sheets/src/api.rs b/tools-src/google-sheets/src/api.rs index 4d7af5ca..c4e7cfbd 100644 --- a/tools-src/google-sheets/src/api.rs +++ b/tools-src/google-sheets/src/api.rs @@ -330,7 +330,13 @@ pub fn add_sheet(spreadsheet_id: &str, title: &str) -> Result