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