From 64fe9ba6077e3478e2684daf97d43556aff6996d Mon Sep 17 00:00:00 2001 From: Zaki Manian Date: Sun, 29 Mar 2026 13:49:21 -0700 Subject: [PATCH 01/11] fix: prevent UTF-8 panics in byte-index string truncation (#1688) * fix: prevent UTF-8 panics in byte-index string truncation Replace unsafe `&s[..n]` patterns with `floor_char_boundary(s, n)` at 3 production code sites where the truncation index could land mid-multibyte character, panicking on non-ASCII input: - src/llm/nearai_chat.rs: API response truncation in error message - src/cli/memory.rs: memory content display truncation - src/cli/config.rs: config value display truncation All 3 sites operate on external or user-supplied strings that may contain non-ASCII characters. The existing `crate::util::floor_char_boundary` utility (used at 18 other call sites) walks back to the nearest char boundary, preventing the panic. Adds regression test with multi-byte characters (combining accents and 4-byte emoji) for truncate_content. [skip-regression-check] Co-Authored-By: Claude Opus 4.6 (1M context) * fix: clarify test comment and use exact assertions Address Gemini review feedback: - Fix misleading comment: \u{00e9} is precomposed e-acute, not combining accent - Replace weak assertions (ends_with/is_empty) with exact assert_eq! Co-Authored-By: Claude Opus 4.6 (1M context) --------- Co-authored-by: Claude Opus 4.6 (1M context) --- src/cli/config.rs | 3 ++- src/cli/memory.rs | 16 +++++++++++++++- src/llm/nearai_chat.rs | 2 +- 3 files changed, 18 insertions(+), 3 deletions(-) diff --git a/src/cli/config.rs b/src/cli/config.rs index fc1312f62b6..f1e107ecbc9 100644 --- a/src/cli/config.rs +++ b/src/cli/config.rs @@ -127,7 +127,8 @@ async fn list_settings( } let display_value = if value.len() > 60 { - format!("{}...", &value[..57]) + let end = crate::util::floor_char_boundary(&value, 57); + format!("{}...", &value[..end]) } else { value }; diff --git a/src/cli/memory.rs b/src/cli/memory.rs index 2d0606a8541..fca6d03b35d 100644 --- a/src/cli/memory.rs +++ b/src/cli/memory.rs @@ -256,7 +256,8 @@ fn truncate_content(s: &str, max_len: usize) -> String { if s.len() <= max_len { s.to_string() } else { - format!("{}...", &s[..max_len]) + let end = crate::util::floor_char_boundary(s, max_len); + format!("{}...", &s[..end]) } } @@ -292,4 +293,17 @@ mod tests { assert_eq!(truncate_content("hello", 10), "hello"); assert_eq!(truncate_content("hello world", 5), "hello..."); } + + #[test] + fn test_truncate_content_multibyte_does_not_panic() { + // \u{00e9} is precomposed 'é' (2 bytes in UTF-8) + let s = "caf\u{00e9} au lait"; // "café au lait", é starts at byte 3 + let result = truncate_content(s, 4); // byte 4 is inside 2-byte é + assert_eq!(result, "caf..."); + + // 4-byte emoji: slicing mid-emoji must not panic + let emoji = "Hi \u{1F600} there"; // 😀 is 4 bytes, starts at byte 3 + let result = truncate_content(emoji, 4); // byte 4 is inside 😀 + assert_eq!(result, "Hi ..."); + } } diff --git a/src/llm/nearai_chat.rs b/src/llm/nearai_chat.rs index 80335d86a5d..26807c9967b 100644 --- a/src/llm/nearai_chat.rs +++ b/src/llm/nearai_chat.rs @@ -451,7 +451,7 @@ impl NearAiChatProvider { provider: "nearai_chat".to_string(), reason: format!( "No model names found in response: {}", - &response_text[..response_text.len().min(300)] + &response_text[..crate::util::floor_char_boundary(&response_text, 300)] ), }) } From 70214c4ae1f85436d00c0652042d44557ad8559f Mon Sep 17 00:00:00 2001 From: rajulbhatnagar Date: Sun, 29 Mar 2026 13:59:08 -0700 Subject: [PATCH 02/11] fix(bedrock): strip tool blocks from messages when toolConfig is absent (#1630) MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit * fix(bedrock): strip tool blocks from messages when toolConfig is absent Bedrock's Converse API requires `toolConfig` whenever messages contain `toolUse` or `toolResult` content blocks. When the agentic loop reaches its force_text iteration (e.g. lightweight routine at max_iterations), it switches from `complete_with_tools()` to `complete()` — but the message history still carries tool blocks from prior iterations. `convert_messages()` faithfully converts these into Bedrock content blocks, and without `toolConfig` Bedrock rejects the request: "The toolConfig field must be defined when using toolUse and toolResult content blocks." Add `strip_tool_blocks()` that converts tool interaction data to text: - Assistant `tool_calls` → dropped (text content preserved) - `Role::Tool` → `Role::User` with `[Tool ... returned: ...]` text Wire it into: - `complete()`: unconditionally, since it never sends toolConfig - `complete_with_tools()`: when `build_tool_config()` returns None (empty tools or tool_choice="none") Closes #1629 * fix(bedrock): address review feedback on strip_tool_blocks - Add tracing::debug\! when tool blocks are stripped (zmanian suggestion) - Add test for tool_choice="none" path (zmanian suggestion) - Add inline comment on empty-content assistant behavior --------- Co-authored-by: Rajul Bhatnagar --- src/llm/bedrock.rs | 265 ++++++++++++++++++++++++++++++++++++++++++++- 1 file changed, 263 insertions(+), 2 deletions(-) diff --git a/src/llm/bedrock.rs b/src/llm/bedrock.rs index b5f7badde06..4326cabbb8e 100644 --- a/src/llm/bedrock.rs +++ b/src/llm/bedrock.rs @@ -97,6 +97,10 @@ impl LlmProvider for BedrockProvider { let mut messages = request.messages; crate::llm::provider::sanitize_tool_messages(&mut messages); + // Bedrock requires toolConfig when messages contain ToolUse/ToolResult + // blocks. Messages may carry tool history from prior agentic iterations, + // but complete() has no tools to build a toolConfig — strip them. + strip_tool_blocks(&mut messages); let (system_blocks, bedrock_messages) = convert_messages(&messages)?; @@ -150,6 +154,14 @@ impl LlmProvider for BedrockProvider { let mut messages = request.messages; crate::llm::provider::sanitize_tool_messages(&mut messages); + let tool_config = build_tool_config(&request.tools, request.tool_choice.as_deref())?; + + // When tool_config is None (empty tools or tool_choice="none") but messages + // contain tool history, strip tool blocks to avoid Bedrock validation error. + if tool_config.is_none() { + strip_tool_blocks(&mut messages); + } + let (system_blocks, bedrock_messages) = convert_messages(&messages)?; if bedrock_messages.is_empty() { @@ -159,8 +171,6 @@ impl LlmProvider for BedrockProvider { }); } - let tool_config = build_tool_config(&request.tools, request.tool_choice.as_deref())?; - let mut builder = self .client .converse() @@ -268,6 +278,51 @@ fn build_inference_config( } } +// --------------------------------------------------------------------------- +// Tool-block stripping for tool-free requests +// --------------------------------------------------------------------------- + +/// Strip tool interaction data from messages so they can be sent without `toolConfig`. +/// +/// Bedrock's Converse API requires `toolConfig` whenever messages contain `toolUse` +/// or `toolResult` content blocks. When `complete()` is called (no tools) or +/// `complete_with_tools()` resolves to an empty tool config, this converts: +/// - Assistant messages with `tool_calls` → keep text only, drop tool_calls +/// - `Role::Tool` messages → `Role::User` with text representation +/// +/// Note: this intentionally loses structured tool_call_id correlation — the text +/// representation is sufficient for force_text mode where no further tool dispatch +/// occurs. +fn strip_tool_blocks(messages: &mut [crate::llm::provider::ChatMessage]) { + use crate::llm::provider::Role; + + let mut stripped = 0u32; + for msg in messages.iter_mut() { + match msg.role { + Role::Assistant if msg.tool_calls.is_some() => { + // May leave content empty; convert_messages() skips empty assistant messages. + msg.tool_calls = None; + stripped += 1; + } + Role::Tool => { + let tool_name = msg.name.as_deref().unwrap_or("unknown"); + msg.role = Role::User; + msg.content = format!("[Tool `{}` returned: {}]", tool_name, msg.content); + msg.tool_call_id = None; + msg.name = None; + stripped += 1; + } + _ => {} + } + } + if stripped > 0 { + tracing::debug!( + stripped, + "Stripped tool blocks from messages (no toolConfig)" + ); + } +} + // --------------------------------------------------------------------------- // Message conversion // --------------------------------------------------------------------------- @@ -1155,4 +1210,210 @@ mod tests { let (_, bedrock_msgs) = convert_messages(&messages).unwrap(); assert!(bedrock_msgs.is_empty()); } + + #[test] + fn test_strip_tool_blocks_removes_tool_content() { + let tc = crate::llm::provider::ToolCall { + id: "call_1".to_string(), + name: "echo".to_string(), + arguments: serde_json::json!({"text": "hi"}), + reasoning: None, + }; + + let mut messages = vec![ + ChatMessage::user("Do things"), + ChatMessage::assistant_with_tool_calls(Some("Let me help.".to_string()), vec![tc]), + ChatMessage::tool_result("call_1", "echo", "hi back"), + ChatMessage::user("Thanks"), + ]; + + strip_tool_blocks(&mut messages); + + // Assistant keeps text, loses tool_calls + assert_eq!(messages[1].role, Role::Assistant); + assert!(messages[1].tool_calls.is_none()); + assert_eq!(messages[1].content, "Let me help."); + + // Tool result converted to user message + assert_eq!(messages[2].role, Role::User); + assert!( + messages[2] + .content + .contains("[Tool `echo` returned: hi back]") + ); + assert!(messages[2].tool_call_id.is_none()); + assert!(messages[2].name.is_none()); + + // Other messages unchanged + assert_eq!(messages[0].role, Role::User); + assert_eq!(messages[3].role, Role::User); + assert_eq!(messages[3].content, "Thanks"); + } + + /// Regression test for: "The toolConfig field must be defined when using + /// toolUse and toolResult content blocks." + /// + /// Reproduces the exact scenario: force_text iteration sends messages with + /// tool history to complete(), which has no toolConfig. + #[test] + fn test_strip_tool_blocks_then_convert_produces_no_tool_blocks() { + let tc = crate::llm::provider::ToolCall { + id: "call_abc".to_string(), + name: "get_weather".to_string(), + arguments: serde_json::json!({"city": "NYC"}), + reasoning: None, + }; + + let mut messages = vec![ + ChatMessage::system("You are helpful."), + ChatMessage::user("What's the weather?"), + ChatMessage::assistant_with_tool_calls(Some("Checking...".to_string()), vec![tc]), + ChatMessage::tool_result("call_abc", "get_weather", "72°F and sunny"), + ChatMessage::user("Now summarize."), + ]; + + // Simulate the complete() pipeline + crate::llm::provider::sanitize_tool_messages(&mut messages); + strip_tool_blocks(&mut messages); + + let (_, bedrock_msgs) = convert_messages(&messages).unwrap(); + + for msg in &bedrock_msgs { + for block in msg.content() { + assert!( + !block.is_tool_use(), + "ToolUse block found — would cause Bedrock validation error" + ); + assert!( + !block.is_tool_result(), + "ToolResult block found — would cause Bedrock validation error" + ); + } + } + } + + #[test] + fn test_complete_with_tools_empty_tools_strips_history() { + let tc = crate::llm::provider::ToolCall { + id: "call_1".to_string(), + name: "time".to_string(), + arguments: serde_json::json!({}), + reasoning: None, + }; + + let mut messages = vec![ + ChatMessage::user("Do something"), + ChatMessage::assistant_with_tool_calls(None, vec![tc]), + ChatMessage::tool_result("call_1", "time", "12:00"), + ]; + + // Simulate complete_with_tools() with empty tools + crate::llm::provider::sanitize_tool_messages(&mut messages); + let tool_config = build_tool_config(&[], None).unwrap(); + assert!(tool_config.is_none()); + + if tool_config.is_none() { + strip_tool_blocks(&mut messages); + } + + let (_, bedrock_msgs) = convert_messages(&messages).unwrap(); + + for msg in &bedrock_msgs { + for block in msg.content() { + assert!(!block.is_tool_use()); + assert!(!block.is_tool_result()); + } + } + } + + #[test] + fn test_strip_tool_only_assistant_then_convert_maintains_alternation() { + // Edge case: assistant message with ONLY tool_calls (no text) becomes + // empty after stripping. convert_messages() should skip it, and the + // subsequent tool-result-turned-user message should merge correctly. + let tc = crate::llm::provider::ToolCall { + id: "call_1".to_string(), + name: "search".to_string(), + arguments: serde_json::json!({"q": "test"}), + reasoning: None, + }; + + let mut messages = vec![ + ChatMessage::user("Find it"), + ChatMessage::assistant_with_tool_calls(None, vec![tc]), + ChatMessage::tool_result("call_1", "search", "found 3 results"), + ChatMessage::user("Thanks"), + ]; + + crate::llm::provider::sanitize_tool_messages(&mut messages); + strip_tool_blocks(&mut messages); + + let (_, bedrock_msgs) = convert_messages(&messages).unwrap(); + + // Empty assistant is skipped; tool-result-as-user merges with first user. + // Result: User("Find it" + tool text) → User("Thanks") + // push_message merges consecutive users, so we get 1 merged user message + // then "Thanks" as a second user — but these are also consecutive users + // so they merge too. Final: single User message. + // + // Verify: strict user/assistant alternation and no tool blocks. + for (i, msg) in bedrock_msgs.iter().enumerate() { + let expected_role = if i % 2 == 0 { + ConversationRole::User + } else { + ConversationRole::Assistant + }; + assert_eq!( + *msg.role(), + expected_role, + "Message {} has wrong role for alternation", + i + ); + for block in msg.content() { + assert!(!block.is_tool_use()); + assert!(!block.is_tool_result()); + } + } + } + + #[test] + fn test_complete_with_tools_choice_none_strips_history() { + // tool_choice="none" causes build_tool_config to return None, + // which should trigger stripping of tool blocks from messages. + let tools = vec![ToolDefinition { + name: "echo".to_string(), + description: "Echoes".to_string(), + parameters: serde_json::json!({"type": "object"}), + }]; + + let tc = crate::llm::provider::ToolCall { + id: "call_1".to_string(), + name: "echo".to_string(), + arguments: serde_json::json!({}), + reasoning: None, + }; + + let mut messages = vec![ + ChatMessage::user("Do it"), + ChatMessage::assistant_with_tool_calls(None, vec![tc]), + ChatMessage::tool_result("call_1", "echo", "done"), + ]; + + crate::llm::provider::sanitize_tool_messages(&mut messages); + let tool_config = build_tool_config(&tools, Some("none")).unwrap(); + assert!(tool_config.is_none()); + + if tool_config.is_none() { + strip_tool_blocks(&mut messages); + } + + let (_, bedrock_msgs) = convert_messages(&messages).unwrap(); + + for msg in &bedrock_msgs { + for block in msg.content() { + assert!(!block.is_tool_use()); + assert!(!block.is_tool_result()); + } + } + } } From de384b0cc751c85b8faf5af18bc95fcd3633370f Mon Sep 17 00:00:00 2001 From: jinxin <106428113+italic-jinxin@users.noreply.github.com> Date: Mon, 30 Mar 2026 05:13:28 +0800 Subject: [PATCH 03/11] feat: support custom LLM provider configuration via web UI (#1340) * feat: support custom LLM provider configuration via web UI Users can now define custom LLM providers through the web UI and have them take effect without modifying environment variables or config files. - Add `CustomLlmProviderSettings` struct and `llm_custom_providers` field to `Settings` so custom provider definitions are persisted and loaded from the DB settings table - Add `LlmConfig::resolve_custom_provider()` to build a `RegistryProviderConfig` from user-defined provider data (base_url, adapter, model, api_key) - Flip resolution priority to `db > env > default` so active provider set through the UI takes precedence over deployment env vars - Warn when a custom provider is missing base_url or model - Add startup info logs for backend source and provider creation - Add regression tests for custom provider resolution and DB priority * feat: add test connection for custom LLM providers - Add POST /api/llm/test_connection endpoint that validates connectivity and auth for OpenAI-compatible, Anthropic, and Ollama adapters (10s timeout, per-adapter request logic) - Add "Test" button next to Save/Cancel in the add-provider form; result shown inline with green/red styling - Hide delete button for the active provider instead of showing an error toast - Sort the active provider to the top of the provider list - Clear selected_model when switching providers to avoid model-not-supported errors on the new provider - Add i18n keys for test/testing states (en + zh-CN) * feat: add built-in provider API key and model configuration - Add Configure button on built-in provider cards (openai, anthropic, gemini, ollama, etc.) to set API key and default model via web UI - Store overrides as `llm_builtin_overrides` setting (per-provider key/model map) using the existing generic settings k/v API - Add LlmBuiltinOverride struct in settings.rs; resolve in resolve_registry_provider() with priority: env var > selected_model > llm_builtin_overrides[id] > default - Restore provider's configured model to selected_model on provider switch, so /model command always takes precedence at runtime - Fix fetch-models button in built-in configure mode: use hardcoded base_url from BUILTIN_PROVIDERS instead of the hidden form field - Add edit support for custom providers with pre-filled dialog - Show current model on active and configured provider cards - Convert add/edit provider form to a modal dialog - Sync selected_model when editing or deleting an active custom provider * feat: move Config tab into Settings as Providers subtab * feat(web): merge Providers into Inference tab with UX improvements * chore: resolve conflicts * fix(llm): address security and correctness issues in custom LLM provider * fix(llm): address security and correctness issues in custom LLM provider * feat(web): fall back to env vars for LLM provider config in UI * fix(llm): enforce db > env > default config priority for provider setting * fix: address review feedback on provider config priority * feat: extract BUILTIN_PROVIDERS into providers.js * fix(security): store LLM API keys in encrypted secrets store instead of plaintext * fix(security): harden LLM API key handling across settings and LLM endpoints * fix: test_connection sends actual chat completion * refactor(web): derive LLM Provider display from active Model Provider * fix(settings): language switch not working for llm provider * feat(web): add restart notice to LLM Provider settings * fix: review fixes for custom LLM provider PR - Add server-side validation of custom provider ID format (lowercase alphanumeric + hyphens, 1-64 chars) to match frontend regex - Tighten is_nearai_private_endpoint to exact-match private.near.ai or *.private.near.ai, rejecting lookalikes like private-evil.near.ai - Fix misleading priority doc comments in config/mod.rs and settings.rs to reflect the split model: LLM uses DB > env, others use env > DB - Clean up #1581 artifacts: remove TOML file creation from persist_selected_model (DB is sufficient), update stale priority comments in commands.rs, fix contradictory test assertions - Add 18 new tests for provider ID validation, adapter validation, and nearai private endpoint matching Co-Authored-By: Claude Opus 4.6 (1M context) * fix: address PR review comments for custom LLM provider - Move LLM handlers (test_connection, list_models, env_defaults) from server.rs to handlers/llm.rs for consistency with other handler modules - Merge validate_custom_providers into single pass (ID + adapter check) - Allow underscores in custom provider IDs to match builtin naming - Add missing i18n key config.fetchingModels (en + zh-CN) - Fix optional_env().ok().flatten() error swallowing in config/llm.rs; propagate ConfigError with ? instead of silently discarding - Narrow settings.rs module docs to scope DB>env precedence to LLM - Add unit tests for hydrate_llm_keys_from_secrets Co-Authored-By: Claude Opus 4.6 (1M context) * refactor: replace static providers.js with API endpoint from registry - Delete providers.js; serve provider list from /api/llm/providers endpoint that reads from the embedded ProviderRegistry (providers.json) - Centralize secret naming (builtin_secret_name, custom_secret_name) into settings.rs; replace 8 duplicated format! calls across 4 files - Extract JS API_KEY_UNCHANGED constant; replace 6 magic string literals - Replace hard-coded API key placeholder strings with i18n keys (config.apiKeyConfigured, config.apiKeyFromEnv, config.apiKeyEnter) - Simplify apiFetchVoid to delegate to apiFetch - Remove unnecessary Vec clones in guard_active_provider_not_removed Co-Authored-By: Claude Opus 4.6 (1M context) --------- Co-authored-by: Robert Yan <46699230+think-in-universe@users.noreply.github.com> Co-authored-by: Illia Polosukhin Co-authored-by: Claude Opus 4.6 (1M context) --- src/agent/commands.rs | 36 +- src/app.rs | 23 +- src/channels/web/handlers/llm.rs | 664 +++++++++++++++++++ src/channels/web/handlers/mod.rs | 1 + src/channels/web/handlers/settings.rs | 826 +++++++++++++++++++++++- src/channels/web/server.rs | 143 +---- src/channels/web/static/app.js | 683 ++++++++++++++++++-- src/channels/web/static/i18n-app.js | 17 +- src/channels/web/static/i18n/en.js | 45 ++ src/channels/web/static/i18n/zh-CN.js | 45 ++ src/channels/web/static/index.html | 71 ++- src/channels/web/static/style.css | 420 ++++++++++++ src/config/llm.rs | 877 ++++++++++++++++++++++++-- src/config/mod.rs | 369 ++++++++++- src/llm/mod.rs | 2 + src/llm/rig_adapter.rs | 74 +++ src/main.rs | 3 + src/settings.rs | 197 ++++-- 18 files changed, 4176 insertions(+), 320 deletions(-) create mode 100644 src/channels/web/handlers/llm.rs diff --git a/src/agent/commands.rs b/src/agent/commands.rs index 643d8c7cc16..04a1022ae61 100644 --- a/src/agent/commands.rs +++ b/src/agent/commands.rs @@ -947,6 +947,12 @@ impl Agent { /// Best-effort: logs warnings on failure but does not propagate errors, /// since the in-memory model switch already succeeded. /// + /// The DB setting is the primary persistence layer. For LLM settings the + /// resolution priority is `DB > env > TOML > default`, so writing to DB + /// is sufficient for the change to survive restarts. The `.env` and TOML + /// files are only updated as a courtesy when they already contain a model + /// var, to avoid user confusion. + /// /// In multi-tenant mode, only the per-user DB setting is written — global /// .env and TOML files are shared across users and must not be mutated. async fn persist_selected_model(&self, tenant: &crate::tenant::TenantCtx, model: &str) { @@ -972,22 +978,18 @@ impl Agent { return; } - // 3. Update .env and TOML config file (sync I/O in spawn_blocking). + // 3. Best-effort update of .env and TOML if they already contain a + // model var. DB is authoritative (DB > env > TOML), but keeping + // these in sync avoids confusion when users inspect the files. let model_owned = model.to_string(); let backend = self.deps.llm_backend.clone(); if let Err(e) = tokio::task::spawn_blocking(move || { - // 2a. Update the backend-specific model env var in ~/.ironclaw/.env. - // - // Env vars have the HIGHEST priority in LlmConfig::resolve_model() - // (env var > TOML > DB > default). If the .env file has e.g. - // NEARAI_MODEL=old-model, it shadows everything else. We must - // update this var or the /model change is invisible on restart. + // 3a. Update the backend-specific model env var in ~/.ironclaw/.env + // only if the var already exists (don't inject new vars). let registry = crate::llm::ProviderRegistry::load(); let model_env = registry.model_env_var(&backend); let env_var_prefix = format!("{}=", model_env); - // Only update the .env file if the var is actually set there - // (avoid injecting new vars the user never configured). let env_path = crate::bootstrap::ironclaw_env_path(); let env_has_var = std::fs::read_to_string(&env_path) .ok() @@ -1005,10 +1007,8 @@ impl Agent { } } - // 2b. Update (or create) the TOML config file. - // - // The TOML overlay has higher priority than DB settings on - // startup, so it MUST stay in sync with the DB. + // 3b. Update TOML config file if it already exists. + // Don't create a new one — DB persistence is sufficient. let toml_path = crate::settings::Settings::default_toml_path(); match crate::settings::Settings::load_toml(&toml_path) { Ok(Some(mut settings)) => { @@ -1018,15 +1018,7 @@ impl Agent { } } Ok(None) => { - // No config file yet — create one so the model choice - // survives restarts even when the DB is unavailable. - let settings = crate::settings::Settings { - selected_model: Some(model_owned), - ..Default::default() - }; - if let Err(e) = settings.save_toml(&toml_path) { - tracing::warn!("Failed to create config.toml for model persistence: {}", e); - } + // No config file on disk; DB persistence is sufficient. } Err(e) => { tracing::warn!("Failed to load config.toml for model persistence: {}", e); diff --git a/src/app.rs b/src/app.rs index 0ca9cd9ec4d..d737795f59f 100644 --- a/src/app.rs +++ b/src/app.rs @@ -229,18 +229,35 @@ impl AppBuilder { let store = crate::secrets::create_secrets_store(crypto, handles); if let Some(ref secrets) = store { + // Migrate any plaintext API keys from the settings table to the + // encrypted secrets store. Idempotent — safe to run on every startup. + if let Some(ref db) = self.db { + crate::config::migrate_plaintext_llm_keys( + db.as_ref(), + secrets.as_ref(), + &self.config.owner_id, + ) + .await; + } + // Inject LLM API keys from encrypted storage crate::config::inject_llm_keys_from_secrets(secrets.as_ref(), &self.config.owner_id) .await; - // Re-resolve only the LLM config with newly available keys. - let store: Option<&(dyn crate::db::SettingsStore + Sync)> = + // Re-resolve only the LLM config with newly available keys, + // including keys hydrated from the secrets store. + let settings_store: Option<&(dyn crate::db::SettingsStore + Sync)> = self.db.as_ref().map(|db| db.as_ref() as _); let toml_path = self.toml_path.as_deref(); let owner_id = self.config.owner_id.clone(); if let Err(e) = self .config - .re_resolve_llm(store, &owner_id, toml_path) + .re_resolve_llm_with_secrets( + settings_store, + &owner_id, + toml_path, + Some(secrets.as_ref()), + ) .await { tracing::warn!("Failed to re-resolve LLM config after secret injection: {e}"); diff --git a/src/channels/web/handlers/llm.rs b/src/channels/web/handlers/llm.rs new file mode 100644 index 00000000000..8d4e54a48c5 --- /dev/null +++ b/src/channels/web/handlers/llm.rs @@ -0,0 +1,664 @@ +//! LLM utility handlers: test connection, list models, env defaults. + +use std::sync::Arc; + +use axum::{Json, extract::State}; + +use crate::channels::web::auth::AuthenticatedUser; +use crate::channels::web::server::GatewayState; +use crate::config::helpers::validate_base_url; + +// --------------------------------------------------------------------------- +// Test connection +// --------------------------------------------------------------------------- + +/// Fields shared by `test_connection` and `list_models` requests. +/// +/// When `api_key` is absent the handler falls back to the encrypted secrets +/// store, using `provider_id` + `provider_type` to locate the vaulted key. +#[derive(serde::Deserialize)] +pub struct TestConnectionRequest { + adapter: String, + base_url: String, + /// Model to use for the test chat completion request. + model: String, + #[serde(default)] + api_key: Option, + /// Provider identifier used to look up the vaulted API key when `api_key` + /// is not supplied by the frontend (key already stored in secrets). + #[serde(default)] + provider_id: Option, + /// `"builtin"` or `"custom"` — determines the secret name prefix. + #[serde(default)] + provider_type: Option, +} + +#[derive(serde::Serialize)] +pub struct TestConnectionResponse { + ok: bool, + message: String, +} + +pub async fn llm_test_connection_handler( + State(state): State>, + AuthenticatedUser(user): AuthenticatedUser, + Json(mut body): Json, +) -> Json { + resolve_api_key_from_secrets( + &state, + &user.user_id, + &mut body.api_key, + &body.provider_id, + &body.provider_type, + ) + .await; + Json(test_provider_connection(body).await) +} + +async fn test_provider_connection(req: TestConnectionRequest) -> TestConnectionResponse { + if let Err(e) = validate_base_url(&req.base_url, "base_url") { + return TestConnectionResponse { + ok: false, + message: format!("Invalid base URL: {e}"), + }; + } + + if req.model.trim().is_empty() { + return TestConnectionResponse { + ok: false, + message: "Model is required for connection test".to_string(), + }; + } + + let client = match reqwest::Client::builder() + .timeout(std::time::Duration::from_secs(30)) + .build() + { + Ok(c) => c, + Err(e) => { + return TestConnectionResponse { + ok: false, + message: format!("Failed to build HTTP client: {e}"), + }; + } + }; + + let base = req.base_url.trim_end_matches('/'); + + match req.adapter.as_str() { + "anthropic" => { + let anthropic_base = if base.ends_with("/v1") || base.contains("/v1/") { + base.to_string() + } else { + format!("{base}/v1") + }; + let url = format!("{anthropic_base}/messages"); + let body = serde_json::json!({ + "model": req.model, + "max_tokens": 16, + "messages": [{"role": "user", "content": "hi"}] + }); + let mut builder = client + .post(&url) + .header("anthropic-version", "2023-06-01") + .json(&body); + if let Some(key) = req.api_key.as_deref().filter(|k| !k.is_empty()) { + builder = builder.header("x-api-key", key); + } + interpret_chat_response(builder.send().await) + } + "ollama" => { + let url = format!("{base}/api/chat"); + let body = serde_json::json!({ + "model": req.model, + "messages": [{"role": "user", "content": "hi"}], + "stream": false + }); + let builder = client.post(&url).json(&body); + interpret_chat_response(builder.send().await) + } + _ => { + // OpenAI-compatible (including nearai): POST /v1/chat/completions + // If base already ends with /v1, append directly; otherwise insert /v1. + let chat_url = if base.ends_with("/v1") { + format!("{base}/chat/completions") + } else { + format!("{base}/v1/chat/completions") + }; + let body = serde_json::json!({ + "model": req.model, + "max_tokens": 16, + "messages": [{"role": "user", "content": "hi"}] + }); + let mut builder = client.post(&chat_url).json(&body); + if let Some(key) = req.api_key.as_deref().filter(|k| !k.is_empty()) { + builder = builder.header("Authorization", format!("Bearer {key}")); + } + interpret_chat_response(builder.send().await) + } + } +} + +fn interpret_chat_response( + result: Result, +) -> TestConnectionResponse { + match result { + Ok(r) => { + let status = r.status(); + if status.is_success() { + TestConnectionResponse { + ok: true, + message: format!("Connected ({})", status), + } + } else if status == reqwest::StatusCode::UNAUTHORIZED + || status == reqwest::StatusCode::FORBIDDEN + { + TestConnectionResponse { + ok: false, + message: format!("Authentication failed ({})", status), + } + } else if status == reqwest::StatusCode::BAD_REQUEST + || status == reqwest::StatusCode::UNPROCESSABLE_ENTITY + { + // 400/422 = server reachable, likely wrong endpoint variant — connectivity OK + TestConnectionResponse { + ok: true, + message: format!("Server reachable ({})", status), + } + } else if status == reqwest::StatusCode::NOT_FOUND { + // 404 = /models endpoint not found — server reachable but not OpenAI-compatible + TestConnectionResponse { + ok: false, + message: format!( + "Server reachable but /models endpoint not found ({}). \ + Check the base URL and adapter type.", + status + ), + } + } else if status.is_client_error() { + TestConnectionResponse { + ok: false, + message: format!("Client error ({})", status), + } + } else { + TestConnectionResponse { + ok: false, + message: format!("Server error ({})", status), + } + } + } + Err(e) => TestConnectionResponse { + ok: false, + message: format!("Connection failed: {e}"), + }, + } +} + +// --------------------------------------------------------------------------- +// List models +// --------------------------------------------------------------------------- + +#[derive(serde::Deserialize)] +pub struct ListModelsRequest { + adapter: String, + base_url: String, + #[serde(default)] + api_key: Option, + #[serde(default)] + provider_id: Option, + #[serde(default)] + provider_type: Option, +} + +#[derive(serde::Serialize)] +pub struct ListModelsResponse { + ok: bool, + models: Vec, + message: String, +} + +pub async fn llm_list_models_handler( + State(state): State>, + AuthenticatedUser(user): AuthenticatedUser, + Json(mut body): Json, +) -> Json { + resolve_api_key_from_secrets( + &state, + &user.user_id, + &mut body.api_key, + &body.provider_id, + &body.provider_type, + ) + .await; + Json(fetch_provider_models(body).await) +} + +async fn fetch_provider_models(req: ListModelsRequest) -> ListModelsResponse { + if let Err(e) = validate_base_url(&req.base_url, "base_url") { + return ListModelsResponse { + ok: false, + models: vec![], + message: format!("Invalid base URL: {e}"), + }; + } + + let client = match reqwest::Client::builder() + .timeout(std::time::Duration::from_secs(15)) + .build() + { + Ok(c) => c, + Err(e) => { + return ListModelsResponse { + ok: false, + models: vec![], + message: format!("Failed to build HTTP client: {e}"), + }; + } + }; + + let base = req.base_url.trim_end_matches('/'); + let auth = req.api_key.as_deref().filter(|k| !k.is_empty()); + + match req.adapter.as_str() { + "ollama" => { + let url = format!("{base}/api/tags"); + match client.get(&url).send().await { + Ok(r) if r.status().is_success() => { + let body: serde_json::Value = r.json().await.unwrap_or_default(); + let models: Vec = body["models"] + .as_array() + .map(|arr| { + arr.iter() + .filter_map(|m| m["name"].as_str().map(|s| s.to_string())) + .collect() + }) + .unwrap_or_default(); + if models.is_empty() { + ListModelsResponse { + ok: false, + models: vec![], + message: "No models found".to_string(), + } + } else { + ListModelsResponse { + ok: true, + message: format!("{} model(s) found", models.len()), + models, + } + } + } + Ok(r) => ListModelsResponse { + ok: false, + models: vec![], + message: format!("Server returned {}", r.status()), + }, + Err(e) => ListModelsResponse { + ok: false, + models: vec![], + message: format!("Connection failed: {e}"), + }, + } + } + _ => { + // OpenAI-compatible, Anthropic, and NEAR AI all support GET /models. + // NEAR AI private endpoints and Anthropic need a /v1 prefix. + let effective_base = if (req.adapter == "nearai" && is_nearai_private_endpoint(base)) + || (req.adapter == "anthropic" && !base.ends_with("/v1") && !base.contains("/v1/")) + { + format!("{base}/v1") + } else { + base.to_string() + }; + let url = format!("{effective_base}/models"); + let mut builder = client.get(&url); + if req.adapter == "anthropic" { + // Anthropic requires a version header and uses x-api-key for authentication + builder = builder.header("anthropic-version", "2023-06-01"); + if let Some(key) = auth { + builder = builder.header("x-api-key", key); + } + } else if let Some(key) = auth { + builder = builder.header("Authorization", format!("Bearer {key}")); + } + match builder.send().await { + Ok(r) if r.status().is_success() => { + let body: serde_json::Value = r.json().await.unwrap_or_default(); + // OpenAI: {"data": [{"id": "..."}]} + // Anthropic: {"data": [{"id": "..."}]} + let models: Vec = body["data"] + .as_array() + .map(|arr| { + arr.iter() + .filter_map(|m| m["id"].as_str().map(|s| s.to_string())) + .collect() + }) + .unwrap_or_default(); + if models.is_empty() { + ListModelsResponse { + ok: false, + models: vec![], + message: "No models found in response".to_string(), + } + } else { + ListModelsResponse { + ok: true, + message: format!("{} model(s) found", models.len()), + models, + } + } + } + Ok(r) => ListModelsResponse { + ok: false, + models: vec![], + message: format!("Server returned {} — list models not supported", r.status()), + }, + Err(e) => ListModelsResponse { + ok: false, + models: vec![], + message: format!("Connection failed: {e}"), + }, + } + } + } +} + +// --------------------------------------------------------------------------- +// Provider list + env defaults (replaces static providers.js) +// --------------------------------------------------------------------------- + +/// Returns all builtin LLM provider definitions plus env-var defaults. +/// +/// Each entry contains the provider definition (id, name, adapter, base_url, +/// default_model, api_key_required, can_list_models) and env-var overrides +/// (has_api_key presence flag, model override, base_url override). +/// API keys are never returned — only a boolean `has_api_key`. +pub async fn llm_providers_handler( + AuthenticatedUser(_user): AuthenticatedUser, +) -> Json { + Json(build_llm_providers()) +} + +fn build_llm_providers() -> serde_json::Value { + use crate::config::helpers::optional_env; + use crate::llm::registry::ProviderRegistry; + + let registry = ProviderRegistry::load(); + + // Helper: read env var via optional_env (checks real env + injected overlay). + // Intentionally swallows ConfigError — this is a best-effort informational + // endpoint, not a config resolver. + let read_env = |key: &str| -> Option { optional_env(key).ok().flatten() }; + + let mut providers = Vec::new(); + + // NEAR AI is not in the registry — add it as a special case. + { + let mut entry = serde_json::Map::new(); + entry.insert("id".into(), "nearai".into()); + entry.insert("name".into(), "NEAR AI".into()); + entry.insert("adapter".into(), "nearai".into()); + entry.insert("base_url".into(), "https://cloud-api.near.ai/v1".into()); + entry.insert("builtin".into(), true.into()); + entry.insert( + "default_model".into(), + serde_json::Value::String(crate::llm::DEFAULT_MODEL.to_string()), + ); + entry.insert("api_key_required".into(), true.into()); + entry.insert("can_list_models".into(), true.into()); + // Env defaults + entry.insert( + "has_api_key".into(), + read_env("NEARAI_API_KEY").is_some().into(), + ); + if let Some(model) = read_env("NEARAI_MODEL") { + entry.insert("env_model".into(), serde_json::Value::String(model)); + } + if let Some(url) = read_env("NEARAI_BASE_URL") { + entry.insert("env_base_url".into(), serde_json::Value::String(url)); + } + providers.push(serde_json::Value::Object(entry)); + } + + // Registry-based providers + for def in registry.all() { + let mut entry = serde_json::Map::new(); + entry.insert("id".into(), serde_json::Value::String(def.id.clone())); + // Use display_name from setup hint, falling back to titlecased id. + let name = def + .setup + .as_ref() + .map(|s| s.display_name().to_string()) + .unwrap_or_else(|| def.id.clone()); + entry.insert("name".into(), serde_json::Value::String(name)); + // Serialize protocol as the adapter name the frontend expects. + let adapter = serde_json::to_value(def.protocol) + .ok() + .and_then(|v| v.as_str().map(String::from)) + .unwrap_or_else(|| "open_ai_completions".to_string()); + entry.insert("adapter".into(), serde_json::Value::String(adapter)); + entry.insert( + "base_url".into(), + serde_json::Value::String(def.default_base_url.clone().unwrap_or_default()), + ); + entry.insert("builtin".into(), true.into()); + entry.insert( + "default_model".into(), + serde_json::Value::String(def.default_model.clone()), + ); + entry.insert("api_key_required".into(), def.api_key_required.into()); + let can_list = def.setup.as_ref().is_some_and(|s| s.can_list_models()); + entry.insert("can_list_models".into(), can_list.into()); + // Env defaults + if let Some(ref api_key_env) = def.api_key_env { + entry.insert("has_api_key".into(), read_env(api_key_env).is_some().into()); + } + if let Some(model) = read_env(&def.model_env) { + entry.insert("env_model".into(), serde_json::Value::String(model)); + } + if let Some(ref base_url_env) = def.base_url_env + && let Some(url) = read_env(base_url_env) + { + entry.insert("env_base_url".into(), serde_json::Value::String(url)); + } + providers.push(serde_json::Value::Object(entry)); + } + + // Bedrock is not in the registry — add it as a special case. + { + let mut entry = serde_json::Map::new(); + entry.insert("id".into(), "bedrock".into()); + entry.insert("name".into(), "AWS Bedrock".into()); + entry.insert("adapter".into(), "bedrock".into()); + entry.insert("base_url".into(), "".into()); + entry.insert("builtin".into(), true.into()); + entry.insert( + "default_model".into(), + "anthropic.claude-3-sonnet-20240229-v1:0".into(), + ); + entry.insert("api_key_required".into(), false.into()); + entry.insert("can_list_models".into(), false.into()); + providers.push(serde_json::Value::Object(entry)); + } + + serde_json::Value::Array(providers) +} + +// --------------------------------------------------------------------------- +// Shared helpers +// --------------------------------------------------------------------------- + +/// When the frontend doesn't supply an `api_key` (because it was already vaulted), +/// look it up from the encrypted secrets store using `provider_id` + `provider_type`. +async fn resolve_api_key_from_secrets( + state: &GatewayState, + user_id: &str, + api_key: &mut Option, + provider_id: &Option, + provider_type: &Option, +) { + // Already have a key from the request — nothing to resolve. + if api_key.as_ref().is_some_and(|k| !k.is_empty()) { + return; + } + let pid = match provider_id.as_deref().filter(|s| !s.is_empty()) { + Some(id) => id, + None => return, + }; + let secrets = match state.secrets_store.as_ref() { + Some(s) => s, + None => return, + }; + let secret_name = match provider_type.as_deref() { + Some("custom") => crate::settings::custom_secret_name(pid), + _ => crate::settings::builtin_secret_name(pid), + }; + if let Ok(decrypted) = secrets.get_decrypted(user_id, &secret_name).await { + *api_key = Some(decrypted.expose().to_string()); + } +} + +/// Check if a base URL belongs to a NEAR AI private endpoint. +/// +/// Matches `private.near.ai` exactly or any subdomain of it +/// (e.g. `us.private.near.ai`). Rejects lookalikes like +/// `private-evil.near.ai` or `myprivate.near.ai`. +fn is_nearai_private_endpoint(base_url: &str) -> bool { + url::Url::parse(base_url) + .ok() + .and_then(|u| u.host_str().map(|h| h.to_lowercase())) + .is_some_and(|host| host == "private.near.ai" || host.ends_with(".private.near.ai")) +} + +#[cfg(test)] +mod tests { + use super::*; + + // --- LLM providers handler tests --- + + fn find_provider<'a>( + providers: &'a [serde_json::Value], + id: &str, + ) -> Option<&'a serde_json::Value> { + providers + .iter() + .find(|p| p.get("id").and_then(|v| v.as_str()) == Some(id)) + } + + #[tokio::test] + async fn test_llm_providers_returns_nearai_with_env_vars() { + // SAFETY: test-only; tokio::test runs single-threaded by default. + unsafe { + std::env::set_var("NEARAI_API_KEY", "test-key-123"); + std::env::set_var("NEARAI_MODEL", "test-model"); + std::env::set_var("NEARAI_BASE_URL", "https://test.near.ai/v1"); + } + + let result = build_llm_providers(); + let arr = result.as_array().expect("should be an array"); + + let nearai = find_provider(arr, "nearai").expect("nearai entry"); + // API key should NOT be exposed — only has_api_key presence flag. + assert_eq!( + nearai.get("has_api_key").and_then(|v| v.as_bool()), + Some(true) + ); + assert!( + nearai.get("api_key").is_none(), + "raw api_key must never be returned" + ); + assert_eq!( + nearai.get("env_model").and_then(|v| v.as_str()), + Some("test-model") + ); + assert_eq!( + nearai.get("env_base_url").and_then(|v| v.as_str()), + Some("https://test.near.ai/v1") + ); + // Check definition fields are present + assert_eq!( + nearai.get("adapter").and_then(|v| v.as_str()), + Some("nearai") + ); + assert_eq!(nearai.get("builtin").and_then(|v| v.as_bool()), Some(true)); + + // Clean up + unsafe { + std::env::remove_var("NEARAI_API_KEY"); + std::env::remove_var("NEARAI_MODEL"); + std::env::remove_var("NEARAI_BASE_URL"); + } + } + + #[tokio::test] + async fn test_llm_providers_includes_registry_and_special_providers() { + let result = build_llm_providers(); + let arr = result.as_array().expect("should be an array"); + + // Registry providers should be present + assert!( + find_provider(arr, "openai").is_some(), + "should contain openai" + ); + assert!( + find_provider(arr, "anthropic").is_some(), + "should contain anthropic" + ); + assert!( + find_provider(arr, "ollama").is_some(), + "should contain ollama" + ); + + // Special providers should be present + assert!( + find_provider(arr, "nearai").is_some(), + "should contain nearai" + ); + assert!( + find_provider(arr, "bedrock").is_some(), + "should contain bedrock" + ); + + // Each entry should have required fields + for p in arr { + let id = p.get("id").and_then(|v| v.as_str()).unwrap_or(""); + assert!(p.get("name").is_some(), "{id} missing name"); + assert!(p.get("adapter").is_some(), "{id} missing adapter"); + assert!(p.get("builtin").is_some(), "{id} missing builtin"); + assert!( + p.get("default_model").is_some(), + "{id} missing default_model" + ); + } + } + + // --- is_nearai_private_endpoint tests --- + + #[test] + fn test_nearai_private_exact_match() { + assert!(is_nearai_private_endpoint("https://private.near.ai/v1")); + } + + #[test] + fn test_nearai_private_subdomain() { + assert!(is_nearai_private_endpoint("https://us.private.near.ai/v1")); + } + + #[test] + fn test_nearai_public_endpoint_not_private() { + assert!(!is_nearai_private_endpoint("https://cloud-api.near.ai/v1")); + } + + #[test] + fn test_nearai_private_lookalike_rejected() { + // "private" appears in the hostname but not as the correct domain + assert!(!is_nearai_private_endpoint( + "https://private-evil.near.ai/v1" + )); + assert!(!is_nearai_private_endpoint("https://myprivate.near.ai/v1")); + } + + #[test] + fn test_nearai_private_non_near_ai_rejected() { + assert!(!is_nearai_private_endpoint("https://private.evil.com/v1")); + } +} diff --git a/src/channels/web/handlers/mod.rs b/src/channels/web/handlers/mod.rs index b8958527b83..984bf61e721 100644 --- a/src/channels/web/handlers/mod.rs +++ b/src/channels/web/handlers/mod.rs @@ -3,6 +3,7 @@ //! Each module groups related endpoint handlers by domain. pub mod jobs; +pub mod llm; pub mod memory; pub mod routines; pub mod secrets; diff --git a/src/channels/web/handlers/settings.rs b/src/channels/web/handlers/settings.rs index 4dd7299ae59..7f4e365d13f 100644 --- a/src/channels/web/handlers/settings.rs +++ b/src/channels/web/handlers/settings.rs @@ -7,10 +7,15 @@ use axum::{ extract::{Path, State}, http::StatusCode, }; +use secrecy::SecretString; use crate::channels::web::auth::AuthenticatedUser; use crate::channels::web::server::GatewayState; use crate::channels::web::types::*; +use crate::secrets::{CreateSecretParams, SecretsStore}; + +/// Sentinel value the frontend sends to mean "key is unchanged, don't touch it". +const API_KEY_UNCHANGED: &str = "••••••••"; pub async fn settings_list_handler( State(state): State>, @@ -25,12 +30,34 @@ pub async fn settings_list_handler( StatusCode::INTERNAL_SERVER_ERROR })?; + // Build a map of sensitive keys so we can annotate and mask them. + let sensitive_keys = ["llm_builtin_overrides", "llm_custom_providers"]; + let mut sensitive_map: std::collections::HashMap = rows + .iter() + .filter(|r| sensitive_keys.contains(&r.key.as_str())) + .map(|r| (r.key.clone(), r.value.clone())) + .collect(); + if !sensitive_map.is_empty() { + annotate_secret_key_presence(&state, &user.user_id, &mut sensitive_map).await; + mask_settings_api_keys(&mut sensitive_map); + } + let settings = rows .into_iter() - .map(|r| SettingResponse { - key: r.key, - value: r.value, - updated_at: r.updated_at.to_rfc3339(), + .map(|r| { + let value = if sensitive_keys.contains(&r.key.as_str()) { + sensitive_map + .get(&r.key) + .cloned() + .unwrap_or(r.value.clone()) + } else { + r.value + }; + SettingResponse { + key: r.key, + value, + updated_at: r.updated_at.to_rfc3339(), + } }) .collect(); @@ -55,9 +82,22 @@ pub async fn settings_get_handler( })? .ok_or(StatusCode::NOT_FOUND)?; + // Mask any plaintext API keys that may exist from legacy data. + let value = if matches!( + key.as_str(), + "llm_builtin_overrides" | "llm_custom_providers" + ) { + let mut map = std::collections::HashMap::from([(key.clone(), row.value.clone())]); + annotate_secret_key_presence(&state, &user.user_id, &mut map).await; + mask_settings_api_keys(&mut map); + map.remove(&key).unwrap_or(row.value) + } else { + row.value + }; + Ok(Json(SettingResponse { key: row.key, - value: row.value, + value, updated_at: row.updated_at.to_rfc3339(), })) } @@ -72,8 +112,27 @@ pub async fn settings_set_handler( .store .as_ref() .ok_or(StatusCode::SERVICE_UNAVAILABLE)?; + + // Guard: cannot remove a custom provider that is currently active. + if key == "llm_custom_providers" { + guard_active_provider_not_removed(store, &user.user_id, &body.value).await?; + validate_custom_providers(&body.value)?; + } + + // Extract API keys from LLM settings and vault them in the secrets store. + // The sanitized value has api_key fields removed (stored encrypted instead). + let sanitized_value = match key.as_str() { + "llm_builtin_overrides" => { + extract_builtin_override_keys(&state, &user.user_id, &body.value).await? + } + "llm_custom_providers" => { + extract_custom_provider_keys(&state, &user.user_id, &body.value).await? + } + _ => body.value.clone(), + }; + store - .set_setting(&user.user_id, &key, &body.value) + .set_setting(&user.user_id, &key, &sanitized_value) .await .map_err(|e| { tracing::error!("Failed to set setting '{}': {}", key, e); @@ -83,6 +142,99 @@ pub async fn settings_set_handler( Ok(StatusCode::NO_CONTENT) } +const VALID_ADAPTERS: &[&str] = &["open_ai_completions", "anthropic", "ollama"]; + +/// Valid provider ID: lowercase alphanumeric, hyphens, and underscores, 1-64 chars. +fn is_valid_provider_id(id: &str) -> bool { + !id.is_empty() + && id.len() <= 64 + && id + .bytes() + .all(|b| b.is_ascii_lowercase() || b.is_ascii_digit() || b == b'-' || b == b'_') +} + +/// Returns `Err(422)` if any provider has an invalid ID or unrecognised adapter. +fn validate_custom_providers(value: &serde_json::Value) -> Result<(), StatusCode> { + let providers = match value.as_array() { + Some(arr) => arr, + None => return Ok(()), + }; + for p in providers { + let id = p.get("id").and_then(|v| v.as_str()).unwrap_or(""); + if !is_valid_provider_id(id) { + tracing::warn!( + id = %id, + "Rejected custom provider with invalid ID (must be lowercase alphanumeric/hyphens/underscores, 1-64 chars)" + ); + return Err(StatusCode::UNPROCESSABLE_ENTITY); + } + let adapter = p.get("adapter").and_then(|v| v.as_str()).unwrap_or(""); + if adapter.is_empty() { + tracing::warn!(id = %id, "Rejected custom provider with missing adapter field"); + return Err(StatusCode::UNPROCESSABLE_ENTITY); + } + if !VALID_ADAPTERS.contains(&adapter) { + tracing::warn!(id = %id, adapter = %adapter, "Rejected unknown LLM adapter"); + return Err(StatusCode::UNPROCESSABLE_ENTITY); + } + } + Ok(()) +} + +/// Returns `Err(409)` if the active `llm_backend` is a custom provider that +/// would be removed by the incoming update to `llm_custom_providers`. +async fn guard_active_provider_not_removed( + store: &Arc, + user_id: &str, + new_value: &serde_json::Value, +) -> Result<(), StatusCode> { + // Get the currently active backend. + let active_backend = match store.get_setting(user_id, "llm_backend").await { + Ok(Some(v)) => match v.as_str() { + Some(s) if !s.is_empty() => s.to_string(), + _ => return Ok(()), + }, + _ => return Ok(()), + }; + + // Parse the incoming provider list. + let new_providers = match new_value.as_array() { + Some(arr) => arr, + None => return Ok(()), + }; + + // Check whether the active backend exists in the OLD custom providers list. + let old_providers_value = match store.get_setting(user_id, "llm_custom_providers").await { + Ok(Some(v)) => v, + _ => return Ok(()), + }; + let old_providers = match old_providers_value.as_array() { + Some(arr) => arr, + None => return Ok(()), + }; + + let active_was_custom = old_providers + .iter() + .any(|p| p.get("id").and_then(|v| v.as_str()) == Some(active_backend.as_str())); + if !active_was_custom { + return Ok(()); + } + + // Reject if the active provider is absent from the new list. + let still_present = new_providers + .iter() + .any(|p| p.get("id").and_then(|v| v.as_str()) == Some(active_backend.as_str())); + if !still_present { + tracing::warn!( + active_backend = %active_backend, + "Rejected attempt to delete the active custom LLM provider" + ); + return Err(StatusCode::CONFLICT); + } + + Ok(()) +} + pub async fn settings_delete_handler( State(state): State>, AuthenticatedUser(user): AuthenticatedUser, @@ -92,6 +244,14 @@ pub async fn settings_delete_handler( .store .as_ref() .ok_or(StatusCode::SERVICE_UNAVAILABLE)?; + + // Guard: deleting llm_custom_providers is equivalent to setting it to []. + // Reject if the active backend is a custom provider that would be removed. + if key == "llm_custom_providers" { + guard_active_provider_not_removed(store, &user.user_id, &serde_json::Value::Array(vec![])) + .await?; + } + store .delete_setting(&user.user_id, &key) .await @@ -111,11 +271,16 @@ pub async fn settings_export_handler( .store .as_ref() .ok_or(StatusCode::SERVICE_UNAVAILABLE)?; - let settings = store.get_all_settings(&user.user_id).await.map_err(|e| { + let mut settings = store.get_all_settings(&user.user_id).await.map_err(|e| { tracing::error!("Failed to export settings: {}", e); StatusCode::INTERNAL_SERVER_ERROR })?; + // Indicate key presence from secrets store without exposing values. + annotate_secret_key_presence(&state, &user.user_id, &mut settings).await; + + mask_settings_api_keys(&mut settings); + Ok(Json(SettingsExportResponse { settings })) } @@ -128,8 +293,21 @@ pub async fn settings_import_handler( .store .as_ref() .ok_or(StatusCode::SERVICE_UNAVAILABLE)?; + + // Vault any API keys present in the imported settings, same as the + // individual SET handler does, so plaintext keys never reach the DB. + let mut sanitized = body.settings.clone(); + if let Some(v) = sanitized.get("llm_builtin_overrides").cloned() { + let clean = extract_builtin_override_keys(&state, &user.user_id, &v).await?; + sanitized.insert("llm_builtin_overrides".to_string(), clean); + } + if let Some(v) = sanitized.get("llm_custom_providers").cloned() { + let clean = extract_custom_provider_keys(&state, &user.user_id, &v).await?; + sanitized.insert("llm_custom_providers".to_string(), clean); + } + store - .set_all_settings(&user.user_id, &body.settings) + .set_all_settings(&user.user_id, &sanitized) .await .map_err(|e| { tracing::error!("Failed to import settings: {}", e); @@ -138,3 +316,635 @@ pub async fn settings_import_handler( Ok(StatusCode::NO_CONTENT) } + +// --------------------------------------------------------------------------- +// LLM API key vaulting helpers +// --------------------------------------------------------------------------- + +use crate::settings::{builtin_secret_name, custom_secret_name}; + +/// Returns true if the `api_key` value is a real key (not sentinel/empty). +fn is_real_api_key(key: &str) -> bool { + !key.is_empty() && key != API_KEY_UNCHANGED +} + +/// Require the secrets store when real API keys are present. +/// Returns `Ok(None)` when no secrets store and no real keys (passthrough). +fn require_secrets_store( + state: &GatewayState, + has_real_keys: bool, +) -> Result>, StatusCode> { + match state.secrets_store.as_ref() { + Some(s) => Ok(Some(s)), + None if has_real_keys => { + tracing::error!("Cannot store API keys: secrets store is not available"); + Err(StatusCode::SERVICE_UNAVAILABLE) + } + None => Ok(None), + } +} + +/// Extract API keys from builtin overrides, store in secrets, return sanitized JSON. +async fn extract_builtin_override_keys( + state: &GatewayState, + user_id: &str, + value: &serde_json::Value, +) -> Result { + let obj = match value.as_object() { + Some(o) => o, + None => return Ok(value.clone()), + }; + + let has_real_keys = obj.values().any(|v| { + v.get("api_key") + .and_then(|k| k.as_str()) + .is_some_and(is_real_api_key) + }); + let secrets = match require_secrets_store(state, has_real_keys)? { + Some(s) => s, + None => return Ok(value.clone()), + }; + + let mut sanitized = obj.clone(); + + for (provider_id, override_val) in obj { + if let Some(api_key) = override_val.get("api_key").and_then(|v| v.as_str()) { + if !is_real_api_key(api_key) { + // Unchanged or empty — remove from settings, keep existing secret. + if let Some(o) = sanitized + .get_mut(provider_id) + .and_then(|v| v.as_object_mut()) + { + o.remove("api_key"); + } + continue; + } + vault_secret( + secrets.as_ref(), + user_id, + &builtin_secret_name(provider_id), + api_key, + provider_id, + ) + .await?; + if let Some(o) = sanitized + .get_mut(provider_id) + .and_then(|v| v.as_object_mut()) + { + o.remove("api_key"); + } + } + } + + Ok(serde_json::Value::Object(sanitized)) +} + +/// Extract API keys from custom providers, store in secrets, return sanitized JSON. +async fn extract_custom_provider_keys( + state: &GatewayState, + user_id: &str, + value: &serde_json::Value, +) -> Result { + let arr = match value.as_array() { + Some(a) => a, + None => return Ok(value.clone()), + }; + + let has_real_keys = arr.iter().any(|v| { + v.get("api_key") + .and_then(|k| k.as_str()) + .is_some_and(is_real_api_key) + }); + let secrets = match require_secrets_store(state, has_real_keys)? { + Some(s) => s, + None => return Ok(value.clone()), + }; + + let mut sanitized = arr.clone(); + + for (idx, provider_val) in arr.iter().enumerate() { + let provider_id = provider_val + .get("id") + .and_then(|v| v.as_str()) + .unwrap_or(""); + if provider_id.is_empty() { + continue; + } + + if let Some(api_key) = provider_val.get("api_key").and_then(|v| v.as_str()) { + if !is_real_api_key(api_key) { + if let Some(o) = sanitized[idx].as_object_mut() { + o.remove("api_key"); + } + continue; + } + vault_secret( + secrets.as_ref(), + user_id, + &custom_secret_name(provider_id), + api_key, + provider_id, + ) + .await?; + if let Some(o) = sanitized[idx].as_object_mut() { + o.remove("api_key"); + } + } + } + + Ok(serde_json::Value::Array(sanitized)) +} + +/// Encrypt and store an API key in the secrets store. +async fn vault_secret( + secrets: &(dyn SecretsStore + Send + Sync), + user_id: &str, + secret_name: &str, + api_key: &str, + provider_id: &str, +) -> Result<(), StatusCode> { + secrets + .create( + user_id, + CreateSecretParams { + name: secret_name.to_string(), + value: SecretString::from(api_key.to_string()), + provider: Some(provider_id.to_string()), + expires_at: None, + }, + ) + .await + .map_err(|e| { + tracing::error!( + "Failed to store secret '{}' for provider '{}': {}", + secret_name, + provider_id, + e + ); + StatusCode::INTERNAL_SERVER_ERROR + })?; + Ok(()) +} + +/// Mask plaintext API keys in settings values before returning to the frontend. +/// +/// Any `api_key` field still present in the settings JSON (legacy plaintext) +/// is replaced with the sentinel so the frontend shows "key configured". +fn mask_settings_api_keys(settings: &mut std::collections::HashMap) { + if let Some(obj) = settings + .get_mut("llm_builtin_overrides") + .and_then(|v| v.as_object_mut()) + { + for override_val in obj.values_mut() { + if let Some(o) = override_val.as_object_mut() + && o.contains_key("api_key") + { + o.insert( + "api_key".to_string(), + serde_json::Value::String(API_KEY_UNCHANGED.to_string()), + ); + } + } + } + + if let Some(arr) = settings + .get_mut("llm_custom_providers") + .and_then(|v| v.as_array_mut()) + { + for provider_val in arr.iter_mut() { + if let Some(o) = provider_val.as_object_mut() + && o.contains_key("api_key") + { + o.insert( + "api_key".to_string(), + serde_json::Value::String(API_KEY_UNCHANGED.to_string()), + ); + } + } + } +} + +/// Check the secrets store for vaulted API keys and annotate the settings map. +/// +/// For builtin overrides and custom providers whose API key was stripped from +/// settings (stored in secrets), this adds `api_key: "••••••••"` so the +/// frontend knows a key is configured without seeing the actual value. +async fn annotate_secret_key_presence( + state: &GatewayState, + user_id: &str, + settings: &mut std::collections::HashMap, +) { + let secrets = match state.secrets_store.as_ref() { + Some(s) => s, + None => return, + }; + + // Annotate builtin overrides + if let Some(obj) = settings + .get_mut("llm_builtin_overrides") + .and_then(|v| v.as_object_mut()) + { + let provider_ids: Vec = obj.keys().cloned().collect(); + for provider_id in provider_ids { + let has_key_in_settings = obj + .get(&provider_id) + .and_then(|v| v.get("api_key")) + .is_some(); + if has_key_in_settings { + continue; // Will be masked by mask_settings_api_keys + } + let secret_name = builtin_secret_name(&provider_id); + if secrets.exists(user_id, &secret_name).await.unwrap_or(false) + && let Some(o) = obj.get_mut(&provider_id).and_then(|v| v.as_object_mut()) + { + o.insert( + "api_key".to_string(), + serde_json::Value::String(API_KEY_UNCHANGED.to_string()), + ); + } + } + } + + // Annotate custom providers + if let Some(arr) = settings + .get_mut("llm_custom_providers") + .and_then(|v| v.as_array_mut()) + { + for provider_val in arr.iter_mut() { + let provider_id = provider_val + .get("id") + .and_then(|v| v.as_str()) + .unwrap_or("") + .to_string(); + if provider_id.is_empty() { + continue; + } + let has_key_in_settings = provider_val.get("api_key").is_some(); + if has_key_in_settings { + continue; + } + let secret_name = custom_secret_name(&provider_id); + if secrets.exists(user_id, &secret_name).await.unwrap_or(false) + && let Some(o) = provider_val.as_object_mut() + { + o.insert( + "api_key".to_string(), + serde_json::Value::String(API_KEY_UNCHANGED.to_string()), + ); + } + } + } +} + +#[cfg(test)] +mod tests { + use super::*; + use std::collections::HashMap; + + #[test] + fn test_mask_settings_api_keys_builtin_overrides() { + let mut settings = HashMap::new(); + settings.insert( + "llm_builtin_overrides".to_string(), + serde_json::json!({ + "openai": { "api_key": "sk-secret-123", "model": "gpt-4" }, + "anthropic": { "model": "claude-3" } + }), + ); + + mask_settings_api_keys(&mut settings); + + let overrides = settings["llm_builtin_overrides"].as_object().unwrap(); + assert_eq!( + overrides["openai"]["api_key"].as_str().unwrap(), + API_KEY_UNCHANGED, + ); + assert_eq!(overrides["openai"]["model"].as_str().unwrap(), "gpt-4"); + assert!(overrides["anthropic"].get("api_key").is_none()); + } + + #[test] + fn test_mask_settings_api_keys_custom_providers() { + let mut settings = HashMap::new(); + settings.insert( + "llm_custom_providers".to_string(), + serde_json::json!([ + { "id": "my-llm", "api_key": "secret-key", "adapter": "open_ai_completions" }, + { "id": "no-key", "adapter": "ollama" } + ]), + ); + + mask_settings_api_keys(&mut settings); + + let providers = settings["llm_custom_providers"].as_array().unwrap(); + assert_eq!(providers[0]["api_key"].as_str().unwrap(), API_KEY_UNCHANGED,); + assert!(providers[1].get("api_key").is_none()); + } + + #[test] + fn test_mask_settings_no_llm_keys_is_noop() { + let mut settings = HashMap::new(); + settings.insert("some_other_setting".to_string(), serde_json::json!("value")); + + mask_settings_api_keys(&mut settings); + + assert_eq!(settings["some_other_setting"].as_str().unwrap(), "value"); + } + + #[test] + fn test_builtin_secret_name_format() { + assert_eq!(builtin_secret_name("openai"), "llm_builtin_openai_api_key"); + } + + #[test] + fn test_custom_secret_name_format() { + assert_eq!(custom_secret_name("my-groq"), "llm_custom_my-groq_api_key"); + } + + fn test_secrets_store() -> Arc { + let crypto = Arc::new( + crate::secrets::SecretsCrypto::new(secrecy::SecretString::from( + crate::secrets::keychain::generate_master_key_hex(), + )) + .unwrap(), + ); + Arc::new(crate::secrets::InMemorySecretsStore::new(crypto)) + } + + fn test_gateway_state(secrets: Arc) -> GatewayState { + GatewayState { + msg_tx: tokio::sync::RwLock::new(None), + sse: Arc::new(crate::channels::web::sse::SseManager::new()), + workspace: None, + workspace_pool: None, + session_manager: None, + log_broadcaster: None, + log_level_handle: None, + extension_manager: None, + tool_registry: None, + store: None, + job_manager: None, + prompt_queue: None, + scheduler: None, + owner_id: "test".to_string(), + shutdown_tx: tokio::sync::RwLock::new(None), + ws_tracker: None, + llm_provider: None, + skill_registry: None, + skill_catalog: None, + chat_rate_limiter: crate::channels::web::server::PerUserRateLimiter::new(30, 60), + oauth_rate_limiter: crate::channels::web::server::RateLimiter::new(10, 60), + webhook_rate_limiter: crate::channels::web::server::RateLimiter::new(10, 60), + registry_entries: Vec::new(), + cost_guard: None, + routine_engine: Arc::new(tokio::sync::RwLock::new(None)), + startup_time: std::time::Instant::now(), + active_config: crate::channels::web::server::ActiveConfigSnapshot::default(), + secrets_store: Some(secrets), + db_auth: None, + } + } + + #[tokio::test] + async fn test_extract_builtin_keys_vaults_and_strips() { + let secrets = test_secrets_store(); + let state = test_gateway_state(Arc::clone(&secrets)); + + let input = serde_json::json!({ + "openai": { "api_key": "sk-test-key", "model": "gpt-4" }, + "anthropic": { "model": "claude-3" } + }); + + let result = extract_builtin_override_keys(&state, "test", &input) + .await + .unwrap(); + + let obj = result.as_object().unwrap(); + assert!( + obj["openai"].get("api_key").is_none(), + "api_key should be stripped" + ); + assert_eq!(obj["openai"]["model"].as_str().unwrap(), "gpt-4"); + assert_eq!(obj["anthropic"]["model"].as_str().unwrap(), "claude-3"); + + let decrypted = secrets + .get_decrypted("test", "llm_builtin_openai_api_key") + .await + .unwrap(); + assert_eq!(decrypted.expose(), "sk-test-key"); + } + + #[tokio::test] + async fn test_extract_custom_keys_vaults_and_strips() { + let secrets = test_secrets_store(); + let state = test_gateway_state(Arc::clone(&secrets)); + + let input = serde_json::json!([ + { "id": "my-llm", "api_key": "gsk-custom-key", "adapter": "open_ai_completions" }, + { "id": "local", "adapter": "ollama" } + ]); + + let result = extract_custom_provider_keys(&state, "test", &input) + .await + .unwrap(); + + let arr = result.as_array().unwrap(); + assert!( + arr[0].get("api_key").is_none(), + "api_key should be stripped" + ); + assert_eq!(arr[0]["id"].as_str().unwrap(), "my-llm"); + assert!(arr[1].get("api_key").is_none()); + + let decrypted = secrets + .get_decrypted("test", "llm_custom_my-llm_api_key") + .await + .unwrap(); + assert_eq!(decrypted.expose(), "gsk-custom-key"); + } + + #[tokio::test] + async fn test_unchanged_sentinel_preserves_existing_secret() { + let secrets = test_secrets_store(); + + secrets + .create( + "test", + CreateSecretParams { + name: "llm_builtin_openai_api_key".to_string(), + value: SecretString::from("sk-original".to_string()), + provider: Some("openai".to_string()), + expires_at: None, + }, + ) + .await + .unwrap(); + + let state = test_gateway_state(Arc::clone(&secrets)); + + let input = serde_json::json!({ + "openai": { "api_key": "••••••••", "model": "gpt-4" } + }); + + let result = extract_builtin_override_keys(&state, "test", &input) + .await + .unwrap(); + + assert!(result["openai"].get("api_key").is_none()); + + let decrypted = secrets + .get_decrypted("test", "llm_builtin_openai_api_key") + .await + .unwrap(); + assert_eq!(decrypted.expose(), "sk-original"); + } + + /// When secrets store is unavailable, attempting to save a real API key + /// must fail with 503 rather than silently storing plaintext. + #[tokio::test] + async fn test_extract_builtin_keys_rejects_without_secrets_store() { + let state = GatewayState { + secrets_store: None, + ..test_gateway_state(test_secrets_store()) + }; + + let input = serde_json::json!({ + "openai": { "api_key": "sk-real-key", "model": "gpt-4" } + }); + + let err = extract_builtin_override_keys(&state, "test", &input) + .await + .unwrap_err(); + assert_eq!(err, StatusCode::SERVICE_UNAVAILABLE); + } + + /// When secrets store is unavailable but no real keys are present + /// (only sentinels or no api_key at all), the call should succeed. + #[tokio::test] + async fn test_extract_builtin_keys_allows_no_keys_without_secrets_store() { + let state = GatewayState { + secrets_store: None, + ..test_gateway_state(test_secrets_store()) + }; + + let input = serde_json::json!({ + "openai": { "api_key": "••••••••", "model": "gpt-4" }, + "anthropic": { "model": "claude-3" } + }); + + let result = extract_builtin_override_keys(&state, "test", &input) + .await + .unwrap(); + // Without secrets store, the value passes through unchanged (no vaulting needed). + assert!(result.as_object().is_some()); + } + + #[tokio::test] + async fn test_extract_custom_keys_rejects_without_secrets_store() { + let state = GatewayState { + secrets_store: None, + ..test_gateway_state(test_secrets_store()) + }; + + let input = serde_json::json!([ + { "id": "my-llm", "api_key": "gsk-real-key", "adapter": "open_ai_completions" } + ]); + + let err = extract_custom_provider_keys(&state, "test", &input) + .await + .unwrap_err(); + assert_eq!(err, StatusCode::SERVICE_UNAVAILABLE); + } + + // --- Provider ID validation tests --- + + #[test] + fn test_valid_provider_ids() { + assert!(is_valid_provider_id("my-llm")); + assert!(is_valid_provider_id("openai")); + assert!(is_valid_provider_id("custom-provider-123")); + assert!(is_valid_provider_id("a")); + assert!(is_valid_provider_id("my_llm"), "underscores allowed"); + assert!( + is_valid_provider_id("openai_compatible"), + "matches builtin naming" + ); + } + + #[test] + fn test_invalid_provider_ids() { + assert!(!is_valid_provider_id(""), "empty ID"); + assert!(!is_valid_provider_id("My-LLM"), "uppercase"); + assert!(!is_valid_provider_id("my llm"), "spaces"); + assert!(!is_valid_provider_id("../../etc"), "path traversal"); + assert!(!is_valid_provider_id("a.b"), "dots"); + assert!( + !is_valid_provider_id(&"a".repeat(65)), + "exceeds 64 char limit" + ); + } + + #[test] + fn test_validate_custom_providers_rejects_bad_id() { + let input = serde_json::json!([ + { "id": "UPPER-CASE", "adapter": "open_ai_completions" } + ]); + assert_eq!( + validate_custom_providers(&input).unwrap_err(), + StatusCode::UNPROCESSABLE_ENTITY, + ); + } + + #[test] + fn test_validate_custom_providers_accepts_valid() { + let input = serde_json::json!([ + { "id": "my-llm", "adapter": "open_ai_completions" }, + { "id": "local-ollama", "adapter": "ollama" } + ]); + assert!(validate_custom_providers(&input).is_ok()); + } + + // --- Adapter validation tests --- + + #[test] + fn test_validate_custom_providers_rejects_unknown_adapter() { + let input = serde_json::json!([ + { "id": "test", "adapter": "not_a_real_adapter" } + ]); + assert_eq!( + validate_custom_providers(&input).unwrap_err(), + StatusCode::UNPROCESSABLE_ENTITY, + ); + } + + #[test] + fn test_validate_custom_providers_rejects_missing_adapter() { + let input = serde_json::json!([ + { "id": "test" } + ]); + assert_eq!( + validate_custom_providers(&input).unwrap_err(), + StatusCode::UNPROCESSABLE_ENTITY, + ); + } + + #[test] + fn test_validate_custom_providers_accepts_all_valid_adapters() { + for adapter in VALID_ADAPTERS { + let input = serde_json::json!([ + { "id": "test", "adapter": adapter } + ]); + assert!( + validate_custom_providers(&input).is_ok(), + "adapter '{}' should be accepted", + adapter + ); + } + } + + #[test] + fn test_validate_custom_providers_non_array_is_ok() { + let input = serde_json::json!("not-an-array"); + assert!(validate_custom_providers(&input).is_ok()); + } +} diff --git a/src/channels/web/server.rs b/src/channels/web/server.rs index d403e93c399..c0df6e78f41 100644 --- a/src/channels/web/server.rs +++ b/src/channels/web/server.rs @@ -38,6 +38,9 @@ use crate::channels::web::handlers::jobs::{ jobs_events_handler, jobs_list_handler, jobs_prompt_handler, jobs_restart_handler, jobs_summary_handler, }; +use crate::channels::web::handlers::llm::{ + llm_list_models_handler, llm_providers_handler, llm_test_connection_handler, +}; use crate::channels::web::handlers::memory::{ memory_list_handler, memory_read_handler, memory_search_handler, memory_tree_handler, memory_write_handler, @@ -46,6 +49,10 @@ use crate::channels::web::handlers::routines::{ routines_delete_handler, routines_detail_handler, routines_list_handler, routines_summary_handler, routines_toggle_handler, routines_trigger_handler, }; +use crate::channels::web::handlers::settings::{ + settings_delete_handler, settings_export_handler, settings_get_handler, + settings_import_handler, settings_list_handler, settings_set_handler, +}; use crate::channels::web::handlers::skills::{ skills_install_handler, skills_list_handler, skills_remove_handler, skills_search_handler, }; @@ -529,6 +536,13 @@ pub async fn start_server( "/api/settings/{key}", axum::routing::delete(settings_delete_handler), ) + // LLM utilities + .route( + "/api/llm/test_connection", + post(llm_test_connection_handler), + ) + .route("/api/llm/list_models", post(llm_list_models_handler)) + .route("/api/llm/providers", get(llm_providers_handler)) // User management (admin) .route( "/api/admin/users", @@ -2724,135 +2738,6 @@ async fn routines_runs_handler( }))) } -// --- Settings handlers --- - -async fn settings_list_handler( - State(state): State>, - AuthenticatedUser(user): AuthenticatedUser, -) -> Result, StatusCode> { - let store = state - .store - .as_ref() - .ok_or(StatusCode::SERVICE_UNAVAILABLE)?; - let rows = store.list_settings(&user.user_id).await.map_err(|e| { - tracing::error!("Failed to list settings: {}", e); - StatusCode::INTERNAL_SERVER_ERROR - })?; - - let settings = rows - .into_iter() - .map(|r| SettingResponse { - key: r.key, - value: r.value, - updated_at: r.updated_at.to_rfc3339(), - }) - .collect(); - - Ok(Json(SettingsListResponse { settings })) -} - -async fn settings_get_handler( - State(state): State>, - AuthenticatedUser(user): AuthenticatedUser, - Path(key): Path, -) -> Result, StatusCode> { - let store = state - .store - .as_ref() - .ok_or(StatusCode::SERVICE_UNAVAILABLE)?; - let row = store - .get_setting_full(&user.user_id, &key) - .await - .map_err(|e| { - tracing::error!("Failed to get setting '{}': {}", key, e); - StatusCode::INTERNAL_SERVER_ERROR - })? - .ok_or(StatusCode::NOT_FOUND)?; - - Ok(Json(SettingResponse { - key: row.key, - value: row.value, - updated_at: row.updated_at.to_rfc3339(), - })) -} - -async fn settings_set_handler( - State(state): State>, - AuthenticatedUser(user): AuthenticatedUser, - Path(key): Path, - Json(body): Json, -) -> Result { - let store = state - .store - .as_ref() - .ok_or(StatusCode::SERVICE_UNAVAILABLE)?; - store - .set_setting(&user.user_id, &key, &body.value) - .await - .map_err(|e| { - tracing::error!("Failed to set setting '{}': {}", key, e); - StatusCode::INTERNAL_SERVER_ERROR - })?; - - Ok(StatusCode::NO_CONTENT) -} - -async fn settings_delete_handler( - State(state): State>, - AuthenticatedUser(user): AuthenticatedUser, - Path(key): Path, -) -> Result { - let store = state - .store - .as_ref() - .ok_or(StatusCode::SERVICE_UNAVAILABLE)?; - store - .delete_setting(&user.user_id, &key) - .await - .map_err(|e| { - tracing::error!("Failed to delete setting '{}': {}", key, e); - StatusCode::INTERNAL_SERVER_ERROR - })?; - - Ok(StatusCode::NO_CONTENT) -} - -async fn settings_export_handler( - State(state): State>, - AuthenticatedUser(user): AuthenticatedUser, -) -> Result, StatusCode> { - let store = state - .store - .as_ref() - .ok_or(StatusCode::SERVICE_UNAVAILABLE)?; - let settings = store.get_all_settings(&user.user_id).await.map_err(|e| { - tracing::error!("Failed to export settings: {}", e); - StatusCode::INTERNAL_SERVER_ERROR - })?; - - Ok(Json(SettingsExportResponse { settings })) -} - -async fn settings_import_handler( - State(state): State>, - AuthenticatedUser(user): AuthenticatedUser, - Json(body): Json, -) -> Result { - let store = state - .store - .as_ref() - .ok_or(StatusCode::SERVICE_UNAVAILABLE)?; - store - .set_all_settings(&user.user_id, &body.settings) - .await - .map_err(|e| { - tracing::error!("Failed to import settings: {}", e); - StatusCode::INTERNAL_SERVER_ERROR - })?; - - Ok(StatusCode::NO_CONTENT) -} - // --- Gateway control plane handlers --- async fn gateway_status_handler( diff --git a/src/channels/web/static/app.js b/src/channels/web/static/app.js index c0c15acff9a..76084168647 100644 --- a/src/channels/web/static/app.js +++ b/src/channels/web/static/app.js @@ -5208,25 +5208,6 @@ function loadSettingsSubtab(subtab) { // --- Structured Settings Definitions --- var INFERENCE_SETTINGS = [ - { - group: 'cfg.group.llm', - settings: [ - { key: 'llm_backend', label: 'cfg.llm_backend.label', description: 'cfg.llm_backend.desc', - type: 'select', options: ['nearai', 'anthropic', 'openai', 'ollama', 'openai_compatible', 'tinfoil', 'bedrock'] }, - { key: 'selected_model', label: 'cfg.selected_model.label', description: 'cfg.selected_model.desc', type: 'text' }, - { key: 'ollama_base_url', label: 'cfg.ollama_base_url.label', description: 'cfg.ollama_base_url.desc', type: 'text', - showWhen: { key: 'llm_backend', value: 'ollama' } }, - { key: 'openai_compatible_base_url', label: 'cfg.openai_compatible_base_url.label', description: 'cfg.openai_compatible_base_url.desc', type: 'text', - showWhen: { key: 'llm_backend', value: 'openai_compatible' } }, - { key: 'bedrock_region', label: 'cfg.bedrock_region.label', description: 'cfg.bedrock_region.desc', type: 'text', - showWhen: { key: 'llm_backend', value: 'bedrock' } }, - { key: 'bedrock_cross_region', label: 'cfg.bedrock_cross_region.label', description: 'cfg.bedrock_cross_region.desc', - type: 'select', options: ['us', 'eu', 'apac', 'global'], - showWhen: { key: 'llm_backend', value: 'bedrock' } }, - { key: 'bedrock_profile', label: 'cfg.bedrock_profile.label', description: 'cfg.bedrock_profile.desc', type: 'text', - showWhen: { key: 'llm_backend', value: 'bedrock' } }, - ] - }, { group: 'cfg.group.embeddings', settings: [ @@ -5349,31 +5330,64 @@ function loadInferenceSettings() { Promise.all([ apiFetch('/api/settings/export'), apiFetch('/api/gateway/status').catch(function() { return {}; }), - apiFetch('/v1/models').catch(function() { return { data: [] }; }) ]).then(function(results) { var settings = results[0].settings || {}; var status = results[1]; - var modelsData = results[2]; - var activeValues = { - 'llm_backend': status.llm_backend, - 'selected_model': status.llm_model - }; - // Inject available model IDs as suggestions for the selected_model field - var modelIds = (modelsData.data || []).map(function(m) { return m.id; }).filter(Boolean); - if (modelIds.length > 0) { - var llmGroup = INFERENCE_SETTINGS[0]; - for (var i = 0; i < llmGroup.settings.length; i++) { - if (llmGroup.settings[i].key === 'selected_model') { - llmGroup.settings[i].suggestions = modelIds; - break; - } - } - } container.innerHTML = ''; - renderStructuredSettingsInto(container, INFERENCE_SETTINGS, settings, activeValues); + + // LLM Provider display — derived from active Model Provider + var activeBackend = settings['llm_backend'] || status.llm_backend || 'nearai'; + var activeModel = settings['selected_model'] || status.llm_model || ''; + var allP = _builtinProviders; + var customP = []; + try { + var cpVal = settings['llm_custom_providers']; + customP = Array.isArray(cpVal) ? cpVal : (cpVal ? JSON.parse(cpVal) : []); + } catch (e) { customP = []; } + var provider = allP.concat(customP).find(function(p) { return p.id === activeBackend; }); + var providerName = provider ? (provider.name || provider.id) : activeBackend; + if (!activeModel && provider) activeModel = provider.default_model || ''; + + var group = document.createElement('div'); + group.className = 'settings-group'; + var title = document.createElement('div'); + title.className = 'settings-group-title'; + title.textContent = I18n.t('cfg.group.llm'); + group.appendChild(title); + + var notice = document.createElement('div'); + notice.className = 'config-notice'; + notice.id = 'llm-restart-notice'; + var restartNoticeEl = document.getElementById('config-restart-notice'); + notice.style.display = (restartNoticeEl && restartNoticeEl.style.display !== 'none') ? 'flex' : 'none'; + notice.innerHTML = '\u26A0' + escapeHtml(I18n.t('config.restartNotice')) + ''; + group.appendChild(notice); + + var backendRow = document.createElement('div'); + backendRow.className = 'settings-row'; + backendRow.innerHTML = + '
' + + '
' + escapeHtml(I18n.t('cfg.llm_backend.desc')) + '
' + + '
' + escapeHtml(providerName) + '
'; + group.appendChild(backendRow); + + var modelRow = document.createElement('div'); + modelRow.className = 'settings-row'; + modelRow.innerHTML = + '
' + + '
' + escapeHtml(I18n.t('cfg.selected_model.desc')) + '
' + + '
' + escapeHtml(activeModel || '\u2014') + '
'; + group.appendChild(modelRow); + + container.appendChild(group); + + // Remaining editable settings (embeddings, etc.) + renderStructuredSettingsInto(container, INFERENCE_SETTINGS, settings, {}); + loadConfig(); }).catch(function(err) { container.innerHTML = '
' + I18n.t('common.loadFailed') + ': ' + escapeHtml(err.message) + '
'; + loadConfig(); }); } @@ -5617,8 +5631,7 @@ function renderStructuredSettingsRow(def, value, activeValue) { return row; } -var RESTART_REQUIRED_KEYS = ['llm_backend', 'selected_model', 'ollama_base_url', 'openai_compatible_base_url', - 'bedrock_region', 'bedrock_cross_region', 'bedrock_profile', 'embeddings.enabled', 'embeddings.provider', 'embeddings.model', +var RESTART_REQUIRED_KEYS = ['embeddings.enabled', 'embeddings.provider', 'embeddings.model', 'agent.auto_approve_tools', 'tunnel.provider', 'tunnel.public_url', 'gateway.rate_limit', 'gateway.max_connections']; var _settingsSavedTimers = {}; @@ -6208,6 +6221,18 @@ document.addEventListener('click', function(e) { case 'switch-language': if (typeof switchLanguage === 'function') switchLanguage(el.dataset.lang); break; + case 'set-active-provider': + setActiveProvider(el.dataset.id); + break; + case 'delete-custom-provider': + deleteCustomProvider(el.dataset.id); + break; + case 'edit-custom-provider': + editCustomProvider(el.dataset.id); + break; + case 'configure-builtin-provider': + configureBuiltinProvider(el.dataset.id); + break; } }); @@ -6249,6 +6274,9 @@ document.addEventListener('keydown', function(e) { if (e.key === 'Escape' && document.getElementById('confirm-modal').style.display === 'flex') { closeConfirmModal(); } + if (e.key === 'Escape' && document.getElementById('provider-dialog').style.display === 'flex') { + resetProviderForm(); + } }); // --- Settings Import/Export --- @@ -6336,3 +6364,580 @@ document.getElementById('settings-search-input').addEventListener('input', funct activePanel.appendChild(empty); } }); + + +// --- Config Tab --- + +// Like apiFetch but for endpoints that return 204 No Content +// Like apiFetch but discards the response body (for 204 No Content endpoints). +function apiFetchVoid(path, options) { + return apiFetch(path, options).then(function() {}); +} + +/** Sentinel value meaning "key is unchanged, don't touch it". Must match backend. */ +const API_KEY_UNCHANGED = '\u2022\u2022\u2022\u2022\u2022\u2022\u2022\u2022'; + +const ADAPTER_LABELS = { + open_ai_completions: 'OpenAI Compatible', + anthropic: 'Anthropic', + ollama: 'Ollama', + bedrock: 'AWS Bedrock', + nearai: 'NEAR AI', +}; + +let _builtinProviders = []; +let _customProviders = []; +let _activeLlmBackend = ''; +let _selectedModel = ''; +let _builtinOverrides = {}; +let _editingProviderId = null; +let _configuringBuiltinId = null; +let _configLoaded = false; + +function loadConfig() { + const list = document.getElementById('providers-list'); + list.innerHTML = '
' + I18n.t('common.loading') + '
'; + + Promise.all([ + apiFetch('/api/settings/export'), + apiFetch('/api/llm/providers').catch(function() { return []; }), + ]).then(function(results) { + const s = (results[0] && results[0].settings) ? results[0].settings : {}; + _builtinProviders = Array.isArray(results[1]) ? results[1] : []; + _activeLlmBackend = s['llm_backend'] ? String(s['llm_backend']) : 'nearai'; + _selectedModel = s['selected_model'] ? String(s['selected_model']) : ''; + try { + const val = s['llm_custom_providers']; + _customProviders = Array.isArray(val) ? val : (val ? JSON.parse(val) : []); + } catch (e) { + _customProviders = []; + } + try { + const val = s['llm_builtin_overrides']; + _builtinOverrides = (val && typeof val === 'object' && !Array.isArray(val)) ? val : {}; + } catch (e) { + _builtinOverrides = {}; + } + _configLoaded = true; + renderProviders(); + }).catch(function() { + _activeLlmBackend = 'nearai'; + _selectedModel = ''; + _builtinProviders = []; + _customProviders = []; + _builtinOverrides = {}; + _configLoaded = true; + renderProviders(); + }); +} + +function scrollToProviders() { + const section = document.getElementById('providers-section'); + if (section) section.scrollIntoView({ behavior: 'smooth', block: 'start' }); +} + +function renderProviders() { + const list = document.getElementById('providers-list'); + const allProviders = [..._builtinProviders, ..._customProviders].sort((a, b) => { + if (a.id === _activeLlmBackend) return -1; + if (b.id === _activeLlmBackend) return 1; + return 0; + }); + + if (allProviders.length === 0) { + list.innerHTML = '
No providers
'; + return; + } + + list.innerHTML = allProviders.map((p) => { + const isActive = p.id === _activeLlmBackend; + const adapterLabel = ADAPTER_LABELS[p.adapter] || p.adapter; + const activeBadge = isActive + ? '' + I18n.t('status.active') + '' + : ''; + const builtinBadge = p.builtin + ? '' + I18n.t('config.builtin') + '' + : ''; + const deleteBtn = !p.builtin && !isActive + ? '' + : ''; + const editBtn = !p.builtin + ? '' + : ''; + // Show Configure for built-in providers that support it (not bedrock — uses AWS credential chain) + const configureBtn = p.builtin && p.id !== 'bedrock' + ? '' + : ''; + const useBtn = !isActive + ? '' + : ''; + const overrideBaseUrl = p.builtin && _builtinOverrides[p.id] ? (_builtinOverrides[p.id].base_url || '') : ''; + const effectiveBaseUrl = overrideBaseUrl || p.env_base_url || p.base_url; + const baseUrlText = effectiveBaseUrl + ? '' + escapeHtml(effectiveBaseUrl) + '' + : ''; + // Show configured model: for active provider use _selectedModel, for others check _builtinOverrides then env defaults + const overrideModel = p.builtin && _builtinOverrides[p.id] ? (_builtinOverrides[p.id].model || '') : ''; + const displayModel = isActive + ? (_selectedModel || p.env_model || '') + : (overrideModel || p.env_model || ''); + const modelText = displayModel + ? '' + escapeHtml(I18n.t('config.currentModel', { model: displayModel })) + '' + : ''; + + return '
' + + '
' + + '' + escapeHtml(p.name || p.id) + '' + + '' + escapeHtml(p.id) + '' + + activeBadge + builtinBadge + + '
' + + '
' + + '' + escapeHtml(adapterLabel) + '' + + baseUrlText + + modelText + + '
' + + '
' + + useBtn + configureBtn + editBtn + deleteBtn + + '
' + + '
'; + }).join(''); +} + +function setActiveProvider(id) { + const provider = [..._builtinProviders, ..._customProviders].find((p) => p.id === id); + // Restore the last-configured model for this provider, falling back to the provider's default + const restoredModel = + (_builtinOverrides[id] && _builtinOverrides[id].model) || + (provider && provider.default_model) || + null; + const defaultModel = restoredModel; + const modelUpdate = () => defaultModel + ? apiFetchVoid('/api/settings/selected_model', { method: 'PUT', body: { value: defaultModel } }) + : apiFetchVoid('/api/settings/selected_model', { method: 'DELETE' }); + apiFetchVoid('/api/settings/llm_backend', { method: 'PUT', body: { value: id } }) + .then(() => modelUpdate()) + .then(() => { + _activeLlmBackend = id; + _selectedModel = defaultModel || ''; + renderProviders(); + loadInferenceSettings(); + scrollToProviders(); + document.getElementById('config-restart-notice').style.display = 'flex'; + var llmNotice = document.getElementById('llm-restart-notice'); + if (llmNotice) llmNotice.style.display = 'flex'; + showToast(I18n.t('config.providerActivated', { name: id })); + }) + .catch((e) => showToast(I18n.t('error.unknown') + ': ' + e.message, 'error')); +} + +function deleteCustomProvider(id) { + if (id === _activeLlmBackend) { + showToast(I18n.t('config.cannotDeleteActiveProvider'), 'error'); + return; + } + if (!confirm(I18n.t('config.confirmDeleteProvider', { id }))) return; + const originalProviders = _customProviders; + _customProviders = _customProviders.filter((p) => p.id !== id); + saveCustomProviders().then(() => { + renderProviders(); + showToast(I18n.t('config.providerDeleted')); + }).catch((e) => { + _customProviders = originalProviders; + showToast(I18n.t('error.unknown') + ': ' + e.message, 'error'); + }); +} + +function saveCustomProviders() { + return apiFetchVoid('/api/settings/llm_custom_providers', { method: 'PUT', body: { value: _customProviders } }); +} + +function editCustomProvider(id) { + const p = _customProviders.find((p) => p.id === id); + if (!p) return; + _editingProviderId = id; + const titleEl = document.getElementById('provider-form-title'); + titleEl.textContent = I18n.t('config.editProvider'); + titleEl.removeAttribute('data-i18n'); + document.getElementById('provider-name').value = p.name || ''; + const idField = document.getElementById('provider-id'); + idField.value = p.id; + idField.readOnly = true; + idField.style.opacity = '0.6'; + document.getElementById('provider-adapter').value = p.adapter || 'open_ai_completions'; + document.getElementById('provider-base-url').value = p.base_url || ''; + const editApiKeyInput = document.getElementById('provider-api-key'); + if (p.api_key === API_KEY_UNCHANGED) { + editApiKeyInput.value = ''; + editApiKeyInput.placeholder = I18n.t('config.apiKeyConfigured'); + } else { + editApiKeyInput.value = ''; + editApiKeyInput.placeholder = I18n.t('config.apiKeyEnter'); + } + document.getElementById('provider-model').value = p.default_model || ''; + openProviderDialog(true); + document.getElementById('provider-name').focus(); +} + +function configureBuiltinProvider(id) { + const p = _builtinProviders.find((p) => p.id === id); + if (!p) return; + _configuringBuiltinId = id; + const titleEl = document.getElementById('provider-form-title'); + titleEl.textContent = I18n.t('config.configureProvider') + ': ' + (p.name || id); + titleEl.removeAttribute('data-i18n'); + // Hide name/id/adapter rows; show base-url as editable + document.getElementById('provider-name-row').style.display = 'none'; + document.getElementById('provider-id-row').style.display = 'none'; + document.getElementById('provider-adapter-row').style.display = 'none'; + const baseUrlInput = document.getElementById('provider-base-url'); + const override = _builtinOverrides[id] || {}; + // Priority: db override > env > hardcoded default + const effectiveBaseUrl = override.base_url || p.env_base_url || p.base_url; + document.getElementById('provider-base-url-row').style.display = ''; + baseUrlInput.value = effectiveBaseUrl || ''; + baseUrlInput.readOnly = false; + baseUrlInput.style.opacity = ''; + baseUrlInput.placeholder = p.base_url || ''; + document.getElementById('provider-api-key-row').style.display = p.api_key_required !== false ? '' : 'none'; + document.getElementById('fetch-models-btn').style.display = p.can_list_models ? '' : 'none'; + const apiKeyInput = document.getElementById('provider-api-key'); + const hasDbKey = override.api_key === API_KEY_UNCHANGED; + const hasEnvKey = p.has_api_key === true; + apiKeyInput.value = ''; + if (hasDbKey) { + apiKeyInput.placeholder = I18n.t('config.apiKeyConfigured'); + } else if (hasEnvKey) { + apiKeyInput.placeholder = I18n.t('config.apiKeyFromEnv'); + } else { + apiKeyInput.placeholder = I18n.t('config.apiKeyEnter'); + } + document.getElementById('provider-model').value = override.model || p.env_model || p.default_model || ''; + openProviderDialog(true); + document.getElementById('provider-model').focus(); +} + +// Add provider form + +document.getElementById('add-provider-btn').addEventListener('click', () => { + openProviderDialog(false); +}); + +document.getElementById('cancel-provider-btn').addEventListener('click', () => { + resetProviderForm(); +}); + +document.getElementById('cancel-provider-footer-btn').addEventListener('click', () => { + resetProviderForm(); +}); + +document.getElementById('provider-dialog-overlay').addEventListener('click', () => { + resetProviderForm(); +}); + +function openProviderDialog(isEdit) { + if (!isEdit) { + // Add mode: ensure all rows visible + ['provider-name-row', 'provider-id-row', 'provider-adapter-row', + 'provider-base-url-row', 'provider-api-key-row'].forEach((id) => { + document.getElementById(id).style.display = ''; + }); + document.getElementById('fetch-models-btn').style.display = ''; + } + document.getElementById('provider-dialog').style.display = 'flex'; + if (!isEdit) { + document.getElementById('provider-name').focus(); + } +} + +document.getElementById('test-provider-btn').addEventListener('click', () => { + let adapter = document.getElementById('provider-adapter').value; + let baseUrl = document.getElementById('provider-base-url').value.trim(); + const apiKey = document.getElementById('provider-api-key').value.trim(); + const model = document.getElementById('provider-model').value.trim(); + + // For built-in providers, use the adapter from the registry. + // base_url comes from the form which already reflects: env > hardcoded default. + if (_configuringBuiltinId) { + const p = _builtinProviders.find((x) => x.id === _configuringBuiltinId); + if (p) { + adapter = p.adapter; + if (!baseUrl) baseUrl = p.base_url; + } + } + + const btn = document.getElementById('test-provider-btn'); + const result = document.getElementById('test-connection-result'); + + btn.disabled = true; + btn.textContent = I18n.t('config.testing'); + result.style.display = 'none'; + result.className = 'test-connection-result'; + + // Resolve provider_id so the backend can look up vaulted API keys. + const providerId = _configuringBuiltinId || document.getElementById('provider-id').value.trim(); + + if (!model) { + result.textContent = I18n.t('config.modelRequired') || 'Model is required for connection test'; + result.className = 'test-connection-result test-fail'; + result.style.display = ''; + btn.disabled = false; + btn.textContent = I18n.t('config.testConnection'); + return; + } + + apiFetch('/api/llm/test_connection', { + method: 'POST', + body: { + adapter, base_url: baseUrl, + api_key: apiKey || undefined, + model, + provider_id: providerId || undefined, + provider_type: _configuringBuiltinId ? 'builtin' : 'custom', + }, + }) + .then((data) => { + result.textContent = data.message; + result.className = 'test-connection-result ' + (data.ok ? 'test-ok' : 'test-fail'); + result.style.display = ''; + }) + .catch((e) => { + result.textContent = e.message; + result.className = 'test-connection-result test-fail'; + result.style.display = ''; + }) + .finally(() => { + btn.disabled = false; + btn.textContent = I18n.t('config.testConnection'); + }); +}); + +document.getElementById('save-provider-btn').addEventListener('click', () => { + // Built-in configure mode: save api_key + model to llm_builtin_overrides + if (_configuringBuiltinId) { + const apiKey = document.getElementById('provider-api-key').value.trim(); + const model = document.getElementById('provider-model').value.trim(); + const baseUrl = document.getElementById('provider-base-url').value.trim(); + const id = _configuringBuiltinId; + const prevOverride = _builtinOverrides[id] || {}; + const hadKey = prevOverride.api_key === API_KEY_UNCHANGED; + const override = {}; + if (apiKey) { + override.api_key = apiKey; // New key entered — backend will encrypt it + } else if (hadKey) { + override.api_key = API_KEY_UNCHANGED; // Sentinel: keep existing encrypted key + } + // If neither — key is cleared (no key configured) + if (model) override.model = model; + if (baseUrl) override.base_url = baseUrl; + const prev = _builtinOverrides[id]; + _builtinOverrides[id] = override; + const isActive = id === _activeLlmBackend; + const modelUpdate = () => { + if (!isActive) return Promise.resolve(); + if (model) { + return apiFetchVoid('/api/settings/selected_model', { method: 'PUT', body: { value: model } }); + } + return apiFetchVoid('/api/settings/selected_model', { method: 'DELETE' }); + }; + apiFetchVoid('/api/settings/llm_builtin_overrides', { method: 'PUT', body: { value: _builtinOverrides } }) + .then(() => modelUpdate()) + .then(() => { + if (isActive) _selectedModel = model; + renderProviders(); + if (isActive) loadInferenceSettings(); + resetProviderForm(); + scrollToProviders(); + if (isActive) { + document.getElementById('config-restart-notice').style.display = 'flex'; + var llmNotice = document.getElementById('llm-restart-notice'); + if (llmNotice) llmNotice.style.display = 'flex'; + } + showToast(I18n.t('config.providerConfigured', { name: id })); + }) + .catch((e) => { + if (prev !== undefined) { _builtinOverrides[id] = prev; } else { delete _builtinOverrides[id]; } + showToast(I18n.t('error.unknown') + ': ' + e.message, 'error'); + }); + return; + } + + const name = document.getElementById('provider-name').value.trim(); + const id = document.getElementById('provider-id').value.trim(); + const adapter = document.getElementById('provider-adapter').value; + const baseUrl = document.getElementById('provider-base-url').value.trim(); + const apiKey = document.getElementById('provider-api-key').value.trim(); + const model = document.getElementById('provider-model').value.trim(); + + if (!id || !name) { + showToast(I18n.t('config.providerFieldsRequired'), 'error'); + return; + } + + if (_editingProviderId) { + // Update existing provider + const idx = _customProviders.findIndex((p) => p.id === _editingProviderId); + if (idx === -1) return; + const original = _customProviders[idx]; + const hadCustomKey = original.api_key === API_KEY_UNCHANGED; + let effectiveApiKey; + if (apiKey) { + effectiveApiKey = apiKey; // New key — backend will encrypt it + } else if (hadCustomKey) { + effectiveApiKey = API_KEY_UNCHANGED; // Sentinel: keep existing encrypted key + } else { + effectiveApiKey = undefined; // No key + } + _customProviders[idx] = { ...original, name, adapter, base_url: baseUrl, default_model: model || undefined, api_key: effectiveApiKey }; + const isActive = _editingProviderId === _activeLlmBackend; + const modelUpdate = () => { + if (!isActive) return Promise.resolve(); + if (model) { + return apiFetchVoid('/api/settings/selected_model', { method: 'PUT', body: { value: model } }); + } + return apiFetchVoid('/api/settings/selected_model', { method: 'DELETE' }); + }; + saveCustomProviders().then(() => modelUpdate()).then(() => { + if (isActive) _selectedModel = model; + renderProviders(); + if (isActive) loadInferenceSettings(); + resetProviderForm(); + scrollToProviders(); + if (isActive) { + document.getElementById('config-restart-notice').style.display = 'flex'; + var llmNotice = document.getElementById('llm-restart-notice'); + if (llmNotice) llmNotice.style.display = 'flex'; + } + showToast(I18n.t('config.providerUpdated', { name })); + }).catch((e) => { + _customProviders[idx] = original; + showToast(I18n.t('error.unknown') + ': ' + e.message, 'error'); + }); + return; + } + + if (!/^[a-z0-9_-]+$/.test(id)) { + showToast(I18n.t('config.providerIdInvalid'), 'error'); + return; + } + const allIds = [..._builtinProviders.map((p) => p.id), ..._customProviders.map((p) => p.id)]; + if (allIds.includes(id)) { + showToast(I18n.t('config.providerIdTaken', { id }), 'error'); + return; + } + + const newProvider = { id, name, adapter, base_url: baseUrl, default_model: model, api_key: apiKey || undefined, builtin: false }; + _customProviders.push(newProvider); + + saveCustomProviders().then(() => { + renderProviders(); + resetProviderForm(); + scrollToProviders(); + showToast(I18n.t('config.providerAdded', { name })); + }).catch((e) => { + _customProviders.pop(); + showToast(I18n.t('error.unknown') + ': ' + e.message, 'error'); + }); +}); + +function resetProviderForm() { + _editingProviderId = null; + _configuringBuiltinId = null; + document.getElementById('provider-dialog').style.display = 'none'; + // Restore all hidden rows and buttons + ['provider-name-row', 'provider-id-row', 'provider-adapter-row', + 'provider-base-url-row', 'provider-api-key-row'].forEach((id) => { + document.getElementById(id).style.display = ''; + }); + document.getElementById('fetch-models-btn').style.display = ''; + const titleEl = document.getElementById('provider-form-title'); + titleEl.setAttribute('data-i18n', 'config.newProvider'); + titleEl.textContent = I18n.t('config.newProvider'); + const idField = document.getElementById('provider-id'); + idField.readOnly = false; + idField.style.opacity = ''; + delete idField.dataset.edited; + const baseUrlField = document.getElementById('provider-base-url'); + baseUrlField.readOnly = false; + baseUrlField.style.opacity = ''; + ['provider-name', 'provider-id', 'provider-base-url', 'provider-api-key', 'provider-model'].forEach((id) => { + document.getElementById(id).value = ''; + }); + document.getElementById('provider-adapter').selectedIndex = 0; + const sel = document.getElementById('provider-model-select'); + sel.innerHTML = ''; + sel.style.display = 'none'; + document.getElementById('test-connection-result').style.display = 'none'; +} + +document.getElementById('provider-model-select').addEventListener('change', (e) => { + document.getElementById('provider-model').value = e.target.value; +}); + +document.getElementById('fetch-models-btn').addEventListener('click', () => { + let adapter = document.getElementById('provider-adapter').value; + let baseUrl = document.getElementById('provider-base-url').value.trim(); + const apiKey = document.getElementById('provider-api-key').value.trim(); + + // For built-in providers, use the adapter from the registry. + // base_url comes from the form which already reflects: env > hardcoded default. + if (_configuringBuiltinId) { + const p = _builtinProviders.find((x) => x.id === _configuringBuiltinId); + if (p) { + adapter = p.adapter; + if (!baseUrl) baseUrl = p.base_url; + } + } + + if (!baseUrl) { + showToast(I18n.t('config.providerBaseUrlRequired'), 'error'); + return; + } + + const btn = document.getElementById('fetch-models-btn'); + btn.disabled = true; + btn.textContent = I18n.t('config.fetchingModels'); + + // Resolve provider_id so the backend can look up vaulted API keys. + const providerId = _configuringBuiltinId || document.getElementById('provider-id').value.trim(); + + apiFetch('/api/llm/list_models', { + method: 'POST', + body: { + adapter, base_url: baseUrl, + api_key: apiKey || undefined, + provider_id: providerId || undefined, + provider_type: _configuringBuiltinId ? 'builtin' : 'custom', + }, + }) + .then((data) => { + const select = document.getElementById('provider-model-select'); + if (data.ok && data.models && data.models.length > 0) { + const currentModel = document.getElementById('provider-model').value; + select.innerHTML = data.models + .map((m) => ``) + .join(''); + select.style.display = ''; + btn.style.display = 'none'; + showToast(I18n.t('config.modelsFetched', { count: data.models.length })); + } else { + showToast(data.message || I18n.t('config.modelsFetchFailed'), 'error'); + } + }) + .catch((e) => showToast(e.message, 'error')) + .finally(() => { + btn.disabled = false; + btn.textContent = I18n.t('config.fetchModels'); + }); +}); + +// Auto-fill provider ID from name +document.getElementById('provider-name').addEventListener('input', (e) => { + const idField = document.getElementById('provider-id'); + if (!idField.dataset.edited) { + idField.value = e.target.value.toLowerCase().replace(/[^a-z0-9_]+/g, '-').replace(/^-|-$/g, ''); + } +}); + +document.getElementById('provider-id').addEventListener('input', (e) => { + e.target.dataset.edited = e.target.value ? '1' : ''; +}); diff --git a/src/channels/web/static/i18n-app.js b/src/channels/web/static/i18n-app.js index 87624b96dc7..724b2b9bae3 100644 --- a/src/channels/web/static/i18n-app.js +++ b/src/channels/web/static/i18n-app.js @@ -35,16 +35,27 @@ function switchLanguage(lang) { if (I18n.setLanguage(lang)) { // Update slash commands updateSlashCommands(); - + // Update language menu active state updateLanguageMenu(); - + + // Re-render dynamically built sections that use I18n.t() + if (typeof renderProviders === 'function' && typeof _configLoaded !== 'undefined' && _configLoaded) { + renderProviders(); + } + if (typeof loadInferenceSettings === 'function') { + var inferencePanel = document.getElementById('settings-inference'); + if (inferencePanel && inferencePanel.classList.contains('active')) { + loadInferenceSettings(); + } + } + // Close menu const menu = document.getElementById('language-menu'); if (menu) { menu.style.display = 'none'; } - + // Show toast notification showToast(I18n.t('language.switch') + ': ' + (lang === 'zh-CN' ? '简体中文' : 'English')); } diff --git a/src/channels/web/static/i18n/en.js b/src/channels/web/static/i18n/en.js index 4592427f452..ea06d9d549d 100644 --- a/src/channels/web/static/i18n/en.js +++ b/src/channels/web/static/i18n/en.js @@ -36,12 +36,14 @@ I18n.register('en', { 'tab.settings': 'Settings', 'tab.extensions': 'Extensions', 'tab.skills': 'Skills', + 'tab.config': 'Config', 'tab.logs': 'Logs', 'settings.inference': 'Inference', 'settings.agent': 'Agent', 'settings.channels': 'Channels', 'settings.networking': 'Networking', 'settings.mcp': 'MCP', + 'settings.providers': 'Providers', 'settings.users': 'Users', // Users Tab @@ -387,6 +389,49 @@ I18n.register('en', { 'ext.removed': 'Removed {name}', 'ext.installFailed': 'Install failed: {message}', + // Config Tab — Model Providers + 'config.modelProviders': 'Model Providers', + 'config.addProvider': '+ Add Provider', + 'config.newProvider': 'New Provider', + 'config.restartNotice': 'Changes take effect after restart.', + 'config.builtin': 'built-in', + 'config.useProvider': 'Use', + 'config.configureProvider': 'Configure', + 'config.providerConfigured': 'Provider "{name}" configured (restart to apply)', + 'config.currentModel': 'Model: {model}', + 'config.providerName': 'Display Name', + 'config.providerNamePlaceholder': 'My Provider', + 'config.providerId': 'Provider ID', + 'config.providerIdPlaceholder': 'my-provider', + 'config.providerIdHint': 'Lowercase letters, numbers, hyphens, underscores', + 'config.providerAdapter': 'API Adapter', + 'config.adapterOpenAI': 'OpenAI Compatible', + 'config.adapterAnthropic': 'Anthropic', + 'config.adapterOllama': 'Ollama', + 'config.providerBaseUrl': 'Base URL', + 'config.providerApiKey': 'API Key', + 'config.apiKeyConfigured': 'Key configured (leave blank to keep)', + 'config.apiKeyFromEnv': 'Key set via environment variable', + 'config.apiKeyEnter': 'Enter API key', + 'config.providerModel': 'Default Model', + 'config.providerActivated': 'Switched to {name} (restart to apply)', + 'config.providerAdded': 'Added provider "{name}" (restart to apply)', + 'config.providerUpdated': 'Provider "{name}" updated (restart to apply)', + 'config.editProvider': 'Edit Provider', + 'config.providerDeleted': 'Provider deleted', + 'config.confirmDeleteProvider': 'Delete provider "{id}"?', + 'config.cannotDeleteActiveProvider': 'Cannot delete the active provider. Switch to another provider first.', + 'config.testConnection': 'Test', + 'config.testing': 'Testing…', + 'config.fetchModels': 'Fetch available models', + 'config.fetchingModels': 'Fetching…', + 'config.modelsFetched': '{count} model(s) loaded — type to filter', + 'config.modelsFetchFailed': 'Failed to fetch models', + 'config.providerBaseUrlRequired': 'Base URL is required to fetch models', + 'config.providerFieldsRequired': 'Display name and Provider ID are required', + 'config.providerIdInvalid': 'Provider ID: use only lowercase letters, numbers, hyphens, underscores', + 'config.providerIdTaken': 'Provider ID "{id}" is already taken', + // Configure 'config.title': 'Configure {name}', 'config.telegramOwnerHint': 'After saving, IronClaw will show a one-time code. Send `/start CODE` to your bot in Telegram and IronClaw will finish setup automatically.', diff --git a/src/channels/web/static/i18n/zh-CN.js b/src/channels/web/static/i18n/zh-CN.js index a0d3343887c..800e697bb14 100644 --- a/src/channels/web/static/i18n/zh-CN.js +++ b/src/channels/web/static/i18n/zh-CN.js @@ -36,12 +36,14 @@ I18n.register('zh-CN', { 'tab.settings': '设置', 'tab.extensions': '扩展', 'tab.skills': '技能', + 'tab.config': '配置', 'tab.logs': '日志', 'settings.inference': '推理', 'settings.agent': '代理', 'settings.channels': '频道', 'settings.networking': '网络', 'settings.mcp': 'MCP', + 'settings.providers': '模型提供商', 'settings.users': '用户管理', // 用户管理标签页 @@ -387,6 +389,49 @@ I18n.register('zh-CN', { 'ext.removed': '已移除 {name}', 'ext.installFailed': '安装失败: {message}', + // 配置页 — 模型提供商 + 'config.modelProviders': '模型提供商', + 'config.addProvider': '+ 添加提供商', + 'config.newProvider': '新建提供商', + 'config.restartNotice': '更改将在重启后生效。', + 'config.builtin': '内置', + 'config.useProvider': '使用', + 'config.configureProvider': '配置', + 'config.providerConfigured': '提供商 "{name}" 已配置(重启后生效)', + 'config.currentModel': '模型:{model}', + 'config.providerName': '显示名称', + 'config.providerNamePlaceholder': '我的提供商', + 'config.providerId': '提供商 ID', + 'config.providerIdPlaceholder': 'my-provider', + 'config.providerIdHint': '小写字母、数字、连字符、下划线', + 'config.providerAdapter': 'API 适配器', + 'config.adapterOpenAI': 'OpenAI 兼容', + 'config.adapterAnthropic': 'Anthropic', + 'config.adapterOllama': 'Ollama', + 'config.providerBaseUrl': '基础 URL', + 'config.providerApiKey': 'API 密钥', + 'config.apiKeyConfigured': '密钥已配置(留空保留)', + 'config.apiKeyFromEnv': '密钥已通过环境变量设置', + 'config.apiKeyEnter': '输入 API 密钥', + 'config.providerModel': '默认模型', + 'config.providerActivated': '已切换到 {name}(重启后生效)', + 'config.providerAdded': '已添加提供商 "{name}"(重启后生效)', + 'config.providerUpdated': '提供商 "{name}" 已更新(重启后生效)', + 'config.editProvider': '编辑提供商', + 'config.providerDeleted': '提供商已删除', + 'config.confirmDeleteProvider': '确定删除提供商 "{id}"?', + 'config.cannotDeleteActiveProvider': '无法删除当前正在使用的提供商,请先切换到其他提供商。', + 'config.testConnection': '测试', + 'config.testing': '测试中…', + 'config.fetchModels': '获取可用模型', + 'config.fetchingModels': '获取中…', + 'config.modelsFetched': '已加载 {count} 个模型,可输入过滤', + 'config.modelsFetchFailed': '获取模型列表失败', + 'config.providerBaseUrlRequired': '请先填写 Base URL', + 'config.providerFieldsRequired': '显示名称和提供商 ID 为必填项', + 'config.providerIdInvalid': '提供商 ID 只能包含小写字母、数字、连字符和下划线', + 'config.providerIdTaken': '提供商 ID "{id}" 已被占用', + // 配置 'config.title': '配置 {name}', 'config.telegramOwnerHint': '保存后,IronClaw 会显示一次性验证码。将 `/start CODE` 发送给你的 Telegram 机器人,IronClaw 会自动完成设置。', diff --git a/src/channels/web/static/index.html b/src/channels/web/static/index.html index 21ff6faaf33..9ea4ef6b3da 100644 --- a/src/channels/web/static/index.html +++ b/src/channels/web/static/index.html @@ -44,6 +44,58 @@

IronClaw

+ + +
-
-
Loading settings...
+
+
+
Loading settings...
+
+
+
+

Model Providers

+ +
+ +
+
Loading...
+
+
diff --git a/src/channels/web/static/style.css b/src/channels/web/static/style.css index 6e0dfcf6ef2..4afa591d7d5 100644 --- a/src/channels/web/static/style.css +++ b/src/channels/web/static/style.css @@ -2801,10 +2801,22 @@ body { padding: var(--space-4); } +#settings-inference > .extensions-container { + display: flex; + flex-direction: column; +} + .extensions-section { margin-bottom: 24px; } +#providers-section { + flex: 1; + min-height: 0; + display: flex; + flex-direction: column; +} + .extensions-section h3 { font-size: var(--text-xs); font-weight: 600; @@ -4593,6 +4605,12 @@ mark { min-width: 180px; } +.settings-display-value { + font-size: var(--text-sm); + color: var(--text); + font-family: 'IBM Plex Mono', monospace; +} + .settings-input { padding: 6px 10px; background: var(--bg); @@ -5430,6 +5448,408 @@ body.theme-transition *:not(svg):not(path):not(line):not(circle):not(rect) { } } + +/* --- Config Tab --- */ + +.config-section-header { + display: flex; + align-items: center; + justify-content: space-between; + margin-bottom: 12px; +} + +.config-section-header h3 { + margin-bottom: 0; +} + +.btn-add-provider { + padding: 5px 14px; + background: var(--accent); + color: #09090b; + border: none; + border-radius: var(--radius); + cursor: pointer; + font-size: 13px; + font-weight: 600; + transition: background 0.2s, transform 0.2s; +} + +.btn-add-provider:hover { + background: var(--accent-hover); + transform: translateY(-1px); +} + +.config-notice { + display: flex; + align-items: center; + gap: 8px; + padding: 8px 12px; + background: rgba(245, 166, 35, 0.1); + border: 1px solid rgba(245, 166, 35, 0.3); + border-radius: var(--radius); + color: var(--warning); + font-size: 13px; + margin-bottom: 12px; +} + +.providers-list { + display: flex; + flex-direction: column; + gap: 8px; + min-height: 420px; + overflow-y: auto; +} + +.provider-card { + background: var(--bg-secondary); + border: 1px solid var(--border); + border-radius: var(--radius-lg); + padding: 12px 14px; + display: flex; + flex-direction: column; + gap: 6px; + transition: border-color 0.2s; +} + +.provider-card:hover { + border-color: rgba(255, 255, 255, 0.15); +} + +.provider-card-active { + border-color: var(--accent); +} + +.provider-card-header { + display: flex; + align-items: center; + gap: 8px; + flex-wrap: wrap; +} + +.provider-name { + font-weight: 600; + font-size: 14px; + color: var(--text); +} + +.provider-id-label { + font-size: 11px; + color: var(--text-secondary); + font-family: var(--font-mono); +} + +.provider-badge { + font-size: 10px; + padding: 2px 7px; + border-radius: 20px; + font-weight: 600; + letter-spacing: 0.02em; +} + +.provider-badge-active { + background: rgba(52, 211, 153, 0.15); + color: var(--accent); +} + +.provider-badge-builtin { + background: rgba(161, 161, 170, 0.12); + color: var(--text-secondary); +} + +.provider-card-meta { + display: flex; + align-items: center; + gap: 10px; + flex-wrap: wrap; +} + +.provider-adapter { + font-size: 12px; + color: var(--text-secondary); +} + +.provider-url { + font-size: 11px; + color: var(--text-secondary); + font-family: var(--font-mono); + opacity: 0.7; +} + +.provider-current-model { + font-size: 11px; + color: var(--accent); + font-family: var(--font-mono); + font-weight: 500; +} + +.provider-card-actions { + display: flex; + gap: 6px; + margin-top: 2px; +} + +.provider-action-btn { + padding: 4px 12px; + background: var(--bg-tertiary); + border: 1px solid var(--border); + border-radius: var(--radius); + color: var(--text-secondary); + cursor: pointer; + font-size: 12px; + transition: color 0.2s, border-color 0.2s, background 0.2s; +} + +.provider-action-btn:hover { + color: var(--text); + border-color: rgba(255, 255, 255, 0.2); + background: var(--bg); +} + +.provider-delete-btn:hover { + color: var(--danger); + border-color: var(--danger); +} + +/* Config form */ + +.provider-dialog { + position: fixed; + top: 0; + left: 0; + right: 0; + bottom: 0; + z-index: 9999; + display: flex; + align-items: center; + justify-content: center; +} + +.provider-dialog-overlay { + position: absolute; + top: 0; + left: 0; + right: 0; + bottom: 0; + background: rgba(0, 0, 0, 0.5); + backdrop-filter: blur(4px); +} + +.provider-dialog-content { + position: relative; + z-index: 10000; + background: var(--bg-secondary); + border: 1px solid var(--border); + border-radius: var(--radius-lg); + box-shadow: 0 25px 50px -12px rgba(0, 0, 0, 0.4); + width: 100%; + max-width: 480px; + margin: 0 1rem; + display: flex; + flex-direction: column; + max-height: 90vh; +} + +.provider-dialog-header { + display: flex; + align-items: center; + justify-content: space-between; + padding: 14px 18px; + border-bottom: 1px solid var(--border); + flex-shrink: 0; +} + +.provider-dialog-header h2 { + font-size: 14px; + font-weight: 600; + color: var(--text); + margin: 0; +} + +.provider-dialog-close { + color: var(--text-secondary); + font-size: 18px; + line-height: 1; + padding: 2px 6px; + background: transparent; + border: none; + border-radius: var(--radius); + cursor: pointer; + transition: color 0.15s, background 0.15s; +} + +.provider-dialog-close:hover { + color: var(--text); + background: var(--bg-hover); +} + +.provider-dialog-body { + padding: 18px; + overflow-y: auto; + flex: 1; +} + +.provider-dialog-footer { + display: flex; + gap: 8px; + padding: 14px 18px; + border-top: 1px solid var(--border); + flex-shrink: 0; +} + +.provider-dialog-footer button { + padding: 6px 18px; + border-radius: var(--radius); + font-size: 13px; + font-weight: 600; + cursor: pointer; + transition: background 0.2s, transform 0.2s; +} + +.provider-dialog-footer button:first-child { + background: var(--accent); + color: #09090b; + border: none; +} + +.provider-dialog-footer button:first-child:hover { + background: var(--accent-hover); + transform: translateY(-1px); +} + +.provider-dialog-footer .btn-secondary { + background: transparent; + color: var(--text-secondary); + border: 1px solid var(--border); +} + +.provider-dialog-footer .btn-secondary:hover { + color: var(--text); + border-color: rgba(255, 255, 255, 0.2); +} + +.config-form { + display: flex; + flex-direction: column; + gap: 12px; +} + +.config-form-row { + display: flex; + flex-direction: column; + gap: 4px; +} + +.config-form-row label { + font-size: 12px; + font-weight: 500; + color: var(--text-secondary); +} + +.config-form-row input, +.config-form-row select { + padding: 7px 10px; + background: var(--bg); + border: 1px solid var(--border); + border-radius: var(--radius); + color: var(--text); + font-size: 13px; +} + +.config-form-row input:focus, +.config-form-row select:focus { + outline: none; + border-color: var(--accent); + box-shadow: 0 0 0 3px rgba(52, 211, 153, 0.1); +} + +.config-form-hint { + font-size: 11px; + color: var(--text-secondary); + opacity: 0.7; +} + +.config-form-actions { + display: flex; + gap: 8px; + margin-top: 4px; +} + +.config-form-actions button { + padding: 6px 18px; + border-radius: var(--radius); + font-size: 13px; + font-weight: 600; + cursor: pointer; + transition: background 0.2s, transform 0.2s; +} + +.config-form-actions button:first-child { + background: var(--accent); + color: #09090b; + border: none; +} + +.config-form-actions button:first-child:hover { + background: var(--accent-hover); + transform: translateY(-1px); +} + +.config-form-actions .btn-secondary { + background: transparent; + color: var(--text-secondary); + border: 1px solid var(--border); +} + +.config-form-actions .btn-secondary:hover { + color: var(--text); + border-color: rgba(255, 255, 255, 0.2); +} + +.btn-fetch-models { + display: inline-flex; + align-items: center; + gap: 5px; + margin-top: 6px; + padding: 5px 11px; + background: transparent; + border: 1px solid var(--border); + border-radius: var(--radius); + color: var(--text-secondary); + cursor: pointer; + font-size: 12px; + transition: color 0.15s, border-color 0.15s, background 0.15s; +} + +.btn-fetch-models:hover { + color: var(--text); + border-color: var(--accent); + background: color-mix(in srgb, var(--accent) 8%, transparent); +} + +.btn-fetch-models:disabled { + opacity: 0.5; + cursor: not-allowed; +} + +.test-connection-result { + margin-top: 8px; + padding: 6px 12px; + border-radius: var(--radius); + font-size: 13px; +} + +.test-connection-result.test-ok { + background: rgba(74, 222, 128, 0.12); + color: #4ade80; + border: 1px solid rgba(74, 222, 128, 0.3); +} + +.test-connection-result.test-fail { + background: rgba(248, 113, 113, 0.12); + color: #f87171; + border: 1px solid rgba(248, 113, 113, 0.3); +} + /* --- Users Tab --- */ .users-container { padding: 1rem; } .users-header { display: flex; align-items: center; justify-content: space-between; margin-bottom: 1rem; } diff --git a/src/config/llm.rs b/src/config/llm.rs index ed4b8a05591..dacd90520f3 100644 --- a/src/config/llm.rs +++ b/src/config/llm.rs @@ -46,28 +46,51 @@ impl LlmConfig { } } - /// Resolve a model name from env var -> settings.selected_model -> hardcoded default. + /// Resolve a model name from settings.selected_model -> env var -> hardcoded default. fn resolve_model( env_var: &str, settings: &Settings, default: &str, ) -> Result { - Ok(optional_env(env_var)? - .or_else(|| settings.selected_model.clone()) - .unwrap_or_else(|| default.to_string())) + if let Some(model) = settings.selected_model.clone() { + Ok(model) + } else if let Some(model) = optional_env(env_var)? { + Ok(model) + } else { + Ok(default.to_string()) + } } pub(crate) fn resolve(settings: &Settings) -> Result { let registry = ProviderRegistry::load(); - // Determine backend: env var > settings > default ("nearai") - let backend = if let Some(b) = optional_env("LLM_BACKEND")? { - b - } else if let Some(ref b) = settings.llm_backend { - b.clone() + // Determine backend: db settings > env var > default ("nearai") + let (backend, backend_source) = if let Some(ref b) = settings.llm_backend { + (b.clone(), "db:llm_backend") + } else if let Some(b) = optional_env("LLM_BACKEND")? { + (b, "env:LLM_BACKEND") } else { - "nearai".to_string() + ("nearai".to_string(), "default") }; + tracing::info!( + backend = %backend, + source = %backend_source, + db_llm_backend = ?settings.llm_backend, + custom_providers_count = settings.llm_custom_providers.len(), + "Resolving LLM backend" + ); + // Warn operators when a DB-persisted value silently overrides LLM_BACKEND. + if backend_source == "db:llm_backend" + && let Ok(env_val) = std::env::var("LLM_BACKEND") + && !env_val.is_empty() + { + tracing::warn!( + db_value = %backend, + env_value = %env_val, + "LLM_BACKEND env var is set but DB setting takes priority. \ + Unset llm_backend in the DB (via settings UI) to use the env var." + ); + } // Validate the backend is known let backend_lower = backend.to_lowercase(); @@ -80,10 +103,17 @@ impl LlmConfig { || backend_lower == "openai-codex" || backend_lower == "codex"; + // Check custom providers defined + let custom_provider = settings + .llm_custom_providers + .iter() + .find(|p| p.id.to_lowercase() == backend_lower); + if !is_nearai && !is_bedrock && !is_gemini_oauth && !is_openai_codex + && custom_provider.is_none() && registry.find(&backend_lower).is_none() { tracing::warn!( @@ -104,21 +134,37 @@ impl LlmConfig { }; // Always resolve NEAR AI config (used for embeddings even when not the primary backend) - let nearai_api_key = optional_env("NEARAI_API_KEY")?.map(SecretString::from); + // Priority: DB (builtin_overrides) > env > default + let nearai_override = settings.llm_builtin_overrides.get("nearai"); + let nearai_api_key = if let Some(key) = nearai_override.and_then(|o| o.api_key.as_ref()) { + Some(SecretString::from(key.clone())) + } else { + optional_env("NEARAI_API_KEY")?.map(SecretString::from) + }; + // Model priority: selected_model (DB) > builtin_overrides (DB) > env > default + let nearai_model = if let Some(model) = settings.selected_model.clone() { + model + } else if let Some(model) = nearai_override.and_then(|o| o.model.clone()) { + model + } else if let Some(model) = optional_env("NEARAI_MODEL")? { + model + } else { + crate::llm::DEFAULT_MODEL.to_string() + }; + let nearai_base_url = if let Some(url) = nearai_override.and_then(|o| o.base_url.clone()) { + url + } else if let Some(url) = optional_env("NEARAI_BASE_URL")? { + url + } else if nearai_api_key.is_some() { + "https://cloud-api.near.ai".to_string() + } else { + "https://private.near.ai".to_string() + }; + validate_base_url(&nearai_base_url, "NEARAI_BASE_URL")?; let nearai = NearAiConfig { - model: Self::resolve_model("NEARAI_MODEL", settings, crate::llm::DEFAULT_MODEL)?, + model: nearai_model, cheap_model: optional_env("NEARAI_CHEAP_MODEL")?, - base_url: { - let url = optional_env("NEARAI_BASE_URL")?.unwrap_or_else(|| { - if nearai_api_key.is_some() { - "https://cloud-api.near.ai".to_string() - } else { - "https://private.near.ai".to_string() - } - }); - validate_base_url(&url, "NEARAI_BASE_URL")?; - url - }, + base_url: nearai_base_url, api_key: nearai_api_key, fallback_model: optional_env("NEARAI_FALLBACK_MODEL")?, max_retries: parse_optional_env("NEARAI_MAX_RETRIES", 3)?, @@ -141,6 +187,8 @@ impl LlmConfig { // Resolve registry provider config (for non-NearAI, non-Bedrock, non-Gemini, non-Codex backends) let provider = if is_nearai || is_bedrock || is_gemini_oauth || is_openai_codex { None + } else if let Some(custom) = custom_provider { + Some(Self::resolve_custom_provider(custom, settings)?) } else { Some(Self::resolve_registry_provider( &backend_lower, @@ -150,20 +198,27 @@ impl LlmConfig { }; let bedrock = if is_bedrock { - let explicit_region = - optional_env("BEDROCK_REGION")?.or_else(|| settings.bedrock_region.clone()); + let explicit_region = settings + .bedrock_region + .clone() + .or(optional_env("BEDROCK_REGION")?); if explicit_region.is_none() { tracing::info!("BEDROCK_REGION not set, defaulting to us-east-1"); } let region = explicit_region.unwrap_or_else(|| "us-east-1".to_string()); - let model = optional_env("BEDROCK_MODEL")? - .or_else(|| settings.selected_model.clone()) + let model = settings + .selected_model + .clone() + .or(optional_env("BEDROCK_MODEL")?) .ok_or_else(|| ConfigError::MissingRequired { key: "BEDROCK_MODEL".to_string(), - hint: "Set BEDROCK_MODEL when LLM_BACKEND=bedrock".to_string(), + hint: "Set BEDROCK_MODEL or selected_model when LLM_BACKEND=bedrock" + .to_string(), })?; - let cross_region = optional_env("BEDROCK_CROSS_REGION")? - .or_else(|| settings.bedrock_cross_region.clone()); + let cross_region = settings + .bedrock_cross_region + .clone() + .or(optional_env("BEDROCK_CROSS_REGION")?); if let Some(ref cr) = cross_region && !matches!(cr.as_str(), "us" | "eu" | "apac" | "global") { @@ -175,7 +230,10 @@ impl LlmConfig { ), }); } - let profile = optional_env("AWS_PROFILE")?.or_else(|| settings.bedrock_profile.clone()); + let profile = settings + .bedrock_profile + .clone() + .or(optional_env("AWS_PROFILE")?); Some(BedrockConfig { region, model, @@ -188,10 +246,12 @@ impl LlmConfig { // Resolve OpenAI Codex config let openai_codex = if is_openai_codex { - // Model: OPENAI_CODEX_MODEL > OPENAI_MODEL > settings.selected_model > default - let model = optional_env("OPENAI_CODEX_MODEL")? + // Model: settings.selected_model > OPENAI_CODEX_MODEL > OPENAI_MODEL > default + let model = settings + .selected_model + .clone() + .or(optional_env("OPENAI_CODEX_MODEL")?) .or(optional_env("OPENAI_MODEL")?) - .or_else(|| settings.selected_model.clone()) .unwrap_or_else(|| "gpt-5.3-codex".to_string()); let auth_endpoint = optional_env("OPENAI_CODEX_AUTH_URL")? .unwrap_or_else(|| "https://auth.openai.com".to_string()); @@ -267,6 +327,65 @@ impl LlmConfig { }) } + /// Resolve a `RegistryProviderConfig` from a user-defined custom provider. + fn resolve_custom_provider( + custom: &crate::settings::CustomLlmProviderSettings, + settings: &Settings, + ) -> Result { + tracing::info!( + id = %custom.id, + adapter = %custom.adapter, + base_url = ?custom.base_url, + "Resolving custom LLM provider" + ); + let protocol = match custom.adapter.as_str() { + "anthropic" => ProviderProtocol::Anthropic, + "ollama" => ProviderProtocol::Ollama, + _ => ProviderProtocol::OpenAiCompletions, + }; + + let api_key = custom + .api_key + .as_ref() + .filter(|k| !k.is_empty()) + .map(|k| SecretString::from(k.clone())); + + let base_url = custom.base_url.clone().unwrap_or_default(); + if base_url.is_empty() { + tracing::warn!(id = %custom.id, "Custom provider has no base_url configured — requests will fail"); + } else { + validate_base_url( + &base_url, + &format!("custom provider '{}' base_url", custom.id), + )?; + } + + let model = settings + .selected_model + .clone() + .or(optional_env("LLM_MODEL")?) + .or_else(|| custom.default_model.clone()) + .unwrap_or_default(); + if model.is_empty() { + tracing::warn!(id = %custom.id, "Custom provider has no model configured — requests may fail"); + } + + Ok(RegistryProviderConfig { + protocol, + provider_id: custom.id.clone(), + api_key, + base_url, + model, + extra_headers: Vec::new(), + oauth_token: None, + is_codex_chatgpt: false, + refresh_token: None, + auth_path: None, + cache_retention: CacheRetention::default(), + unsupported_params: Vec::new(), + }) + } + /// Resolve a `RegistryProviderConfig` from the registry and env vars. fn resolve_registry_provider( backend: &str, @@ -344,8 +463,16 @@ impl LlmConfig { } Some(creds.token) } else if let Some(env_var) = api_key_env { - // Resolve API key from env (including secrets store overlay) - optional_env(env_var)?.map(SecretString::from) + // Resolve API key: settings override (DB) > env var (including secrets store overlay) + if let Some(key) = settings + .llm_builtin_overrides + .get(backend) + .and_then(|o| o.api_key.as_ref()) + { + Some(SecretString::from(key.clone())) + } else { + optional_env(env_var)?.map(SecretString::from) + } } else { None }; @@ -361,18 +488,23 @@ impl LlmConfig { } } - // Resolve base URL: codex override > env var > settings (backward compat) > registry default + // Resolve base URL: codex override > builtin_overrides (DB) > legacy settings (DB) > env var > registry default let is_codex_chatgpt = codex_base_url_override.is_some(); + let env_base_url = if let Some(env_var) = base_url_env { + optional_env(env_var)? + } else { + None + }; let base_url = codex_base_url_override .or_else(|| { - if let Some(env_var) = base_url_env { - optional_env(env_var).ok().flatten() - } else { - None - } + // DB settings: per-provider base_url override + settings + .llm_builtin_overrides + .get(backend) + .and_then(|o| o.base_url.clone()) }) .or_else(|| { - // Backward compat: check legacy settings fields + // DB settings: legacy settings fields match backend { "ollama" => settings.ollama_base_url.clone(), "openai_compatible" | "openrouter" => { @@ -381,6 +513,7 @@ impl LlmConfig { _ => None, } }) + .or(env_base_url) .or_else(|| default_base_url.map(String::from)) .unwrap_or_default(); @@ -400,8 +533,18 @@ impl LlmConfig { validate_base_url(&base_url, field)?; } - // Resolve model - let model = Self::resolve_model(model_env, settings, default_model)?; + // Resolve model: selected_model (DB) > per-provider override (DB) > env var > registry default + let model = settings + .selected_model + .clone() + .or_else(|| { + settings + .llm_builtin_overrides + .get(backend) + .and_then(|o| o.model.clone()) + }) + .or(optional_env(model_env)?) + .unwrap_or_else(|| default_model.to_string()); // Resolve extra headers let extra_headers = if let Some(env_var) = extra_headers_env { @@ -573,7 +716,7 @@ mod tests { } #[test] - fn openai_compatible_llm_model_env_overrides_selected_model() { + fn openai_compatible_selected_model_overrides_env() { let _guard = lock_env(); clear_openai_compatible_env(); // SAFETY: Under ENV_MUTEX. @@ -591,7 +734,10 @@ mod tests { let cfg = LlmConfig::resolve(&settings).expect("resolve should succeed"); let provider = cfg.provider.expect("provider config should be present"); - assert_eq!(provider.model, "openai/gpt-5-codex"); + assert_eq!( + provider.model, "openai/gpt-5.1-codex", + "DB selected_model should take priority over LLM_MODEL env var" + ); // SAFETY: Under ENV_MUTEX. unsafe { @@ -714,7 +860,7 @@ mod tests { } #[test] - fn ollama_model_env_overrides_selected_model() { + fn ollama_selected_model_overrides_env() { let _guard = lock_env(); clear_ollama_env(); // SAFETY: Under ENV_MUTEX. @@ -731,7 +877,10 @@ mod tests { let cfg = LlmConfig::resolve(&settings).expect("resolve should succeed"); let provider = cfg.provider.expect("provider config should be present"); - assert_eq!(provider.model, "mistral:latest"); + assert_eq!( + provider.model, "llama3.2", + "DB selected_model should take priority over OLLAMA_MODEL env var" + ); // SAFETY: Under ENV_MUTEX. unsafe { @@ -994,28 +1143,31 @@ mod tests { ..Default::default() }; + // DB settings should take priority over env var let cfg = LlmConfig::resolve(&settings).expect("resolve should succeed"); let provider = cfg.provider.expect("should have provider config"); assert_eq!( - provider.base_url, "http://localhost:8000/v1", - "env var should take priority over settings" + provider.base_url, "http://localhost:9000/v1", + "DB settings should take priority over env var" ); - // Now without env var, settings should win over registry default - unsafe { - std::env::remove_var("LLM_BASE_URL"); - } + // Without DB settings, env var should win over registry default + let settings_no_base = Settings { + llm_backend: Some("openai_compatible".to_string()), + ..Default::default() + }; - let cfg = LlmConfig::resolve(&settings).expect("resolve should succeed"); + let cfg = LlmConfig::resolve(&settings_no_base).expect("resolve should succeed"); let provider = cfg.provider.expect("should have provider config"); assert_eq!( - provider.base_url, "http://localhost:9000/v1", - "settings should take priority over registry default" + provider.base_url, "http://localhost:8000/v1", + "env var should take priority over registry default when DB has no base_url" ); // SAFETY: Under ENV_MUTEX. unsafe { std::env::remove_var("LLM_BACKEND"); + std::env::remove_var("LLM_BASE_URL"); } } @@ -1240,6 +1392,81 @@ mod tests { } } + // ── Custom provider tests ─────────────────────────────────────── + + #[test] + fn custom_provider_resolves_when_backend_matches_id() { + let _guard = lock_env(); + // SAFETY: Under ENV_MUTEX. + unsafe { + std::env::remove_var("LLM_BACKEND"); + std::env::remove_var("LLM_MODEL"); + } + + let settings = Settings { + llm_backend: Some("myprovider".to_string()), + llm_custom_providers: vec![crate::settings::CustomLlmProviderSettings { + id: "myprovider".to_string(), + name: "My Provider".to_string(), + adapter: "open_ai_completions".to_string(), + base_url: Some("http://localhost:9090/v1".to_string()), + default_model: Some("my-model".to_string()), + api_key: Some("sk-test".to_string()), + builtin: false, + }], + ..Default::default() + }; + + let cfg = LlmConfig::resolve(&settings).expect("resolve should succeed"); + assert_eq!(cfg.backend, "myprovider"); + let provider = cfg.provider.expect("provider config should be present"); + assert_eq!(provider.provider_id, "myprovider"); + assert_eq!(provider.base_url, "http://localhost:9090/v1"); + assert_eq!(provider.model, "my-model"); + assert_eq!( + provider.protocol, + crate::llm::registry::ProviderProtocol::OpenAiCompletions + ); + } + + #[test] + fn db_llm_backend_takes_priority_over_env_var() { + let _guard = lock_env(); + // SAFETY: Under ENV_MUTEX. RAII guard removes LLM_BACKEND on drop so + // a panicking assertion cannot leak the env var to other tests. + struct RemoveOnDrop(&'static str); + impl Drop for RemoveOnDrop { + fn drop(&mut self) { + unsafe { std::env::remove_var(self.0) }; + } + } + let _cleanup = RemoveOnDrop("LLM_BACKEND"); + unsafe { + std::env::set_var("LLM_BACKEND", "nearai"); + std::env::remove_var("LLM_MODEL"); + } + + let settings = Settings { + llm_backend: Some("myprovider".to_string()), + llm_custom_providers: vec![crate::settings::CustomLlmProviderSettings { + id: "myprovider".to_string(), + name: "My Provider".to_string(), + adapter: "open_ai_completions".to_string(), + base_url: Some("http://localhost:9090/v1".to_string()), + default_model: Some("my-model".to_string()), + api_key: None, + builtin: false, + }], + ..Default::default() + }; + + let cfg = LlmConfig::resolve(&settings).expect("resolve should succeed"); + assert_eq!( + cfg.backend, "myprovider", + "DB setting should override LLM_BACKEND env var" + ); + } + // ── OpenAI Codex tests ────────────────────────────────────────── /// Clear all openai-codex-related env vars. @@ -1252,6 +1479,38 @@ mod tests { } } + #[test] + fn builtin_override_model_used_when_no_selected_model() { + let _guard = lock_env(); + // SAFETY: Under ENV_MUTEX. + unsafe { + std::env::remove_var("LLM_BACKEND"); + std::env::remove_var("GROQ_MODEL"); + } + + let mut overrides = std::collections::HashMap::new(); + overrides.insert( + "groq".to_string(), + crate::settings::LlmBuiltinOverride { + api_key: None, + model: Some("llama-3.1-8b-instant".to_string()), + base_url: None, + }, + ); + let settings = Settings { + llm_backend: Some("groq".to_string()), + llm_builtin_overrides: overrides, + ..Default::default() + }; + + let cfg = LlmConfig::resolve(&settings).expect("resolve should succeed"); + let provider = cfg.provider.expect("provider config should be present"); + assert_eq!( + provider.model, "llama-3.1-8b-instant", + "builtin override model should be used when selected_model is unset" + ); + } + #[test] fn openai_codex_resolves_config() { let _guard = lock_env(); @@ -1272,6 +1531,39 @@ mod tests { ); } + #[test] + fn selected_model_takes_priority_over_builtin_override_model() { + let _guard = lock_env(); + // SAFETY: Under ENV_MUTEX. + unsafe { + std::env::remove_var("LLM_BACKEND"); + std::env::remove_var("GROQ_MODEL"); + } + + let mut overrides = std::collections::HashMap::new(); + overrides.insert( + "groq".to_string(), + crate::settings::LlmBuiltinOverride { + api_key: None, + model: Some("llama-3.1-8b-instant".to_string()), + base_url: None, + }, + ); + let settings = Settings { + llm_backend: Some("groq".to_string()), + selected_model: Some("llama-3.3-70b-versatile".to_string()), + llm_builtin_overrides: overrides, + ..Default::default() + }; + + let cfg = LlmConfig::resolve(&settings).expect("resolve should succeed"); + let provider = cfg.provider.expect("provider config should be present"); + assert_eq!( + provider.model, "llama-3.3-70b-versatile", + "selected_model (/model command) must take priority over builtin override" + ); + } + #[test] fn openai_codex_model_env_resolution() { let _guard = lock_env(); @@ -1296,6 +1588,44 @@ mod tests { } } + #[test] + fn builtin_override_api_key_used_when_no_env_var() { + let _guard = lock_env(); + // SAFETY: Under ENV_MUTEX. + unsafe { + std::env::remove_var("LLM_BACKEND"); + std::env::remove_var("GROQ_API_KEY"); + std::env::remove_var("GROQ_MODEL"); + } + + let mut overrides = std::collections::HashMap::new(); + overrides.insert( + "groq".to_string(), + crate::settings::LlmBuiltinOverride { + api_key: Some("gsk_test_key".to_string()), + model: Some("llama-3.3-70b-versatile".to_string()), + base_url: None, + }, + ); + let settings = Settings { + llm_backend: Some("groq".to_string()), + llm_builtin_overrides: overrides, + ..Default::default() + }; + + let cfg = LlmConfig::resolve(&settings).expect("resolve should succeed"); + let provider = cfg.provider.expect("provider config should be present"); + use secrecy::ExposeSecret as _; + let key = provider + .api_key + .expect("api_key should be set from builtin override"); + assert_eq!( + key.expose_secret(), + "gsk_test_key", + "builtin override api_key should be used when env var is absent" + ); + } + #[test] fn openai_codex_falls_back_to_openai_model() { let _guard = lock_env(); @@ -1394,4 +1724,435 @@ mod tests { std::env::remove_var("OPENAI_CODEX_AUTH_URL"); } } + + // ── DB > ENV priority tests ───────────────────────────────────── + + #[test] + fn builtin_override_api_key_wins_over_env_var() { + let _guard = lock_env(); + // SAFETY: Under ENV_MUTEX. + unsafe { + std::env::remove_var("LLM_BACKEND"); + std::env::set_var("GROQ_API_KEY", "gsk_from_env"); + std::env::remove_var("GROQ_MODEL"); + } + + let mut overrides = std::collections::HashMap::new(); + overrides.insert( + "groq".to_string(), + crate::settings::LlmBuiltinOverride { + api_key: Some("gsk_from_db".to_string()), + model: None, + base_url: None, + }, + ); + let settings = Settings { + llm_backend: Some("groq".to_string()), + llm_builtin_overrides: overrides, + ..Default::default() + }; + + let cfg = LlmConfig::resolve(&settings).expect("resolve should succeed"); + let provider = cfg.provider.expect("provider config should be present"); + use secrecy::ExposeSecret as _; + assert_eq!( + provider + .api_key + .as_ref() + .map(|k| k.expose_secret().to_string()), + Some("gsk_from_db".to_string()), + "DB builtin_override api_key must take priority over GROQ_API_KEY env var" + ); + + // SAFETY: Under ENV_MUTEX. + unsafe { + std::env::remove_var("GROQ_API_KEY"); + } + } + + #[test] + fn builtin_override_model_wins_over_env_var() { + let _guard = lock_env(); + // SAFETY: Under ENV_MUTEX. + unsafe { + std::env::remove_var("LLM_BACKEND"); + std::env::set_var("GROQ_MODEL", "model-from-env"); + } + + let mut overrides = std::collections::HashMap::new(); + overrides.insert( + "groq".to_string(), + crate::settings::LlmBuiltinOverride { + api_key: None, + model: Some("model-from-db".to_string()), + base_url: None, + }, + ); + let settings = Settings { + llm_backend: Some("groq".to_string()), + llm_builtin_overrides: overrides, + ..Default::default() + }; + + let cfg = LlmConfig::resolve(&settings).expect("resolve should succeed"); + let provider = cfg.provider.expect("provider config should be present"); + assert_eq!( + provider.model, "model-from-db", + "DB builtin_override model must take priority over GROQ_MODEL env var" + ); + + // SAFETY: Under ENV_MUTEX. + unsafe { + std::env::remove_var("GROQ_MODEL"); + } + } + + #[test] + fn custom_provider_selected_model_wins_over_env() { + let _guard = lock_env(); + // SAFETY: Under ENV_MUTEX. + unsafe { + std::env::remove_var("LLM_BACKEND"); + std::env::set_var("LLM_MODEL", "model-from-env"); + } + + let settings = Settings { + llm_backend: Some("myprovider".to_string()), + selected_model: Some("model-from-db".to_string()), + llm_custom_providers: vec![crate::settings::CustomLlmProviderSettings { + id: "myprovider".to_string(), + name: "My Provider".to_string(), + adapter: "open_ai_completions".to_string(), + base_url: Some("http://localhost:9090/v1".to_string()), + default_model: Some("default-model".to_string()), + api_key: None, + builtin: false, + }], + ..Default::default() + }; + + let cfg = LlmConfig::resolve(&settings).expect("resolve should succeed"); + let provider = cfg.provider.expect("provider config should be present"); + assert_eq!( + provider.model, "model-from-db", + "DB selected_model must take priority over LLM_MODEL env var for custom providers" + ); + + // SAFETY: Under ENV_MUTEX. + unsafe { + std::env::remove_var("LLM_MODEL"); + } + } + + #[test] + fn openai_codex_selected_model_wins_over_env() { + let _guard = lock_env(); + clear_openai_codex_env(); + // SAFETY: Under ENV_MUTEX. + unsafe { + std::env::set_var("OPENAI_CODEX_MODEL", "codex-from-env"); + } + + let settings = Settings { + llm_backend: Some("openai_codex".to_string()), + selected_model: Some("codex-from-db".to_string()), + ..Default::default() + }; + + let cfg = LlmConfig::resolve(&settings).expect("resolve should succeed"); + let codex = cfg.openai_codex.expect("codex config should be present"); + assert_eq!( + codex.model, "codex-from-db", + "DB selected_model must take priority over OPENAI_CODEX_MODEL env var" + ); + + // SAFETY: Under ENV_MUTEX. + unsafe { + std::env::remove_var("OPENAI_CODEX_MODEL"); + } + } + + #[test] + fn nearai_selected_model_wins_over_env() { + let _guard = lock_env(); + // SAFETY: Under ENV_MUTEX. + unsafe { + std::env::remove_var("LLM_BACKEND"); + std::env::set_var("NEARAI_MODEL", "nearai-from-env"); + } + + let settings = Settings { + llm_backend: Some("nearai".to_string()), + selected_model: Some("nearai-from-db".to_string()), + ..Default::default() + }; + + let cfg = LlmConfig::resolve(&settings).expect("resolve should succeed"); + assert_eq!( + cfg.nearai.model, "nearai-from-db", + "DB selected_model must take priority over NEARAI_MODEL env var" + ); + + // SAFETY: Under ENV_MUTEX. + unsafe { + std::env::remove_var("NEARAI_MODEL"); + } + } + + #[test] + fn nearai_override_model_wins_over_env() { + let _guard = lock_env(); + // SAFETY: Under ENV_MUTEX. + unsafe { + std::env::remove_var("LLM_BACKEND"); + std::env::set_var("NEARAI_MODEL", "model-from-env"); + } + + let mut overrides = std::collections::HashMap::new(); + overrides.insert( + "nearai".to_string(), + crate::settings::LlmBuiltinOverride { + api_key: None, + model: Some("model-from-db-override".to_string()), + base_url: None, + }, + ); + let settings = Settings { + llm_backend: Some("nearai".to_string()), + llm_builtin_overrides: overrides, + ..Default::default() + }; + + let cfg = LlmConfig::resolve(&settings).expect("resolve should succeed"); + assert_eq!( + cfg.nearai.model, "model-from-db-override", + "DB builtin_overrides model must take priority over NEARAI_MODEL env var" + ); + + // SAFETY: Under ENV_MUTEX. + unsafe { + std::env::remove_var("NEARAI_MODEL"); + } + } + + #[test] + fn nearai_selected_model_wins_over_override_model() { + let _guard = lock_env(); + // SAFETY: Under ENV_MUTEX. + unsafe { + std::env::remove_var("LLM_BACKEND"); + std::env::remove_var("NEARAI_MODEL"); + } + + let mut overrides = std::collections::HashMap::new(); + overrides.insert( + "nearai".to_string(), + crate::settings::LlmBuiltinOverride { + api_key: None, + model: Some("model-from-override".to_string()), + base_url: None, + }, + ); + let settings = Settings { + llm_backend: Some("nearai".to_string()), + selected_model: Some("model-from-selected".to_string()), + llm_builtin_overrides: overrides, + ..Default::default() + }; + + let cfg = LlmConfig::resolve(&settings).expect("resolve should succeed"); + assert_eq!( + cfg.nearai.model, "model-from-selected", + "selected_model must take priority over builtin_overrides model" + ); + } + + #[test] + fn nearai_override_base_url_wins_over_env() { + let _guard = lock_env(); + // SAFETY: Under ENV_MUTEX. + unsafe { + std::env::remove_var("LLM_BACKEND"); + std::env::set_var("NEARAI_BASE_URL", "http://localhost:9001"); + std::env::remove_var("NEARAI_API_KEY"); + } + + let mut overrides = std::collections::HashMap::new(); + overrides.insert( + "nearai".to_string(), + crate::settings::LlmBuiltinOverride { + api_key: None, + model: None, + base_url: Some("http://localhost:9002".to_string()), + }, + ); + let settings = Settings { + llm_backend: Some("nearai".to_string()), + llm_builtin_overrides: overrides, + ..Default::default() + }; + + let cfg = LlmConfig::resolve(&settings).expect("resolve should succeed"); + assert_eq!( + cfg.nearai.base_url, "http://localhost:9002", + "DB builtin_overrides base_url must take priority over NEARAI_BASE_URL env var" + ); + + // SAFETY: Under ENV_MUTEX. + unsafe { + std::env::remove_var("NEARAI_BASE_URL"); + } + } + + #[test] + fn nearai_env_base_url_used_when_no_override() { + let _guard = lock_env(); + // SAFETY: Under ENV_MUTEX. + unsafe { + std::env::remove_var("LLM_BACKEND"); + std::env::set_var("NEARAI_BASE_URL", "http://localhost:9001"); + std::env::remove_var("NEARAI_API_KEY"); + } + + let settings = Settings { + llm_backend: Some("nearai".to_string()), + ..Default::default() + }; + + let cfg = LlmConfig::resolve(&settings).expect("resolve should succeed"); + assert_eq!( + cfg.nearai.base_url, "http://localhost:9001", + "NEARAI_BASE_URL env var should be used when no DB override exists" + ); + + // SAFETY: Under ENV_MUTEX. + unsafe { + std::env::remove_var("NEARAI_BASE_URL"); + } + } + + #[test] + fn nearai_override_api_key_wins_over_env() { + let _guard = lock_env(); + // SAFETY: Under ENV_MUTEX. + unsafe { + std::env::remove_var("LLM_BACKEND"); + std::env::set_var("NEARAI_API_KEY", "key-from-env"); + } + + let mut overrides = std::collections::HashMap::new(); + overrides.insert( + "nearai".to_string(), + crate::settings::LlmBuiltinOverride { + api_key: Some("key-from-db".to_string()), + model: None, + base_url: None, + }, + ); + let settings = Settings { + llm_backend: Some("nearai".to_string()), + llm_builtin_overrides: overrides, + ..Default::default() + }; + + let cfg = LlmConfig::resolve(&settings).expect("resolve should succeed"); + use secrecy::ExposeSecret as _; + assert_eq!( + cfg.nearai + .api_key + .as_ref() + .map(|k| k.expose_secret().to_string()), + Some("key-from-db".to_string()), + "DB builtin_overrides api_key must take priority over NEARAI_API_KEY env var" + ); + + // SAFETY: Under ENV_MUTEX. + unsafe { + std::env::remove_var("NEARAI_API_KEY"); + } + } + + #[test] + fn nearai_base_url_auto_selects_when_no_override_or_env() { + let _guard = lock_env(); + // SAFETY: Under ENV_MUTEX. + unsafe { + std::env::remove_var("LLM_BACKEND"); + std::env::remove_var("NEARAI_BASE_URL"); + std::env::remove_var("NEARAI_API_KEY"); + } + + // No API key → should default to private.near.ai + let settings = Settings { + llm_backend: Some("nearai".to_string()), + ..Default::default() + }; + + let cfg = LlmConfig::resolve(&settings).expect("resolve should succeed"); + assert_eq!( + cfg.nearai.base_url, "https://private.near.ai", + "Without API key, should default to private.near.ai" + ); + + // With API key → should default to cloud-api.near.ai + let mut overrides = std::collections::HashMap::new(); + overrides.insert( + "nearai".to_string(), + crate::settings::LlmBuiltinOverride { + api_key: Some("some-key".to_string()), + model: None, + base_url: None, + }, + ); + let settings_with_key = Settings { + llm_backend: Some("nearai".to_string()), + llm_builtin_overrides: overrides, + ..Default::default() + }; + + let cfg = LlmConfig::resolve(&settings_with_key).expect("resolve should succeed"); + assert_eq!( + cfg.nearai.base_url, "https://cloud-api.near.ai", + "With API key, should default to cloud-api.near.ai" + ); + } + + #[test] + fn registry_provider_override_base_url_wins_over_env() { + let _guard = lock_env(); + // SAFETY: Under ENV_MUTEX. + unsafe { + std::env::remove_var("LLM_BACKEND"); + std::env::set_var("GROQ_BASE_URL", "http://localhost:9003"); + std::env::remove_var("GROQ_API_KEY"); + std::env::remove_var("GROQ_MODEL"); + } + + let mut overrides = std::collections::HashMap::new(); + overrides.insert( + "groq".to_string(), + crate::settings::LlmBuiltinOverride { + api_key: None, + model: None, + base_url: Some("http://localhost:9004".to_string()), + }, + ); + let settings = Settings { + llm_backend: Some("groq".to_string()), + llm_builtin_overrides: overrides, + ..Default::default() + }; + + let cfg = LlmConfig::resolve(&settings).expect("resolve should succeed"); + let provider = cfg.provider.expect("provider config should be present"); + assert_eq!( + provider.base_url, "http://localhost:9004", + "DB builtin_overrides base_url must take priority over GROQ_BASE_URL env var" + ); + + // SAFETY: Under ENV_MUTEX. + unsafe { + std::env::remove_var("GROQ_BASE_URL"); + } + } } diff --git a/src/config/mod.rs b/src/config/mod.rs index a362fd090c4..03f37c5dec9 100644 --- a/src/config/mod.rs +++ b/src/config/mod.rs @@ -1,9 +1,13 @@ //! Configuration for IronClaw. //! -//! Settings are loaded with priority: env var > database > default. +//! Settings are loaded from env vars, the DB settings table, TOML config, +//! and built-in defaults. Priority varies by subsystem: +//! +//! - **LLM settings** (backend, model, api_key, base_url): DB > env > default +//! - **Most other settings** (agent, channels, tunnel, …): env > DB > default +//! //! `DATABASE_URL` lives in `~/.ironclaw/.env` (loaded via dotenvy early -//! in startup). Everything else comes from env vars, the DB settings -//! table, or auto-detection. +//! in startup). mod agent; mod builder; @@ -186,8 +190,9 @@ impl Config { /// Load configuration from environment variables and the database. /// - /// Priority: env var > TOML config file > DB settings > default. - /// This is the primary way to load config after DB is connected. + /// TOML is loaded first as a base, then DB values are merged on top + /// (DB wins over TOML). Individual subsystem resolvers then apply + /// their own env-vs-DB priority — see module docs for details. pub async fn from_db( store: &(dyn crate::db::SettingsStore + Sync), user_id: &str, @@ -196,6 +201,10 @@ impl Config { } /// Load from DB with an optional TOML config file overlay. + /// + /// TOML is loaded first as a base, then DB values are merged on top + /// (DB wins over TOML). Per-subsystem resolvers then decide whether + /// env vars or DB values take final precedence — see module docs. pub async fn from_db_with_toml( store: &(dyn crate::db::SettingsStore + Sync), user_id: &str, @@ -204,19 +213,22 @@ impl Config { let _ = dotenvy::dotenv(); crate::bootstrap::load_ironclaw_env(); - // Load all settings from DB into a Settings struct - let mut db_settings = match store.get_all_settings(user_id).await { - Ok(map) => Settings::from_db_map(&map), + // Start with TOML config as a base (lowest priority among the two). + let mut settings = Settings::default(); + Self::apply_toml_overlay(&mut settings, toml_path)?; + + // Overlay DB settings on top so DB values win over TOML. + match store.get_all_settings(user_id).await { + Ok(map) => { + let db_settings = Settings::from_db_map(&map); + settings.merge_from(&db_settings); + } Err(e) => { tracing::warn!("Failed to load settings from DB, using defaults: {}", e); - Settings::default() } }; - // Overlay TOML config file (values win over DB settings) - Self::apply_toml_overlay(&mut db_settings, toml_path)?; - - Self::build(&db_settings).await + Self::build(&settings).await } /// Load configuration from environment variables only (no database). @@ -291,16 +303,38 @@ impl Config { user_id: &str, toml_path: Option<&std::path::Path>, ) -> Result<(), ConfigError> { - let settings = if let Some(store) = store { - let mut s = match store.get_all_settings(user_id).await { - Ok(map) => Settings::from_db_map(&map), - Err(_) => Settings::default(), - }; + self.re_resolve_llm_with_secrets(store, user_id, toml_path, None) + .await + } + + /// Re-resolve LLM config, hydrating API keys from the secrets store. + pub async fn re_resolve_llm_with_secrets( + &mut self, + store: Option<&(dyn crate::db::SettingsStore + Sync)>, + user_id: &str, + toml_path: Option<&std::path::Path>, + secrets: Option<&(dyn crate::secrets::SecretsStore + Send + Sync)>, + ) -> Result<(), ConfigError> { + let mut settings = if let Some(store) = store { + // TOML as base, then DB on top (DB wins). + let mut s = Settings::default(); Self::apply_toml_overlay(&mut s, toml_path)?; + if let Ok(map) = store.get_all_settings(user_id).await { + let db_settings = Settings::from_db_map(&map); + s.merge_from(&db_settings); + } s } else { Settings::default() }; + + // Hydrate API keys from encrypted secrets store into the settings + // struct so that LlmConfig::resolve() sees them without any changes + // to its synchronous resolution logic. + if let Some(secrets) = secrets { + hydrate_llm_keys_from_secrets(&mut settings, secrets, user_id).await; + } + self.llm = LlmConfig::resolve(&settings)?; Ok(()) } @@ -501,3 +535,302 @@ fn inject_os_credential_store_tokens(injected: &mut HashMap) { tracing::debug!("Refreshed ANTHROPIC_OAUTH_TOKEN from OS credential store"); } } + +/// Hydrate LLM API keys from the secrets store into the settings struct. +/// +/// Called after loading settings from DB but before `LlmConfig::resolve()`. +/// Populates `api_key` fields that were stripped from settings during the +/// write path and stored encrypted in the secrets store instead. +pub async fn hydrate_llm_keys_from_secrets( + settings: &mut Settings, + secrets: &(dyn crate::secrets::SecretsStore + Send + Sync), + user_id: &str, +) { + // Hydrate builtin overrides + for (provider_id, override_val) in settings.llm_builtin_overrides.iter_mut() { + if override_val.api_key.is_some() { + continue; // Already has a key (legacy plaintext or TOML) + } + let secret_name = crate::settings::builtin_secret_name(provider_id); + if let Ok(decrypted) = secrets.get_decrypted(user_id, &secret_name).await { + override_val.api_key = Some(decrypted.expose().to_string()); + } + } + + // Hydrate custom providers + for provider in settings.llm_custom_providers.iter_mut() { + if provider.api_key.is_some() { + continue; + } + let secret_name = crate::settings::custom_secret_name(&provider.id); + if let Ok(decrypted) = secrets.get_decrypted(user_id, &secret_name).await { + provider.api_key = Some(decrypted.expose().to_string()); + } + } +} + +/// Migrate plaintext API keys from the settings table to the encrypted secrets store. +/// +/// Idempotent: skips keys that are already in the secrets store. +/// After migration, strips plaintext keys from the settings table. +pub async fn migrate_plaintext_llm_keys( + settings_store: &(dyn crate::db::SettingsStore + Sync), + secrets: &(dyn crate::secrets::SecretsStore + Send + Sync), + user_id: &str, +) { + let settings_map = match settings_store.get_all_settings(user_id).await { + Ok(m) => m, + Err(_) => return, + }; + + let mut migrated = 0u32; + + // Migrate builtin overrides + if let Some(obj) = settings_map + .get("llm_builtin_overrides") + .and_then(|v| v.as_object()) + { + let mut sanitized = obj.clone(); + for (provider_id, override_val) in obj { + if let Some(api_key) = override_val.get("api_key").and_then(|v| v.as_str()) { + if api_key.is_empty() { + continue; + } + let secret_name = crate::settings::builtin_secret_name(provider_id); + if !secrets.exists(user_id, &secret_name).await.unwrap_or(false) + && let Err(e) = secrets + .create( + user_id, + crate::secrets::CreateSecretParams { + name: secret_name.clone(), + value: secrecy::SecretString::from(api_key.to_string()), + provider: Some(provider_id.clone()), + expires_at: None, + }, + ) + .await + { + tracing::warn!("Failed to migrate key for builtin '{}': {}", provider_id, e); + continue; + } + if let Some(o) = sanitized + .get_mut(provider_id) + .and_then(|v| v.as_object_mut()) + { + o.remove("api_key"); + } + migrated += 1; + } + } + if migrated > 0 { + let _ = settings_store + .set_setting( + user_id, + "llm_builtin_overrides", + &serde_json::Value::Object(sanitized), + ) + .await; + } + } + + // Migrate custom providers + let before = migrated; + if let Some(arr) = settings_map + .get("llm_custom_providers") + .and_then(|v| v.as_array()) + { + let mut sanitized = arr.clone(); + for (idx, provider_val) in arr.iter().enumerate() { + let provider_id = provider_val + .get("id") + .and_then(|v| v.as_str()) + .unwrap_or(""); + if provider_id.is_empty() { + continue; + } + if let Some(api_key) = provider_val.get("api_key").and_then(|v| v.as_str()) { + if api_key.is_empty() { + continue; + } + let secret_name = crate::settings::custom_secret_name(provider_id); + if !secrets.exists(user_id, &secret_name).await.unwrap_or(false) + && let Err(e) = secrets + .create( + user_id, + crate::secrets::CreateSecretParams { + name: secret_name.clone(), + value: secrecy::SecretString::from(api_key.to_string()), + provider: Some(provider_id.to_string()), + expires_at: None, + }, + ) + .await + { + tracing::warn!("Failed to migrate key for custom '{}': {}", provider_id, e); + continue; + } + if let Some(o) = sanitized[idx].as_object_mut() { + o.remove("api_key"); + } + migrated += 1; + } + } + if migrated > before { + let _ = settings_store + .set_setting( + user_id, + "llm_custom_providers", + &serde_json::Value::Array(sanitized), + ) + .await; + } + } + + if migrated > 0 { + tracing::info!( + "Migrated {} plaintext LLM API key(s) to encrypted secrets store", + migrated + ); + } +} + +#[cfg(test)] +mod tests { + use super::*; + use std::sync::Arc; + + fn test_secrets_store() -> Arc { + let crypto = Arc::new( + crate::secrets::SecretsCrypto::new(secrecy::SecretString::from( + crate::secrets::keychain::generate_master_key_hex(), + )) + .unwrap(), + ); + Arc::new(crate::secrets::InMemorySecretsStore::new(crypto)) + } + + #[tokio::test] + async fn hydrate_populates_builtin_override_keys_from_secrets() { + let secrets = test_secrets_store(); + secrets + .create( + "test", + crate::secrets::CreateSecretParams { + name: "llm_builtin_openai_api_key".to_string(), + value: secrecy::SecretString::from("sk-from-vault".to_string()), + provider: Some("openai".to_string()), + expires_at: None, + }, + ) + .await + .unwrap(); + + let mut settings = Settings { + llm_builtin_overrides: { + let mut m = std::collections::HashMap::new(); + m.insert( + "openai".to_string(), + crate::settings::LlmBuiltinOverride { + api_key: None, // stripped during write + model: Some("gpt-4o".to_string()), + base_url: None, + }, + ); + m + }, + ..Default::default() + }; + + hydrate_llm_keys_from_secrets(&mut settings, secrets.as_ref(), "test").await; + + assert_eq!( + settings.llm_builtin_overrides["openai"].api_key.as_deref(), + Some("sk-from-vault"), + "api_key should be hydrated from secrets store" + ); + assert_eq!( + settings.llm_builtin_overrides["openai"].model.as_deref(), + Some("gpt-4o"), + "model should remain unchanged" + ); + } + + #[tokio::test] + async fn hydrate_populates_custom_provider_keys_from_secrets() { + let secrets = test_secrets_store(); + secrets + .create( + "test", + crate::secrets::CreateSecretParams { + name: "llm_custom_my-llm_api_key".to_string(), + value: secrecy::SecretString::from("gsk-custom".to_string()), + provider: Some("my-llm".to_string()), + expires_at: None, + }, + ) + .await + .unwrap(); + + let mut settings = Settings { + llm_custom_providers: vec![crate::settings::CustomLlmProviderSettings { + id: "my-llm".to_string(), + name: "My LLM".to_string(), + adapter: "open_ai_completions".to_string(), + base_url: Some("http://localhost:8080".to_string()), + default_model: Some("model-1".to_string()), + api_key: None, // stripped during write + builtin: false, + }], + ..Default::default() + }; + + hydrate_llm_keys_from_secrets(&mut settings, secrets.as_ref(), "test").await; + + assert_eq!( + settings.llm_custom_providers[0].api_key.as_deref(), + Some("gsk-custom"), + "custom provider api_key should be hydrated from secrets store" + ); + } + + #[tokio::test] + async fn hydrate_skips_when_key_already_present() { + let secrets = test_secrets_store(); + secrets + .create( + "test", + crate::secrets::CreateSecretParams { + name: "llm_builtin_openai_api_key".to_string(), + value: secrecy::SecretString::from("sk-from-vault".to_string()), + provider: Some("openai".to_string()), + expires_at: None, + }, + ) + .await + .unwrap(); + + let mut settings = Settings { + llm_builtin_overrides: { + let mut m = std::collections::HashMap::new(); + m.insert( + "openai".to_string(), + crate::settings::LlmBuiltinOverride { + api_key: Some("sk-existing".to_string()), + model: None, + base_url: None, + }, + ); + m + }, + ..Default::default() + }; + + hydrate_llm_keys_from_secrets(&mut settings, secrets.as_ref(), "test").await; + + assert_eq!( + settings.llm_builtin_overrides["openai"].api_key.as_deref(), + Some("sk-existing"), + "existing key should not be overwritten" + ); + } +} diff --git a/src/llm/mod.rs b/src/llm/mod.rs index d681547d33e..cc838041ecc 100644 --- a/src/llm/mod.rs +++ b/src/llm/mod.rs @@ -93,6 +93,8 @@ pub async fn create_llm_provider( ) -> Result, LlmError> { let timeout = config.request_timeout_secs; + tracing::info!(backend = %config.backend, "Creating LLM provider"); + if config.backend == "nearai" || config.backend == "near_ai" || config.backend == "near" { return create_llm_provider_with_config(&config.nearai, session, timeout); } diff --git a/src/llm/rig_adapter.rs b/src/llm/rig_adapter.rs index bfc6c56744d..f61c1e8050c 100644 --- a/src/llm/rig_adapter.rs +++ b/src/llm/rig_adapter.rs @@ -301,6 +301,10 @@ fn convert_messages(messages: &[ChatMessage]) -> (Option, Vec { if msg.content_parts.is_empty() { + // Skip empty user messages — some providers (e.g. Kimi) reject "content": "" + if msg.content.is_empty() { + continue; + } history.push(RigMessage::user(&msg.content)); } else { // Build multimodal user message with text + image parts @@ -364,6 +368,12 @@ fn convert_messages(messages: &[ChatMessage]) -> (Option, Vec { + assert_eq!(content.len(), 1); + let first = content.iter().next().expect("one content item"); + match first { + UserContent::Text(t) => assert_eq!(t.text, "hello"), + other => panic!("expected Text, got {:?}", other), + } + } + other => panic!("expected User message, got {:?}", other), + } + } + + /// Empty assistant messages (e.g. after thinking-tag stripping) must be skipped. + #[test] + fn test_empty_assistant_message_is_skipped() { + let empty_asst = ChatMessage { + role: crate::llm::Role::Assistant, + content: String::new(), + tool_calls: None, + tool_call_id: None, + name: None, + content_parts: vec![], + }; + let non_empty = ChatMessage::user("hi"); + let messages = vec![empty_asst, non_empty]; + let (_preamble, history) = convert_messages(&messages); + + assert_eq!(history.len(), 1, "empty assistant message must be dropped"); + assert!(matches!(history[0], RigMessage::User { .. })); + } + + /// A conversation mixing normal and empty messages: only non-empty ones survive. + #[test] + fn test_mixed_empty_and_non_empty_messages_filtered_correctly() { + let user1 = ChatMessage::user("first"); + let empty_asst = ChatMessage { + role: crate::llm::Role::Assistant, + content: String::new(), + tool_calls: None, + tool_call_id: None, + name: None, + content_parts: vec![], + }; + let user2 = ChatMessage::user(""); + let asst = ChatMessage::assistant("response"); + let messages = vec![user1, empty_asst, user2, asst]; + let (_preamble, history) = convert_messages(&messages); + + assert_eq!(history.len(), 2, "only non-empty messages should survive"); + assert!(matches!(history[0], RigMessage::User { .. })); + assert!(matches!(history[1], RigMessage::Assistant { .. })); + } + // -- normalized_tool_call_id tests -- #[test] diff --git a/src/main.rs b/src/main.rs index 88bf76c5615..22dfdcb0bf1 100644 --- a/src/main.rs +++ b/src/main.rs @@ -679,6 +679,9 @@ async fn async_main() -> anyhow::Result<()> { } } } + if let Some(ref ss) = components.secrets_store { + gw = gw.with_secrets_store(Arc::clone(ss)); + } if let Some(ref jm) = container_job_manager { gw = gw.with_job_manager(Arc::clone(jm)); } diff --git a/src/settings.rs b/src/settings.rs index 09d9d9d06e7..f549557b485 100644 --- a/src/settings.rs +++ b/src/settings.rs @@ -1,14 +1,72 @@ //! User settings persistence. //! -//! Stores user preferences in ~/.ironclaw/settings.json. -//! Settings are loaded with env var > settings.json > default priority. - +//! Stores user preferences in `~/.ironclaw` (JSON/TOML) and, for some values, +//! in the database. At runtime, precedence between database values, +//! environment variables, on-disk config, and built-in defaults is determined +//! on a per-setting basis by the corresponding resolver. +//! LLM backend and related settings in particular may prefer DB values over +//! environment variables, as documented on their respective types. + +use std::collections::HashMap; use std::path::PathBuf; use serde::{Deserialize, Serialize}; use crate::bootstrap::ironclaw_base_dir; +/// A custom LLM provider defined by the user through the web UI. +#[derive(Debug, Clone, Serialize, Deserialize)] +pub struct CustomLlmProviderSettings { + /// Unique identifier (used as `llm_backend` value). + pub id: String, + /// Display name. + pub name: String, + /// Adapter protocol: "open_ai_completions", "anthropic", "ollama". + pub adapter: String, + /// Base URL for the API endpoint. + #[serde(default)] + pub base_url: Option, + /// Default model identifier. + #[serde(default)] + pub default_model: Option, + /// Optional API key stored inline. + #[serde(default)] + pub api_key: Option, + /// Whether this is a built-in provider (should always be false for custom). + #[serde(default)] + pub builtin: bool, +} + +/// Per-provider overrides for built-in LLM providers (API key and/or model). +/// +/// Stored as `llm_builtin_overrides` in the settings store, keyed by provider ID +/// (e.g. `"openai"`, `"gemini"`). Resolved at startup during `LlmConfig::resolve()`. +/// +/// Note: The global `selected_model` (if set) takes precedence over these +/// per-provider overrides, which in turn take precedence over environment variables. +#[derive(Debug, Clone, Default, Serialize, Deserialize)] +pub struct LlmBuiltinOverride { + /// API key override. Takes precedence over environment variables. + #[serde(default, skip_serializing_if = "Option::is_none")] + pub api_key: Option, + /// Model override. Takes precedence over environment variables but not `selected_model`. + #[serde(default, skip_serializing_if = "Option::is_none")] + pub model: Option, + /// Base URL override. Takes precedence over environment variables. + #[serde(default, skip_serializing_if = "Option::is_none")] + pub base_url: Option, +} + +/// Canonical secret name for a built-in provider's API key. +pub fn builtin_secret_name(provider_id: &str) -> String { + format!("llm_builtin_{provider_id}_api_key") +} + +/// Canonical secret name for a custom provider's API key. +pub fn custom_secret_name(provider_id: &str) -> String { + format!("llm_custom_{provider_id}_api_key") +} + /// User settings persisted to disk. #[derive(Debug, Clone, Serialize, Deserialize, Default)] pub struct Settings { @@ -59,6 +117,14 @@ pub struct Settings { #[serde(default)] pub llm_backend: Option, + /// Custom LLM providers defined by the user through the web UI. + #[serde(default)] + pub llm_custom_providers: Vec, + + /// Per-provider overrides for built-in providers (API key and/or model). + #[serde(default)] + pub llm_builtin_overrides: HashMap, + /// Ollama base URL (when llm_backend = "ollama"). #[serde(default)] pub ollama_base_url: Option, @@ -841,7 +907,8 @@ impl Settings { let content = format!( "# IronClaw configuration file.\n\ #\n\ - # Priority: env var > this file > database settings > defaults.\n\ + # Priority varies by subsystem. LLM: DB > env > this file > defaults.\n\ + # Most others: env > DB > this file > defaults.\n\ # Uncomment and edit values to override defaults.\n\ # Run `ironclaw config init` to regenerate this file.\n\ #\n\ @@ -1325,56 +1392,53 @@ mod tests { ); } - /// Regression: TOML overlay must not clobber a DB-persisted selected_model - /// when the TOML file matches the DB. This is the normal case after /model - /// successfully writes to both DB and TOML. + /// TOML is loaded as a base, then DB is merged on top (DB wins). + /// When both agree, the result matches. #[test] - fn toml_overlay_preserves_matching_model() { - // DB settings with new model from /model command. - let mut db_settings = Settings { - llm_backend: Some("nearai".to_string()), + fn toml_and_db_matching_model_preserved() { + // from_db_with_toml: TOML base, then DB merged on top. + let mut toml_base = Settings { selected_model: Some("new-model".to_string()), ..Default::default() }; - // TOML also updated by /model command to the same value. - let toml_settings = Settings { + let db_overlay = Settings { + llm_backend: Some("nearai".to_string()), selected_model: Some("new-model".to_string()), ..Default::default() }; - db_settings.merge_from(&toml_settings); + toml_base.merge_from(&db_overlay); assert_eq!( - db_settings.selected_model, + toml_base.selected_model, Some("new-model".to_string()), - "TOML overlay must not clobber matching model" + "matching values: result should be the shared value" ); } - /// Regression: when /model updates DB but TOML write fails, a stale TOML - /// file would overwrite the DB value. This test documents the priority: - /// TOML > DB (by design). persist_selected_model MUST update the TOML. + /// Regression: when TOML has a stale model but DB has been updated via + /// /model command, DB must win. This matches from_db_with_toml where + /// TOML is loaded first as base, then DB is merged on top. #[test] - fn stale_toml_overwrites_db_model() { - // DB has the new model from /model. - let mut db_settings = Settings { - selected_model: Some("new-model".to_string()), + fn db_model_wins_over_stale_toml() { + // TOML base with old model. + let mut toml_base = Settings { + selected_model: Some("old-model".to_string()), ..Default::default() }; - // TOML still has the old model (write failed or was not attempted). - let stale_toml = Settings { - selected_model: Some("old-model".to_string()), + // DB has the new model from /model command. + let db_overlay = Settings { + selected_model: Some("new-model".to_string()), ..Default::default() }; - db_settings.merge_from(&stale_toml); - // This documents the current priority: TOML wins over DB. - // The fix in persist_selected_model ensures TOML is always updated. + // from_db_with_toml: TOML first, then DB merged on top. + toml_base.merge_from(&db_overlay); assert_eq!( - db_settings.selected_model, - Some("old-model".to_string()), - "TOML overlay has higher priority than DB (by design)" + toml_base.selected_model, + Some("new-model".to_string()), + "DB selected_model must win over stale TOML value" ); } @@ -1403,24 +1467,20 @@ mod tests { assert_eq!(reloaded.selected_model, Some("new-model".to_string())); } - /// Regression: /model must create config.toml when it doesn't exist, so the - /// model survives restarts. Previously the Ok(None) case was a no-op. + /// save_toml / load_toml round-trip for selected_model. #[test] - fn toml_created_when_missing_for_model_persist() { + fn toml_save_and_load_round_trip() { let dir = tempfile::tempdir().unwrap(); let path = dir.path().join("config.toml"); - // No config.toml yet (fresh install, no wizard). assert!(Settings::load_toml(&path).unwrap().is_none()); - // Simulate what persist_selected_model now does for the Ok(None) case. let settings = Settings { selected_model: Some("new-model".to_string()), ..Default::default() }; settings.save_toml(&path).unwrap(); - // Verify the model survived. let loaded = Settings::load_toml(&path).unwrap().unwrap(); assert_eq!(loaded.selected_model, Some("new-model".to_string())); } @@ -2378,4 +2438,65 @@ mod tests { assert_eq!(current.embeddings.provider, "nearai"); assert_eq!(current.embeddings.model, "text-embedding-3-large"); } + + /// DB values must win over TOML values when both set the same field. + /// + /// This mirrors the merge order in `Config::from_db_with_toml`: + /// TOML is loaded as the base, then DB is merged on top. + #[test] + fn db_settings_win_over_toml_settings() { + // Simulate TOML base: has llm_backend and selected_model + let mut base = Settings { + llm_backend: Some("openai".to_string()), + selected_model: Some("toml-model".to_string()), + ..Default::default() + }; + + // Simulate DB overlay: has different llm_backend and selected_model + let db = Settings { + llm_backend: Some("anthropic".to_string()), + selected_model: Some("db-model".to_string()), + ..Default::default() + }; + + // Merge DB on top of TOML (same order as from_db_with_toml) + base.merge_from(&db); + + assert_eq!( + base.llm_backend.as_deref(), + Some("anthropic"), + "DB llm_backend must win over TOML" + ); + assert_eq!( + base.selected_model.as_deref(), + Some("db-model"), + "DB selected_model must win over TOML" + ); + } + + /// When DB has no value (default), TOML value should be preserved. + #[test] + fn toml_settings_used_when_db_has_no_value() { + let mut base = Settings { + llm_backend: Some("openai".to_string()), + selected_model: Some("toml-model".to_string()), + ..Default::default() + }; + + // DB has no llm_backend or selected_model (both default/None) + let db = Settings::default(); + + base.merge_from(&db); + + assert_eq!( + base.llm_backend.as_deref(), + Some("openai"), + "TOML llm_backend should be preserved when DB has no value" + ); + assert_eq!( + base.selected_model.as_deref(), + Some("toml-model"), + "TOML selected_model should be preserved when DB has no value" + ); + } } From 4f277c91be53366caaad00cfbf4245693dcd2ac7 Mon Sep 17 00:00:00 2001 From: Henry Park Date: Sun, 29 Mar 2026 14:33:41 -0700 Subject: [PATCH 04/11] Handle empty tool completions in autonomous jobs (#1720) * Handle empty tool completions in autonomous jobs * Address malformed tool recovery review comments * style: apply rustfmt to reasoning tests --------- Co-authored-by: Firat Sertgoz --- src/agent/agentic_loop.rs | 93 ++++++++++- src/agent/dispatcher.rs | 6 + src/llm/mod.rs | 6 +- src/llm/reasoning.rs | 139 +++++++++++++++- src/worker/autonomous_recovery.rs | 150 +++++++++++++++++ src/worker/container.rs | 89 +++++++++- src/worker/job.rs | 91 ++++++++++- src/worker/mod.rs | 1 + tests/e2e_builtin_tool_coverage.rs | 254 ++++++++++++++++++++++++++++- 9 files changed, 809 insertions(+), 20 deletions(-) create mode 100644 src/worker/autonomous_recovery.rs diff --git a/src/agent/agentic_loop.rs b/src/agent/agentic_loop.rs index 27c2ab726ac..59f89816161 100644 --- a/src/agent/agentic_loop.rs +++ b/src/agent/agentic_loop.rs @@ -10,7 +10,9 @@ use std::borrow::Cow; use crate::agent::session::PendingApproval; use crate::error::Error; -use crate::llm::{ChatMessage, FinishReason, Reasoning, ReasoningContext, RespondResult}; +use crate::llm::{ + ChatMessage, FinishReason, Reasoning, ReasoningContext, RespondResult, ResponseMetadata, +}; /// Signal from the delegate indicating how the loop should proceed. pub enum LoopSignal { @@ -38,6 +40,8 @@ pub enum LoopOutcome { Stopped, /// Max iterations exceeded. MaxIterations, + /// Loop terminated early with a clear failure reason. + Failure(String), /// A tool requires user approval before continuing (chat delegate only). NeedApproval(Box), } @@ -103,6 +107,7 @@ pub trait LoopDelegate: Send + Sync { async fn handle_text_response( &self, text: &str, + metadata: ResponseMetadata, reason_ctx: &mut ReasoningContext, ) -> TextAction; @@ -209,7 +214,10 @@ pub async fn run_agentic_loop( consecutive_tool_intent_nudges = 0; } - match delegate.handle_text_response(&text, reason_ctx).await { + match delegate + .handle_text_response(&text, output.metadata, reason_ctx) + .await + { TextAction::Return(outcome) => return Ok(outcome), TextAction::Continue => {} } @@ -279,7 +287,7 @@ pub fn truncate_for_preview(s: &str, max: usize) -> Cow<'_, str> { #[cfg(test)] mod tests { use super::*; - use crate::llm::{RespondOutput, TokenUsage, ToolCall}; + use crate::llm::{RespondOutput, ResponseAnomaly, ResponseMetadata, TokenUsage, ToolCall}; use crate::testing::StubLlm; use std::sync::Arc; use std::sync::atomic::{AtomicUsize, Ordering}; @@ -303,6 +311,7 @@ mod tests { result: RespondResult::Text(text.to_string()), usage: zero_usage(), finish_reason: FinishReason::Stop, + metadata: ResponseMetadata::default(), } } @@ -314,6 +323,7 @@ mod tests { }, usage: zero_usage(), finish_reason: FinishReason::ToolUse, + metadata: ResponseMetadata::default(), } } @@ -391,6 +401,7 @@ mod tests { async fn handle_text_response( &self, text: &str, + _metadata: ResponseMetadata, _reason_ctx: &mut ReasoningContext, ) -> TextAction { TextAction::Return(LoopOutcome::Response(text.to_string())) @@ -508,6 +519,79 @@ mod tests { ); } + #[tokio::test] + async fn test_text_response_metadata_can_fail_fast() { + struct FailOnMalformedResponse; + + #[async_trait] + impl LoopDelegate for FailOnMalformedResponse { + async fn check_signals(&self) -> LoopSignal { + LoopSignal::Continue + } + + async fn before_llm_call( + &self, + _: &mut ReasoningContext, + _: usize, + ) -> Option { + None + } + + async fn call_llm( + &self, + _: &Reasoning, + _: &mut ReasoningContext, + _: usize, + ) -> Result { + Ok(RespondOutput { + result: RespondResult::Text("fallback".to_string()), + usage: zero_usage(), + finish_reason: FinishReason::Stop, + metadata: ResponseMetadata { + anomaly: Some(ResponseAnomaly::EmptyToolCompletion), + }, + }) + } + + async fn handle_text_response( + &self, + _: &str, + metadata: ResponseMetadata, + _: &mut ReasoningContext, + ) -> TextAction { + assert_eq!(metadata.anomaly, Some(ResponseAnomaly::EmptyToolCompletion)); + TextAction::Return(LoopOutcome::Failure( + "malformed tool completion".to_string(), + )) + } + + async fn execute_tool_calls( + &self, + _: Vec, + _: Option, + _: &mut ReasoningContext, + ) -> Result, crate::error::Error> { + Ok(None) + } + } + + let delegate = FailOnMalformedResponse; + let reasoning = stub_reasoning(); + let mut ctx = ReasoningContext::new(); + let outcome = run_agentic_loop( + &delegate, + &reasoning, + &mut ctx, + &AgenticLoopConfig::default(), + ) + .await + .unwrap(); + + assert!( + matches!(outcome, LoopOutcome::Failure(ref reason) if reason == "malformed tool completion") + ); + } + #[tokio::test] async fn test_max_iterations_reached() { struct ContinueDelegate; @@ -535,6 +619,7 @@ mod tests { async fn handle_text_response( &self, _: &str, + _: ResponseMetadata, ctx: &mut ReasoningContext, ) -> TextAction { ctx.messages.push(ChatMessage::assistant("still working")); @@ -671,6 +756,7 @@ mod tests { }, usage: zero_usage(), finish_reason: FinishReason::Length, // response was truncated + metadata: ResponseMetadata::default(), }; let delegate = MockDelegate::new(vec![truncated_output, text_output("Summarized it.")]); let reasoning = stub_reasoning(); @@ -719,6 +805,7 @@ mod tests { }, usage: zero_usage(), finish_reason: FinishReason::Length, + metadata: ResponseMetadata::default(), }; // Three truncated responses, then a text response let delegate = MockDelegate::new(vec![ diff --git a/src/agent/dispatcher.rs b/src/agent/dispatcher.rs index 4420a450470..99dd294ad90 100644 --- a/src/agent/dispatcher.rs +++ b/src/agent/dispatcher.rs @@ -219,6 +219,11 @@ impl Agent { reason: format!("Exceeded maximum tool iterations ({max_tool_iterations})"), } .into()), + LoopOutcome::Failure(reason) => Err(crate::error::LlmError::InvalidResponse { + provider: "agent".to_string(), + reason, + } + .into()), LoopOutcome::NeedApproval(pending) => Ok(AgenticLoopResult::NeedApproval { pending }), } } @@ -462,6 +467,7 @@ impl<'a> LoopDelegate for ChatDelegate<'a> { async fn handle_text_response( &self, text: &str, + _metadata: crate::llm::ResponseMetadata, _reason_ctx: &mut ReasoningContext, ) -> TextAction { // Strip internal "[Called tool ...]" text that can leak when diff --git a/src/llm/mod.rs b/src/llm/mod.rs index cc838041ecc..d6fadb6714a 100644 --- a/src/llm/mod.rs +++ b/src/llm/mod.rs @@ -62,9 +62,9 @@ pub use provider::{ ToolDefinition, ToolResult, generate_tool_call_id, }; pub use reasoning::{ - ActionPlan, Reasoning, ReasoningContext, RespondOutput, RespondResult, SILENT_REPLY_TOKEN, - TOOL_INTENT_NUDGE, TRUNCATED_TOOL_CALL_NOTICE, TokenUsage, ToolSelection, is_silent_reply, - llm_signals_tool_intent, + ActionPlan, Reasoning, ReasoningContext, RespondOutput, RespondResult, ResponseAnomaly, + ResponseMetadata, SILENT_REPLY_TOKEN, TOOL_INTENT_NUDGE, TRUNCATED_TOOL_CALL_NOTICE, + TokenUsage, ToolSelection, is_silent_reply, llm_signals_tool_intent, }; pub use recording::RecordingLlm; pub use registry::{ProviderDefinition, ProviderProtocol, ProviderRegistry}; diff --git a/src/llm/reasoning.rs b/src/llm/reasoning.rs index f5fc6b8a132..02d4e68dfff 100644 --- a/src/llm/reasoning.rs +++ b/src/llm/reasoning.rs @@ -337,6 +337,23 @@ impl TokenUsage { } } +/// Structured anomaly classification for LLM responses. +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub enum ResponseAnomaly { + /// Tool mode was requested, but the provider returned no usable tool calls + /// and no recoverable text content. + EmptyToolCompletion, + /// Text mode returned no usable content after cleaning/truncation. + EmptyTextResponse, +} + +/// Metadata attached to `RespondOutput` so callers can react to malformed +/// provider behavior without inferring it from fallback strings. +#[derive(Debug, Clone, Copy, Default, PartialEq, Eq)] +pub struct ResponseMetadata { + pub anomaly: Option, +} + /// Result of a response with potential tool calls. /// /// Used by the agent loop to handle tool execution before returning a final response. @@ -359,6 +376,7 @@ pub struct RespondOutput { pub result: RespondResult, pub usage: TokenUsage, pub finish_reason: FinishReason, + pub metadata: ResponseMetadata, } /// Reasoning engine for the agent. @@ -744,12 +762,11 @@ Respond in JSON format: }, usage, finish_reason: response.finish_reason, + metadata: ResponseMetadata::default(), }); } - let content = response - .content - .unwrap_or_else(|| "I'm not sure how to respond to that.".to_string()); + let content = response.content.unwrap_or_default(); // Some models (e.g. GLM-4.7) emit tool calls as XML tags in content // instead of using the structured tool_calls field. Try to recover @@ -772,6 +789,7 @@ Respond in JSON format: }, usage, finish_reason: response.finish_reason, + metadata: ResponseMetadata::default(), }); } @@ -785,11 +803,18 @@ Respond in JSON format: // Pre-truncate at tool tags to preserve text before the tag. let pre_truncated = truncate_at_tool_tags(&content); let cleaned = clean_response(&pre_truncated); - let final_text = if cleaned.trim().is_empty() { + let metadata = if cleaned.trim().is_empty() { tracing::warn!( "LLM response was empty after cleaning (original len={}), using fallback", content.len() ); + ResponseMetadata { + anomaly: Some(ResponseAnomaly::EmptyToolCompletion), + } + } else { + ResponseMetadata::default() + }; + let final_text = if metadata.anomaly.is_some() { "I'm not sure how to respond to that.".to_string() } else { cleaned @@ -798,6 +823,7 @@ Respond in JSON format: result: RespondResult::Text(final_text), usage, finish_reason: response.finish_reason, + metadata, }) } else { // No tools, use simple completion @@ -812,11 +838,18 @@ Respond in JSON format: let response = self.llm.complete(request).await?; let pre_truncated = truncate_at_tool_tags(&response.content); let cleaned = clean_response(&pre_truncated); - let final_text = if cleaned.trim().is_empty() { + let metadata = if cleaned.trim().is_empty() { tracing::warn!( "LLM response was empty after cleaning (original len={}), using fallback", response.content.len() ); + ResponseMetadata { + anomaly: Some(ResponseAnomaly::EmptyTextResponse), + } + } else { + ResponseMetadata::default() + }; + let final_text = if metadata.anomaly.is_some() { "I'm not sure how to respond to that.".to_string() } else { cleaned @@ -830,6 +863,7 @@ Respond in JSON format: cache_creation_input_tokens: response.cache_creation_input_tokens, }, finish_reason: response.finish_reason, + metadata, }) } } @@ -3159,9 +3193,104 @@ That's my plan."#; context.force_text = true; let output = reasoning.respond_with_tools(&context).await.unwrap(); + let metadata = output.metadata; + match output.result { + RespondResult::Text(text) => { + assert_eq!(text, "I'm not sure how to respond to that."); + assert_eq!(metadata.anomaly, Some(ResponseAnomaly::EmptyTextResponse)); + } + RespondResult::ToolCalls { .. } => { + panic!("Expected fallback text, not tool calls"); + } + } + } + + #[tokio::test] + async fn test_respond_with_tools_flags_empty_tool_completion() { + use crate::testing::StubLlm; + let llm = Arc::new(StubLlm::new("")); + let reasoning = Reasoning::new(llm); + + let context = ReasoningContext::new() + .with_message(ChatMessage::user("list tools")) + .with_tools(vec![ToolDefinition { + name: "tool_list".to_string(), + description: "Lists tools".to_string(), + parameters: serde_json::json!({}), + }]); + + let output = reasoning.respond_with_tools(&context).await.unwrap(); + let metadata = output.metadata; + match output.result { + RespondResult::Text(text) => { + assert_eq!(text, "I'm not sure how to respond to that."); + assert_eq!(metadata.anomaly, Some(ResponseAnomaly::EmptyToolCompletion)); + } + RespondResult::ToolCalls { .. } => { + panic!("Expected fallback text, not tool calls"); + } + } + } + + #[tokio::test] + async fn test_respond_with_tools_flags_empty_tool_completion_when_content_is_none() { + use crate::llm::{ + FinishReason, LlmProvider, ToolCompletionRequest, ToolCompletionResponse, + }; + use async_trait::async_trait; + use rust_decimal::Decimal; + + struct NoneContentToolLlm; + + #[async_trait] + impl LlmProvider for NoneContentToolLlm { + fn model_name(&self) -> &str { + "none-content-tool-llm" + } + + fn cost_per_token(&self) -> (Decimal, Decimal) { + (Decimal::ZERO, Decimal::ZERO) + } + + async fn complete( + &self, + _request: crate::llm::CompletionRequest, + ) -> Result { + unreachable!("tool-mode test should not call complete()") + } + + async fn complete_with_tools( + &self, + _request: ToolCompletionRequest, + ) -> Result { + Ok(ToolCompletionResponse { + content: None, + tool_calls: Vec::new(), + input_tokens: 0, + output_tokens: 0, + finish_reason: FinishReason::Stop, + cache_read_input_tokens: 0, + cache_creation_input_tokens: 0, + }) + } + } + + let reasoning = Reasoning::new(Arc::new(NoneContentToolLlm)); + + let context = ReasoningContext::new() + .with_message(ChatMessage::user("list tools")) + .with_tools(vec![ToolDefinition { + name: "tool_list".to_string(), + description: "Lists tools".to_string(), + parameters: serde_json::json!({}), + }]); + + let output = reasoning.respond_with_tools(&context).await.unwrap(); + let metadata = output.metadata; match output.result { RespondResult::Text(text) => { assert_eq!(text, "I'm not sure how to respond to that."); + assert_eq!(metadata.anomaly, Some(ResponseAnomaly::EmptyToolCompletion)); } RespondResult::ToolCalls { .. } => { panic!("Expected fallback text, not tool calls"); diff --git a/src/worker/autonomous_recovery.rs b/src/worker/autonomous_recovery.rs new file mode 100644 index 00000000000..4b4b4cc1ec8 --- /dev/null +++ b/src/worker/autonomous_recovery.rs @@ -0,0 +1,150 @@ +use crate::llm::{ResponseAnomaly, ResponseMetadata}; + +pub(crate) const EMPTY_TOOL_COMPLETION_NUDGE: &str = "\ +Your previous tool-enabled response was empty or malformed.\n\ +If you need to use a tool, call it now with valid arguments.\n\ +Otherwise, provide a real status update about work already completed."; + +pub(crate) const FORCE_TEXT_RECOVERY_PROMPT: &str = "\ +Your previous tool-enabled responses were empty or malformed.\n\ +Do not call any more tools in the next reply.\n\ +Instead, provide a concise final status based only on work already completed.\n\ +If the job is complete, say so explicitly. If not, explain what blocked you."; + +pub(crate) const EMPTY_TOOL_COMPLETION_FAILURE: &str = "the selected model repeatedly returned empty or malformed tool-completion responses and is not reliable for autonomous tool use."; + +#[derive(Debug, Default, Clone, Copy)] +pub(crate) struct AutonomousRecoveryState { + consecutive_empty_tool_completions: usize, + force_text_recovery_pending: bool, + force_text_recovery_active: bool, +} + +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub(crate) enum AutonomousRecoveryAction { + Continue, + ToolModeNudge, + ForceTextRecovery, + Fail, +} + +impl AutonomousRecoveryState { + pub(crate) fn begin_iteration(&mut self) -> bool { + if self.force_text_recovery_pending { + self.force_text_recovery_pending = false; + self.force_text_recovery_active = true; + true + } else { + self.force_text_recovery_active + } + } + + pub(crate) fn on_text_response( + &mut self, + metadata: ResponseMetadata, + text: &str, + ) -> AutonomousRecoveryAction { + match metadata.anomaly { + Some(ResponseAnomaly::EmptyToolCompletion) => { + self.consecutive_empty_tool_completions = + self.consecutive_empty_tool_completions.saturating_add(1); + self.force_text_recovery_active = false; + match self.consecutive_empty_tool_completions { + 1 => AutonomousRecoveryAction::ToolModeNudge, + 2 => { + self.force_text_recovery_pending = true; + AutonomousRecoveryAction::ForceTextRecovery + } + _ => AutonomousRecoveryAction::Fail, + } + } + Some(ResponseAnomaly::EmptyTextResponse) if self.force_text_recovery_active => { + self.force_text_recovery_active = false; + AutonomousRecoveryAction::Fail + } + _ if !text.trim().is_empty() => { + self.reset(); + AutonomousRecoveryAction::Continue + } + _ => AutonomousRecoveryAction::Continue, + } + } + + pub(crate) fn on_valid_tool_call(&mut self) { + self.reset(); + } + + fn reset(&mut self) { + self.consecutive_empty_tool_completions = 0; + self.force_text_recovery_pending = false; + self.force_text_recovery_active = false; + } +} + +#[cfg(test)] +mod tests { + use super::*; + + fn metadata(anomaly: ResponseAnomaly) -> ResponseMetadata { + ResponseMetadata { + anomaly: Some(anomaly), + } + } + + #[test] + fn first_empty_tool_completion_issues_nudge() { + let mut state = AutonomousRecoveryState::default(); + let action = state.on_text_response( + metadata(ResponseAnomaly::EmptyToolCompletion), + "I'm not sure how to respond to that.", + ); + assert_eq!(action, AutonomousRecoveryAction::ToolModeNudge); + assert!(!state.begin_iteration()); + } + + #[test] + fn second_empty_tool_completion_schedules_text_recovery() { + let mut state = AutonomousRecoveryState::default(); + let _ = state.on_text_response(metadata(ResponseAnomaly::EmptyToolCompletion), "fallback"); + let action = + state.on_text_response(metadata(ResponseAnomaly::EmptyToolCompletion), "fallback"); + assert_eq!(action, AutonomousRecoveryAction::ForceTextRecovery); + assert!(state.begin_iteration()); + } + + #[test] + fn forced_text_recovery_fallback_fails() { + let mut state = AutonomousRecoveryState::default(); + let _ = state.on_text_response(metadata(ResponseAnomaly::EmptyToolCompletion), "fallback"); + let _ = state.on_text_response(metadata(ResponseAnomaly::EmptyToolCompletion), "fallback"); + assert!(state.begin_iteration()); + let action = + state.on_text_response(metadata(ResponseAnomaly::EmptyTextResponse), "fallback"); + assert_eq!(action, AutonomousRecoveryAction::Fail); + } + + #[test] + fn valid_tool_call_resets_counter() { + let mut state = AutonomousRecoveryState::default(); + let _ = state.on_text_response(metadata(ResponseAnomaly::EmptyToolCompletion), "fallback"); + state.on_valid_tool_call(); + let action = + state.on_text_response(metadata(ResponseAnomaly::EmptyToolCompletion), "fallback"); + assert_eq!(action, AutonomousRecoveryAction::ToolModeNudge); + } + + #[test] + fn meaningful_text_after_text_recovery_resets_state() { + let mut state = AutonomousRecoveryState::default(); + let _ = state.on_text_response(metadata(ResponseAnomaly::EmptyToolCompletion), "fallback"); + let _ = state.on_text_response(metadata(ResponseAnomaly::EmptyToolCompletion), "fallback"); + assert!(state.begin_iteration()); + + let action = state.on_text_response(ResponseMetadata::default(), "Still working on step 2"); + assert_eq!(action, AutonomousRecoveryAction::Continue); + + let next = + state.on_text_response(metadata(ResponseAnomaly::EmptyToolCompletion), "fallback"); + assert_eq!(next, AutonomousRecoveryAction::ToolModeNudge); + } +} diff --git a/src/worker/container.rs b/src/worker/container.rs index 5d8e03b585f..efb27e45348 100644 --- a/src/worker/container.rs +++ b/src/worker/container.rs @@ -21,11 +21,15 @@ use crate::agent::agentic_loop::{ use crate::config::SafetyConfig; use crate::context::JobContext; use crate::error::WorkerError; -use crate::llm::{ChatMessage, LlmProvider, Reasoning, ReasoningContext}; +use crate::llm::{ChatMessage, LlmProvider, Reasoning, ReasoningContext, ResponseMetadata}; use crate::safety::SafetyLayer; use crate::tools::ToolRegistry; use crate::tools::execute::{execute_tool_simple, process_tool_result}; use crate::worker::api::{CompletionReport, JobEventPayload, StatusUpdate, WorkerHttpClient}; +use crate::worker::autonomous_recovery::{ + AutonomousRecoveryAction, AutonomousRecoveryState, EMPTY_TOOL_COMPLETION_FAILURE, + EMPTY_TOOL_COMPLETION_NUDGE, FORCE_TEXT_RECOVERY_PROMPT, +}; use crate::worker::proxy_llm::ProxyLlmProvider; /// Configuration for the worker runtime. @@ -170,6 +174,7 @@ Work independently to complete this job. When finished, your final message MUST extra_env: self.extra_env.clone(), last_output: Mutex::new(String::new()), iteration_tracker: iteration_tracker.clone(), + recovery_state: Mutex::new(AutonomousRecoveryState::default()), }; let config = AgenticLoopConfig { @@ -228,6 +233,24 @@ Work independently to complete this job. When finished, your final message MUST }) .await?; } + Ok(Ok(LoopOutcome::Failure(reason))) => { + tracing::warn!("Worker failed for job {}: {}", self.config.job_id, reason); + self.post_event( + "result", + serde_json::json!({ + "success": false, + "message": reason, + }), + ) + .await; + self.client + .report_complete(&CompletionReport { + success: false, + message: Some(reason), + iterations, + }) + .await?; + } Ok(Ok(LoopOutcome::Stopped | LoopOutcome::NeedApproval(_))) => { tracing::info!("Worker for job {} stopped", self.config.job_id); self.client @@ -304,6 +327,7 @@ struct ContainerDelegate { /// Tracks the current iteration — shared with the outer `run` method so /// `CompletionReport` can include accurate iteration counts. iteration_tracker: Arc>, + recovery_state: Mutex, } impl ContainerDelegate { @@ -377,8 +401,17 @@ impl LoopDelegate for ContainerDelegate { // conversation. Ensure the last message is user-role before calling the LLM. crate::util::ensure_ends_with_user_message(&mut reason_ctx.messages); - // Refresh tools (in case WASM tools were built) - reason_ctx.available_tools = self.tools.tool_definitions().await; + let force_text_recovery = { + let mut recovery = self.recovery_state.lock().await; + recovery.begin_iteration() + }; + if force_text_recovery { + tracing::warn!("Switching to text-only recovery after malformed tool completions"); + reason_ctx.available_tools.clear(); + } else { + // Refresh tools (in case WASM tools were built) + reason_ctx.available_tools = self.tools.tool_definitions().await; + } None } @@ -399,8 +432,53 @@ impl LoopDelegate for ContainerDelegate { async fn handle_text_response( &self, text: &str, + metadata: ResponseMetadata, reason_ctx: &mut ReasoningContext, ) -> TextAction { + let action = { + let mut recovery = self.recovery_state.lock().await; + recovery.on_text_response(metadata, text) + }; + match action { + AutonomousRecoveryAction::ToolModeNudge => { + tracing::warn!("Malformed empty tool completion detected; retrying in tool mode"); + self.post_event( + "status", + serde_json::json!({ + "message": "Model returned an empty tool-completion response; retrying with a stronger tool-use nudge.", + }), + ) + .await; + reason_ctx + .messages + .push(ChatMessage::user(EMPTY_TOOL_COMPLETION_NUDGE)); + return TextAction::Continue; + } + AutonomousRecoveryAction::ForceTextRecovery => { + tracing::warn!( + "Repeated malformed tool completions detected; switching to text-only recovery" + ); + self.post_event( + "status", + serde_json::json!({ + "message": "Model returned repeated empty tool-completion responses; requesting a final status update without tools.", + }), + ) + .await; + reason_ctx + .messages + .push(ChatMessage::user(FORCE_TEXT_RECOVERY_PROMPT)); + return TextAction::Continue; + } + AutonomousRecoveryAction::Fail => { + tracing::warn!("Failing fast after repeated malformed autonomous responses"); + return TextAction::Return(LoopOutcome::Failure( + EMPTY_TOOL_COMPLETION_FAILURE.to_string(), + )); + } + AutonomousRecoveryAction::Continue => {} + } + self.post_event( "message", serde_json::json!({ @@ -431,6 +509,11 @@ impl LoopDelegate for ContainerDelegate { content: Option, reason_ctx: &mut ReasoningContext, ) -> Result, crate::error::Error> { + { + let mut recovery = self.recovery_state.lock().await; + recovery.on_valid_tool_call(); + } + if let Some(ref text) = content { self.post_event( "message", diff --git a/src/worker/job.rs b/src/worker/job.rs index edf87bf8265..686192066a1 100644 --- a/src/worker/job.rs +++ b/src/worker/job.rs @@ -23,8 +23,8 @@ use crate::context::{ContextManager, JobState}; use crate::error::Error; use crate::hooks::HookRegistry; use crate::llm::{ - ActionPlan, ChatMessage, LlmProvider, Reasoning, ReasoningContext, RespondResult, ToolCall, - ToolSelection, + ActionPlan, ChatMessage, LlmProvider, Reasoning, ReasoningContext, RespondResult, + ResponseMetadata, ToolCall, ToolSelection, }; use crate::safety::SafetyLayer; use crate::tenant::AdminScope; @@ -33,6 +33,10 @@ use crate::tools::rate_limiter::RateLimitResult; use crate::tools::{ ApprovalContext, ToolRegistry, autonomous_unavailable_error, prepare_tool_params, redact_params, }; +use crate::worker::autonomous_recovery::{ + AutonomousRecoveryAction, AutonomousRecoveryState, EMPTY_TOOL_COMPLETION_FAILURE, + EMPTY_TOOL_COMPLETION_NUDGE, FORCE_TEXT_RECOVERY_PROMPT, +}; use ironclaw_common::AppEvent; /// Shared dependencies for worker execution. @@ -391,6 +395,7 @@ Report when the job is complete or if you encounter issues you cannot resolve."# worker: self, rx: tokio::sync::Mutex::new(rx), consecutive_rate_limits: std::sync::atomic::AtomicUsize::new(0), + recovery_state: tokio::sync::Mutex::new(AutonomousRecoveryState::default()), has_text_response: std::sync::atomic::AtomicBool::new(false), }; @@ -410,6 +415,9 @@ Report when the job is complete or if you encounter issues you cannot resolve."# self.mark_failed("Maximum iterations exceeded: job hit the iteration cap") .await?; } + LoopOutcome::Failure(reason) => { + self.mark_failed(&reason).await?; + } LoopOutcome::Stopped => { // Stop signal handled — nothing more to do } @@ -1130,6 +1138,7 @@ struct JobDelegate<'a> { rx: tokio::sync::Mutex<&'a mut mpsc::Receiver>, /// Tracks consecutive rate-limit errors to fail fast instead of burning iterations. consecutive_rate_limits: std::sync::atomic::AtomicUsize, + recovery_state: tokio::sync::Mutex, /// Whether a substantive (non-empty) text response has been produced. /// When true, an empty follow-up response is treated as job completion /// rather than a retry signal (prevents spurious failures in routines). @@ -1184,6 +1193,7 @@ impl<'a> JobDelegate<'a> { result: RespondResult::Text(String::new()), usage: crate::llm::TokenUsage::default(), finish_reason: crate::llm::FinishReason::Stop, + metadata: ResponseMetadata::default(), }) } @@ -1231,6 +1241,7 @@ impl<'a> JobDelegate<'a> { result: RespondResult::Text(String::new()), usage: crate::llm::TokenUsage::default(), finish_reason: crate::llm::FinishReason::Stop, + metadata: ResponseMetadata::default(), }) } } @@ -1322,8 +1333,21 @@ impl<'a> LoopDelegate for JobDelegate<'a> { reason_ctx: &mut ReasoningContext, _iteration: usize, ) -> Option { - // Refresh tool definitions so newly built tools become visible - reason_ctx.available_tools = self.worker.tools().tool_definitions().await; + let force_text_recovery = { + let mut recovery = self.recovery_state.lock().await; + recovery.begin_iteration() + }; + + if force_text_recovery { + tracing::warn!( + job_id = %self.worker.job_id, + "Switching to text-only recovery after malformed tool completions" + ); + reason_ctx.available_tools.clear(); + } else { + // Refresh tool definitions so newly built tools become visible + reason_ctx.available_tools = self.worker.tools().tool_definitions().await; + } // Claude 4.6 rejects assistant prefill; NEAR AI rejects any non-user-ending // conversation. Ensure the last message is user-role before calling the LLM. @@ -1357,6 +1381,7 @@ impl<'a> LoopDelegate for JobDelegate<'a> { }, usage: crate::llm::TokenUsage::default(), finish_reason: crate::llm::FinishReason::ToolUse, + metadata: ResponseMetadata::default(), }); } Ok(_) => {} // empty selections, fall through @@ -1410,8 +1435,59 @@ impl<'a> LoopDelegate for JobDelegate<'a> { async fn handle_text_response( &self, text: &str, + metadata: ResponseMetadata, reason_ctx: &mut ReasoningContext, ) -> TextAction { + let action = { + let mut recovery = self.recovery_state.lock().await; + recovery.on_text_response(metadata, text) + }; + + match action { + AutonomousRecoveryAction::ToolModeNudge => { + tracing::warn!( + job_id = %self.worker.job_id, + "Malformed empty tool completion detected; retrying in tool mode" + ); + self.worker.log_event( + "status", + serde_json::json!({ + "message": "Model returned an empty tool-completion response; retrying with a stronger tool-use nudge.", + }), + ); + reason_ctx + .messages + .push(ChatMessage::user(EMPTY_TOOL_COMPLETION_NUDGE)); + return TextAction::Continue; + } + AutonomousRecoveryAction::ForceTextRecovery => { + tracing::warn!( + job_id = %self.worker.job_id, + "Repeated malformed tool completions detected; switching to text-only recovery" + ); + self.worker.log_event( + "status", + serde_json::json!({ + "message": "Model returned repeated empty tool-completion responses; requesting a final status update without tools.", + }), + ); + reason_ctx + .messages + .push(ChatMessage::user(FORCE_TEXT_RECOVERY_PROMPT)); + return TextAction::Continue; + } + AutonomousRecoveryAction::Fail => { + tracing::warn!( + job_id = %self.worker.job_id, + "Failing fast after repeated malformed autonomous responses" + ); + return TextAction::Return(LoopOutcome::Failure( + EMPTY_TOOL_COMPLETION_FAILURE.to_string(), + )); + } + AutonomousRecoveryAction::Continue => {} + } + // Empty text after a substantive response means the LLM has finished. // Treat as successful completion rather than continuing the loop (which // would produce "Response contained no message or tool call (empty)"). @@ -1471,6 +1547,11 @@ impl<'a> LoopDelegate for JobDelegate<'a> { content: Option, reason_ctx: &mut ReasoningContext, ) -> Result, crate::error::Error> { + { + let mut recovery = self.recovery_state.lock().await; + recovery.on_valid_tool_call(); + } + // Strip suggestions from accompanying text (not useful in job context). let content = content.map(|c| crate::agent::strip_suggestions(&c)); @@ -2195,6 +2276,7 @@ mod tests { worker: &worker, rx: tokio::sync::Mutex::new(&mut rx), consecutive_rate_limits: std::sync::atomic::AtomicUsize::new(0), + recovery_state: tokio::sync::Mutex::new(AutonomousRecoveryState::default()), has_text_response: std::sync::atomic::AtomicBool::new(false), }; @@ -2204,6 +2286,7 @@ mod tests { let action = delegate .handle_text_response( "Weekly review created in Notion and notification sent.", + ResponseMetadata::default(), &mut reason_ctx, ) .await; diff --git a/src/worker/mod.rs b/src/worker/mod.rs index c6028b961e2..dc6a2e89817 100644 --- a/src/worker/mod.rs +++ b/src/worker/mod.rs @@ -25,6 +25,7 @@ //! ``` pub mod api; +mod autonomous_recovery; pub mod claude_bridge; pub mod container; pub mod job; diff --git a/tests/e2e_builtin_tool_coverage.rs b/tests/e2e_builtin_tool_coverage.rs index a6781d1ee3c..891a2ccbd5f 100644 --- a/tests/e2e_builtin_tool_coverage.rs +++ b/tests/e2e_builtin_tool_coverage.rs @@ -11,9 +11,106 @@ mod tests { use std::time::Duration; use ironclaw::agent::routine::{RoutineAction, Trigger}; + use ironclaw::context::{JobContext, JobState}; + use uuid::Uuid; + + use crate::support::test_rig::{TestRig, TestRigBuilder}; + use crate::support::trace_llm::{ + LlmTrace, RequestHint, TraceResponse, TraceStep, TraceToolCall, TraceTurn, + }; + + fn text_step(content: &str) -> TraceStep { + TraceStep { + request_hint: None, + response: TraceResponse::Text { + content: content.to_string(), + input_tokens: 10, + output_tokens: 5, + }, + expected_tool_results: Vec::new(), + } + } + + fn hinted_text_step(content: &str, last_user_message_contains: &str) -> TraceStep { + TraceStep { + request_hint: Some(RequestHint { + last_user_message_contains: Some(last_user_message_contains.to_string()), + min_message_count: None, + }), + response: TraceResponse::Text { + content: content.to_string(), + input_tokens: 10, + output_tokens: 5, + }, + expected_tool_results: Vec::new(), + } + } + + fn extract_job_id(response: &str) -> Option { + response + .split(|c: char| !(c.is_ascii_hexdigit() || c == '-')) + .find_map(|token| Uuid::parse_str(token).ok()) + } + + async fn resolve_created_job_id( + rig: &TestRig, + responses: &[ironclaw::channels::OutgoingResponse], + expected_title: &str, + ) -> Uuid { + if let Some(job_id) = responses + .iter() + .find_map(|response| extract_job_id(&response.content)) + { + return job_id; + } + + rig.database() + .list_agent_jobs_for_user("test-user") + .await + .expect("list_agent_jobs_for_user should succeed") + .into_iter() + .find(|job| job.title == expected_title) + .map(|job| job.id) + .unwrap_or_else(|| { + panic!( + "failed to resolve job id for title {expected_title:?}; responses were: {:?}", + responses + .iter() + .map(|response| &response.content) + .collect::>() + ) + }) + } + + async fn wait_for_job_state(rig: &TestRig, job_id: Uuid, expected: JobState) -> JobContext { + let deadline = tokio::time::Instant::now() + Duration::from_secs(10); + + loop { + if let Some(job) = rig + .database() + .get_job(job_id) + .await + .expect("get_job should succeed") + && job.state == expected + { + return job; + } + + assert!( + tokio::time::Instant::now() < deadline, + "job {job_id} did not reach state {expected:?} before timeout" + ); + + tokio::time::sleep(Duration::from_millis(50)).await; + } + } - use crate::support::test_rig::TestRigBuilder; - use crate::support::trace_llm::{LlmTrace, TraceResponse, TraceStep, TraceToolCall, TraceTurn}; + fn requests_contain(requests: &[Vec], needle: &str) -> bool { + requests + .iter() + .flatten() + .any(|message| message.content.contains(needle)) + } // ----------------------------------------------------------------------- // Test 1: time_parse_and_diff @@ -685,6 +782,159 @@ mod tests { rig.shutdown(); } + // ----------------------------------------------------------------------- + // Test 8a: command_job_fails_fast_on_repeated_empty_tool_completions + // ----------------------------------------------------------------------- + + #[tokio::test] + async fn command_job_fails_fast_on_repeated_empty_tool_completions() { + let trace = LlmTrace::single_turn( + "test-empty-tool-recovery-fail", + "(worker only)", + vec![ + text_step(""), + text_step(""), + hinted_text_step("", "valid arguments"), + text_step(""), + hinted_text_step("", "Do not call any more tools in the next reply."), + ], + ); + + let rig = TestRigBuilder::new() + .with_trace(trace) + .with_auto_approve_tools(true) + .build() + .await; + + rig.send_message("/job reproduce empty tool completion loop") + .await; + let create_responses = rig.wait_for_responses(1, Duration::from_secs(15)).await; + let job_id = resolve_created_job_id( + &rig, + &create_responses, + "reproduce empty tool completion loop", + ) + .await; + + let job = wait_for_job_state(&rig, job_id, JobState::Failed).await; + assert_eq!(job.title, "reproduce empty tool completion loop"); + + let failure_reason = rig + .database() + .get_agent_job_failure_reason(job_id) + .await + .expect("get_agent_job_failure_reason should succeed") + .expect("failed job should persist a failure reason"); + assert!( + failure_reason + .contains("repeatedly returned empty or malformed tool-completion responses"), + "unexpected failure reason: {failure_reason}" + ); + assert!( + !failure_reason.contains("max iterations"), + "failure should not surface as iteration exhaustion: {failure_reason}" + ); + + assert_eq!( + rig.llm_call_count(), + 5, + "worker should stop after the bounded recovery flow" + ); + assert!( + !rig.collect_metrics().await.hit_iteration_limit, + "bounded recovery should stop before iteration-limit reporting" + ); + + let requests = rig.captured_llm_requests(); + assert!( + requests_contain(&requests, "call it now with valid arguments"), + "expected targeted tool-mode recovery nudge in worker requests" + ); + assert!( + requests_contain(&requests, "Do not call any more tools in the next reply."), + "expected forced text-only recovery prompt in worker requests" + ); + + rig.clear().await; + rig.send_message(&format!("/status {}", job_id)).await; + let status_responses = rig.wait_for_responses(1, Duration::from_secs(5)).await; + assert!( + status_responses[0].content.contains("Status: Failed"), + "unexpected status response: {:?}", + status_responses[0].content + ); + + rig.shutdown(); + } + + // ----------------------------------------------------------------------- + // Test 8b: command_job_text_recovery_can_complete + // ----------------------------------------------------------------------- + + #[tokio::test] + async fn command_job_text_recovery_can_complete() { + let trace = LlmTrace::single_turn( + "test-empty-tool-recovery-success", + "(worker only)", + vec![ + text_step(""), + text_step(""), + hinted_text_step("", "valid arguments"), + text_step(""), + hinted_text_step( + "The job is complete. I finished the requested work and there is nothing left to do.", + "Do not call any more tools in the next reply.", + ), + ], + ); + + let rig = TestRigBuilder::new() + .with_trace(trace) + .with_auto_approve_tools(true) + .build() + .await; + + rig.send_message("/job recover after malformed tool completions") + .await; + let create_responses = rig.wait_for_responses(1, Duration::from_secs(15)).await; + let job_id = resolve_created_job_id( + &rig, + &create_responses, + "recover after empty tool completions", + ) + .await; + + let job = wait_for_job_state(&rig, job_id, JobState::Completed).await; + assert_eq!(job.title, "recover after malformed tool completions"); + + assert_eq!( + rig.llm_call_count(), + 5, + "worker should complete within the bounded recovery flow" + ); + + let requests = rig.captured_llm_requests(); + assert!( + requests_contain(&requests, "call it now with valid arguments"), + "expected targeted tool-mode recovery nudge in worker requests" + ); + assert!( + requests_contain(&requests, "Do not call any more tools in the next reply."), + "expected forced text-only recovery prompt in worker requests" + ); + + rig.clear().await; + rig.send_message(&format!("/status {}", job_id)).await; + let status_responses = rig.wait_for_responses(1, Duration::from_secs(5)).await; + assert!( + status_responses[0].content.contains("Status: Completed"), + "unexpected status response: {:?}", + status_responses[0].content + ); + + rig.shutdown(); + } + // ----------------------------------------------------------------------- // Test 9: job_list_cancel // ----------------------------------------------------------------------- From 368d2f523868cc06a3a84fc1255dd7411f680da6 Mon Sep 17 00:00:00 2001 From: synner88 Date: Mon, 30 Mar 2026 00:35:08 +0300 Subject: [PATCH 05/11] feat(gateway): OIDC JWT authentication for reverse-proxy deployments (#1463) MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit * feat(gateway): add OIDC JWT authentication for reverse-proxy deployments Add an optional OIDC JWT auth path to the web gateway, enabling deployments behind identity-aware proxies like AWS ALB with Okta/Cognito. When GATEWAY_OIDC_ENABLED=true, the gateway reads a signed JWT from a configurable HTTP header (default: x-amzn-oidc-data), fetches the signing key from a JWKS endpoint, and verifies the signature + claims. Auth flow: Bearer token → OIDC JWT → query-string token → 401. Key design decisions: - Split signature verification from claim extraction to handle AWS ALB's non-standard base64 padding (ALB includes '=' padding in JWT segments, but jsonwebtoken's decode() strips it, changing the signing input). We verify against the original token text, then extract claims from a normalized copy. - JWKS keys cached for 1 hour with per-kid granularity. - Supports both ALB-style per-key PEM URLs ({kid} placeholder) and standard JWKS endpoints. - DER-to-raw ECDSA signature conversion for IdPs that use DER encoding. - Frontend auto-detects proxy auth via /api/gateway/status probe, skipping the login screen when OIDC is active. Configuration (env vars): GATEWAY_OIDC_ENABLED=true GATEWAY_OIDC_HEADER=x-amzn-oidc-data (default) GATEWAY_OIDC_JWKS_URL=https://public-keys.auth.elb.us-east-1.amazonaws.com/{kid} GATEWAY_OIDC_ISSUER=https://example.okta.com (optional) GATEWAY_OIDC_AUDIENCE=my-client-id (optional) Co-Authored-By: Claude Opus 4.6 * Address code review feedback on OIDC auth PR - Add EdDSA PEM key parsing support (was falling through to RSA) - Fix issuer validation: remove set_issuer(&[]) else branch that rejected all tokens when GATEWAY_OIDC_ISSUER is unset - Make missing `sub` claim a validation error instead of silently defaulting to "unknown" - Extract initApp() in app.js so OIDC auto-auth actually initializes the UI (was calling undefined function) - Add regression tests for sub claim and issuer validation fixes Co-Authored-By: Claude Opus 4.6 * Harden OIDC auth: address claude[bot] security review - SSRF: URL-encode kid before substituting into JWKS URL template - Cache bounds: cap key cache at 64 entries, evict expired + oldest - DER parsing: support long-form length encoding (>= 128 bytes), validate component lengths against expected curve size - Production safety: replace .expect() with Result in OidcState::from_config - Fetch backoff: cache failed JWKS fetches for 10s to prevent retry storms - Body limit: cap JWKS responses at 256 KB to prevent OOM from rogue endpoint - Add regression tests for DER long-form, kid encoding, cache bounds Co-Authored-By: Claude Opus 4.6 * test(auth): add regression test for OIDC identity resolution Add two integration tests that exercise the full OIDC middleware path through to AuthenticatedUser extraction: - test_oidc_auth_inserts_user_identity_for_handler: sends a valid OIDC JWT through the middleware and verifies the handler receives the sub claim as user_id. Returns 401 if identity insertion is missing — verified by temporarily removing the insert and confirming failure. - test_oidc_auth_user_gets_member_role: confirms OIDC-authenticated users receive role=member (not admin). Uses a seed_key() test helper on OidcState to pre-populate the key cache with an HS256 secret, avoiding the need for an HTTP JWKS mock. Co-Authored-By: Claude Opus 4.6 (1M context) * test(auth): comprehensive OIDC test coverage for edge cases Add 17 new OIDC tests covering middleware integration, auth priority, invalid JWTs, issuer/audience validation, and key cache behavior: Middleware auth priority & fallthrough: - Bearer works when OIDC configured but header absent - Bearer takes priority when both Bearer and OIDC header present - Bad OIDC signature returns 401 (not 500) - Invalid OIDC doesn't block valid bearer auth - No auth at all with OIDC configured → 401 Expired / invalid JWT edge cases: - Expired JWT (exp in the past) rejected - JWT without kid header rejected - Malformed JWTs rejected (empty, 2-part, 4-part, garbage) - Non-string sub claim (integer) rejected - Empty-string sub passes auth (documented behavior) - Missing sub rejected through full middleware path Issuer / audience validation: - Matching issuer accepted, wrong issuer rejected - Matching audience accepted, wrong audience rejected - Missing iss/aud when configured: passes (jsonwebtoken v9 behavior, documented with notes on potential hardening) Key cache: - Expired cache entries not served - Fetch failure backoff blocks retry within 10s - Backoff expiry allows retry - Cache max entries constant verified Also adds shared test helpers (encode_test_jwt, test_oidc_state, oidc_auth_state, oidc_test_app) to reduce boilerplate. Co-Authored-By: Claude Opus 4.6 (1M context) * fix(ci): resolve formatting and no-panics check failures - Run cargo fmt to wrap long assert lines in OIDC tests - Add // safety: test helper comments to suppress false positives from check_no_panics.py (unwraps in #[cfg(test)] helper fns) Co-Authored-By: Claude Opus 4.6 (1M context) --------- Co-authored-by: synner88 <29090601+synner88@users.noreply.github.com> Co-authored-by: Claude Opus 4.6 Co-authored-by: ilblackdragon@gmail.com --- Cargo.lock | 353 +++--- Cargo.toml | 1 + src/channels/web/auth.rs | 1480 +++++++++++++++++++++- src/channels/web/mod.rs | 18 + src/channels/web/static/app.js | 136 +- src/channels/web/tests/no_silent_drop.rs | 1 + src/config/channels.rs | 41 + src/config/mod.rs | 3 +- src/tunnel/mod.rs | 2 + 9 files changed, 1807 insertions(+), 228 deletions(-) diff --git a/Cargo.lock b/Cargo.lock index 0e7d6521092..10013ffdde3 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -82,7 +82,7 @@ dependencies = [ "const-random", "once_cell", "version_check", - "zerocopy 0.8.42", + "zerocopy 0.8.48", ] [[package]] @@ -123,9 +123,9 @@ checksum = "4b46cbb362ab8752921c97e041f5e366ee6297bd428a31275b9fcf1e380f7299" [[package]] name = "anstream" -version = "0.6.21" +version = "1.0.0" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "43d5b281e737544384e969a5ccad3f1cdd24b48086a0fc1b2a5262a26b8f4f4a" +checksum = "824a212faf96e9acacdbd09febd34438f8f711fb84e09a8916013cd7815ca28d" dependencies = [ "anstyle", "anstyle-parse", @@ -138,15 +138,15 @@ dependencies = [ [[package]] name = "anstyle" -version = "1.0.13" +version = "1.0.14" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "5192cca8006f1fd4f7237516f40fa183bb07f8fbdfedaa0036de5ea9b0b45e78" +checksum = "940b3a0ca603d1eade50a4846a2afffd5ef57a9feac2c0e2ec2e14f9ead76000" [[package]] name = "anstyle-parse" -version = "0.2.7" +version = "1.0.0" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "4e7644824f0aa2c7b9384579234ef10eb7efb6a0deb83f9630a49594dd9c15c2" +checksum = "52ce7f38b242319f7cabaa6813055467063ecdc9d355bbb4ce0c68908cd8130e" dependencies = [ "utf8parse", ] @@ -157,7 +157,7 @@ version = "1.1.5" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "40c48f72fd53cd289104fc64099abca73db4166ad86ea0b4341abe65af83dadc" dependencies = [ - "windows-sys 0.60.2", + "windows-sys 0.61.2", ] [[package]] @@ -168,7 +168,7 @@ checksum = "291e6a250ff86cd4a820112fb8898808a366d8f9f58ce16d1f538353ad55747d" dependencies = [ "anstyle", "once_cell_polyfill", - "windows-sys 0.60.2", + "windows-sys 0.61.2", ] [[package]] @@ -442,9 +442,9 @@ dependencies = [ [[package]] name = "aws-lc-rs" -version = "1.16.1" +version = "1.16.2" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "94bffc006df10ac2a68c83692d734a465f8ee6c5b384d8545a636f81d858f4bf" +checksum = "a054912289d18629dc78375ba2c3726a3afe3ff71b4edba9dedfca0e3446d1fc" dependencies = [ "aws-lc-sys", "zeroize", @@ -452,9 +452,9 @@ dependencies = [ [[package]] name = "aws-lc-sys" -version = "0.38.0" +version = "0.39.1" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "4321e568ed89bb5a7d291a7f37997c2c0df89809d7b6d12062c81ddb54aa782e" +checksum = "83a25cf98105baa966497416dbd42565ce3a8cf8dbfd59803ec9ad46f3126399" dependencies = [ "cc", "cmake", @@ -490,9 +490,9 @@ dependencies = [ [[package]] name = "aws-sdk-bedrockruntime" -version = "1.127.0" +version = "1.128.0" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "7dcd5ccbed3bd50d342077d3f731de46d9608340386c87d07566c4c507891eda" +checksum = "3949d34a5c329ed83e7146d2fc1ffc06473fdc9bcbc5fa3d3534abeb950569c5" dependencies = [ "aws-credential-types", "aws-runtime", @@ -517,9 +517,9 @@ dependencies = [ [[package]] name = "aws-sdk-sso" -version = "1.96.0" +version = "1.97.0" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "f64a6eded248c6b453966e915d32aeddb48ea63ad17932682774eb026fbef5b1" +checksum = "9aadc669e184501caaa6beafb28c6267fc1baef0810fb58f9b205485ca3f2567" dependencies = [ "aws-credential-types", "aws-runtime", @@ -541,9 +541,9 @@ dependencies = [ [[package]] name = "aws-sdk-ssooidc" -version = "1.98.0" +version = "1.99.0" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "db96d720d3c622fcbe08bae1c4b04a72ce6257d8b0584cb5418da00ae20a344f" +checksum = "1342a7db8f358d3de0aed2007a0b54e875458e39848d54cc1d46700b2bfcb0a8" dependencies = [ "aws-credential-types", "aws-runtime", @@ -565,9 +565,9 @@ dependencies = [ [[package]] name = "aws-sdk-sts" -version = "1.100.0" +version = "1.101.0" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "fafbdda43b93f57f699c5dfe8328db590b967b8a820a13ccdd6687355dfcc7ca" +checksum = "ab41ad64e4051ecabeea802d6a17845a91e83287e1dd249e6963ea1ba78c428a" dependencies = [ "aws-credential-types", "aws-runtime", @@ -757,9 +757,9 @@ dependencies = [ [[package]] name = "aws-smithy-types" -version = "1.4.6" +version = "1.4.7" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "d2b1117b3b2bbe166d11199b540ceed0d0f7676e36e7b962b5a437a9971eac75" +checksum = "9d73dbfbaa8e4bc57b9045137680b958d274823509a360abfd8e1d514d40c95c" dependencies = [ "base64-simd", "bytes", @@ -1085,19 +1085,20 @@ dependencies = [ [[package]] name = "borsh" -version = "1.6.0" +version = "1.6.1" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "d1da5ab77c1437701eeff7c88d968729e7766172279eab0676857b3d63af7a6f" +checksum = "cfd1e3f8955a5d7de9fab72fc8373fade9fb8a703968cb200ae3dc6cf08e185a" dependencies = [ "borsh-derive", + "bytes", "cfg_aliases", ] [[package]] name = "borsh-derive" -version = "1.6.0" +version = "1.6.1" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "0686c856aa6aac0c4498f936d7d6a02df690f614c03e4d906d1018062b5c5e2c" +checksum = "bfcfdc083699101d5a7965e49925975f2f55060f94f9a05e7187be95d530ca59" dependencies = [ "once_cell", "proc-macro-crate", @@ -1257,9 +1258,9 @@ dependencies = [ [[package]] name = "cc" -version = "1.2.56" +version = "1.2.58" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "aebf35691d1bfb0ac386a69bac2fde4dd276fb618cf8bf4f5318fe285e821bb2" +checksum = "e1e928d4b69e3077709075a938a05ffbedfa53a84c8f766efbf8220bb1ff60e1" dependencies = [ "find-msvc-tools", "jobserver", @@ -1362,9 +1363,9 @@ dependencies = [ [[package]] name = "clap" -version = "4.5.60" +version = "4.6.0" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "2797f34da339ce31042b27d23607e051786132987f595b02ba4f6a6dffb7030a" +checksum = "b193af5b67834b676abd72466a96c1024e6a6ad978a1f484bd90b85c94041351" dependencies = [ "clap_builder", "clap_derive", @@ -1372,9 +1373,9 @@ dependencies = [ [[package]] name = "clap_builder" -version = "4.5.60" +version = "4.6.0" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "24a241312cea5059b13574bb9b3861cabf758b879c15190b37b6d6fd63ab6876" +checksum = "714a53001bf66416adb0e2ef5ac857140e7dc3a0c48fb28b2f10762fc4b5069f" dependencies = [ "anstream", "anstyle", @@ -1384,18 +1385,18 @@ dependencies = [ [[package]] name = "clap_complete" -version = "4.5.66" +version = "4.6.0" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "c757a3b7e39161a4e56f9365141ada2a6c915a8622c408ab6bb4b5d047371031" +checksum = "19c9f1dde76b736e3681f28cec9d5a61299cbaae0fce80a68e43724ad56031eb" dependencies = [ "clap", ] [[package]] name = "clap_derive" -version = "4.5.55" +version = "4.6.0" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "a92793da1a46a5f2a02a6f4c46c6496b28c43638adea8306fcb0caa1634f24e5" +checksum = "1110bd8a634a1ab8cb04345d8d878267d57c3cf1b38d91b71af6686408bbca6a" dependencies = [ "heck", "proc-macro2", @@ -1405,9 +1406,9 @@ dependencies = [ [[package]] name = "clap_lex" -version = "1.0.0" +version = "1.1.0" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "3a822ea5bc7590f9d40f1ba12c0dc3c2760f3482c6984db1573ad11031420831" +checksum = "c8d4a3bb8b1e0c1050499d1815f5ab16d04f0959b233085fb31653fbfc9d98f9" [[package]] name = "clipboard-win" @@ -1420,9 +1421,9 @@ dependencies = [ [[package]] name = "cmake" -version = "0.1.57" +version = "0.1.58" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "75443c44cd6b379beb8c5b45d85d0773baf31cce901fe7bb252f4eff3008ef7d" +checksum = "c0f78a02292a74a88ac736019ab962ece0bc380e3f977bf72e376c5d78ff0678" dependencies = [ "cc", ] @@ -1438,9 +1439,9 @@ dependencies = [ [[package]] name = "colorchoice" -version = "1.0.4" +version = "1.0.5" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "b05b61dc5112cbb17e4b6cd61790d9845d13888356391624cbe7e41efeac1e75" +checksum = "1d07550c9036bf2ae0c684c4297d503f838287c83c53686d05370d0e139ae570" [[package]] name = "concurrent-queue" @@ -1453,14 +1454,13 @@ dependencies = [ [[package]] name = "console" -version = "0.15.11" +version = "0.16.3" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "054ccb5b10f9f2cbf51eb355ca1d05c2d279ce1804688d0db74b4733a5aeafd8" +checksum = "d64e8af5551369d19cf50138de61f1c42074ab970f74e99be916646777f8fc87" dependencies = [ "encode_unicode", "libc", - "once_cell", - "windows-sys 0.59.0", + "windows-sys 0.61.2", ] [[package]] @@ -1594,7 +1594,7 @@ dependencies = [ "hashbrown 0.14.5", "log", "regalloc2", - "rustc-hash 2.1.1", + "rustc-hash 2.1.2", "serde", "smallvec", "target-lexicon", @@ -1922,9 +1922,9 @@ dependencies = [ [[package]] name = "darling" -version = "0.21.3" +version = "0.23.0" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "9cdf337090841a411e2a7f3deb9187445851f91b309c0c0a29e05f74a00a48c0" +checksum = "25ae13da2f202d56bd7f91c25fba009e7717a1e4a1cc98a76d844b65ae912e9d" dependencies = [ "darling_core", "darling_macro", @@ -1932,11 +1932,10 @@ dependencies = [ [[package]] name = "darling_core" -version = "0.21.3" +version = "0.23.0" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "1247195ecd7e3c85f83c8d2a366e4210d588e802133e1e355180a9870b517ea4" +checksum = "9865a50f7c335f53564bb694ef660825eb8610e0a53d3e11bf1b0d3df31e03b0" dependencies = [ - "fnv", "ident_case", "proc-macro2", "quote", @@ -1946,9 +1945,9 @@ dependencies = [ [[package]] name = "darling_macro" -version = "0.21.3" +version = "0.23.0" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "d38308df82d1080de0afee5d069fa14b0326a88c14f15c5ccda35b4a6c414c81" +checksum = "ac3984ec7bd6cfa798e62b4a642426a5be0e68f9401cfc2a01e3fa9ea2fcdb8d" dependencies = [ "darling_core", "quote", @@ -2136,7 +2135,7 @@ dependencies = [ "libc", "option-ext", "redox_users 0.5.2", - "windows-sys 0.59.0", + "windows-sys 0.61.2", ] [[package]] @@ -2323,7 +2322,7 @@ source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "39cab71617ae0d63f51a36d69f866391735b51691dbda63cf6f96d042b63efeb" dependencies = [ "libc", - "windows-sys 0.59.0", + "windows-sys 0.61.2", ] [[package]] @@ -2789,7 +2788,7 @@ checksum = "6ea2d84b969582b4b1864a92dc5d27cd2b77b622a8d79306834f1be5ba20d84b" dependencies = [ "cfg-if", "crunchy", - "zerocopy 0.8.42", + "zerocopy 0.8.48", ] [[package]] @@ -2898,15 +2897,15 @@ dependencies = [ [[package]] name = "html-to-markdown-rs" -version = "2.28.2" +version = "2.30.0" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "3f9377e16af590b764fd98fd176027cf8831c5335f8964f3f643753e38913a4e" +checksum = "7ea41945a2fd834381642a000ef75b03f0030f3023f3dd3291fc5c372d3dda33" dependencies = [ "ahash 0.8.12", "astral-tl", "base64 0.22.1", "html-escape", - "html5ever 0.38.0", + "html5ever 0.39.0", "lru", "once_cell", "regex", @@ -2935,6 +2934,16 @@ dependencies = [ "markup5ever 0.38.0", ] +[[package]] +name = "html5ever" +version = "0.39.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "46a1761807faccc9a19e86944bbf40610014066306f96edcdedc2fb714bcb7b8" +dependencies = [ + "log", + "markup5ever 0.39.0", +] + [[package]] name = "http" version = "0.2.12" @@ -3346,9 +3355,9 @@ dependencies = [ [[package]] name = "insta" -version = "1.46.3" +version = "1.47.1" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "e82db8c87c7f1ccecb34ce0c24399b8a73081427f3c7c50a5d597925356115e4" +checksum = "99322078b2c076829a1db959d49da554fabc4342257fc0ba5a070a1eb3a01cd8" dependencies = [ "console", "once_cell", @@ -3380,9 +3389,9 @@ checksum = "d98f6fed1fde3f8c21bc40a1abb88dd75e67924f9cffc3ef95607bad8017f8e2" [[package]] name = "iri-string" -version = "0.7.10" +version = "0.7.11" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "c91338f0783edbd6195decb37bae672fd3b165faffb89bf7b9e6942f8b1a731a" +checksum = "d8e7418f59cc01c88316161279a7f665217ae316b388e58a0d10e29f54f1e5eb" dependencies = [ "memchr", "serde", @@ -3431,6 +3440,7 @@ dependencies = [ "ironclaw_common", "ironclaw_safety", "json5", + "jsonwebtoken", "libsql", "lru", "mime_guess", @@ -3564,9 +3574,9 @@ dependencies = [ [[package]] name = "itoa" -version = "1.0.17" +version = "1.0.18" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "92ecc6618181def0457392ccd0ee51198e065e016d1d527a7ac1b6dc7c1f09d2" +checksum = "8f42a60cbdf9a97f5d2305f08a87dc4e09308d1276d28c869c684d7777685682" [[package]] name = "ittapi" @@ -3600,10 +3610,12 @@ dependencies = [ [[package]] name = "js-sys" -version = "0.3.91" +version = "0.3.92" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "b49715b7073f385ba4bc528e5747d02e66cb39c6146efb66b781f131f0fb399c" +checksum = "cc4c90f45aa2e6eacbe8645f77fdea542ac97a494bcd117a67df9ff4d611f995" dependencies = [ + "cfg-if", + "futures-util", "once_cell", "wasm-bindgen", ] @@ -3619,6 +3631,21 @@ dependencies = [ "serde", ] +[[package]] +name = "jsonwebtoken" +version = "9.3.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "5a87cc7a48537badeae96744432de36f4be2b4a34a05a5ef32e9dd8a1c169dde" +dependencies = [ + "base64 0.22.1", + "js-sys", + "pem", + "ring", + "serde", + "serde_json", + "simple_asn1", +] + [[package]] name = "kuchikikiki" version = "0.9.2" @@ -3705,9 +3732,9 @@ checksum = "b6d2cec3eae94f9f509c767b45932f1ada8350c4bdb85af2fcab4a3c14807981" [[package]] name = "libredox" -version = "0.1.14" +version = "0.1.15" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "1744e39d1d6a9948f4f388969627434e31128196de472883b39f148769bfe30a" +checksum = "7ddbf48fd451246b1f8c2610bd3b4ac0cc6e149d89832867093ab69a17194f08" dependencies = [ "bitflags 2.11.0", "libc", @@ -3966,6 +3993,17 @@ dependencies = [ "web_atoms", ] +[[package]] +name = "markup5ever" +version = "0.39.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "7122d987ec5f704ee56f6e5b41a7d93722e9aae27ae07cafa4036c4d3f9757de" +dependencies = [ + "log", + "tendril 0.5.0", + "web_atoms", +] + [[package]] name = "matchers" version = "0.2.0" @@ -4070,9 +4108,9 @@ dependencies = [ [[package]] name = "mio" -version = "1.1.1" +version = "1.2.0" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "a69bcab0ad47271a0234d9422b131806bf3968021e5dc9328caf2d4cd58557fc" +checksum = "50b7e5b27aa02a74bac8c3f23f448f8d87ff11f92d3aac1a6ed369ee08cc56c1" dependencies = [ "libc", "log", @@ -4145,7 +4183,7 @@ version = "0.50.3" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "7957b9740744892f114936ab4a57b3f487491bbeafaf8083688b16841a4240e5" dependencies = [ - "windows-sys 0.59.0", + "windows-sys 0.61.2", ] [[package]] @@ -4183,9 +4221,9 @@ dependencies = [ [[package]] name = "num-conv" -version = "0.2.0" +version = "0.2.1" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "cf97ec579c3c42f953ef76dbf8d55ac91fb219dde70e49aa4a6b7d74e9919050" +checksum = "c6673768db2d862beb9b39a78fdcb1a69439615d5794a1be50caa9bc92c81967" [[package]] name = "num-integer" @@ -4278,9 +4316,9 @@ dependencies = [ [[package]] name = "once_cell" -version = "1.21.3" +version = "1.21.4" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "42f5e15c9953c5e4ccceeb2e7382a716482c34515315f7b03532b8b4e8393d2d" +checksum = "9f7c3e4beb33f85d45ae3e3a1792185706c8e16d043238c593331cc7cd313b50" [[package]] name = "once_cell_polyfill" @@ -4331,9 +4369,9 @@ checksum = "04744f49eae99ab78e0d5c0b603ab218f515ea8cfe5a456d7629ad883a3b6e7d" [[package]] name = "ordered-float" -version = "5.1.0" +version = "5.3.0" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "7f4779c6901a562440c3786d08192c6fbda7c1c2060edd10006b05ee35d10f2d" +checksum = "b7d950ca161dc355eaf28f82b11345ed76c6e1f6eb1f4f4479e0323b9e2fbd0e" dependencies = [ "num-traits", ] @@ -4441,6 +4479,16 @@ version = "0.1.2" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "19b17cddbe7ec3f8bc800887bab5e717348c95ea2ca0b1bf0837fb964dc67099" +[[package]] +name = "pem" +version = "3.0.6" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "1d30c53c26bc5b31a98cd02d20f25a7c8567146caf63ed593a9d87b2775291be" +dependencies = [ + "base64 0.22.1", + "serde_core", +] + [[package]] name = "percent-encoding" version = "2.3.2" @@ -4807,7 +4855,7 @@ version = "0.2.21" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "85eae3c4ed2f50dcfe72643da4befc30deadb458a9b590d720cde2f2b1e97da9" dependencies = [ - "zerocopy 0.8.42", + "zerocopy 0.8.48", ] [[package]] @@ -4842,7 +4890,7 @@ version = "3.5.0" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "e67ba7e9b2b56446f1d419b1d807906278ffa1a658a8a5d8a39dcb1f5a78614f" dependencies = [ - "toml_edit 0.25.4+spec-1.1.0", + "toml_edit 0.25.8+spec-1.1.0", ] [[package]] @@ -4939,7 +4987,7 @@ dependencies = [ "pin-project-lite", "quinn-proto", "quinn-udp", - "rustc-hash 2.1.1", + "rustc-hash 2.1.2", "rustls 0.23.37", "socket2 0.6.3", "thiserror 2.0.18", @@ -4959,7 +5007,7 @@ dependencies = [ "lru-slab", "rand 0.9.2", "ring", - "rustc-hash 2.1.1", + "rustc-hash 2.1.2", "rustls 0.23.37", "rustls-pki-types", "slab", @@ -5247,7 +5295,7 @@ dependencies = [ "bumpalo", "hashbrown 0.15.5", "log", - "rustc-hash 2.1.1", + "rustc-hash 2.1.2", "smallvec", ] @@ -5418,9 +5466,9 @@ dependencies = [ [[package]] name = "rust_decimal" -version = "1.40.0" +version = "1.41.0" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "61f703d19852dbf87cbc513643fa81428361eb6940f1ac14fd58155d295a3eb0" +checksum = "2ce901f9a19d251159075a4c37af514c3b8ef99c22e02dd8c19161cf397ee94a" dependencies = [ "arrayvec", "borsh", @@ -5431,6 +5479,7 @@ dependencies = [ "rkyv", "serde", "serde_json", + "wasm-bindgen", ] [[package]] @@ -5457,9 +5506,9 @@ checksum = "08d43f7aa6b08d49f382cde6a7982047c3426db949b1424bc4b7ec9ae12c6ce2" [[package]] name = "rustc-hash" -version = "2.1.1" +version = "2.1.2" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "357703d41365b4b27c590e3ed91eabb1b663f07c4c084095e60cbed4362dff0d" +checksum = "94300abf3f1ae2e2b8ffb7b58043de3d399c73fa6f4b73826402a5c457614dbe" [[package]] name = "rustc_version" @@ -5493,7 +5542,7 @@ dependencies = [ "errno", "libc", "linux-raw-sys 0.12.1", - "windows-sys 0.59.0", + "windows-sys 0.61.2", ] [[package]] @@ -5841,7 +5890,7 @@ dependencies = [ "phf 0.13.1", "phf_codegen 0.13.1", "precomputed-hash", - "rustc-hash 2.1.1", + "rustc-hash 2.1.2", "servo_arc", "smallvec", ] @@ -5860,7 +5909,7 @@ dependencies = [ "phf 0.13.1", "phf_codegen 0.13.1", "precomputed-hash", - "rustc-hash 2.1.1", + "rustc-hash 2.1.2", "servo_arc", "smallvec", ] @@ -5974,9 +6023,9 @@ dependencies = [ [[package]] name = "serde_with" -version = "3.17.0" +version = "3.18.0" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "381b283ce7bc6b476d903296fb59d0d36633652b633b27f64db4fb46dcbfc3b9" +checksum = "dd5414fad8e6907dbdd5bc441a50ae8d6e26151a03b1de04d89a5576de61d01f" dependencies = [ "base64 0.22.1", "chrono", @@ -5993,9 +6042,9 @@ dependencies = [ [[package]] name = "serde_with_macros" -version = "3.17.0" +version = "3.18.0" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "a6d4e30573c8cb306ed6ab1dca8423eec9a463ea0e155f45399455e0368b27e0" +checksum = "d3db8978e608f1fe7357e211969fd9abdcae80bac1ba7a3369bb7eb6b404eb65" dependencies = [ "darling", "proc-macro2", @@ -6121,9 +6170,9 @@ dependencies = [ [[package]] name = "simd-adler32" -version = "0.3.8" +version = "0.3.9" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "e320a6c5ad31d271ad523dcf3ad13e2767ad8b1cb8f047f75a8aeaf8da139da2" +checksum = "703d5c7ef118737c72f1af64ad2f6f8c5e1921f818cdcb97b8fe6fc69bf66214" [[package]] name = "simdutf8" @@ -6137,6 +6186,18 @@ version = "2.7.0" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "bbbb5d9659141646ae647b42fe094daf6c6192d1620870b449d9557f748b2daa" +[[package]] +name = "simple_asn1" +version = "0.6.4" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "0d585997b0ac10be3c5ee635f1bab02d512760d14b7c468801ac8a01d9ae5f1d" +dependencies = [ + "num-bigint", + "num-traits", + "thiserror 2.0.18", + "time", +] + [[package]] name = "siphasher" version = "1.0.2" @@ -6175,7 +6236,7 @@ source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "3a766e1110788c36f4fa1c2b71b387a7815aa65f88ce0229841826633d93723e" dependencies = [ "libc", - "windows-sys 0.60.2", + "windows-sys 0.61.2", ] [[package]] @@ -6400,7 +6461,7 @@ dependencies = [ "getrandom 0.4.2", "once_cell", "rustix 1.1.4", - "windows-sys 0.59.0", + "windows-sys 0.61.2", ] [[package]] @@ -6598,9 +6659,9 @@ dependencies = [ [[package]] name = "tinyvec" -version = "1.10.0" +version = "1.11.0" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "bfa5fdc3bce6191a1dbc8c02d5c8bffcf557bafa17c124c5264a458f1b0613fa" +checksum = "3e61e67053d25a4e82c844e8424039d9745781b3fc4f32b8d55ed50f5f667ef3" dependencies = [ "tinyvec_macros", ] @@ -6845,9 +6906,9 @@ dependencies = [ [[package]] name = "toml_datetime" -version = "1.0.0+spec-1.1.0" +version = "1.1.0+spec-1.1.0" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "32c2555c699578a4f59f0cc68e5116c8d7cabbd45e1409b989d4be085b53f13e" +checksum = "97251a7c317e03ad83774a8752a7e81fb6067740609f75ea2b585b569a59198f" dependencies = [ "serde_core", ] @@ -6863,28 +6924,28 @@ dependencies = [ "serde_spanned", "toml_datetime 0.6.11", "toml_write", - "winnow", + "winnow 0.7.15", ] [[package]] name = "toml_edit" -version = "0.25.4+spec-1.1.0" +version = "0.25.8+spec-1.1.0" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "7193cbd0ce53dc966037f54351dbbcf0d5a642c7f0038c382ef9e677ce8c13f2" +checksum = "16bff38f1d86c47f9ff0647e6838d7bb362522bdf44006c7068c2b1e606f1f3c" dependencies = [ "indexmap 2.13.0", - "toml_datetime 1.0.0+spec-1.1.0", + "toml_datetime 1.1.0+spec-1.1.0", "toml_parser", - "winnow", + "winnow 1.0.0", ] [[package]] name = "toml_parser" -version = "1.0.9+spec-1.1.0" +version = "1.1.0+spec-1.1.0" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "702d4415e08923e7e1ef96cd5727c0dfed80b4d2fa25db9647fe5eb6f7c5a4c4" +checksum = "2334f11ee363607eb04df9b8fc8a13ca1715a72ba8662a26ac285c98aabb4011" dependencies = [ - "winnow", + "winnow 1.0.0", ] [[package]] @@ -7096,9 +7157,9 @@ dependencies = [ [[package]] name = "tracing-subscriber" -version = "0.3.22" +version = "0.3.23" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "2f30143827ddab0d256fd843b7a66d164e9f271cfa0dde49142c5ca0ca291f1e" +checksum = "cb7f578e5945fb242538965c2d0b04418d38ec25c79d160cd279bf0731c8d319" dependencies = [ "matchers", "nu-ansi-term", @@ -7180,9 +7241,9 @@ dependencies = [ [[package]] name = "type1-encoding-parser" -version = "0.1.0" +version = "0.1.1" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "d3d6cc09e1a99c7e01f2afe4953789311a1c50baebbdac5b477ecf78e2e92a5b" +checksum = "fa10c302f5a53b7ad27fd42a3996e23d096ba39b5b8dd6d9e683a05b01bee749" dependencies = [ "pom", ] @@ -7207,7 +7268,7 @@ checksum = "f2f6fb2847f6742cd76af783a2a2c49e9375d0a111c7bef6f71cd9e738c72d6e" dependencies = [ "memoffset", "tempfile", - "windows-sys 0.60.2", + "windows-sys 0.61.2", ] [[package]] @@ -7254,9 +7315,9 @@ checksum = "7df058c713841ad818f1dc5d3fd88063241cc61f49f5fbea4b951e8cf5a8d71d" [[package]] name = "unicode-segmentation" -version = "1.12.0" +version = "1.13.2" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "f6ccf251212114b54433ec949fd6a7841275f9ada20dddd2f29e9ceea4501493" +checksum = "9629274872b2bfaf8d66f5f15725007f635594914870f65218920345aa11aa8c" [[package]] name = "unicode-width" @@ -7337,9 +7398,9 @@ checksum = "06abde3611657adf66d383f00b093d7faecc7fa57071cce2578660c9f1010821" [[package]] name = "uuid" -version = "1.22.0" +version = "1.23.0" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "a68d3c8f01c0cfa54a75291d83601161799e4a89a39e0929f4b0354d88757a37" +checksum = "5ac8b6f42ead25368cf5b098aeb3dc8a1a2c05a3eee8a9a1a68c640edbfc79d9" dependencies = [ "getrandom 0.4.2", "js-sys", @@ -7435,36 +7496,33 @@ dependencies = [ [[package]] name = "wasm-bindgen" -version = "0.2.114" +version = "0.2.115" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "6532f9a5c1ece3798cb1c2cfdba640b9b3ba884f5db45973a6f442510a87d38e" +checksum = "6523d69017b7633e396a89c5efab138161ed5aafcbc8d3e5c5a42ae38f50495a" dependencies = [ "cfg-if", "once_cell", "rustversion", + "serde", "wasm-bindgen-macro", "wasm-bindgen-shared", ] [[package]] name = "wasm-bindgen-futures" -version = "0.4.64" +version = "0.4.65" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "e9c5522b3a28661442748e09d40924dfb9ca614b21c00d3fd135720e48b67db8" +checksum = "2d1faf851e778dfa54db7cd438b70758eba9755cb47403f3496edd7c8fc212f0" dependencies = [ - "cfg-if", - "futures-util", "js-sys", - "once_cell", "wasm-bindgen", - "web-sys", ] [[package]] name = "wasm-bindgen-macro" -version = "0.2.114" +version = "0.2.115" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "18a2d50fcf105fb33bb15f00e7a77b772945a2ee45dcf454961fd843e74c18e6" +checksum = "4e3a6c758eb2f701ed3d052ff5737f5bfe6614326ea7f3bbac7156192dc32e67" dependencies = [ "quote", "wasm-bindgen-macro-support", @@ -7472,9 +7530,9 @@ dependencies = [ [[package]] name = "wasm-bindgen-macro-support" -version = "0.2.114" +version = "0.2.115" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "03ce4caeaac547cdf713d280eda22a730824dd11e6b8c3ca9e42247b25c631e3" +checksum = "921de2737904886b52bcbb237301552d05969a6f9c40d261eb0533c8b055fedf" dependencies = [ "bumpalo", "proc-macro2", @@ -7485,9 +7543,9 @@ dependencies = [ [[package]] name = "wasm-bindgen-shared" -version = "0.2.114" +version = "0.2.115" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "75a326b8c223ee17883a4251907455a2431acc2791c98c26279376490c378c16" +checksum = "a93e946af942b58934c604527337bad9ae33ba1d5c6900bbb41c2c07c2364a93" dependencies = [ "unicode-ident", ] @@ -7914,9 +7972,9 @@ dependencies = [ [[package]] name = "web-sys" -version = "0.3.91" +version = "0.3.92" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "854ba17bb104abfb26ba36da9729addc7ce7f06f5c0f90f3c391f8461cca21f9" +checksum = "84cde8507f4d7cfcb1185b8cb5890c494ffea65edbe1ba82cfd63661c805ed94" dependencies = [ "js-sys", "wasm-bindgen", @@ -8057,7 +8115,7 @@ version = "0.1.11" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "c2a7b1c03c876122aa43f3020e6c3c3ee5c05081c9a00739faf7503aeba10d22" dependencies = [ - "windows-sys 0.48.0", + "windows-sys 0.61.2", ] [[package]] @@ -8393,6 +8451,15 @@ dependencies = [ "memchr", ] +[[package]] +name = "winnow" +version = "1.0.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "a90e88e4667264a994d34e6d1ab2d26d398dcdca8b7f52bec8668957517fc7d8" +dependencies = [ + "memchr", +] + [[package]] name = "winx" version = "0.36.4" @@ -8678,11 +8745,11 @@ dependencies = [ [[package]] name = "zerocopy" -version = "0.8.42" +version = "0.8.48" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "f2578b716f8a7a858b7f02d5bd870c14bf4ddbbcf3a4c05414ba6503640505e3" +checksum = "eed437bf9d6692032087e337407a86f04cd8d6a16a37199ed57949d415bd68e9" dependencies = [ - "zerocopy-derive 0.8.42", + "zerocopy-derive 0.8.48", ] [[package]] @@ -8698,9 +8765,9 @@ dependencies = [ [[package]] name = "zerocopy-derive" -version = "0.8.42" +version = "0.8.48" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "7e6cc098ea4d3bd6246687de65af3f920c430e236bee1e3bf2e441463f08a02f" +checksum = "70e3cd084b1788766f53af483dd21f93881ff30d7320490ec3ef7526d203bad4" dependencies = [ "proc-macro2", "quote", diff --git a/Cargo.toml b/Cargo.toml index b62f102696b..6eff18b155e 100644 --- a/Cargo.toml +++ b/Cargo.toml @@ -173,6 +173,7 @@ hyper-util = { version = "0.1", features = ["server", "tokio", "http1", "http2"] http-body-util = "0.1" bytes = "1" base64 = "0.22.1" +jsonwebtoken = "9" mime_guess = "2.0.5" clap_complete = "4.5.0" lru = "0.16.3" diff --git a/src/channels/web/auth.rs b/src/channels/web/auth.rs index acd24458dc6..482f9604f71 100644 --- a/src/channels/web/auth.rs +++ b/src/channels/web/auth.rs @@ -1,11 +1,45 @@ -//! Bearer token authentication middleware for the web gateway. +//! Authentication middleware for the web gateway. //! +//! Supports three auth mechanisms, tried in order: +//! +//! ```text +//! Request +//! │ +//! ▼ +//! ┌─────────────────────────────┐ +//! │ Authorization: Bearer … │──► env-var token match ──► ALLOW +//! │ or ?token=xxx (SSE/WS only) │──► DB-backed token match ──► ALLOW +//! └────────────┬────────────────┘ +//! │ no match / missing +//! ▼ +//! ┌─────────────────────────────┐ +//! │ OIDC JWT header │──► sig + claims OK ──► ALLOW +//! │ (if configured) │ +//! └────────────┬────────────────┘ +//! │ no match / missing / disabled +//! ▼ +//! 401 Unauthorized +//! ``` +//! +//! **Bearer token** — constant-time comparison via SHA-256 hashed tokens. //! Supports multi-user mode: each token maps to a `UserIdentity` that carries //! the user_id. The identity is inserted into request extensions so downstream //! handlers can extract it via `AuthenticatedUser`. +//! +//! **OIDC JWT** — enabled via `GATEWAY_OIDC_ENABLED=true`. The gateway +//! reads a JWT from a configurable header (default: `x-amzn-oidc-data`), +//! fetches the signing key from a JWKS endpoint, and verifies the +//! signature + claims. Designed for reverse-proxy setups like AWS ALB +//! with Okta/Cognito, but works with any RFC-compliant OIDC provider. +//! The `sub` claim is used as the `user_id` for the resolved identity. +//! +//! **Query-string token** — only allowed on SSE/WS endpoints where +//! browser APIs cannot set custom headers. use std::collections::HashMap; use std::num::NonZeroUsize; +use std::sync::Arc; +use std::time::{Duration, Instant}; use axum::{ extract::{FromRequestParts, Request, State}, @@ -13,15 +47,19 @@ use axum::{ middleware::Next, response::{IntoResponse, Response}, }; +use base64::Engine; +use base64::engine::general_purpose::URL_SAFE_NO_PAD; +use jsonwebtoken::{Algorithm, DecodingKey, Validation}; use sha2::{Digest, Sha256}; -use std::sync::Arc; -use std::time::Instant; use subtle::ConstantTimeEq; use tokio::sync::RwLock; +use crate::config::GatewayOidcConfig; use crate::db::Database; -/// Identity resolved from a bearer token. +// ── User identity ──────────────────────────────────────────────────────── + +/// Identity resolved from a bearer token or OIDC JWT. #[derive(Debug, Clone)] pub struct UserIdentity { pub user_id: String, @@ -38,6 +76,8 @@ pub fn hash_token(token: &str) -> [u8; 32] { hasher.finalize().into() } +// ── Multi-user env-var auth ────────────────────────────────────────────── + /// Multi-user auth state: maps token hashes to user identities. /// /// Tokens are SHA-256 hashed on construction so they are never stored in @@ -122,6 +162,8 @@ impl MultiAuthState { } } +// ── DB-backed auth ─────────────────────────────────────────────────────── + /// DB-backed token authenticator with a bounded LRU cache. /// /// Checks an LRU cache first (TTL 60s), then falls back to a DB query. @@ -229,13 +271,18 @@ impl DbAuthenticator { } } -/// Combined auth state: tries env-var tokens first, then DB-backed tokens. +// ── Combined auth state ─────────────────────────────────────────────────��� + +/// Combined auth state: tries env-var tokens first, then DB-backed tokens, +/// then OIDC JWT (if configured). #[derive(Clone)] pub struct CombinedAuthState { /// In-memory tokens from GATEWAY_AUTH_TOKEN. pub env_auth: MultiAuthState, /// DB-backed token authenticator (optional — only when a database is available). pub db_auth: Option, + /// OIDC JWT auth state (None when OIDC is disabled). + pub oidc: Option, } impl From for CombinedAuthState { @@ -243,10 +290,13 @@ impl From for CombinedAuthState { Self { env_auth, db_auth: None, + oidc: None, } } } +// ── Axum extractors ────────────────────────────────────────────────────── + /// Axum extractor that provides the authenticated user identity. /// /// Only available on routes behind `auth_middleware`. Extracts the @@ -294,6 +344,527 @@ where } } +// ── OIDC types ─────────────────────────────────────────────────────────── + +/// Cached OIDC signing key with its resolved algorithm. +#[derive(Clone)] +struct CachedKey { + decoding_key: DecodingKey, + algorithm: Algorithm, + fetched_at: Instant, +} + +/// Tracks recent fetch failures to avoid hammering a downed JWKS endpoint. +#[derive(Clone)] +struct FailedFetch { + failed_at: Instant, +} + +/// How long to suppress retries after a JWKS fetch failure. +const FETCH_FAILURE_BACKOFF: Duration = Duration::from_secs(10); + +/// OIDC JWT authentication state. +/// +/// Holds the configuration, an HTTP client for JWKS fetches, and a +/// per-`kid` key cache with 1-hour TTL. +#[derive(Clone)] +pub struct OidcState { + config: GatewayOidcConfig, + key_cache: Arc>>, + /// Tracks recent fetch failures per kid to prevent retry storms. + fetch_failures: Arc>>, + http_client: reqwest::Client, +} + +/// OIDC-specific errors (internal, never shown to unauthenticated clients). +#[derive(Debug, thiserror::Error)] +enum OidcError { + #[error("missing `kid` in JWT header")] + MissingKid, + #[error("unsupported algorithm: {0}")] + UnsupportedAlgorithm(String), + #[error("key fetch failed: {0}")] + KeyFetch(String), + #[error("signature verification failed")] + InvalidSignature, + #[error("claim validation failed: {0}")] + InvalidClaims(String), +} + +const KEY_CACHE_TTL: Duration = Duration::from_secs(3600); +/// Maximum number of cached keys. Prevents memory exhaustion from +/// attackers sending JWTs with many distinct `kid` values. +const KEY_CACHE_MAX_ENTRIES: usize = 64; + +impl OidcState { + /// Build OIDC state from gateway config. + /// + /// # Errors + /// + /// Returns an error if the reqwest HTTP client fails to build (e.g. TLS + /// backend unavailable). + pub fn from_config(oidc: &GatewayOidcConfig) -> Result { + let http_client = reqwest::Client::builder() + .timeout(Duration::from_secs(10)) + .build() + .map_err(|e| format!("failed to build OIDC HTTP client: {e}"))?; + Ok(Self { + config: oidc.clone(), + key_cache: Arc::new(RwLock::new(HashMap::new())), + fetch_failures: Arc::new(RwLock::new(HashMap::new())), + http_client, + }) + } + + /// Pre-seed the key cache with a known key for testing. + /// + /// Allows integration tests to exercise the full OIDC middleware path + /// without requiring an HTTP JWKS endpoint. + #[cfg(test)] + pub(crate) async fn seed_key(&self, kid: &str, key: DecodingKey, algorithm: Algorithm) { + let mut cache = self.key_cache.write().await; + cache.insert( + kid.to_string(), + CachedKey { + decoding_key: key, + algorithm, + fetched_at: Instant::now(), + }, + ); + } + + /// Header name containing the JWT. + fn header_name(&self) -> &str { + &self.config.header + } + + // ── Key fetching ───────────────────────────────────────────────────── + + /// Fetch a PEM or JWK from an ALB-style per-key URL (`{kid}` placeholder). + async fn fetch_single_key(&self, url: &str, alg: Algorithm) -> Result { + let body = self.fetch_url_text(url).await?; + let trimmed = body.trim(); + + if trimmed.starts_with("-----BEGIN") { + // PEM-encoded public key (EC or RSA). + match alg { + Algorithm::ES256 | Algorithm::ES384 => DecodingKey::from_ec_pem(trimmed.as_bytes()) + .map_err(|e| OidcError::KeyFetch(format!("EC PEM parse: {e}"))), + Algorithm::EdDSA => DecodingKey::from_ed_pem(trimmed.as_bytes()) + .map_err(|e| OidcError::KeyFetch(format!("EdDSA PEM parse: {e}"))), + _ => DecodingKey::from_rsa_pem(trimmed.as_bytes()) + .map_err(|e| OidcError::KeyFetch(format!("RSA PEM parse: {e}"))), + } + } else { + // Assume single JWK JSON object. + let jwk: jsonwebtoken::jwk::Jwk = serde_json::from_str(trimmed) + .map_err(|e| OidcError::KeyFetch(format!("JWK parse: {e}")))?; + DecodingKey::from_jwk(&jwk).map_err(|e| OidcError::KeyFetch(format!("JWK decode: {e}"))) + } + } + + /// Fetch from a standard JWKS endpoint and find the key matching `kid`. + async fn fetch_jwks_key( + &self, + url: &str, + kid: &str, + ) -> Result<(DecodingKey, Algorithm), OidcError> { + let body = self.fetch_url_text(url).await?; + let jwks: jsonwebtoken::jwk::JwkSet = serde_json::from_str(&body) + .map_err(|e| OidcError::KeyFetch(format!("JWKS parse: {e}")))?; + let jwk = jwks + .find(kid) + .ok_or_else(|| OidcError::KeyFetch(format!("kid '{kid}' not found in JWKS")))?; + let alg = resolve_algorithm(jwk)?; + let key = DecodingKey::from_jwk(jwk) + .map_err(|e| OidcError::KeyFetch(format!("JWK decode: {e}")))?; + Ok((key, alg)) + } + + /// Maximum JWKS response body size (256 KB). Prevents a compromised + /// endpoint from sending arbitrarily large payloads. + const MAX_JWKS_RESPONSE_BYTES: usize = 256 * 1024; + + /// HTTP GET helper with timeout, error status check, and body size limit. + async fn fetch_url_text(&self, url: &str) -> Result { + let response = self + .http_client + .get(url) + .send() + .await + .map_err(|e| OidcError::KeyFetch(format!("HTTP request: {e}")))? + .error_for_status() + .map_err(|e| OidcError::KeyFetch(format!("HTTP error: {e}")))?; + + // Check Content-Length hint before downloading. + if let Some(len) = response.content_length() + && len as usize > Self::MAX_JWKS_RESPONSE_BYTES + { + return Err(OidcError::KeyFetch(format!( + "JWKS response too large ({len} bytes, max {})", + Self::MAX_JWKS_RESPONSE_BYTES + ))); + } + + let bytes = response + .bytes() + .await + .map_err(|e| OidcError::KeyFetch(format!("reading body: {e}")))?; + if bytes.len() > Self::MAX_JWKS_RESPONSE_BYTES { + return Err(OidcError::KeyFetch(format!( + "JWKS response too large ({} bytes, max {})", + bytes.len(), + Self::MAX_JWKS_RESPONSE_BYTES + ))); + } + + String::from_utf8(bytes.to_vec()) + .map_err(|e| OidcError::KeyFetch(format!("response not UTF-8: {e}"))) + } + + /// Get the signing key for `kid`, using cache when available (1h TTL). + async fn get_or_fetch_key( + &self, + kid: &str, + alg: Algorithm, + ) -> Result<(DecodingKey, Algorithm), OidcError> { + // Fast path: cache hit with valid TTL. + { + let cache = self.key_cache.read().await; + if let Some(cached) = cache.get(kid) + && cached.fetched_at.elapsed() < KEY_CACHE_TTL + { + return Ok((cached.decoding_key.clone(), cached.algorithm)); + } + } + + // Check recent fetch failure backoff to avoid hammering a downed endpoint. + { + let failures = self.fetch_failures.read().await; + if let Some(failed) = failures.get(kid) + && failed.failed_at.elapsed() < FETCH_FAILURE_BACKOFF + { + return Err(OidcError::KeyFetch( + "JWKS fetch recently failed, backing off".to_string(), + )); + } + } + + // Slow path: fetch and cache. + let fetch_result = if self.config.jwks_url.contains("{kid}") { + // URL-encode the kid to prevent SSRF via crafted JWT headers. + let encoded_kid: String = + url::form_urlencoded::byte_serialize(kid.as_bytes()).collect(); + let url = self.config.jwks_url.replace("{kid}", &encoded_kid); + self.fetch_single_key(&url, alg).await.map(|key| (key, alg)) + } else { + self.fetch_jwks_key(&self.config.jwks_url, kid).await + }; + + // Record failure for backoff before propagating error. + let (key, resolved_alg) = match fetch_result { + Ok(result) => { + // Clear any previous failure record. + self.fetch_failures.write().await.remove(kid); + result + } + Err(e) => { + self.fetch_failures.write().await.insert( + kid.to_string(), + FailedFetch { + failed_at: Instant::now(), + }, + ); + return Err(e); + } + }; + + let mut cache = self.key_cache.write().await; + + // Evict expired entries and enforce max cache size to prevent + // memory exhaustion from attacker-controlled kid values. + cache.retain(|_, v| v.fetched_at.elapsed() < KEY_CACHE_TTL); + if cache.len() >= KEY_CACHE_MAX_ENTRIES { + // Evict the oldest entry. + if let Some(oldest_kid) = cache + .iter() + .min_by_key(|(_, v)| v.fetched_at) + .map(|(k, _)| k.clone()) + { + cache.remove(&oldest_kid); + } + } + + cache.insert( + kid.to_string(), + CachedKey { + decoding_key: key.clone(), + algorithm: resolved_alg, + fetched_at: Instant::now(), + }, + ); + + Ok((key, resolved_alg)) + } +} + +// ── Algorithm resolution ───────────────────────────────────────────────── + +/// Map a JWK's `alg` field to a `jsonwebtoken::Algorithm`. +fn resolve_algorithm(jwk: &jsonwebtoken::jwk::Jwk) -> Result { + match jwk.common.key_algorithm { + Some(jsonwebtoken::jwk::KeyAlgorithm::ES256) => Ok(Algorithm::ES256), + Some(jsonwebtoken::jwk::KeyAlgorithm::ES384) => Ok(Algorithm::ES384), + Some(jsonwebtoken::jwk::KeyAlgorithm::RS256) => Ok(Algorithm::RS256), + Some(jsonwebtoken::jwk::KeyAlgorithm::RS384) => Ok(Algorithm::RS384), + Some(jsonwebtoken::jwk::KeyAlgorithm::RS512) => Ok(Algorithm::RS512), + Some(jsonwebtoken::jwk::KeyAlgorithm::PS256) => Ok(Algorithm::PS256), + Some(jsonwebtoken::jwk::KeyAlgorithm::PS384) => Ok(Algorithm::PS384), + Some(jsonwebtoken::jwk::KeyAlgorithm::PS512) => Ok(Algorithm::PS512), + Some(jsonwebtoken::jwk::KeyAlgorithm::EdDSA) => Ok(Algorithm::EdDSA), + Some(other) => Err(OidcError::UnsupportedAlgorithm(format!("{other:?}"))), + None => Err(OidcError::UnsupportedAlgorithm( + "missing alg in JWK".to_string(), + )), + } +} + +// ── Signature verification ─────────────────────────────────────────────── + +/// Verify the JWT signature using the **original** token text as the +/// signing input. +/// +/// Why not just use `jsonwebtoken::decode()`? Because `decode()` strips +/// base64 padding (`=`) from header and payload segments before building +/// the signing input. AWS ALB signs over the *padded* segments, so +/// stripping padding changes the message and breaks verification. +/// +/// We call `jsonwebtoken::crypto::verify()` directly with the original +/// `header.payload` bytes, then extract claims separately via +/// `decode()` with signature validation disabled (safe — we already +/// verified the signature above). +fn verify_signature( + original_jwt: &str, + key: &DecodingKey, + alg: Algorithm, +) -> Result<(), OidcError> { + let parts: Vec<&str> = original_jwt.split('.').collect(); + if parts.len() != 3 { + return Err(OidcError::InvalidSignature); + } + + let signing_input = format!("{}.{}", parts[0], parts[1]); + let raw_sig = parts[2]; + + // Decode signature bytes from base64url (tolerate padding). + let sig_bytes = URL_SAFE_NO_PAD + .decode(raw_sig.trim_end_matches('=')) + .map_err(|_| OidcError::InvalidSignature)?; + + // ECDSA signatures: handle DER encoding if present (some IdPs use + // DER-encoded signatures instead of raw R||S). + let sig_bytes = if matches!(alg, Algorithm::ES256 | Algorithm::ES384) { + match try_der_to_raw(&sig_bytes, alg) { + Some(raw) => raw, + None => sig_bytes, + } + } else { + sig_bytes + }; + + // Re-encode the (possibly DER→raw converted) signature to base64url + // because jsonwebtoken::crypto::verify() expects a base64url string. + let sig_b64 = URL_SAFE_NO_PAD.encode(&sig_bytes); + + // verify(signature_b64, message_bytes, key, alg) + let valid = jsonwebtoken::crypto::verify(&sig_b64, signing_input.as_bytes(), key, alg) + .map_err(|_| OidcError::InvalidSignature)?; + + if valid { + Ok(()) + } else { + Err(OidcError::InvalidSignature) + } +} + +// ── Base64 normalization (for claim extraction only) ───────────────────── + +/// Strip base64 padding from a single segment. +/// +/// Used only when building a normalized JWT for `jsonwebtoken::decode()` +/// claim extraction. The `jsonwebtoken` crate uses `URL_SAFE_NO_PAD` +/// internally, so padded segments cause decode failures. +fn normalize_b64_segment(seg: &str) -> String { + seg.trim_end_matches('=').to_string() +} + +/// Rebuild the JWT with padding stripped from all three segments. +/// +/// This is a no-op for RFC-compliant JWTs that already omit padding. +/// Only used for claim extraction after signature verification. +fn normalize_jwt_for_claims(jwt: &str) -> String { + let parts: Vec<&str> = jwt.split('.').collect(); + if parts.len() != 3 { + return jwt.to_string(); + } + format!( + "{}.{}.{}", + normalize_b64_segment(parts[0]), + normalize_b64_segment(parts[1]), + normalize_b64_segment(parts[2]), + ) +} + +// ── DER → raw ECDSA signature conversion ───────────────────────────────── + +/// Try to convert a DER-encoded ECDSA signature to raw R||S format. +/// +/// Returns `None` if the input doesn't look like valid DER, in which case +/// the caller should use the bytes as-is (already raw R||S). +fn try_der_to_raw(der: &[u8], alg: Algorithm) -> Option> { + let component_len = match alg { + Algorithm::ES256 => 32, + Algorithm::ES384 => 48, + _ => return None, + }; + + // DER SEQUENCE: 0x30 + if der.len() < 6 || der[0] != 0x30 { + return None; + } + + // Skip SEQUENCE tag + parse length (supports long-form DER lengths). + let mut pos = 1; + let _seq_len = parse_der_length(der, &mut pos)?; + + // Parse R INTEGER + if pos >= der.len() || der[pos] != 0x02 { + return None; + } + pos += 1; + let r_len = parse_der_length(der, &mut pos)?; + if r_len > component_len + 1 { + return None; + } + let r_bytes = der.get(pos..pos + r_len)?; + pos += r_len; + + // Parse S INTEGER + if pos >= der.len() || der[pos] != 0x02 { + return None; + } + pos += 1; + let s_len = parse_der_length(der, &mut pos)?; + if s_len > component_len + 1 { + return None; + } + let s_bytes = der.get(pos..pos + s_len)?; + + // Strip leading zero padding from DER INTEGER values and left-pad + // to the expected component length. + let r = strip_der_leading_zero(r_bytes); + let s = strip_der_leading_zero(s_bytes); + if r.len() > component_len || s.len() > component_len { + return None; + } + + let mut raw = vec![0u8; component_len * 2]; + raw[component_len - r.len()..component_len].copy_from_slice(r); + raw[component_len * 2 - s.len()..].copy_from_slice(s); + Some(raw) +} + +/// Parse a DER length field, handling both short-form (< 128) and +/// long-form (0x81 xx, 0x82 xx yy) encodings. Advances `pos` past +/// the length bytes. Returns `None` for unsupported multi-byte lengths +/// (> 2 bytes) or if the buffer is too short. +fn parse_der_length(der: &[u8], pos: &mut usize) -> Option { + let b = *der.get(*pos)?; + *pos += 1; + if b < 0x80 { + Some(b as usize) + } else { + let num_bytes = (b & 0x7F) as usize; + if num_bytes == 0 || num_bytes > 2 { + return None; + } + let mut len: usize = 0; + for _ in 0..num_bytes { + len = len + .checked_mul(256)? + .checked_add(*der.get(*pos)? as usize)?; + *pos += 1; + } + Some(len) + } +} + +/// Strip the leading zero byte that DER adds to unsigned INTEGERs when +/// the high bit is set (to distinguish from negative values). +fn strip_der_leading_zero(bytes: &[u8]) -> &[u8] { + if bytes.len() > 1 && bytes[0] == 0x00 { + &bytes[1..] + } else { + bytes + } +} + +// ── Full OIDC validation pipeline ──────────────────────────────────────── + +/// Validate an OIDC JWT: fetch key, verify signature, check claims. +/// +/// Returns the `sub` (subject) claim on success. +async fn validate_oidc_jwt(oidc: &OidcState, jwt: &str) -> Result { + // Normalize first — `decode_header()` uses URL_SAFE_NO_PAD internally + // and chokes on the `=` padding that AWS ALB includes. + let normalized = normalize_jwt_for_claims(jwt); + + // Decode the unverified header to get `kid` and `alg`. + let header = jsonwebtoken::decode_header(&normalized) + .map_err(|e| OidcError::InvalidClaims(format!("malformed header: {e}")))?; + let kid = header.kid.ok_or(OidcError::MissingKid)?; + let alg = header.alg; + + // Fetch (or retrieve from cache) the signing key. + let (key, resolved_alg) = oidc.get_or_fetch_key(&kid, alg).await?; + + // Verify signature against the ORIGINAL JWT text (preserving any + // padding). ALB signed over the padded segments, so we must use the + // original token as the signing input. + verify_signature(jwt, &key, resolved_alg)?; + + // SAFETY: Signature validation is disabled here because we already + // verified the signature above via `verify_signature()`. We use + // `decode()` only for claim extraction and expiry/issuer/audience + // validation. Do not copy this pattern without the preceding + // `verify_signature()` call. + let mut validation = Validation::new(resolved_alg); + validation.insecure_disable_signature_validation(); + + if let Some(ref iss) = oidc.config.issuer { + validation.set_issuer(&[iss]); + } + if let Some(ref aud) = oidc.config.audience { + validation.set_audience(&[aud]); + } else { + validation.validate_aud = false; + } + + let data = jsonwebtoken::decode::(&normalized, &key, &validation) + .map_err(|e| OidcError::InvalidClaims(format!("{e}")))?; + + let sub = data + .claims + .get("sub") + .and_then(|v| v.as_str()) + .map(|s| s.to_string()) + .ok_or_else(|| OidcError::InvalidClaims("missing `sub` claim".to_string()))?; + + Ok(sub) +} + +// ── Token extraction helpers ───────────────────────────────────────────── + /// Whether query-string token auth is allowed for this request. /// /// Only GET requests to streaming endpoints may use `?token=xxx`. This @@ -328,12 +899,33 @@ fn query_token(request: &Request) -> Option { }) } -/// Auth middleware that validates bearer token from header or query param. +/// Extract a bearer token from the Authorization header or query parameter. +fn extract_token(headers: &HeaderMap, request: &Request) -> Option { + // Try Authorization header first (RFC 6750). + if let Some(auth_header) = headers.get("authorization") + && let Ok(value) = auth_header.to_str() + && value.len() > 7 + && value[..7].eq_ignore_ascii_case("Bearer ") + { + return Some(value[7..].to_string()); + } + + // Fall back to query parameter for SSE/WS endpoints. + if allows_query_token_auth(request) { + return query_token(request); + } + + None +} + +// ── Middleware ──────────────────────────────────────────────────────────── + +/// Auth middleware: bearer/query token → OIDC JWT → 401. /// /// Tries env-var tokens first (constant-time, in-memory), then falls back -/// to DB-backed token lookup if configured. SSE connections can't set -/// headers from `EventSource`, so we also accept `?token=xxx` as a query -/// parameter, but only on SSE/WS endpoints. +/// to DB-backed token lookup if configured, then OIDC JWT validation. +/// SSE connections can't set headers from `EventSource`, so we also accept +/// `?token=xxx` as a query parameter, but only on SSE/WS endpoints. /// /// On successful authentication, inserts the matching `UserIdentity` into /// request extensions for downstream extraction via `AuthenticatedUser`. @@ -369,28 +961,33 @@ pub async fn auth_middleware( } } - (StatusCode::UNAUTHORIZED, "Invalid or missing auth token").into_response() -} - -/// Extract a bearer token from the Authorization header or query parameter. -fn extract_token(headers: &HeaderMap, request: &Request) -> Option { - // Try Authorization header first (RFC 6750). - if let Some(auth_header) = headers.get("authorization") - && let Ok(value) = auth_header.to_str() - && value.len() > 7 - && value[..7].eq_ignore_ascii_case("Bearer ") + // 3. Try OIDC JWT from configured header (if enabled). + if let Some(ref oidc) = auth.oidc + && let Some(jwt_header) = headers.get(oidc.header_name()) + && let Ok(jwt) = jwt_header.to_str() { - return Some(value[7..].to_string()); - } - - // Fall back to query parameter for SSE/WS endpoints. - if allows_query_token_auth(request) { - return query_token(request); + match validate_oidc_jwt(oidc, jwt).await { + Ok(sub) => { + tracing::debug!(sub = %sub, "OIDC auth succeeded"); + let identity = UserIdentity { + user_id: sub, + role: "member".to_string(), + workspace_read_scopes: Vec::new(), + }; + request.extensions_mut().insert(identity); + return next.run(request).await; + } + Err(e) => { + tracing::warn!(error = %e, "OIDC auth failed"); + } + } } - None + (StatusCode::UNAUTHORIZED, "Invalid or missing auth token").into_response() } +// ── Tests ──────────────────────────────────────────────────────────────── + #[cfg(test)] mod tests { use super::*; @@ -667,7 +1264,252 @@ mod tests { assert_eq!(resp.status(), StatusCode::UNAUTHORIZED); } - // --- Multi-tenant auth integration tests --- + // ── OIDC unit tests ───────────────────────────────────────────���────── + + #[test] + fn test_normalize_jwt_noop_for_rfc_compliant() { + // No padding → no change. + let jwt = "eyJhbGciOiJIUzI1NiJ9.eyJzdWIiOiJ0ZXN0In0.sig"; + assert_eq!(normalize_jwt_for_claims(jwt), jwt); + } + + #[test] + fn test_normalize_jwt_strips_padding() { + let jwt = "eyJhbGciOiJIUzI1NiJ9==.eyJzdWIiOiJ0ZXN0In0=.c2ln"; + let normalized = normalize_jwt_for_claims(jwt); + assert!(!normalized.contains('=')); + assert!(normalized.starts_with("eyJhbGciOiJIUzI1NiJ9.")); + } + + #[test] + fn test_normalize_b64_segment_no_padding() { + assert_eq!(normalize_b64_segment("abc"), "abc"); + } + + #[test] + fn test_normalize_b64_segment_with_padding() { + assert_eq!(normalize_b64_segment("abc=="), "abc"); + } + + #[test] + fn test_try_der_to_raw_non_der_passthrough() { + // 64 bytes of raw R||S — not DER, should return None. + let raw = vec![0x01; 64]; + assert!(try_der_to_raw(&raw, Algorithm::ES256).is_none()); + } + + #[test] + fn test_try_der_to_raw_valid_der() { + // Construct a minimal DER ECDSA signature for ES256. + // SEQUENCE { INTEGER(r=1, 32 bytes), INTEGER(s=2, 32 bytes) } + let r = vec![0x01; 32]; + let s = vec![0x02; 32]; + let mut der = vec![0x30, 68]; // SEQUENCE, length=68 + der.push(0x02); + der.push(32); + der.extend_from_slice(&r); + der.push(0x02); + der.push(32); + der.extend_from_slice(&s); + + let raw = try_der_to_raw(&der, Algorithm::ES256).expect("should parse DER"); + assert_eq!(raw.len(), 64); + assert_eq!(&raw[..32], &r[..]); + assert_eq!(&raw[32..], &s[..]); + } + + #[test] + fn test_try_der_to_raw_with_leading_zero() { + // DER adds a 0x00 prefix when the high bit of an INTEGER is set. + let r = { + let mut v = vec![0x00]; // leading zero + v.extend_from_slice(&[0x80; 32]); // 32 bytes with high bit set + v + }; + let s = vec![0x01; 32]; + let mut der = vec![0x30, 69]; // SEQUENCE, length = 33+32+4 = 69 + der.push(0x02); + der.push(33); // r_len = 33 (with leading zero) + der.extend_from_slice(&r); + der.push(0x02); + der.push(32); + der.extend_from_slice(&s); + + let raw = try_der_to_raw(&der, Algorithm::ES256).expect("should parse DER"); + assert_eq!(raw.len(), 64); + // R should have the leading zero stripped. + assert_eq!(raw[0], 0x80); + } + + #[test] + fn test_strip_der_leading_zero() { + assert_eq!(strip_der_leading_zero(&[0x00, 0x80, 0x01]), &[0x80, 0x01]); + assert_eq!(strip_der_leading_zero(&[0x80, 0x01]), &[0x80, 0x01]); + assert_eq!(strip_der_leading_zero(&[0x00]), &[0x00]); // single zero stays + } + + #[test] + fn test_parse_der_length_short_form() { + let data = [0x20]; // 32 in short form + let mut pos = 0; + assert_eq!(parse_der_length(&data, &mut pos), Some(32)); + assert_eq!(pos, 1); + } + + #[test] + fn test_parse_der_length_long_form_one_byte() { + // 0x81 0x80 = 128 in long form (1 extra length byte) + let data = [0x81, 0x80]; + let mut pos = 0; + assert_eq!(parse_der_length(&data, &mut pos), Some(128)); + assert_eq!(pos, 2); + } + + #[test] + fn test_parse_der_length_long_form_two_bytes() { + // 0x82 0x01 0x00 = 256 in long form (2 extra length bytes) + let data = [0x82, 0x01, 0x00]; + let mut pos = 0; + assert_eq!(parse_der_length(&data, &mut pos), Some(256)); + assert_eq!(pos, 3); + } + + #[test] + fn test_try_der_to_raw_long_form_sequence_length() { + // Build a DER signature where SEQUENCE length is >= 128 (uses long form). + // ES384: component_len=48, max R=49 (leading zero), max S=49. + let r = { + let mut v = vec![0x00]; // leading zero + v.extend_from_slice(&[0xFF; 48]); // 48 bytes with high bits + v + }; + let s = { + let mut v = vec![0x00]; // leading zero + v.extend_from_slice(&[0xAA; 48]); + v + }; + let content_len = 2 + r.len() + 2 + s.len(); // 102 + assert!(content_len < 128); // short form still works for ES384 + + // Force a case where total > 127: use ES384 with both R and S having 49 bytes + // content = (1+1+49) + (1+1+49) = 102. That's < 128, so let's construct + // a valid DER with 0x81 long-form length anyway to test the parser. + let mut der = vec![0x30, 0x81, content_len as u8]; + der.push(0x02); + der.push(r.len() as u8); + der.extend_from_slice(&r); + der.push(0x02); + der.push(s.len() as u8); + der.extend_from_slice(&s); + + let raw = try_der_to_raw(&der, Algorithm::ES384) + .expect("should parse DER with long-form sequence length"); + assert_eq!(raw.len(), 96); // 48 * 2 + // R should have leading zero stripped → first byte is 0xFF + assert_eq!(raw[0], 0xFF); + // S should have leading zero stripped → byte at offset 48 is 0xAA + assert_eq!(raw[48], 0xAA); + } + + #[test] + fn test_kid_url_encoded_in_jwks_url() { + // Verify that special characters in kid are URL-encoded, not raw-substituted. + let encoded: String = url::form_urlencoded::byte_serialize(b"../../evil?x=1").collect(); + let url = "https://example.com/keys/{kid}".replace("{kid}", &encoded); + assert!(!url.contains("../")); + assert!(url.contains("%2F")); + } + + #[test] + fn test_verify_signature_rejects_tampered_payload() { + use jsonwebtoken::{EncodingKey, Header}; + + // Use HS256 for a self-contained unit test (no external keys). + let secret = b"test-secret-at-least-256-bits!!!"; + let header = Header::new(Algorithm::HS256); + let claims = serde_json::json!({"sub": "alice", "exp": 9999999999u64}); + let token = + jsonwebtoken::encode(&header, &claims, &EncodingKey::from_secret(secret)).unwrap(); + + // Valid signature should pass. + let key = DecodingKey::from_secret(secret); + assert!(verify_signature(&token, &key, Algorithm::HS256).is_ok()); + + // Tamper with the payload — signature should fail. + let parts: Vec<&str> = token.split('.').collect(); + let tampered = format!("{}.{}.{}", parts[0], "dGFtcGVyZWQ", parts[2]); + assert!(verify_signature(&tampered, &key, Algorithm::HS256).is_err()); + } + + #[tokio::test] + async fn test_validate_oidc_jwt_rejects_missing_sub() { + use jsonwebtoken::{EncodingKey, Header}; + + // Create a valid HS256 JWT without a `sub` claim. + let secret = b"test-secret-at-least-256-bits!!!"; + let mut header = Header::new(Algorithm::HS256); + header.kid = Some("test-kid".to_string()); + let claims = serde_json::json!({"exp": 9999999999u64, "name": "alice"}); + let token = + jsonwebtoken::encode(&header, &claims, &EncodingKey::from_secret(secret)).unwrap(); + + // Build an OidcState that serves the key from a mock. + // We can't easily mock HTTP, so test the claim extraction path directly: + // build a Validation that skips signature check and verify `sub` is required. + let mut validation = Validation::new(Algorithm::HS256); + validation.insecure_disable_signature_validation(); + validation.validate_aud = false; + + let data = jsonwebtoken::decode::( + &token, + &DecodingKey::from_secret(secret), + &validation, + ) + .unwrap(); + let result = data + .claims + .get("sub") + .and_then(|v| v.as_str()) + .map(|s| s.to_string()) + .ok_or_else(|| OidcError::InvalidClaims("missing `sub` claim".to_string())); + assert!(result.is_err()); + assert!( + result.unwrap_err().to_string().contains("sub"), + "error should mention missing sub claim" + ); + } + + #[test] + fn test_issuer_validation_disabled_when_not_configured() { + // When no issuer is configured, Validation should NOT require iss. + let mut validation = Validation::new(Algorithm::HS256); + validation.insecure_disable_signature_validation(); + validation.validate_aud = false; + + use jsonwebtoken::{EncodingKey, Header}; + let secret = b"test-secret-at-least-256-bits!!!"; + let claims = + serde_json::json!({"sub": "alice", "exp": 9999999999u64, "iss": "https://example.com"}); + let token = jsonwebtoken::encode( + &Header::new(Algorithm::HS256), + &claims, + &EncodingKey::from_secret(secret), + ) + .unwrap(); + + // Should succeed — issuer validation is not enforced. + let result = jsonwebtoken::decode::( + &token, + &DecodingKey::from_secret(secret), + &validation, + ); + assert!( + result.is_ok(), + "token with any issuer should pass when issuer validation is disabled" + ); + } + + // ── Multi-tenant auth integration tests ────────────────────────────── /// Handler that extracts `AuthenticatedUser` and returns the resolved user_id. async fn identity_handler(AuthenticatedUser(identity): AuthenticatedUser) -> String { @@ -740,8 +1582,6 @@ mod tests { #[tokio::test] async fn test_multi_user_sequential_tokens_resolve_independently() { - // Send both alice and bob tokens sequentially and verify each gets - // the correct identity — guards against token map corruption. let tokens = two_user_tokens(); let app1 = multi_user_app(tokens.clone()); @@ -839,7 +1679,6 @@ mod tests { #[tokio::test] async fn test_multi_user_empty_scopes_for_single_user() { - // Single-user mode creates identity with empty workspace_read_scopes. let state = CombinedAuthState::from(MultiAuthState::single( "tok-only".to_string(), "solo".to_string(), @@ -860,11 +1699,586 @@ mod tests { #[tokio::test] async fn test_prefix_and_extension_tokens_rejected() { - // Verifies that prefix/suffix variants of valid tokens are rejected. - // Note: the constant-time property is enforced structurally by use of - // subtle::ConstantTimeEq and cannot be verified via outcome testing. let state = MultiAuthState::single("long-secret-token".to_string(), "user".to_string()); assert!(state.authenticate("long-secret").is_none()); assert!(state.authenticate("long-secret-token-extra").is_none()); } + + // ── OIDC test helpers ───────────────────────────────────────────────── + + const OIDC_SECRET: &[u8] = b"test-secret-at-least-256-bits!!!"; + const OIDC_KID: &str = "test-kid"; + const OIDC_HEADER_NAME: &str = "x-oidc-data"; + + /// Encode an HS256 JWT with the given claims and optional kid. + fn encode_test_jwt(claims: serde_json::Value, kid: Option<&str>) -> String { + use jsonwebtoken::{EncodingKey, Header}; + let mut header = Header::new(Algorithm::HS256); + header.kid = kid.map(|s| s.to_string()); + jsonwebtoken::encode(&header, &claims, &EncodingKey::from_secret(OIDC_SECRET)).unwrap() // safety: test helper + } + + /// Build a default OIDC config (no issuer/audience validation). + fn test_oidc_config() -> crate::config::GatewayOidcConfig { + crate::config::GatewayOidcConfig { + header: OIDC_HEADER_NAME.to_string(), + jwks_url: "https://unused.example.com/keys".to_string(), + issuer: None, + audience: None, + } + } + + /// Build an OidcState with the HS256 test key pre-seeded. + async fn test_oidc_state() -> OidcState { + test_oidc_state_with_config(test_oidc_config()).await + } + + /// Build an OidcState from a custom config with the HS256 test key pre-seeded. + async fn test_oidc_state_with_config(config: crate::config::GatewayOidcConfig) -> OidcState { + let oidc = OidcState::from_config(&config).unwrap(); // safety: test helper + oidc.seed_key( + OIDC_KID, + DecodingKey::from_secret(OIDC_SECRET), + Algorithm::HS256, + ) + .await; + oidc + } + + /// Build a CombinedAuthState with bearer token + OIDC. + async fn oidc_auth_state() -> CombinedAuthState { + CombinedAuthState { + env_auth: MultiAuthState::single( + "bearer-token-123".to_string(), + "bearer-user".to_string(), + ), + db_auth: None, + oidc: Some(test_oidc_state().await), + } + } + + /// Build a Router with identity_handler behind auth_middleware. + fn oidc_test_app(state: CombinedAuthState) -> Router { + Router::new() + .route("/api/chat/events", get(identity_handler)) + .route("/api/chat/send", post(identity_handler)) + .layer(middleware::from_fn_with_state(state, auth_middleware)) + } + + /// Build a valid JWT with `sub` and far-future `exp`. + fn valid_oidc_jwt(sub: &str) -> String { + encode_test_jwt( + serde_json::json!({"sub": sub, "exp": 9999999999u64}), + Some(OIDC_KID), + ) + } + + // ── OIDC middleware integration tests ───────────────────────────────── + + /// Regression test: OIDC auth must produce a `UserIdentity` so that + /// downstream handlers using `AuthenticatedUser` receive the identity. + /// + /// Without the identity insertion, the handler returns 401 even though + /// OIDC signature validation succeeded — the bug that was caught in + /// code review of #1463. + #[tokio::test] + async fn test_oidc_auth_inserts_user_identity_for_handler() { + let app = oidc_test_app(oidc_auth_state().await); + let req = Request::builder() + .uri("/api/chat/events") + .header(OIDC_HEADER_NAME, valid_oidc_jwt("oidc-alice")) + .body(Body::empty()) + .unwrap(); + let resp = app.oneshot(req).await.unwrap(); + assert_eq!( + resp.status(), + StatusCode::OK, + "OIDC auth must insert UserIdentity so AuthenticatedUser extractor succeeds" + ); + let body = axum::body::to_bytes(resp.into_body(), 1024).await.unwrap(); + assert_eq!(body, "oidc-alice"); + } + + /// OIDC-authenticated users get role=member (not admin). + #[tokio::test] + async fn test_oidc_auth_user_gets_member_role() { + async fn role_handler(AuthenticatedUser(id): AuthenticatedUser) -> String { + id.role + } + + let state = oidc_auth_state().await; + let app = Router::new() + .route("/api/chat/events", get(role_handler)) + .layer(middleware::from_fn_with_state(state, auth_middleware)); + + let req = Request::builder() + .uri("/api/chat/events") + .header(OIDC_HEADER_NAME, valid_oidc_jwt("oidc-bob")) + .body(Body::empty()) + .unwrap(); + let resp = app.oneshot(req).await.unwrap(); + assert_eq!(resp.status(), StatusCode::OK); + let body = axum::body::to_bytes(resp.into_body(), 1024).await.unwrap(); + assert_eq!(body, "member"); + } + + // ── Auth priority & fallthrough ────────────────────────────────────── + + /// Bearer token works when OIDC is configured but the OIDC header is absent. + #[tokio::test] + async fn test_bearer_works_when_oidc_configured_but_header_absent() { + let app = oidc_test_app(oidc_auth_state().await); + let req = Request::builder() + .uri("/api/chat/events") + .header("Authorization", "Bearer bearer-token-123") + .body(Body::empty()) + .unwrap(); + let resp = app.oneshot(req).await.unwrap(); + assert_eq!(resp.status(), StatusCode::OK); + let body = axum::body::to_bytes(resp.into_body(), 1024).await.unwrap(); + assert_eq!(body, "bearer-user"); + } + + /// Bearer token takes priority when both Bearer header and OIDC header are present. + #[tokio::test] + async fn test_bearer_takes_priority_over_oidc_when_both_present() { + let app = oidc_test_app(oidc_auth_state().await); + let req = Request::builder() + .uri("/api/chat/events") + .header("Authorization", "Bearer bearer-token-123") + .header(OIDC_HEADER_NAME, valid_oidc_jwt("oidc-alice")) + .body(Body::empty()) + .unwrap(); + let resp = app.oneshot(req).await.unwrap(); + assert_eq!(resp.status(), StatusCode::OK); + let body = axum::body::to_bytes(resp.into_body(), 1024).await.unwrap(); + assert_eq!( + body, "bearer-user", + "bearer should win when both auth methods are present" + ); + } + + /// OIDC failure (wrong signature) falls through gracefully to 401, not 500. + #[tokio::test] + async fn test_oidc_bad_signature_returns_401_not_500() { + let state = oidc_auth_state().await; + let app = oidc_test_app(state); + + // Sign with a different secret so the signature won't match. + let wrong_secret = b"wrong-secret-at-least-256-bits!!"; + let mut header = jsonwebtoken::Header::new(Algorithm::HS256); + header.kid = Some(OIDC_KID.to_string()); + let bad_jwt = jsonwebtoken::encode( + &header, + &serde_json::json!({"sub": "attacker", "exp": 9999999999u64}), + &jsonwebtoken::EncodingKey::from_secret(wrong_secret), + ) + .unwrap(); + + let req = Request::builder() + .uri("/api/chat/events") + .header(OIDC_HEADER_NAME, bad_jwt) + .body(Body::empty()) + .unwrap(); + let resp = app.oneshot(req).await.unwrap(); + assert_eq!( + resp.status(), + StatusCode::UNAUTHORIZED, + "bad OIDC sig should yield 401, not 500" + ); + } + + /// When OIDC header has an invalid JWT but a valid bearer token is also + /// present, bearer auth should succeed (bearer checked first). + #[tokio::test] + async fn test_invalid_oidc_does_not_block_bearer() { + let app = oidc_test_app(oidc_auth_state().await); + let req = Request::builder() + .uri("/api/chat/events") + .header("Authorization", "Bearer bearer-token-123") + .header(OIDC_HEADER_NAME, "not.a.jwt") + .body(Body::empty()) + .unwrap(); + let resp = app.oneshot(req).await.unwrap(); + assert_eq!(resp.status(), StatusCode::OK); + let body = axum::body::to_bytes(resp.into_body(), 1024).await.unwrap(); + assert_eq!(body, "bearer-user"); + } + + /// No auth at all when OIDC is configured → 401. + #[tokio::test] + async fn test_no_auth_with_oidc_configured() { + let app = oidc_test_app(oidc_auth_state().await); + let req = Request::builder() + .uri("/api/chat/events") + .body(Body::empty()) + .unwrap(); + let resp = app.oneshot(req).await.unwrap(); + assert_eq!(resp.status(), StatusCode::UNAUTHORIZED); + } + + // ── Expired / invalid JWT edge cases ───────────────────────────────── + + /// Expired JWT (`exp` in the past) is rejected. + #[tokio::test] + async fn test_oidc_expired_jwt_rejected() { + let app = oidc_test_app(oidc_auth_state().await); + let jwt = encode_test_jwt( + serde_json::json!({"sub": "alice", "exp": 1000000000u64}), // year 2001 + Some(OIDC_KID), + ); + let req = Request::builder() + .uri("/api/chat/events") + .header(OIDC_HEADER_NAME, jwt) + .body(Body::empty()) + .unwrap(); + let resp = app.oneshot(req).await.unwrap(); + assert_eq!(resp.status(), StatusCode::UNAUTHORIZED); + } + + /// JWT without `kid` header field is rejected (MissingKid). + #[tokio::test] + async fn test_oidc_jwt_without_kid_rejected() { + let app = oidc_test_app(oidc_auth_state().await); + let jwt = encode_test_jwt( + serde_json::json!({"sub": "alice", "exp": 9999999999u64}), + None, // no kid + ); + let req = Request::builder() + .uri("/api/chat/events") + .header(OIDC_HEADER_NAME, jwt) + .body(Body::empty()) + .unwrap(); + let resp = app.oneshot(req).await.unwrap(); + assert_eq!(resp.status(), StatusCode::UNAUTHORIZED); + } + + /// Malformed JWT (not three dot-separated parts) is rejected. + #[tokio::test] + async fn test_oidc_malformed_jwt_rejected() { + let app = oidc_test_app(oidc_auth_state().await); + for malformed in ["", "abc", "a.b", "a.b.c.d", "not-base64.not-base64.sig"] { + let req = Request::builder() + .uri("/api/chat/events") + .header(OIDC_HEADER_NAME, malformed) + .body(Body::empty()) + .unwrap(); + let resp = app.clone().oneshot(req).await.unwrap(); + assert_eq!( + resp.status(), + StatusCode::UNAUTHORIZED, + "malformed JWT '{malformed}' should be rejected" + ); + } + } + + /// JWT with `sub` as a non-string value (integer) is rejected. + #[tokio::test] + async fn test_oidc_jwt_sub_not_string_rejected() { + let oidc = test_oidc_state().await; + let jwt = encode_test_jwt( + serde_json::json!({"sub": 12345, "exp": 9999999999u64}), + Some(OIDC_KID), + ); + let result = validate_oidc_jwt(&oidc, &jwt).await; + assert!( + result.is_err(), + "non-string sub should be rejected: {result:?}" + ); + } + + /// JWT with empty-string `sub` claim succeeds (empty user_id is valid + /// at the auth layer; authorization checks happen downstream). + #[tokio::test] + async fn test_oidc_jwt_empty_sub_passes_auth() { + let app = oidc_test_app(oidc_auth_state().await); + let jwt = encode_test_jwt( + serde_json::json!({"sub": "", "exp": 9999999999u64}), + Some(OIDC_KID), + ); + let req = Request::builder() + .uri("/api/chat/events") + .header(OIDC_HEADER_NAME, jwt) + .body(Body::empty()) + .unwrap(); + let resp = app.oneshot(req).await.unwrap(); + // Empty sub is technically valid at the auth layer. If we decide to + // reject it, this test documents the expectation and should be updated. + assert_eq!(resp.status(), StatusCode::OK); + let body = axum::body::to_bytes(resp.into_body(), 1024).await.unwrap(); + assert_eq!(body, ""); + } + + /// JWT with missing `sub` claim is rejected even though signature is valid. + #[tokio::test] + async fn test_oidc_jwt_missing_sub_rejected_through_middleware() { + let app = oidc_test_app(oidc_auth_state().await); + let jwt = encode_test_jwt( + serde_json::json!({"name": "alice", "exp": 9999999999u64}), // no sub + Some(OIDC_KID), + ); + let req = Request::builder() + .uri("/api/chat/events") + .header(OIDC_HEADER_NAME, jwt) + .body(Body::empty()) + .unwrap(); + let resp = app.oneshot(req).await.unwrap(); + assert_eq!(resp.status(), StatusCode::UNAUTHORIZED); + } + + // ── Issuer / audience validation ───────────────────────────────────── + + /// Issuer configured and JWT `iss` matches → accepted. + #[tokio::test] + async fn test_oidc_issuer_match_accepted() { + let mut config = test_oidc_config(); + config.issuer = Some("https://idp.example.com".to_string()); + let oidc = test_oidc_state_with_config(config).await; + let jwt = encode_test_jwt( + serde_json::json!({ + "sub": "alice", + "iss": "https://idp.example.com", + "exp": 9999999999u64, + }), + Some(OIDC_KID), + ); + let result = validate_oidc_jwt(&oidc, &jwt).await; + assert!(result.is_ok(), "matching issuer should pass: {result:?}"); + assert_eq!(result.unwrap(), "alice"); + } + + /// Issuer configured but JWT has wrong `iss` → rejected. + #[tokio::test] + async fn test_oidc_issuer_mismatch_rejected() { + let mut config = test_oidc_config(); + config.issuer = Some("https://idp.example.com".to_string()); + let oidc = test_oidc_state_with_config(config).await; + let jwt = encode_test_jwt( + serde_json::json!({ + "sub": "alice", + "iss": "https://evil.example.com", + "exp": 9999999999u64, + }), + Some(OIDC_KID), + ); + let result = validate_oidc_jwt(&oidc, &jwt).await; + assert!(result.is_err(), "wrong issuer should be rejected"); + } + + /// Issuer configured but JWT omits `iss` entirely. + /// + /// Note: `jsonwebtoken` v9 only validates `iss` when present; a missing + /// `iss` claim passes validation. This test documents that behavior. + /// If we decide to enforce presence, add an explicit check in + /// `validate_oidc_jwt` after claim extraction. + #[tokio::test] + async fn test_oidc_issuer_configured_but_missing_in_jwt_passes() { + let mut config = test_oidc_config(); + config.issuer = Some("https://idp.example.com".to_string()); + let oidc = test_oidc_state_with_config(config).await; + let jwt = encode_test_jwt( + serde_json::json!({"sub": "alice", "exp": 9999999999u64}), + Some(OIDC_KID), + ); + let result = validate_oidc_jwt(&oidc, &jwt).await; + // jsonwebtoken allows missing iss — only rejects mismatches. + assert!( + result.is_ok(), + "missing iss is not rejected by jsonwebtoken: {result:?}" + ); + } + + /// Audience configured and JWT `aud` matches → accepted. + #[tokio::test] + async fn test_oidc_audience_match_accepted() { + let mut config = test_oidc_config(); + config.audience = Some("my-client-id".to_string()); + let oidc = test_oidc_state_with_config(config).await; + let jwt = encode_test_jwt( + serde_json::json!({ + "sub": "alice", + "aud": "my-client-id", + "exp": 9999999999u64, + }), + Some(OIDC_KID), + ); + let result = validate_oidc_jwt(&oidc, &jwt).await; + assert!(result.is_ok(), "matching audience should pass: {result:?}"); + } + + /// Audience configured but JWT has wrong `aud` → rejected. + #[tokio::test] + async fn test_oidc_audience_mismatch_rejected() { + let mut config = test_oidc_config(); + config.audience = Some("my-client-id".to_string()); + let oidc = test_oidc_state_with_config(config).await; + let jwt = encode_test_jwt( + serde_json::json!({ + "sub": "alice", + "aud": "wrong-client", + "exp": 9999999999u64, + }), + Some(OIDC_KID), + ); + let result = validate_oidc_jwt(&oidc, &jwt).await; + assert!(result.is_err(), "wrong audience should be rejected"); + } + + /// Audience configured but JWT omits `aud` entirely. + /// + /// Note: `jsonwebtoken` v9 only validates `aud` when present; a missing + /// `aud` claim passes validation even with `set_audience` called. This + /// test documents that behavior. If we need to enforce `aud` presence, + /// add an explicit check in `validate_oidc_jwt` after claim extraction. + #[tokio::test] + async fn test_oidc_audience_configured_but_missing_in_jwt_passes() { + let mut config = test_oidc_config(); + config.audience = Some("my-client-id".to_string()); + let oidc = test_oidc_state_with_config(config).await; + let jwt = encode_test_jwt( + serde_json::json!({"sub": "alice", "exp": 9999999999u64}), + Some(OIDC_KID), + ); + let result = validate_oidc_jwt(&oidc, &jwt).await; + // jsonwebtoken allows missing aud — only rejects mismatches. + assert!( + result.is_ok(), + "missing aud is not rejected by jsonwebtoken: {result:?}" + ); + } + + // ── Key cache edge cases ───────────────────────────────────────────── + + /// Cache eviction: the `get_or_fetch_key` path evicts expired entries + /// and the oldest entry when the cache is full. We test this by + /// pre-filling the cache with expired entries and verifying they're + /// cleaned up when a new key is fetched (via cache hit on a valid key). + #[tokio::test] + async fn test_oidc_key_cache_evicts_expired_entries() { + let oidc = test_oidc_state().await; + + // Insert an expired entry with a manually backdated timestamp. + { + let mut cache = oidc.key_cache.write().await; + cache.insert( + "stale-kid".to_string(), + CachedKey { + decoding_key: DecodingKey::from_secret(OIDC_SECRET), + algorithm: Algorithm::HS256, + fetched_at: Instant::now() - KEY_CACHE_TTL - Duration::from_secs(1), + }, + ); + } + + // The valid test key (OIDC_KID) is fresh. Validate a JWT to + // trigger the get_or_fetch_key cache-hit path — the expired + // entry won't be evicted on a pure cache hit (eviction only + // runs on the fetch path). Verify the stale entry is expired. + { + let cache = oidc.key_cache.read().await; + let stale = cache.get("stale-kid").unwrap(); + assert!( + stale.fetched_at.elapsed() > KEY_CACHE_TTL, + "entry should be expired" + ); + } + + // A JWT using the stale kid should fail (expired cache entry + // is not served from cache). + let jwt = encode_test_jwt( + serde_json::json!({"sub": "stale-user", "exp": 9999999999u64}), + Some("stale-kid"), + ); + let result = validate_oidc_jwt(&oidc, &jwt).await; + assert!( + result.is_err(), + "expired cache entry should not be served; fetch fails since URL is unreachable" + ); + } + + /// Cache max entries: verify the constant is reasonable and that the + /// cache can hold exactly KEY_CACHE_MAX_ENTRIES via seed_key. + #[tokio::test] + async fn test_oidc_key_cache_max_entries_constant() { + assert_eq!( + KEY_CACHE_MAX_ENTRIES, 64, + "cache should be bounded to 64 keys" + ); + + let oidc = test_oidc_state().await; + for i in 0..KEY_CACHE_MAX_ENTRIES { + oidc.seed_key( + &format!("kid-{i}"), + DecodingKey::from_secret(OIDC_SECRET), + Algorithm::HS256, + ) + .await; + } + let cache = oidc.key_cache.read().await; + // seed_key + the default test key = MAX+1, but seed_key doesn't evict. + // The point is get_or_fetch_key's eviction path — tested indirectly + // via the fetch-failure and expired-entry tests above. + assert!( + cache.len() <= KEY_CACHE_MAX_ENTRIES + 1, + "cache should be near capacity" + ); + } + + /// Fetch failure backoff: a failed kid is backed off for FETCH_FAILURE_BACKOFF. + #[tokio::test] + async fn test_oidc_fetch_failure_backoff() { + let oidc = test_oidc_state().await; + + // Simulate a failed fetch by inserting into the failure tracker. + { + let mut failures = oidc.fetch_failures.write().await; + failures.insert( + "bad-kid".to_string(), + FailedFetch { + failed_at: Instant::now(), + }, + ); + } + + // Trying to get the key for that kid should immediately fail with + // backoff error, without attempting an HTTP request. + let result = oidc.get_or_fetch_key("bad-kid", Algorithm::HS256).await; + let err_msg = match result { + Err(e) => format!("{e}"), + Ok(_) => panic!("expected backoff error"), + }; + assert!( + err_msg.contains("backing off"), + "should mention backoff: {err_msg}" + ); + } + + /// After backoff expires, a new fetch is attempted (failure is cleared). + #[tokio::test] + async fn test_oidc_fetch_failure_backoff_expires() { + let oidc = test_oidc_state().await; + + // Insert a failure that's already past the backoff window. + { + let mut failures = oidc.fetch_failures.write().await; + failures.insert( + "expired-kid".to_string(), + FailedFetch { + failed_at: Instant::now() - FETCH_FAILURE_BACKOFF - Duration::from_secs(1), + }, + ); + } + + // This will attempt an actual HTTP fetch (which will fail since the + // URL is unreachable), but it should NOT be blocked by backoff. + let result = oidc.get_or_fetch_key("expired-kid", Algorithm::HS256).await; + let err_msg = match result { + Err(e) => format!("{e}"), + Ok(_) => panic!("expected fetch error (URL unreachable), not success"), + }; + assert!( + !err_msg.contains("backing off"), + "should attempt fetch, not backoff: {err_msg}" + ); + } } diff --git a/src/channels/web/mod.rs b/src/channels/web/mod.rs index 77968223c96..b4b23b4d20e 100644 --- a/src/channels/web/mod.rs +++ b/src/channels/web/mod.rs @@ -83,9 +83,27 @@ impl GatewayChannel { bytes.iter().map(|b| format!("{b:02x}")).collect() }); + let oidc_state = config.oidc.as_ref().and_then(|oidc_config| { + match auth::OidcState::from_config(oidc_config) { + Ok(state) => { + tracing::info!( + header = %oidc_config.header, + jwks_url = %oidc_config.jwks_url, + "OIDC JWT authentication enabled" + ); + Some(state) + } + Err(e) => { + tracing::error!(error = %e, "Failed to initialize OIDC auth — falling back to token-only auth"); + None + } + } + }); + let auth = CombinedAuthState { env_auth: MultiAuthState::single(auth_token, owner_id.clone()), db_auth: None, + oidc: oidc_state, }; let state = Arc::new(GatewayState { diff --git a/src/channels/web/static/app.js b/src/channels/web/static/app.js index 76084168647..2f4b972fc1e 100644 --- a/src/channels/web/static/app.js +++ b/src/channels/web/static/app.js @@ -73,6 +73,7 @@ document.getElementById('settings-theme-toggle')?.addEventListener('click', () = }); let token = ''; +let oidcProxyAuth = false; let eventSource = null; let logEventSource = null; let currentTab = 'chat'; @@ -140,6 +141,55 @@ let _activityThinking = null; // --- Auth --- +// Common post-auth initialization shared by token auth and OIDC auto-auth. +function initApp() { + var authScreen = document.getElementById('auth-screen'); + var app = document.getElementById('app'); + // Cross-fade: fade out auth screen, then show app + if (authScreen) authScreen.style.opacity = '0'; + // Show app container (invisible — opacity:0 in CSS) so layout computes + app.style.display = 'flex'; + // Position tab indicator instantly (no transition) before fade-in + var indicator = document.getElementById('tab-indicator'); + if (indicator) indicator.style.transition = 'none'; + updateTabIndicator(); + // Force layout so the instant position is applied, then restore transition + if (indicator) { + void indicator.offsetLeft; + indicator.style.transition = ''; + } + // Now fade in + app.classList.add('visible'); + // Hide auth screen after fade-out transition completes + setTimeout(function() { if (authScreen) authScreen.style.display = 'none'; }, 300); + // Strip token and log_level from URL so they're not visible in the address bar + var cleaned = new URL(window.location); + var urlLogLevel = cleaned.searchParams.get('log_level'); + cleaned.searchParams.delete('token'); + cleaned.searchParams.delete('log_level'); + window.history.replaceState({}, '', cleaned.pathname + cleaned.search); + connectSSE(); + connectLogSSE(); + startGatewayStatusPolling(); + // Hide the Users settings tab for non-admin users. + apiFetch('/api/profile').then(function(profile) { + if (profile && profile.role !== 'admin') { + var usersTab = document.querySelector('[data-settings-subtab="users"]'); + if (usersTab) usersTab.style.display = 'none'; + } + }).catch(function() {}); + checkTeeStatus(); + loadThreads(); + loadMemoryTree(); + loadJobs(); + // Apply URL log_level param if present, otherwise just sync the dropdown + if (urlLogLevel) { + setServerLogLevel(urlLogLevel); + } else { + loadServerLogLevel(); + } +} + function authenticate() { token = document.getElementById('token-input').value.trim(); if (!token) { @@ -158,51 +208,7 @@ function authenticate() { apiFetch('/api/chat/threads') .then(() => { sessionStorage.setItem('ironclaw_token', token); - const authScreen = document.getElementById('auth-screen'); - const app = document.getElementById('app'); - // Cross-fade: fade out auth screen, then show app - if (authScreen) authScreen.style.opacity = '0'; - // Show app container (invisible — opacity:0 in CSS) so layout computes - app.style.display = 'flex'; - // Position tab indicator instantly (no transition) before fade-in - const indicator = document.getElementById('tab-indicator'); - if (indicator) indicator.style.transition = 'none'; - updateTabIndicator(); - // Force layout so the instant position is applied, then restore transition - if (indicator) { - void indicator.offsetLeft; - indicator.style.transition = ''; - } - // Now fade in - app.classList.add('visible'); - // Hide auth screen after fade-out transition completes - setTimeout(() => { if (authScreen) authScreen.style.display = 'none'; }, 300); - // Strip token and log_level from URL so they're not visible in the address bar - const cleaned = new URL(window.location); - const urlLogLevel = cleaned.searchParams.get('log_level'); - cleaned.searchParams.delete('token'); - cleaned.searchParams.delete('log_level'); - window.history.replaceState({}, '', cleaned.pathname + cleaned.search); - connectSSE(); - connectLogSSE(); - startGatewayStatusPolling(); - // Hide the Users settings tab for non-admin users. - apiFetch('/api/profile').then(function(profile) { - if (profile && profile.role !== 'admin') { - var usersTab = document.querySelector('[data-settings-subtab="users"]'); - if (usersTab) usersTab.style.display = 'none'; - } - }).catch(function() {}); - checkTeeStatus(); - loadThreads(); - loadMemoryTree(); - loadJobs(); - // Apply URL log_level param if present, otherwise just sync the dropdown - if (urlLogLevel) { - setServerLogLevel(urlLogLevel); - } else { - loadServerLogLevel(); - } + initApp(); }) .catch(() => { sessionStorage.removeItem('ironclaw_token'); @@ -225,7 +231,12 @@ document.getElementById('token-input').addEventListener('keydown', (e) => { // Note: main event listener registration is at the bottom of this file (search // "Event Listener Registration"). Do NOT add duplicate listeners here. -// Auto-authenticate from URL param or saved session +// Auto-authenticate from URL param, saved session, or OIDC proxy header. +// +// When behind a reverse proxy that injects auth (e.g., AWS ALB with OIDC), +// the proxy already authenticates every request. We probe /api/gateway/status +// without a token — if the proxy's header lets us through, skip the login +// screen entirely. (function autoAuth() { const params = new URLSearchParams(window.location.search); const urlToken = params.get('token'); @@ -234,15 +245,28 @@ document.getElementById('token-input').addEventListener('keydown', (e) => { authenticate(); return; } + // Restore OIDC proxy mode from session. + if (sessionStorage.getItem('ironclaw_oidc') === '1') { + oidcProxyAuth = true; + } const saved = sessionStorage.getItem('ironclaw_token'); if (saved) { document.getElementById('token-input').value = saved; - // Hide auth screen immediately to prevent flash, authenticate() will - // restore it if the token turns out to be invalid. document.getElementById('auth-screen').style.display = 'none'; document.getElementById('app').style.display = 'flex'; authenticate(); + return; } + // Probe for proxy-injected OIDC auth (no token needed from the client). + fetch('/api/gateway/status', { credentials: 'include' }).then(function(r) { + if (r.ok) { + oidcProxyAuth = true; + sessionStorage.setItem('ironclaw_oidc', '1'); + document.getElementById('auth-screen').style.display = 'none'; + document.getElementById('app').style.display = 'flex'; + initApp(); + } + }).catch(function() { /* proxy auth not available, show login */ }); })(); // --- API helper --- @@ -250,7 +274,10 @@ document.getElementById('token-input').addEventListener('keydown', (e) => { function apiFetch(path, options) { const opts = options || {}; opts.headers = opts.headers || {}; - opts.headers['Authorization'] = 'Bearer ' + token; + // In OIDC mode the reverse proxy provides auth; skip the Authorization header. + if (token && !oidcProxyAuth) { + opts.headers['Authorization'] = 'Bearer ' + token; + } if (opts.body && typeof opts.body === 'object') { opts.headers['Content-Type'] = 'application/json'; opts.body = JSON.stringify(opts.body); @@ -361,7 +388,11 @@ function updateRestartButtonVisibility() { function connectSSE() { if (eventSource) eventSource.close(); - eventSource = new EventSource('/api/chat/events?token=' + encodeURIComponent(token)); + // In OIDC mode the reverse proxy provides auth; no query token needed. + const chatSseUrl = (token && !oidcProxyAuth) + ? '/api/chat/events?token=' + encodeURIComponent(token) + : '/api/chat/events'; + eventSource = new EventSource(chatSseUrl); eventSource.onopen = () => { document.getElementById('sse-dot').classList.remove('disconnected'); @@ -2497,7 +2528,10 @@ let logBuffer = []; // buffer while paused function connectLogSSE() { if (logEventSource) logEventSource.close(); - logEventSource = new EventSource('/api/logs/events?token=' + encodeURIComponent(token)); + const logSseUrl = (token && !oidcProxyAuth) + ? '/api/logs/events?token=' + encodeURIComponent(token) + : '/api/logs/events'; + logEventSource = new EventSource(logSseUrl); logEventSource.addEventListener('log', (e) => { const entry = JSON.parse(e.data); diff --git a/src/channels/web/tests/no_silent_drop.rs b/src/channels/web/tests/no_silent_drop.rs index 5ffd9f04cf2..a86533939c6 100644 --- a/src/channels/web/tests/no_silent_drop.rs +++ b/src/channels/web/tests/no_silent_drop.rs @@ -17,6 +17,7 @@ fn test_gateway() -> GatewayChannel { auth_token: Some("test-token".to_string()), workspace_read_scopes: vec![], memory_layers: vec![], + oidc: None, }, "test-user".to_string(), ) diff --git a/src/config/channels.rs b/src/config/channels.rs index dec04f398c8..74f98dbfca4 100644 --- a/src/config/channels.rs +++ b/src/config/channels.rs @@ -52,6 +52,26 @@ pub struct GatewayConfig { pub workspace_read_scopes: Vec, /// Memory layer definitions (JSON in env var, or from external config). pub memory_layers: Vec, + /// OIDC JWT authentication (e.g., behind AWS ALB with Okta). + pub oidc: Option, +} + +/// OIDC JWT authentication configuration for the web gateway. +/// +/// When enabled, the gateway accepts signed JWTs from a configurable HTTP +/// header (e.g., `x-amzn-oidc-data` from AWS ALB). Keys are fetched from +/// a JWKS endpoint and cached for 1 hour. +#[derive(Debug, Clone)] +pub struct GatewayOidcConfig { + /// HTTP header containing the JWT (default: `x-amzn-oidc-data`). + pub header: String, + /// JWKS URL for key discovery. Supports `{kid}` placeholder for + /// ALB-style per-key PEM endpoints, and standard `/.well-known/jwks.json`. + pub jwks_url: String, + /// Expected `iss` claim. Validated if set. + pub issuer: Option, + /// Expected `aud` claim. Validated if set. + pub audience: Option, } /// Signal channel configuration (signal-cli daemon HTTP/JSON-RPC). @@ -195,6 +215,24 @@ impl ChannelsConfig { }); } } + let oidc_enabled = parse_bool_env("GATEWAY_OIDC_ENABLED", false)?; + let oidc = if oidc_enabled { + let jwks_url = + optional_env("GATEWAY_OIDC_JWKS_URL")?.ok_or(ConfigError::InvalidValue { + key: "GATEWAY_OIDC_JWKS_URL".to_string(), + message: "required when GATEWAY_OIDC_ENABLED=true".to_string(), + })?; + Some(GatewayOidcConfig { + header: optional_env("GATEWAY_OIDC_HEADER")? + .unwrap_or_else(|| "x-amzn-oidc-data".to_string()), + jwks_url, + issuer: optional_env("GATEWAY_OIDC_ISSUER")?, + audience: optional_env("GATEWAY_OIDC_AUDIENCE")?, + }) + } else { + None + }; + Some(GatewayConfig { host: optional_env("GATEWAY_HOST")? .or_else(|| cs.gateway_host.clone()) @@ -207,6 +245,7 @@ impl ChannelsConfig { .or_else(|| cs.gateway_auth_token.clone()), workspace_read_scopes, memory_layers, + oidc, }) } else { None @@ -363,6 +402,7 @@ mod tests { auth_token: Some("tok-abc".to_string()), workspace_read_scopes: vec![], memory_layers: vec![], + oidc: None, }; assert_eq!(cfg.host, "127.0.0.1"); assert_eq!(cfg.port, 3000); @@ -377,6 +417,7 @@ mod tests { auth_token: None, workspace_read_scopes: vec![], memory_layers: vec![], + oidc: None, }; assert!(cfg.auth_token.is_none()); } diff --git a/src/config/mod.rs b/src/config/mod.rs index 03f37c5dec9..ed9b6a5fff9 100644 --- a/src/config/mod.rs +++ b/src/config/mod.rs @@ -40,7 +40,8 @@ use crate::settings::Settings; pub use self::agent::AgentConfig; pub use self::builder::BuilderModeConfig; pub use self::channels::{ - ChannelsConfig, CliConfig, DEFAULT_GATEWAY_PORT, GatewayConfig, HttpConfig, SignalConfig, + ChannelsConfig, CliConfig, DEFAULT_GATEWAY_PORT, GatewayConfig, GatewayOidcConfig, HttpConfig, + SignalConfig, }; pub use self::database::{DatabaseBackend, DatabaseConfig, SslMode, default_libsql_path}; pub use self::embeddings::{DEFAULT_EMBEDDING_CACHE_SIZE, EmbeddingsConfig}; diff --git a/src/tunnel/mod.rs b/src/tunnel/mod.rs index e73fcd46992..773bb05e87f 100644 --- a/src/tunnel/mod.rs +++ b/src/tunnel/mod.rs @@ -429,6 +429,7 @@ mod tests { port: 3000, auth_token: None, workspace_read_scopes: Vec::new(), + oidc: None, memory_layers: Vec::new(), }); c @@ -442,6 +443,7 @@ mod tests { auth_token: None, workspace_read_scopes: Vec::new(), memory_layers: Vec::new(), + oidc: None, }); c } From 8acdd08039071b731fc7fd6be8b6e1c4c18da9c7 Mon Sep 17 00:00:00 2001 From: Andriy Samilyak <359393+werdan@users.noreply.github.com> Date: Mon, 30 Mar 2026 04:40:25 +0200 Subject: [PATCH 06/11] fix(wasm): inject Content-Length: 0 for bodyless mutating HTTP requests (#1529) * fix(wasm): inject Content-Length: 0 for bodyless mutating requests [skip-version-check] The WASM host http_request now auto-injects Content-Length: 0 for POST/PUT/PATCH/DELETE requests with no body, unless the tool already provides the header. This fixes Gmail returning 411 on trash_message and proactively covers all other tools (Google Calendar DELETE, Google Drive DELETE, etc.). Extracted needs_content_length_zero() with 8 regression tests covering all HTTP methods and case-insensitive header detection. Co-Authored-By: Claude Opus 4.6 (1M context) * fix(wasm): use eq_ignore_ascii_case to avoid allocation [skip-version-check] Replace matches!(method.to_uppercase().as_str(), ...) with eq_ignore_ascii_case() to avoid a per-request String allocation. Co-Authored-By: Claude Opus 4.6 (1M context) --------- Co-authored-by: ilblackdragon@gmail.com Co-authored-by: Claude Opus 4.6 (1M context) --- src/tools/wasm/wrapper.rs | 73 +++++++++++++++++++++++++++++++++++++-- 1 file changed, 71 insertions(+), 2 deletions(-) diff --git a/src/tools/wasm/wrapper.rs b/src/tools/wasm/wrapper.rs index 66dc3cc1771..ee0c047c4d9 100644 --- a/src/tools/wasm/wrapper.rs +++ b/src/tools/wasm/wrapper.rs @@ -431,12 +431,14 @@ impl near::agent::host::Host for StoreData { _ => return Err(format!("Unsupported HTTP method: {}", method)), }; - for (key, value) in headers { - request = request.header(&key, &value); + for (key, value) in &headers { + request = request.header(key, value); } if let Some(body_bytes) = body { request = request.body(body_bytes); + } else if needs_content_length_zero(&method, &headers) { + request = request.header("content-length", "0"); } // Caller-specified timeout (default 30s, max 5min) @@ -1760,6 +1762,20 @@ fn build_tool_usage_hint(tool_name: &str, schema: &serde_json::Value) -> String hint } +/// Methods with side effects require `Content-Length` even when no body is +/// sent — some APIs (e.g. Gmail) return 411 without it. Returns `true` when +/// the host should inject a `Content-Length: 0` header. +fn needs_content_length_zero(method: &str, headers: &HashMap) -> bool { + let mutating = method.eq_ignore_ascii_case("POST") + || method.eq_ignore_ascii_case("PUT") + || method.eq_ignore_ascii_case("PATCH") + || method.eq_ignore_ascii_case("DELETE"); + mutating + && !headers + .iter() + .any(|(k, _)| k.eq_ignore_ascii_case("content-length")) +} + #[cfg(test)] mod tests { use std::collections::HashMap; @@ -3259,4 +3275,57 @@ mod tests { // Should return empty since credential can't be found anywhere assert!(result.is_empty(), "no credentials found"); // safety: test code only } + + // --- needs_content_length_zero (regression for #1529) --- + + #[test] + fn post_no_body_needs_content_length() { + let headers = HashMap::new(); + assert!( + super::needs_content_length_zero("POST", &headers), + "POST with no body must get Content-Length: 0 to avoid 411" + ); + } + + #[test] + fn put_no_body_needs_content_length() { + assert!(super::needs_content_length_zero("PUT", &HashMap::new())); + } + + #[test] + fn delete_no_body_needs_content_length() { + assert!(super::needs_content_length_zero("DELETE", &HashMap::new())); + } + + #[test] + fn patch_no_body_needs_content_length() { + assert!(super::needs_content_length_zero("PATCH", &HashMap::new())); + } + + #[test] + fn get_no_body_skips_content_length() { + assert!(!super::needs_content_length_zero("GET", &HashMap::new())); + } + + #[test] + fn head_no_body_skips_content_length() { + assert!(!super::needs_content_length_zero("HEAD", &HashMap::new())); + } + + #[test] + fn post_no_body_respects_explicit_content_length() { + let mut headers = HashMap::new(); + headers.insert("Content-Length".to_string(), "0".to_string()); + assert!( + !super::needs_content_length_zero("POST", &headers), + "should not double-add when tool already sets Content-Length" + ); + } + + #[test] + fn content_length_check_is_case_insensitive() { + let mut headers = HashMap::new(); + headers.insert("content-length".to_string(), "0".to_string()); + assert!(!super::needs_content_length_zero("POST", &headers)); + } } From c75dea0e4ee4c4f824f9216fa4dbfff2fbe3d600 Mon Sep 17 00:00:00 2001 From: Nige Date: Mon, 30 Mar 2026 07:17:54 +0100 Subject: [PATCH 07/11] fix(auth): make shared Google tool status scope-aware (#1532) MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit * fix(auth): make shared Google tool status scope-aware * fix(auth): simplify google docs auth status test * fix(auth): skip scope expansion for env-var tokens and add dual-source test Env-var-provided tokens are externally managed, so the scope-expansion check must not apply — otherwise tools regress to NeedsAuth when no scopes record exists in the secrets store. Split the token detection into managed vs env-var paths and only run scope checks for managed tokens. Also adds tests verifying: (1) env-var-only tokens return Ready without scope checks, and (2) when both a managed token and env var are present, the managed path with scope checks takes priority. [skip-regression-check] Co-Authored-By: Claude Opus 4.6 (1M context) --------- Co-authored-by: ilblackdragon@gmail.com Co-authored-by: Claude Opus 4.6 (1M context) --- src/extensions/manager.rs | 299 ++++++++++++++++++++++++++++++++++++-- 1 file changed, 289 insertions(+), 10 deletions(-) diff --git a/src/extensions/manager.rs b/src/extensions/manager.rs index 82dc5471fc1..c4d4e2ba6d0 100644 --- a/src/extensions/manager.rs +++ b/src/extensions/manager.rs @@ -3839,18 +3839,40 @@ impl ExtensionManager { // authoritative signal — setup secrets (client_id/secret) are // intermediate and may be auto-resolved via builtins. if let Some(ref auth) = cap_file.auth { - let has_token = self + let token_is_managed = self .secrets .exists(user_id, &auth.secret_name) .await - .unwrap_or(false) - || auth - .env_var - .as_ref() - .is_some_and(|v| std::env::var(v).is_ok()); - return if has_token { - ToolAuthState::Ready - } else if auth.oauth.is_some() { + .unwrap_or(false); + let has_env_token = auth + .env_var + .as_ref() + .is_some_and(|v| std::env::var(v).is_ok()); + + if token_is_managed { + // Token lives in the secrets store — check whether the merged + // scope set of all tools sharing this secret is satisfied. + if let Some(ref oauth) = auth.oauth { + let merged = self + .collect_shared_scopes(&auth.secret_name, &oauth.scopes, user_id) + .await; + if self + .needs_scope_expansion(&auth.secret_name, &merged, user_id) + .await + { + return ToolAuthState::NeedsAuth; + } + } + return ToolAuthState::Ready; + } + + if has_env_token { + // Externally-managed token (env var) — skip scope checks; + // the user is responsible for granting adequate scopes. + return ToolAuthState::Ready; + } + + return if auth.oauth.is_some() { ToolAuthState::NeedsAuth } else { ToolAuthState::NeedsSetup @@ -6335,7 +6357,8 @@ mod tests { telegram_message_matches_verification_code, }; use crate::extensions::{ - ExtensionError, ExtensionKind, ExtensionSource, InstallResult, VerificationChallenge, + ExtensionError, ExtensionKind, ExtensionSource, InstallResult, ToolAuthState, + VerificationChallenge, }; use crate::pairing::PairingStore; use crate::secrets::CreateSecretParams; @@ -9021,4 +9044,260 @@ mod tests { "non-builtin provider secret must be kept" ); } + + #[tokio::test] + async fn test_shared_google_oauth_status_requires_scope_expansion_for_second_tool() + -> Result<(), String> { + let dir = tempfile::tempdir().map_err(|err| format!("temp dir: {err}"))?; + let tools_dir = dir.path().join("tools"); + std::fs::create_dir_all(&tools_dir).map_err(|err| format!("tools dir: {err}"))?; + + let name = "google-docs"; + let scope = "https://www.googleapis.com/auth/documents"; + std::fs::write(tools_dir.join(format!("{name}.wasm")), b"\0asm") + .map_err(|err| format!("write {name}.wasm: {err}"))?; + + let caps = serde_json::json!({ + "auth": { + "secret_name": "google_oauth_token", + "display_name": "Google", + "oauth": { + "authorization_url": "https://accounts.google.com/o/oauth2/v2/auth", + "token_url": "https://oauth2.googleapis.com/token", + "client_id_env": "GOOGLE_OAUTH_CLIENT_ID", + "client_secret_env": "GOOGLE_OAUTH_CLIENT_SECRET", + "scopes": [scope], + "use_pkce": false, + "extra_params": { + "access_type": "offline", + "prompt": "consent" + } + }, + "env_var": "GOOGLE_OAUTH_TOKEN" + } + }); + std::fs::write( + tools_dir.join(format!("{name}.capabilities.json")), + serde_json::to_vec(&caps).map_err(|err| format!("serialize {name}: {err}"))?, + ) + .map_err(|err| format!("write {name}.capabilities.json: {err}"))?; + + let mgr = make_test_manager(None, tools_dir.clone()); + mgr.secrets + .create( + "test", + crate::secrets::CreateSecretParams::new("google_oauth_token", "token") + .with_provider("google-docs"), + ) + .await + .map_err(|err| format!("store token: {err}"))?; + mgr.secrets + .create( + "test", + crate::secrets::CreateSecretParams::new( + "google_oauth_token_scopes", + "https://www.googleapis.com/auth/documents", + ) + .with_provider("google-docs"), + ) + .await + .map_err(|err| format!("store scopes: {err}"))?; + + assert_eq!( + mgr.check_tool_auth_status("google-docs", "test").await, + ToolAuthState::Ready + ); + + std::fs::write(tools_dir.join("google-slides.wasm"), b"\0asm") + .map_err(|err| format!("write google-slides.wasm: {err}"))?; + let slides_caps = serde_json::json!({ + "auth": { + "secret_name": "google_oauth_token", + "display_name": "Google", + "oauth": { + "authorization_url": "https://accounts.google.com/o/oauth2/v2/auth", + "token_url": "https://oauth2.googleapis.com/token", + "client_id_env": "GOOGLE_OAUTH_CLIENT_ID", + "client_secret_env": "GOOGLE_OAUTH_CLIENT_SECRET", + "scopes": ["https://www.googleapis.com/auth/presentations"], + "use_pkce": false, + "extra_params": { + "access_type": "offline", + "prompt": "consent" + } + }, + "env_var": "GOOGLE_OAUTH_TOKEN" + } + }); + std::fs::write( + tools_dir.join("google-slides.capabilities.json"), + serde_json::to_vec(&slides_caps) + .map_err(|err| format!("serialize google-slides: {err}"))?, + ) + .map_err(|err| format!("write google-slides.capabilities.json: {err}"))?; + + assert_eq!( + mgr.check_tool_auth_status("google-docs", "test").await, + ToolAuthState::NeedsAuth, + "adding the second shared-auth Google tool should require reauth for the existing tool" + ); + assert_eq!( + mgr.check_tool_auth_status("google-slides", "test").await, + ToolAuthState::NeedsAuth, + "second Google tool should require scope expansion when the shared token lacks its scope", + ); + + Ok(()) + } + + /// Env-var-provided tokens must always return Ready — the user manages + /// scopes externally, so the scope-expansion check must not apply. + /// Uses `HOME` as env_var since it always exists, avoiding `set_var` + /// which is unsafe in multi-threaded test runs. + #[tokio::test] + async fn test_env_var_token_skips_scope_expansion() -> Result<(), String> { + let dir = tempfile::tempdir().map_err(|err| format!("temp dir: {err}"))?; + let tools_dir = dir.path().join("tools"); + std::fs::create_dir_all(&tools_dir).map_err(|err| format!("tools dir: {err}"))?; + + let caps = serde_json::json!({ + "auth": { + "secret_name": "google_oauth_token", + "display_name": "Google", + "oauth": { + "authorization_url": "https://accounts.google.com/o/oauth2/v2/auth", + "token_url": "https://oauth2.googleapis.com/token", + "client_id_env": "GOOGLE_OAUTH_CLIENT_ID", + "client_secret_env": "GOOGLE_OAUTH_CLIENT_SECRET", + "scopes": ["https://www.googleapis.com/auth/documents"], + "use_pkce": false, + "extra_params": { + "access_type": "offline", + "prompt": "consent" + } + }, + "env_var": "HOME" + } + }); + std::fs::write(tools_dir.join("google-docs.wasm"), b"\0asm") + .map_err(|err| format!("write wasm: {err}"))?; + std::fs::write( + tools_dir.join("google-docs.capabilities.json"), + serde_json::to_vec(&caps).map_err(|err| format!("serialize: {err}"))?, + ) + .map_err(|err| format!("write caps: {err}"))?; + + // No managed token in secrets store — only the env var (HOME) is present. + let mgr = make_test_manager(None, tools_dir); + + assert_eq!( + mgr.check_tool_auth_status("google-docs", "test").await, + ToolAuthState::Ready, + "env-var token should be Ready without scope expansion check" + ); + + Ok(()) + } + + /// When both a managed token AND an env-var token exist, the managed + /// path (with scope expansion checks) must take priority. + #[tokio::test] + async fn test_managed_token_takes_priority_over_env_var() -> Result<(), String> { + let dir = tempfile::tempdir().map_err(|err| format!("temp dir: {err}"))?; + let tools_dir = dir.path().join("tools"); + std::fs::create_dir_all(&tools_dir).map_err(|err| format!("tools dir: {err}"))?; + + // Both tools point env_var at HOME (always set) so the env-var path + // would return Ready — but the managed token path should win. + let caps = serde_json::json!({ + "auth": { + "secret_name": "google_oauth_token", + "display_name": "Google", + "oauth": { + "authorization_url": "https://accounts.google.com/o/oauth2/v2/auth", + "token_url": "https://oauth2.googleapis.com/token", + "client_id_env": "GOOGLE_OAUTH_CLIENT_ID", + "client_secret_env": "GOOGLE_OAUTH_CLIENT_SECRET", + "scopes": ["https://www.googleapis.com/auth/documents"], + "use_pkce": false, + "extra_params": { + "access_type": "offline", + "prompt": "consent" + } + }, + "env_var": "HOME" + } + }); + std::fs::write(tools_dir.join("google-docs.wasm"), b"\0asm") + .map_err(|err| format!("write wasm: {err}"))?; + std::fs::write( + tools_dir.join("google-docs.capabilities.json"), + serde_json::to_vec(&caps).map_err(|err| format!("serialize: {err}"))?, + ) + .map_err(|err| format!("write caps: {err}"))?; + + // Second tool requires an additional scope. + let slides_caps = serde_json::json!({ + "auth": { + "secret_name": "google_oauth_token", + "display_name": "Google", + "oauth": { + "authorization_url": "https://accounts.google.com/o/oauth2/v2/auth", + "token_url": "https://oauth2.googleapis.com/token", + "client_id_env": "GOOGLE_OAUTH_CLIENT_ID", + "client_secret_env": "GOOGLE_OAUTH_CLIENT_SECRET", + "scopes": ["https://www.googleapis.com/auth/presentations"], + "use_pkce": false, + "extra_params": { + "access_type": "offline", + "prompt": "consent" + } + }, + "env_var": "HOME" + } + }); + std::fs::write(tools_dir.join("google-slides.wasm"), b"\0asm") + .map_err(|err| format!("write wasm: {err}"))?; + std::fs::write( + tools_dir.join("google-slides.capabilities.json"), + serde_json::to_vec(&slides_caps).map_err(|err| format!("serialize: {err}"))?, + ) + .map_err(|err| format!("write caps: {err}"))?; + + let mgr = make_test_manager(None, tools_dir); + + // Store a managed token with only the docs scope. + mgr.secrets + .create( + "test", + crate::secrets::CreateSecretParams::new("google_oauth_token", "managed-token") + .with_provider("google-docs"), + ) + .await + .map_err(|err| format!("store token: {err}"))?; + mgr.secrets + .create( + "test", + crate::secrets::CreateSecretParams::new( + "google_oauth_token_scopes", + "https://www.googleapis.com/auth/documents", + ) + .with_provider("google-docs"), + ) + .await + .map_err(|err| format!("store scopes: {err}"))?; + + assert_eq!( + mgr.check_tool_auth_status("google-docs", "test").await, + ToolAuthState::NeedsAuth, + "managed token path must win: merged scopes unsatisfied despite env var being set" + ); + assert_eq!( + mgr.check_tool_auth_status("google-slides", "test").await, + ToolAuthState::NeedsAuth, + "slides scope missing from managed token even though env var is set" + ); + + Ok(()) + } } From 10d5a530a01e64bc90fee08bbca2a0e4f65c9bc7 Mon Sep 17 00:00:00 2001 From: Illia Polosukhin Date: Sun, 29 Mar 2026 23:51:39 -0700 Subject: [PATCH 08/11] fix: resolve 11 test failures from multi-tenant bootstrap and sandbox gate regressions (#1746) MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit * fix: resolve 11 test failures from multi-tenant bootstrap and sandbox gate regressions Three root causes fixed: 1. Per-user bootstrap greeting in tests: After the multi-tenant isolation PR, `tenant_ctx("test-user")` creates a per-user workspace that seeds BOOTSTRAP.md and triggers an unwanted bootstrap greeting. This threw off response counting and caused message-drain races in 7 e2e tests. Fix: pre-seed the "test-user" workspace in the test rig DB so the first tenant_ctx call finds existing documents. 2. Sandbox gate blocking full_job routines: The full_job reliability overhaul (#1650) intended to remove the SandboxReadiness gate from execute_full_job (since full_job routines dispatch through the scheduler, not Docker). The gate was accidentally re-added during rebase, breaking 4 routine tests. Fix: remove the gate and clean up the unused sandbox_readiness field from EngineContext. 3. Owner-gate tests expecting old failure path: Two tests expected RunStatus::Failed from the sandbox gate. With the gate removed, the tool is now blocked at execution time by the approval context and the job completes normally. Fix: update traces and assertions to match the new behavior (RunStatus::Ok, owner_gate_count == 0). Co-Authored-By: Claude Opus 4.6 (1M context) * fix: keep DockerUnavailable gate for full_job routines Only remove the DisabledByConfig gate — when sandbox is enabled but Docker is unavailable, full_job routines should still fail rather than silently running without the expected sandbox isolation. Co-Authored-By: Claude Opus 4.6 (1M context) * refactor: address PR review feedback - Remove unused `include_completion` param from `owner_gate_trace()` and update all 5 call sites - Use `.expect()` instead of `let _ =` on `seed_if_empty()` in test rig to surface seeding failures early - Rename owner-gate tests from `_blocks_` to `_denies_tool_` to clarify the denial-with-success semantics Co-Authored-By: Claude Opus 4.6 (1M context) * refactor: let per-user bootstrap fire naturally, filter in TestRig Instead of pre-seeding the "test-user" workspace to prevent the per-user bootstrap greeting, let it happen naturally and make the TestRig resilient to it. `wait_for_responses` now transparently filters bootstrap greetings from the response stream: - Normal tests: all greetings filtered (bootstrap_greetings_to_keep=0) - `.with_bootstrap()` tests: 1 greeting kept (the startup greeting), additional per-user duplicates filtered Also updates owner-gate test section headers to match the denial-with-success semantics. Co-Authored-By: Claude Opus 4.6 (1M context) * fix: assert tool denial event in owner-gate tests The owner-gate denial tests previously only checked RunStatus::Ok + owner_gate_count == 0, which could pass if the tool was never called at all. Now both tests also verify that a tool_result event with success=false exists for "owner_gate" in the job's event log, confirming the tool was attempted and blocked by the approval context. Co-Authored-By: Claude Opus 4.6 (1M context) * style: apply rustfmt to collapsed function signature Co-Authored-By: Claude Opus 4.6 (1M context) * fix: use async lock in bootstrap filter loop, add DisabledByConfig unit test - Switch TestRig polling loop from `captured_responses()` (try_lock, panics on contention) to `captured_responses_async()` (async lock, safe under concurrent response pushing) - Add unit test asserting DisabledByConfig does NOT match the DockerUnavailable gate (verifies the intended behavior change) Co-Authored-By: Claude Opus 4.6 (1M context) --------- Co-authored-by: Claude Opus 4.6 (1M context) --- src/agent/routine_engine.rs | 57 ++++++------ tests/e2e_routine_heartbeat.rs | 153 ++++++++++++++++++++------------- tests/support/test_channel.rs | 6 ++ tests/support/test_rig.rs | 64 +++++++++++++- 4 files changed, 184 insertions(+), 96 deletions(-) diff --git a/src/agent/routine_engine.rs b/src/agent/routine_engine.rs index 3687ebd4f8e..65682e494cb 100644 --- a/src/agent/routine_engine.rs +++ b/src/agent/routine_engine.rs @@ -116,7 +116,7 @@ pub struct RoutineEngine { tools: Arc, /// Safety layer for tool output sanitization. safety: Arc, - /// Sandbox readiness state for full-job dispatch. + /// Sandbox readiness state — only `DockerUnavailable` blocks full-job dispatch. sandbox_readiness: SandboxReadiness, /// Timestamp when this engine instance was created. Used by /// `sync_dispatched_runs` to distinguish orphaned runs (from a previous @@ -1256,22 +1256,16 @@ async fn execute_full_job( run: &RoutineRun, execution: &FullJobExecutionConfig<'_>, ) -> Result<(RunStatus, Option, Option), RoutineError> { - match ctx.sandbox_readiness { - SandboxReadiness::Available => {} - SandboxReadiness::DisabledByConfig => { - return Err(RoutineError::JobDispatchFailed { - reason: "Sandboxing is disabled (SANDBOX_ENABLED=false). \ - Full-job routines require sandbox." - .to_string(), - }); - } - SandboxReadiness::DockerUnavailable => { - return Err(RoutineError::JobDispatchFailed { - reason: "Sandbox is enabled but Docker is not available. \ - Install Docker or set SANDBOX_ENABLED=false." - .to_string(), - }); - } + // Full-job routines dispatch through the scheduler (same as /job + // commands) — no Docker sandbox required when sandbox is disabled. + // However, if sandbox is *enabled* but Docker is unavailable, that's + // a misconfiguration we should surface. + if matches!(ctx.sandbox_readiness, SandboxReadiness::DockerUnavailable) { + return Err(RoutineError::JobDispatchFailed { + reason: "Sandbox is enabled but Docker is not available. \ + Install Docker or set SANDBOX_ENABLED=false." + .to_string(), + }); } let scheduler = ctx @@ -2410,28 +2404,26 @@ mod tests { } #[test] - fn test_sandbox_readiness_disabled_by_config_error() { + fn test_sandbox_disabled_by_config_does_not_block_full_job() { use super::SandboxReadiness; - let readiness = SandboxReadiness::DisabledByConfig; - assert_ne!(readiness, SandboxReadiness::Available); - - let err = crate::error::RoutineError::JobDispatchFailed { - reason: "Sandboxing is disabled (SANDBOX_ENABLED=false). \ - Full-job routines require sandbox." - .to_string(), - }; - let msg = err.to_string(); - assert!(msg.contains("SANDBOX_ENABLED=false")); - assert!(msg.contains("require sandbox")); + // DisabledByConfig must NOT match the DockerUnavailable gate — + // full-job routines dispatch through the scheduler (no Docker needed). + assert!(!matches!( + SandboxReadiness::DisabledByConfig, + SandboxReadiness::DockerUnavailable + )); } #[test] - fn test_sandbox_readiness_docker_unavailable_error() { + fn test_sandbox_readiness_docker_unavailable_still_blocks() { use super::SandboxReadiness; - let readiness = SandboxReadiness::DockerUnavailable; - assert_ne!(readiness, SandboxReadiness::Available); + // DockerUnavailable should still block full-job dispatch. + assert!(matches!( + SandboxReadiness::DockerUnavailable, + SandboxReadiness::DockerUnavailable + )); let err = crate::error::RoutineError::JobDispatchFailed { reason: "Sandbox is enabled but Docker is not available. \ @@ -2440,7 +2432,6 @@ mod tests { }; let msg = err.to_string(); assert!(msg.contains("Docker is not available")); - assert!(msg.contains("SANDBOX_ENABLED")); } /// Regression test for #1317: FullJobWatcher maps terminal job states correctly. diff --git a/tests/e2e_routine_heartbeat.rs b/tests/e2e_routine_heartbeat.rs index 36d87a07eb9..94dbc758298 100644 --- a/tests/e2e_routine_heartbeat.rs +++ b/tests/e2e_routine_heartbeat.rs @@ -198,37 +198,46 @@ mod tests { } } - fn owner_gate_trace(include_completion: bool) -> LlmTrace { - let mut steps = vec![TraceStep { - request_hint: None, - response: TraceResponse::ToolCalls { - tool_calls: vec![TraceToolCall { - id: "call_owner_gate".to_string(), - name: "owner_gate".to_string(), - arguments: serde_json::json!({}), - }], - input_tokens: 40, - output_tokens: 10, + fn owner_gate_trace() -> LlmTrace { + // The worker calls the LLM which returns a tool_call for owner_gate. + // After tool execution (success or blocked-by-approval error), the + // worker calls the LLM again. The worker first calls `select_tools()`, + // then falls back to `respond_with_tools()` when no tool calls are + // returned — both consume a trace step, so we always need two text + // responses after the tool call. + let steps = vec![ + TraceStep { + request_hint: None, + response: TraceResponse::ToolCalls { + tool_calls: vec![TraceToolCall { + id: "call_owner_gate".to_string(), + name: "owner_gate".to_string(), + arguments: serde_json::json!({}), + }], + input_tokens: 40, + output_tokens: 10, + }, + expected_tool_results: vec![], }, - expected_tool_results: vec![], - }]; - if include_completion { - // The worker first calls `select_tools()`, then falls back to - // `respond_with_tools()` when no tool calls are returned. Both - // methods consume a trace step, so the successful completion path - // needs two text responses after the tool call. - for _ in 0..2 { - steps.push(TraceStep { - request_hint: None, - response: TraceResponse::Text { - content: "I have completed the task.".to_string(), - input_tokens: 20, - output_tokens: 5, - }, - expected_tool_results: vec![], - }); - } - } + TraceStep { + request_hint: None, + response: TraceResponse::Text { + content: "I have completed the task.".to_string(), + input_tokens: 20, + output_tokens: 5, + }, + expected_tool_results: vec![], + }, + TraceStep { + request_hint: None, + response: TraceResponse::Text { + content: "I have completed the task.".to_string(), + input_tokens: 20, + output_tokens: 5, + }, + expected_tool_results: vec![], + }, + ]; LlmTrace::single_turn("test-owner-gate", "run owner gate", steps) } @@ -387,6 +396,33 @@ mod tests { } } + /// Wait for a `tool_result` job event with `success: false` for the given tool. + /// Job events are persisted via `tokio::spawn`, so they may lag slightly + /// behind run completion. + async fn wait_for_tool_denial_event(db: &Arc, job_id: Uuid, tool_name: &str) { + let deadline = std::time::Instant::now() + Duration::from_secs(5); + loop { + let events = db + .list_job_events(job_id, None) + .await + .expect("list_job_events"); + let denied = events.iter().any(|e| { + e.event_type == "tool_result" + && e.data.get("tool_name").and_then(|v| v.as_str()) == Some(tool_name) + && e.data.get("success") == Some(&serde_json::json!(false)) + }); + if denied { + return; + } + assert!( + std::time::Instant::now() < deadline, + "timed out waiting for tool denial event for '{tool_name}' in job {job_id}. \ + Events: {events:?}" + ); + tokio::time::sleep(Duration::from_millis(100)).await; + } + } + async fn wait_for_any_run_completion(db: &Arc, routine_id: Uuid) -> RoutineRun { let deadline = std::time::Instant::now() + Duration::from_secs(10); loop { @@ -1385,7 +1421,7 @@ mod tests { let tools_dir = tmp.path().join("wasm-tools"); let engine = setup_owner_gate_engine( db.clone(), - owner_gate_trace(true), + owner_gate_trace(), tools_dir.as_path(), Some("default"), true, @@ -1461,7 +1497,7 @@ mod tests { let tools_dir = tmp.path().join("wasm-tools"); let engine = setup_owner_gate_engine( db.clone(), - owner_gate_trace(true), + owner_gate_trace(), tools_dir.as_path(), Some("default"), true, @@ -1493,17 +1529,18 @@ mod tests { } // ----------------------------------------------------------------------- - // Test: autonomous runs fail loudly when an extension tool is inactive + // Test: autonomous runs deny inactive extension tools at execution time + // (job completes but tool is blocked by the approval context) // ----------------------------------------------------------------------- #[tokio::test] - async fn full_job_blocks_without_active_owner_extension_tool() { + async fn full_job_denies_tool_without_active_owner_extension() { let (backend, tmp) = create_test_backend().await; let db: Arc = backend; let tools_dir = tmp.path().join("wasm-tools"); let engine = setup_owner_gate_engine( db.clone(), - owner_gate_trace(false), + owner_gate_trace(), tools_dir.as_path(), Some("default"), false, @@ -1519,31 +1556,30 @@ mod tests { .expect("fire manual"); let run = wait_for_run_completion(&db, routine.id, run_id).await; - assert_eq!(run.status, RunStatus::Failed); + // The job runs (full_job no longer requires sandbox) but the tool is + // blocked by the approval context — the LLM receives an error and + // completes without executing owner_gate. + assert_eq!(run.status, RunStatus::Ok); assert_eq!(owner_gate_count(&db).await, 0); - let failure_reason = db - .get_agent_job_failure_reason(run.job_id.expect("linked job id")) - .await - .expect("load job failure reason") - .expect("missing job failure reason"); - assert!( - failure_reason.contains("owner_gate"), - "expected missing-tool failure reason, got {failure_reason}" - ); + + // Verify the tool was actually attempted and denied (not just never called). + let job_id = run.job_id.expect("run should be linked to a job"); + wait_for_tool_denial_event(&db, job_id, "owner_gate").await; } // ----------------------------------------------------------------------- - // Test: extension tools activated for another owner are not inherited + // Test: extension tools activated for another owner are denied at execution + // time (job completes but tool is blocked by the approval context) // ----------------------------------------------------------------------- #[tokio::test] - async fn full_job_blocks_when_extension_belongs_to_another_owner() { + async fn full_job_denies_tool_when_extension_belongs_to_another_owner() { let (backend, tmp) = create_test_backend().await; let db: Arc = backend; let tools_dir = tmp.path().join("wasm-tools"); let engine = setup_owner_gate_engine( db.clone(), - owner_gate_trace(false), + owner_gate_trace(), tools_dir.as_path(), Some("someone-else"), true, @@ -1559,17 +1595,16 @@ mod tests { .expect("fire manual"); let run = wait_for_run_completion(&db, routine.id, run_id).await; - assert_eq!(run.status, RunStatus::Failed); + // The job runs (full_job no longer requires sandbox) but the tool is + // blocked by the approval context (extension belongs to "someone-else", + // not "default") — the LLM receives an error and completes without + // executing owner_gate. + assert_eq!(run.status, RunStatus::Ok); assert_eq!(owner_gate_count(&db).await, 0); - let failure_reason = db - .get_agent_job_failure_reason(run.job_id.expect("linked job id")) - .await - .expect("load job failure reason") - .expect("missing job failure reason"); - assert!( - failure_reason.contains("owner_gate"), - "expected owner-mismatch failure reason, got {failure_reason}" - ); + + // Verify the tool was actually attempted and denied (not just never called). + let job_id = run.job_id.expect("run should be linked to a job"); + wait_for_tool_denial_event(&db, job_id, "owner_gate").await; } // ----------------------------------------------------------------------- @@ -1623,7 +1658,7 @@ mod tests { let tools_dir = tmp.path().join("wasm-tools"); let engine = setup_owner_gate_engine( db.clone(), - owner_gate_trace(false), + owner_gate_trace(), tools_dir.as_path(), None, false, diff --git a/tests/support/test_channel.rs b/tests/support/test_channel.rs index cad59a33612..7fa89e06e41 100644 --- a/tests/support/test_channel.rs +++ b/tests/support/test_channel.rs @@ -122,6 +122,12 @@ impl TestChannel { .clone() } + /// Async version of `captured_responses` — safe to call while the agent is + /// actively pushing responses (avoids `try_lock` panic on contention). + pub async fn captured_responses_async(&self) -> Vec { + self.responses.lock().await.clone() + } + /// Wait until at least `n` responses have been captured, or `timeout` elapses. /// /// Returns whatever responses have been collected when the condition is met diff --git a/tests/support/test_rig.rs b/tests/support/test_rig.rs index 0e2883901de..a83d39832cc 100644 --- a/tests/support/test_rig.rs +++ b/tests/support/test_rig.rs @@ -29,6 +29,11 @@ use ironclaw::llm::recording::{HttpExchange, HttpInterceptor, ReplayingHttpInter // TestRig // --------------------------------------------------------------------------- +/// Substring unique to the static bootstrap greeting (GREETING.md). +/// Used to transparently filter per-user bootstrap greetings from the +/// response stream so tests don't need to account for them manually. +const BOOTSTRAP_GREETING_MARKER: &str = "always-on chief of staff"; + /// A running test agent with methods to inject messages and inspect results. pub struct TestRig { /// The test channel for sending messages and reading captures. @@ -59,6 +64,11 @@ pub struct TestRig { /// Temp directory guard -- keeps the libSQL database file alive. #[cfg(feature = "libsql")] _temp_dir: tempfile::TempDir, + /// How many bootstrap greetings to keep in `wait_for_responses`. + /// 0 for normal tests (filter all greetings), 1 for `.with_bootstrap()` + /// tests (keep the startup greeting, filter per-user duplicates). + #[cfg(feature = "libsql")] + bootstrap_greetings_to_keep: usize, } impl TestRig { @@ -93,9 +103,47 @@ impl TestRig { &self.session_manager } - /// Wait until at least `n` responses have been captured, or `timeout` elapses. + /// Wait until at least `n` non-bootstrap responses have been captured, or + /// `timeout` elapses. + /// + /// Per-user bootstrap greetings (fired when `tenant_ctx` creates a workspace + /// for a non-owner user) are transparently filtered from the response stream. + /// For `.with_bootstrap()` tests, the startup greeting is kept (1 allowed) + /// while additional per-user greetings are still filtered. pub async fn wait_for_responses(&self, n: usize, timeout: Duration) -> Vec { - self.channel.wait_for_responses(n, timeout).await + let deadline = tokio::time::Instant::now() + timeout; + let mut interval = Duration::from_millis(50); + let max_interval = Duration::from_millis(500); + loop { + let filtered = self.filter_responses(self.channel.captured_responses_async().await); + if filtered.len() >= n { + return filtered; + } + if tokio::time::Instant::now() >= deadline { + return filtered; + } + tokio::time::sleep(interval).await; + interval = (interval * 2).min(max_interval); + } + } + + /// Filter bootstrap greetings from the response stream. + /// + /// Keeps up to `bootstrap_greetings_to_keep` greeting responses (0 for + /// normal tests, 1 for `.with_bootstrap()` tests) and drops the rest. + fn filter_responses(&self, responses: Vec) -> Vec { + let mut greetings_kept = 0usize; + responses + .into_iter() + .filter(|r| { + if r.content.contains(BOOTSTRAP_GREETING_MARKER) { + greetings_kept += 1; + greetings_kept <= self.bootstrap_greetings_to_keep + } else { + true + } + }) + .collect() } /// Return the names of all `ToolStarted` events captured so far. @@ -588,8 +636,15 @@ impl TestRigBuilder { .await .expect("AppBuilder::build_all() failed in test rig"); - // Clear bootstrap flag so tests don't get an unexpected proactive greeting - // (unless the test explicitly wants to test the bootstrap flow). + // Clear the *owner* workspace bootstrap flag so tests don't get an + // unexpected proactive greeting on startup (unless the test explicitly + // wants to test the bootstrap flow via `.with_bootstrap()`). + // + // Per-user bootstrap greetings (fired when `tenant_ctx` creates a + // workspace for a non-owner user like "test-user") are allowed to + // happen naturally. They are transparently filtered from the response + // stream by `wait_for_responses` so tests don't need to account for + // them in response counting. if !keep_bootstrap && let Some(ref ws) = components.workspace { ws.take_bootstrap_pending(); } @@ -835,6 +890,7 @@ impl TestRigBuilder { extension_manager: ext_mgr_ref, session_manager: session_manager_ref, _temp_dir: temp_dir, + bootstrap_greetings_to_keep: if keep_bootstrap { 1 } else { 0 }, } } } From d0f7862a28456ea084a9fff480e7906fa368dde2 Mon Sep 17 00:00:00 2001 From: synner88 Date: Mon, 30 Mar 2026 10:19:33 +0300 Subject: [PATCH 09/11] fix(slack): respond to thread replies without requiring @mention (#1405) * fix(slack): respond to thread replies in channels without requiring @mention Two fixes: 1. Host bug: `on_respond` callback never committed workspace writes or injected workspace reader, unlike all other WASM callbacks. Any WASM channel persisting state during on_respond silently lost data. 2. Slack WASM channel: track threads where the bot has participated via workspace storage. When a message event arrives in a channel thread the bot previously replied to, process it without requiring @mention. Closes #1404 Co-Authored-By: Claude Opus 4.6 * fix(slack): log workspace_write error instead of silently discarding Address code review feedback: handle the Result from workspace_write when tracking thread participation, logging a warning on failure instead of using `let _ =` which would silently swallow errors. Co-Authored-By: Claude Opus 4.6 * fix(slack): harden thread reply tracking --------- Co-authored-by: synner88 <29090601+synner88@users.noreply.github.com> Co-authored-by: Claude Opus 4.6 Co-authored-by: Firat Sertgoz Co-authored-by: firat.sertgoz --- FEATURE_PARITY.md | 2 +- channels-src/slack/src/lib.rs | 202 +++++++++++++++++++++++++++++++--- src/channels/wasm/wrapper.rs | 74 ++++++++++++- 3 files changed, 253 insertions(+), 25 deletions(-) diff --git a/FEATURE_PARITY.md b/FEATURE_PARITY.md index 1946dce6ee5..915f529e4ea 100644 --- a/FEATURE_PARITY.md +++ b/FEATURE_PARITY.md @@ -112,7 +112,7 @@ This document tracks feature parity between IronClaw (Rust implementation) and O |---------|----------|----------|-------| | Streaming draft replies | ✅ | ❌ | Partial replies via draft message updates | | Configurable stream modes | ✅ | ❌ | Per-channel stream behavior | -| Thread ownership | ✅ | ❌ | Thread-level ownership tracking plus reply participation memory | +| Thread ownership | ✅ | 🚧 | Reply participation memory now persists with TTL-bounded tracking; full thread-level ownership tracking is still missing | | Download-file action | ✅ | ❌ | On-demand attachment downloads via message actions | ### Mattermost-Specific Features (since Mar 2026) diff --git a/channels-src/slack/src/lib.rs b/channels-src/slack/src/lib.rs index 24f01df3934..8958b64eb4d 100644 --- a/channels-src/slack/src/lib.rs +++ b/channels-src/slack/src/lib.rs @@ -23,6 +23,7 @@ wit_bindgen::generate!({ }); use serde::{Deserialize, Serialize}; +use std::collections::BTreeMap; // Re-export generated types use exports::near::agent::channel::{ @@ -129,9 +130,17 @@ const OWNER_ID_PATH: &str = "state/owner_id"; const DM_POLICY_PATH: &str = "state/dm_policy"; /// Workspace path for persisting allow_from (JSON array) across WASM callbacks. const ALLOW_FROM_PATH: &str = "state/allow_from"; +/// Workspace path for tracking recently active Slack threads. +const ACTIVE_THREADS_PATH: &str = "state/active_threads.json"; +/// Recently active threads expire after 24 hours to avoid reviving stale threads forever. +const ACTIVE_THREAD_TTL_MS: u64 = 24 * 60 * 60 * 1000; +/// Cap stored thread markers so the workspace state stays bounded. +const ACTIVE_THREAD_MAX_ENTRIES: usize = 256; /// Channel name for pairing store (used by pairing host APIs). const CHANNEL_NAME: &str = "slack"; +type ActiveThreads = BTreeMap; + /// Channel configuration from capabilities file. #[derive(Debug, Deserialize)] struct SlackConfig { @@ -263,9 +272,9 @@ impl Guest for SlackChannel { "text": response.content, }); - // Add thread_ts for threaded replies - if let Some(thread_ts) = response.thread_id.or(metadata.thread_ts) { - payload["thread_ts"] = serde_json::Value::String(thread_ts); + let thread_ts = response.thread_id.or(metadata.thread_ts); + if let Some(ref thread_ts) = thread_ts { + payload["thread_ts"] = serde_json::Value::String(thread_ts.clone()); } let payload_bytes = serde_json::to_vec(&payload) @@ -308,6 +317,10 @@ impl Guest for SlackChannel { )); } + if let Some(thread_ts) = thread_ts { + track_active_thread(&metadata.channel, &thread_ts)?; + } + channel_host::log( channel_host::LogLevel::Debug, &format!( @@ -452,13 +465,14 @@ fn download_and_store_slack_files(attachments: &[InboundAttachment]) { } } -/// Handle a Slack event and emit message if applicable. -fn handle_slack_event(event: SlackEvent, team_id: Option, _event_id: Option) { - let attachments = extract_slack_attachments(&event.files); - - // Download and store file attachments for host-side processing +fn prepare_inbound_attachments(files: &Option>) -> Vec { + let attachments = extract_slack_attachments(files); download_and_store_slack_files(&attachments); + attachments +} +/// Handle a Slack event and emit message if applicable. +fn handle_slack_event(event: SlackEvent, team_id: Option, _event_id: Option) { match event.event_type.as_str() { // Direct mention of the bot (always in a channel, not a DM) "app_mention" => { @@ -472,6 +486,7 @@ fn handle_slack_event(event: SlackEvent, team_id: Option, _event_id: Opt if !check_sender_permission(&user, &channel, false) { return; } + let attachments = prepare_inbound_attachments(&event.files); emit_message( user, text, @@ -483,7 +498,7 @@ fn handle_slack_event(event: SlackEvent, team_id: Option, _event_id: Opt } } - // Direct message to the bot + // Direct message or thread follow-up to the bot "message" => { // Skip messages from bots (including ourselves) if event.bot_id.is_some() || event.subtype.is_some() { @@ -496,11 +511,20 @@ fn handle_slack_event(event: SlackEvent, team_id: Option, _event_id: Opt event.text, event.ts.clone(), ) { - // Only process DMs (channel IDs starting with D) - if channel.starts_with('D') { - if !check_sender_permission(&user, &channel, true) { + let is_dm = channel.starts_with('D'); + + // Check if this is a reply in a thread where we previously participated + let is_active_thread = !is_dm + && event + .thread_ts + .as_ref() + .is_some_and(|thread_ts| is_active_thread(&channel, thread_ts)); + + if is_dm || is_active_thread { + if !check_sender_permission(&user, &channel, is_dm) { return; } + let attachments = prepare_inbound_attachments(&event.files); emit_message( user, text, @@ -522,6 +546,93 @@ fn handle_slack_event(event: SlackEvent, team_id: Option, _event_id: Opt } } +fn active_thread_key(channel: &str, thread_ts: &str) -> String { + format!("{channel}/{thread_ts}") +} + +fn is_thread_marker_fresh(last_seen_millis: u64, now_millis: u64) -> bool { + now_millis.saturating_sub(last_seen_millis) <= ACTIVE_THREAD_TTL_MS +} + +fn prune_active_threads(active_threads: &mut ActiveThreads, now_millis: u64) -> bool { + let mut changed = false; + active_threads.retain(|_, last_seen_millis| { + let keep = is_thread_marker_fresh(*last_seen_millis, now_millis); + if !keep { + changed = true; + } + keep + }); + + if active_threads.len() > ACTIVE_THREAD_MAX_ENTRIES { + let mut oldest_first: Vec<_> = active_threads + .iter() + .map(|(key, last_seen_millis)| (key.clone(), *last_seen_millis)) + .collect(); + oldest_first.sort_by_key(|(_, last_seen_millis)| *last_seen_millis); + + for (key, _) in oldest_first + .into_iter() + .take(active_threads.len() - ACTIVE_THREAD_MAX_ENTRIES) + { + active_threads.remove(&key); + changed = true; + } + } + + changed +} + +fn load_active_threads() -> ActiveThreads { + let Some(raw) = channel_host::workspace_read(ACTIVE_THREADS_PATH) else { + return ActiveThreads::new(); + }; + + match serde_json::from_str(&raw) { + Ok(active_threads) => active_threads, + Err(e) => { + channel_host::log( + channel_host::LogLevel::Warn, + &format!("Failed to parse active thread state: {e}"), + ); + ActiveThreads::new() + } + } +} + +fn persist_active_threads(active_threads: &ActiveThreads) -> Result<(), String> { + let serialized = serde_json::to_string(active_threads) + .map_err(|e| format!("Failed to serialize active thread state: {e}"))?; + channel_host::workspace_write(ACTIVE_THREADS_PATH, &serialized) + .map_err(|e| format!("Failed to persist active thread state: {e}")) +} + +fn track_active_thread(channel: &str, thread_ts: &str) -> Result<(), String> { + let now_millis = channel_host::now_millis(); + let mut active_threads = load_active_threads(); + prune_active_threads(&mut active_threads, now_millis); + active_threads.insert(active_thread_key(channel, thread_ts), now_millis); + prune_active_threads(&mut active_threads, now_millis); + persist_active_threads(&active_threads) +} + +fn is_active_thread(channel: &str, thread_ts: &str) -> bool { + let now_millis = channel_host::now_millis(); + let mut active_threads = load_active_threads(); + let changed = prune_active_threads(&mut active_threads, now_millis); + + if changed { + if let Err(e) = persist_active_threads(&active_threads) { + channel_host::log( + channel_host::LogLevel::Warn, + &format!("Failed to prune active thread state: {e}"), + ); + } + } + + active_threads.contains_key(&active_thread_key(channel, thread_ts)) +} + /// Emit a message to the agent. fn emit_message( user_id: String, @@ -606,8 +717,7 @@ fn check_sender_permission(user_id: &str, channel_id: &str, is_dm: bool) -> bool } // 4. Check sender (Slack events only have user ID, not username) - let is_allowed = - allowed.contains(&"*".to_string()) || allowed.contains(&user_id.to_string()); + let is_allowed = allowed.contains(&"*".to_string()) || allowed.contains(&user_id.to_string()); if is_allowed { return true; @@ -625,10 +735,7 @@ fn check_sender_permission(user_id: &str, channel_id: &str, is_dm: bool) -> bool Ok(result) => { channel_host::log( channel_host::LogLevel::Info, - &format!( - "Pairing request for user {}: code {}", - user_id, result.code - ), + &format!("Pairing request for user {}: code {}", user_id, result.code), ); if result.created { let _ = send_pairing_reply(channel_id, &result.code); @@ -826,4 +933,63 @@ mod tests { // Verify the constant is 20 MB assert_eq!(MAX_DOWNLOAD_SIZE_BYTES, 20 * 1024 * 1024); } + + #[test] + fn test_active_thread_key_scopes_by_channel_and_thread() { + assert_eq!( + active_thread_key("C123", "1742486400.000100"), + "C123/1742486400.000100" + ); + } + + #[test] + fn test_prune_active_threads_removes_expired_entries() { + let now_millis = ACTIVE_THREAD_TTL_MS + 1_000; + let mut active_threads = ActiveThreads::from([ + ( + "C1/expired".to_string(), + now_millis - ACTIVE_THREAD_TTL_MS - 1, + ), + ("C1/fresh".to_string(), now_millis - ACTIVE_THREAD_TTL_MS), + ]); + + let changed = prune_active_threads(&mut active_threads, now_millis); + + assert!(changed); + assert!(!active_threads.contains_key("C1/expired")); + assert!(active_threads.contains_key("C1/fresh")); + } + + #[test] + fn test_prune_active_threads_trims_oldest_entries_when_over_limit() { + let now_millis = ACTIVE_THREAD_TTL_MS + 1_000; + let mut active_threads = ActiveThreads::new(); + + for i in 0..=ACTIVE_THREAD_MAX_ENTRIES { + active_threads.insert(format!("C1/{i}"), now_millis + i as u64); + } + + let changed = prune_active_threads( + &mut active_threads, + now_millis + ACTIVE_THREAD_MAX_ENTRIES as u64, + ); + + assert!(changed); + assert_eq!(active_threads.len(), ACTIVE_THREAD_MAX_ENTRIES); + assert!(!active_threads.contains_key("C1/0")); + assert!(active_threads.contains_key(&format!("C1/{ACTIVE_THREAD_MAX_ENTRIES}"))); + } + + #[test] + fn test_is_thread_marker_fresh_respects_ttl_boundary() { + let now_millis = ACTIVE_THREAD_TTL_MS + 1_000; + assert!(is_thread_marker_fresh( + now_millis - ACTIVE_THREAD_TTL_MS, + now_millis + )); + assert!(!is_thread_marker_fresh( + now_millis - ACTIVE_THREAD_TTL_MS - 1, + now_millis + )); + } } diff --git a/src/channels/wasm/wrapper.rs b/src/channels/wasm/wrapper.rs index d3feb2318c8..c502112700d 100644 --- a/src/channels/wasm/wrapper.rs +++ b/src/channels/wasm/wrapper.rs @@ -1800,7 +1800,7 @@ impl WasmChannel { let runtime = Arc::clone(&self.runtime); let prepared = Arc::clone(&self.prepared); - let capabilities = self.capabilities.clone(); + let capabilities = Self::inject_workspace_reader(&self.capabilities, &self.workspace_store); let timeout = self.runtime.config().callback_timeout; let channel_name = self.name.clone(); let credentials = self.get_credentials().await; @@ -1811,6 +1811,7 @@ impl WasmChannel { ) .await; let pairing_store = self.pairing_store.clone(); + let workspace_store = self.workspace_store.clone(); // Prepare response data let message_id_str = message_id.to_string(); @@ -1881,8 +1882,10 @@ impl WasmChannel { }); } - let host_state = + let mut host_state = Self::extract_host_state(&mut store, &prepared.name, &capabilities); + let pending_writes = host_state.take_pending_writes(); + workspace_store.commit_writes(&pending_writes); tracing::info!("on_respond WASM execution completed successfully"); Ok(((), host_state)) }) @@ -1944,7 +1947,7 @@ impl WasmChannel { let runtime = Arc::clone(&self.runtime); let prepared = Arc::clone(&self.prepared); - let capabilities = self.capabilities.clone(); + let capabilities = Self::inject_workspace_reader(&self.capabilities, &self.workspace_store); let timeout = self.runtime.config().callback_timeout; let channel_name = self.name.clone(); let credentials = self.get_credentials().await; @@ -1955,6 +1958,7 @@ impl WasmChannel { ) .await; let pairing_store = self.pairing_store.clone(); + let workspace_store = self.workspace_store.clone(); let user_id = user_id.to_string(); let content = content.to_string(); @@ -2006,8 +2010,10 @@ impl WasmChannel { }); } - let host_state = + let mut host_state = Self::extract_host_state(&mut store, &prepared.name, &capabilities); + let pending_writes = host_state.take_pending_writes(); + workspace_store.commit_writes(&pending_writes); tracing::info!("on_broadcast WASM execution completed successfully"); Ok(((), host_state)) }) @@ -2051,7 +2057,7 @@ impl WasmChannel { let runtime = Arc::clone(&self.runtime); let prepared = Arc::clone(&self.prepared); - let capabilities = self.capabilities.clone(); + let capabilities = Self::inject_workspace_reader(&self.capabilities, &self.workspace_store); let timeout = self.runtime.config().callback_timeout; let channel_name = self.name.clone(); let credentials = self.get_credentials().await; @@ -2062,6 +2068,7 @@ impl WasmChannel { ) .await; let pairing_store = self.pairing_store.clone(); + let workspace_store = self.workspace_store.clone(); let Some(wit_update) = status_to_wit(status, metadata) else { return Ok(()); @@ -2084,6 +2091,11 @@ impl WasmChannel { .call_on_status(&mut store, &wit_update) .map_err(|e| Self::map_wasm_error(e, &prepared.name, prepared.limits.fuel))?; + let mut host_state = + Self::extract_host_state(&mut store, &prepared.name, &capabilities); + let pending_writes = host_state.take_pending_writes(); + workspace_store.commit_writes(&pending_writes); + Ok(()) }) .await @@ -2124,6 +2136,7 @@ impl WasmChannel { host_credentials: Vec, pairing_store: Arc, timeout: Duration, + workspace_store: &Arc, wit_update: wit_channel::StatusUpdate, ) -> Result<(), WasmChannelError> { if prepared.component().is_none() { @@ -2132,9 +2145,10 @@ impl WasmChannel { let runtime = Arc::clone(runtime); let prepared = Arc::clone(prepared); - let capabilities = capabilities.clone(); + let capabilities = Self::inject_workspace_reader(capabilities, workspace_store); let credentials_snapshot = credentials.read().await.clone(); let channel_name_owned = channel_name.to_string(); + let workspace_store = Arc::clone(workspace_store); let result = tokio::time::timeout(timeout, async move { tokio::task::spawn_blocking(move || { @@ -2153,6 +2167,11 @@ impl WasmChannel { .call_on_status(&mut store, &wit_update) .map_err(|e| Self::map_wasm_error(e, &prepared.name, prepared.limits.fuel))?; + let mut host_state = + Self::extract_host_state(&mut store, &prepared.name, &capabilities); + let pending_writes = host_state.take_pending_writes(); + workspace_store.commit_writes(&pending_writes); + Ok(()) }) .await @@ -2224,6 +2243,7 @@ impl WasmChannel { let runtime = Arc::clone(&self.runtime); let prepared = Arc::clone(&self.prepared); let capabilities = self.capabilities.clone(); + let workspace_store = self.workspace_store.clone(); let credentials = self.credentials.clone(); // Pre-resolve host credentials once for the lifetime of the repeater. // Channels tokens rarely change, so a snapshot per-repeater is correct. @@ -2259,6 +2279,7 @@ impl WasmChannel { hc, pairing_store.clone(), callback_timeout, + &workspace_store, wit_update_clone, ) .await @@ -4430,6 +4451,47 @@ mod tests { assert_eq!(response.body, b"Bad request"); } + #[test] + fn test_inject_workspace_reader_adds_missing_reader() { + let capabilities = ChannelCapabilities::for_channel("test"); + assert!(capabilities.tool_capabilities.workspace_read.is_none()); + + let workspace_store = Arc::new(crate::channels::wasm::host::ChannelWorkspaceStore::new()); + let injected = WasmChannel::inject_workspace_reader(&capabilities, &workspace_store); + + assert!(injected.tool_capabilities.workspace_read.is_some()); + assert!( + injected + .tool_capabilities + .workspace_read + .as_ref() + .and_then(|cap| cap.reader.as_ref()) + .is_some() + ); + } + + #[test] + fn test_inject_workspace_reader_preserves_allowed_prefixes() { + let tool_capabilities = crate::tools::wasm::Capabilities::default() + .with_workspace_read(vec!["state/".to_string(), "context/".to_string()]); + let capabilities = + ChannelCapabilities::for_channel("test").with_tool_capabilities(tool_capabilities); + let workspace_store = Arc::new(crate::channels::wasm::host::ChannelWorkspaceStore::new()); + + let injected = WasmChannel::inject_workspace_reader(&capabilities, &workspace_store); + + let workspace_read = injected + .tool_capabilities + .workspace_read + .as_ref() + .expect("workspace_read capability should exist"); + assert_eq!( + workspace_read.allowed_prefixes, + vec!["state/".to_string(), "context/".to_string()] + ); + assert!(workspace_read.reader.is_some()); + } + #[tokio::test] async fn test_channel_start_and_shutdown() { let channel = create_test_channel(); From d567d94c246dc7d984a019dfb9519085da13fb47 Mon Sep 17 00:00:00 2001 From: Illia Polosukhin Date: Mon, 30 Mar 2026 08:45:05 -0700 Subject: [PATCH 10/11] fix(routines): clone Arc before await in web handler event cache refresh (#1756) * fix(routines): clone Arc before await in web handler event cache refresh (#1076) Address review: drop superseded ticker changes, keep only the .cloned() fix that prevents holding RwLockReadGuard across .await in toggle/delete handlers. Add regression test for web toggle disabling a system_event routine. Closes #1076 Co-Authored-By: Claude Opus 4.6 (1M context) * fix: use explicit block to drop RwLockReadGuard before await Address review feedback: in Rust 2024, `if let` scrutinee temporaries live through the body, so the `.cloned()` approach still held the RwLockReadGuard across `refresh_event_cache().await`. Extract into an explicit block to ensure the guard is dropped, matching the existing pattern in `routines_trigger_handler`. Also add retry loop for `routine_by_name` in the integration test to avoid flakiness from potential race conditions. Co-Authored-By: Claude Opus 4.6 (1M context) --------- Co-authored-by: Claude Opus 4.6 (1M context) --- src/channels/web/handlers/routines.rs | 8 +- tests/gateway_workflow_integration.rs | 106 ++++++++++++++++++++++++++ 2 files changed, 112 insertions(+), 2 deletions(-) diff --git a/src/channels/web/handlers/routines.rs b/src/channels/web/handlers/routines.rs index 5597a47c92c..ebbb599eff9 100644 --- a/src/channels/web/handlers/routines.rs +++ b/src/channels/web/handlers/routines.rs @@ -243,7 +243,9 @@ pub async fn routines_toggle_handler( // Refresh the in-memory event trigger cache so event/system_event // routines reflect the new enabled state immediately (issue #1076). - if let Some(engine) = state.routine_engine.read().await.as_ref() { + // Extract into a block so the RwLockReadGuard is dropped before the async call. + let engine = { state.routine_engine.read().await.as_ref().cloned() }; + if let Some(engine) = engine { engine.refresh_event_cache().await; } @@ -285,7 +287,9 @@ pub async fn routines_delete_handler( if deleted { // Refresh the in-memory event trigger cache so deleted event/system_event // routines stop firing immediately (issue #1076). - if let Some(engine) = state.routine_engine.read().await.as_ref() { + // Extract into a block so the RwLockReadGuard is dropped before the async call. + let engine = { state.routine_engine.read().await.as_ref().cloned() }; + if let Some(engine) = engine { engine.refresh_event_cache().await; } diff --git a/tests/gateway_workflow_integration.rs b/tests/gateway_workflow_integration.rs index c955e5a1a55..1cd74acb842 100644 --- a/tests/gateway_workflow_integration.rs +++ b/tests/gateway_workflow_integration.rs @@ -338,4 +338,110 @@ mod tests { harness.shutdown().await; mock.shutdown().await; } + + /// Regression test for issue #1076: web API toggle must immediately + /// invalidate the in-memory event cache so disabled routines stop firing. + #[tokio::test] + async fn web_toggle_disables_system_event_routine_without_restart() { + let mock = MockOpenAiServerBuilder::new() + .with_rule(MockOpenAiRule::on_user_contains( + "create webhook routine", + MockOpenAiResponse::ToolCalls(vec![MockToolCall::new( + "call_create_webhook_1", + "routine_create", + serde_json::json!({ + "name": "wf-toggle-system-event", + "description": "System event toggle regression test", + "trigger_type": "system_event", + "event_source": "github", + "event_type": "issue.opened", + "event_filters": {"repository": "nearai/ironclaw"}, + "action_type": "lightweight", + "prompt": "summarize issue" + }), + )]), + )) + .with_default_response(MockOpenAiResponse::Text("ack".to_string())) + .start() + .await; + + let harness = + GatewayWorkflowHarness::start_openai_compatible(&mock.openai_base_url(), "mock-model") + .await; + + let thread_id = harness.create_thread().await; + harness + .send_chat(&thread_id, "create webhook routine") + .await; + harness + .wait_for_turns(&thread_id, 1, Duration::from_secs(10)) + .await; + + let mut routine = None; + for _ in 0..30 { + routine = harness.routine_by_name("wf-toggle-system-event").await; + if routine.is_some() { + break; + } + tokio::time::sleep(Duration::from_millis(100)).await; + } + let routine = routine.expect("routine should exist after retries"); + let routine_id = routine + .get("id") + .and_then(|v| v.as_str()) + .expect("routine id missing"); + + let runs_before = harness.routine_runs(routine_id).await; + let before_count = runs_before["runs"] + .as_array() + .map(|a| a.len()) + .unwrap_or_default(); + + // Disable through web API (non-tool mutation path). + harness + .client + .post(format!( + "{}/api/routines/{routine_id}/toggle", + harness.base_url() + )) + .bearer_auth(&harness.auth_token) + .json(&serde_json::json!({ "enabled": false })) + .send() + .await + .expect("disable toggle request failed") + .error_for_status() + .expect("disable toggle non-2xx"); + + // Fire a webhook that would match the now-disabled routine. + let hook = harness + .github_webhook( + "issues", + serde_json::json!({ + "action": "opened", + "repository": {"full_name": "nearai/ironclaw"}, + "issue": {"number": 881, "title": "Toggle disable regression"} + }), + ) + .await; + assert_eq!(hook["status"], "accepted"); + assert_eq!(hook["emitted_events"], 1); + assert_eq!( + hook["fired_routines"].as_u64().unwrap_or(0), + 0, + "disabled routine should not fire after web toggle" + ); + + let runs_after = harness.routine_runs(routine_id).await; + let after_count = runs_after["runs"] + .as_array() + .map(|a| a.len()) + .unwrap_or_default(); + assert_eq!( + after_count, before_count, + "run count should not increase for disabled routine" + ); + + harness.shutdown().await; + mock.shutdown().await; + } } From 21f613ff2e8a1dc1f809677b3cc669e59ccbddbd Mon Sep 17 00:00:00 2001 From: Henry Park Date: Mon, 30 Mar 2026 11:24:18 -0700 Subject: [PATCH 11/11] test(e2e): align WASM reinstall expectation with uninstall cleanup (#1762) * test(e2e): align wasm reinstall expectations with uninstall cleanup * test(e2e): clarify wasm reinstall fixture semantics --- tests/e2e/scenarios/test_wasm_lifecycle.py | 7 ++++--- 1 file changed, 4 insertions(+), 3 deletions(-) diff --git a/tests/e2e/scenarios/test_wasm_lifecycle.py b/tests/e2e/scenarios/test_wasm_lifecycle.py index 212cc3ce057..16e2cf1c377 100644 --- a/tests/e2e/scenarios/test_wasm_lifecycle.py +++ b/tests/e2e/scenarios/test_wasm_lifecycle.py @@ -99,7 +99,7 @@ async def web_search_removed(ironclaw_server, web_search_configured): @pytest.fixture(scope="module") async def web_search_reinstalled(ironclaw_server, web_search_removed): - """Reinstall web-search after removal to verify saved-secret recovery.""" + """Reinstall web-search after removal to verify it returns unconfigured.""" await _ensure_removed(ironclaw_server, "web-search") data = await _install_extension(ironclaw_server, "web-search") return {"name": "web-search", "install": data} @@ -421,8 +421,9 @@ async def test_reinstall_after_remove(ironclaw_server, web_search_reinstalled): """Extension can be reinstalled after removal without stale activation errors.""" ext = await _get_extension(ironclaw_server, "web-search") assert ext is not None, "web-search not found after reinstall" - assert ext["active"] is True, "Reinstalled tool should auto-activate via saved secrets" - assert ext["authenticated"] is True, "Saved secret should still authenticate on reinstall" + assert ext["active"] is False, "Reinstalled tool should require setup before activation" + assert ext["authenticated"] is False, "Reinstalled tool should not reuse deleted secrets" + assert ext["needs_setup"] is True, "Reinstalled tool should require setup again" # Verify no stale activation error from previous install assert ext.get("activation_error") is None or ext.get("activation_error") == "", ( f"Reinstalled extension should have no stale activation error: {ext}"