From 0b33ca99262925558760fcfa2930bee60fe65997 Mon Sep 17 00:00:00 2001 From: Achieve Date: Sat, 28 Mar 2026 22:10:39 +0800 Subject: [PATCH 01/23] fix(oauth): tighten legacy state validation and fallback handling (#1701) * fix(oauth): tighten legacy state validation and fallback handling * style: fix formatting * refactor: separate validation checks for clearer error messages --- src/cli/oauth_defaults.rs | 126 ++++++++++++++++++++++++++++++++++++-- 1 file changed, 122 insertions(+), 4 deletions(-) diff --git a/src/cli/oauth_defaults.rs b/src/cli/oauth_defaults.rs index 384d5833834..5628f3d689b 100644 --- a/src/cli/oauth_defaults.rs +++ b/src/cli/oauth_defaults.rs @@ -569,6 +569,42 @@ pub async fn sweep_expired_flows(registry: &PendingOAuthRegistry) { const HOSTED_STATE_PREFIX: &str = "ic2"; const HOSTED_STATE_CHECKSUM_BYTES: usize = 12; +/// Maximum length for a legacy flow ID or instance name. +const LEGACY_STATE_MAX_LEN: usize = 128; +/// Minimum length for a legacy flow ID. +const LEGACY_STATE_MIN_LEN: usize = 8; + +/// Validate that a legacy state component (flow_id or instance_name) contains +/// only safe characters: alphanumeric, dash, underscore. +fn is_valid_legacy_state_component(s: &str) -> bool { + !s.is_empty() + && s.len() <= LEGACY_STATE_MAX_LEN + && s.bytes() + .all(|b| b.is_ascii_alphanumeric() || b == b'-' || b == b'_') +} + +fn validate_legacy_flow_id(flow_id: &str) -> Result<(), String> { + if flow_id.len() < LEGACY_STATE_MIN_LEN { + return Err(format!( + "Legacy OAuth flow_id too short ({} chars, minimum {LEGACY_STATE_MIN_LEN})", + flow_id.len() + )); + } + if flow_id.len() > LEGACY_STATE_MAX_LEN { + return Err(format!( + "Legacy OAuth flow_id too long ({} chars, maximum {LEGACY_STATE_MAX_LEN})", + flow_id.len() + )); + } + if !flow_id + .bytes() + .all(|b| b.is_ascii_alphanumeric() || b == b'-' || b == b'_') + { + return Err("Legacy OAuth flow_id contains invalid characters".to_string()); + } + Ok(()) +} + #[derive(Debug, Clone, PartialEq, Eq)] pub struct DecodedHostedOAuthState { pub flow_id: String, @@ -653,6 +689,17 @@ pub fn decode_hosted_oauth_state(state: &str) -> Result Result Date: Sat, 28 Mar 2026 22:13:28 +0800 Subject: [PATCH 02/23] fix(web): redact database error details from API responses (#1711) --- src/channels/web/handlers/jobs.rs | 62 +++++++++++++++---------------- 1 file changed, 30 insertions(+), 32 deletions(-) diff --git a/src/channels/web/handlers/jobs.rs b/src/channels/web/handlers/jobs.rs index 35adeec68a1..b171561c767 100644 --- a/src/channels/web/handlers/jobs.rs +++ b/src/channels/web/handlers/jobs.rs @@ -15,6 +15,14 @@ use crate::channels::web::auth::AuthenticatedUser; use crate::channels::web::server::GatewayState; use crate::channels::web::types::*; +fn db_error(context: &str, e: impl std::fmt::Display) -> (StatusCode, String) { + tracing::error!(%e, context, "Database error in jobs handler"); + ( + StatusCode::INTERNAL_SERVER_ERROR, + "Internal database error".to_string(), + ) +} + pub async fn jobs_list_handler( State(state): State>, AuthenticatedUser(user): AuthenticatedUser, @@ -213,10 +221,7 @@ pub async fn jobs_detail_handler( } Ok(None) => {} Err(e) => { - return Err(( - StatusCode::INTERNAL_SERVER_ERROR, - format!("Database error: {}", e), - )); + return Err(db_error("jobs_handler", e)); } } @@ -257,10 +262,7 @@ pub async fn jobs_detail_handler( })) } Ok(None) => Err((StatusCode::NOT_FOUND, "Job not found".to_string())), - Err(e) => Err(( - StatusCode::INTERNAL_SERVER_ERROR, - format!("Database error: {}", e), - )), + Err(e) => Err(db_error("jobs_handler", e)), } } @@ -304,10 +306,7 @@ pub async fn jobs_cancel_handler( } Ok(None) => {} Err(e) => { - return Err(( - StatusCode::INTERNAL_SERVER_ERROR, - format!("Database error: {}", e), - )); + return Err(db_error("jobs_handler", e)); } } } @@ -350,10 +349,7 @@ pub async fn jobs_cancel_handler( } Ok(None) => {} Err(e) => { - return Err(( - StatusCode::INTERNAL_SERVER_ERROR, - format!("Database error: {}", e), - )); + return Err(db_error("jobs_handler", e)); } } } @@ -471,10 +467,7 @@ pub async fn jobs_restart_handler( } Ok(None) => {} Err(e) => { - return Err(( - StatusCode::INTERNAL_SERVER_ERROR, - format!("Database error: {}", e), - )); + return Err(db_error("jobs_handler", e)); } } @@ -530,10 +523,7 @@ pub async fn jobs_restart_handler( }))) } Ok(None) => Err((StatusCode::NOT_FOUND, "Job not found".to_string())), - Err(e) => Err(( - StatusCode::INTERNAL_SERVER_ERROR, - format!("Database error: {}", e), - )), + Err(e) => Err(db_error("jobs_handler", e)), } } @@ -609,10 +599,7 @@ pub async fn jobs_prompt_handler( return Err((StatusCode::NOT_FOUND, "Job not found".to_string())); } Err(e) => { - return Err(( - StatusCode::INTERNAL_SERVER_ERROR, - format!("Database error: {}", e), - )); + return Err(db_error("jobs_handler", e)); } } } @@ -667,10 +654,7 @@ pub async fn jobs_events_handler( return Err((StatusCode::NOT_FOUND, "Job not found".to_string())); } Err(e) => { - return Err(( - StatusCode::INTERNAL_SERVER_ERROR, - format!("Database error: {}", e), - )); + return Err(db_error("jobs_handler", e)); } } @@ -823,3 +807,17 @@ pub async fn job_files_read_handler( content, })) } + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn test_db_error_does_not_leak_details() { + let (status, body) = db_error("test_context", "relation \"jobs\" does not exist"); + assert_eq!(status, StatusCode::INTERNAL_SERVER_ERROR); + assert_eq!(body, "Internal database error"); + assert!(!body.contains("relation")); + assert!(!body.contains("does not exist")); + } +} From 9ce3a9fc53707aa80cd8c8a6276e63adc4f61280 Mon Sep 17 00:00:00 2001 From: Achieve Date: Sat, 28 Mar 2026 23:31:27 +0800 Subject: [PATCH 03/23] feat(discord): implement on_broadcast via DM channel creation (#1693) MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit - Implement broadcast_dm() that creates a DM channel with the target user (POST /users/@me/channels, cached by Discord) and sends the message to it - Extract DISCORD_API_BASE constant for all Discord REST API URLs - Extract send_channel_message() shared helper to deduplicate message posting between on_respond and broadcast_dm - Add snowflake validation on user_id before API calls - Fix pre-existing clippy redundant_closure warning - Use typed DmChannelResponse struct instead of serde_json::Value Closes no specific issue — completes the previously stubbed on_broadcast. Co-authored-by: Claude Opus 4.6 (1M context) --- channels-src/discord/src/lib.rs | 141 ++++++++++++++++++++++++++------ 1 file changed, 118 insertions(+), 23 deletions(-) diff --git a/channels-src/discord/src/lib.rs b/channels-src/discord/src/lib.rs index cdb6c515077..249d1e5b341 100644 --- a/channels-src/discord/src/lib.rs +++ b/channels-src/discord/src/lib.rs @@ -28,6 +28,9 @@ use std::{cmp::Ordering, collections::HashMap}; use ed25519_dalek::{Signature, Verifier, VerifyingKey}; use serde::{Deserialize, Serialize}; +/// Discord REST API v10 base URL. +const DISCORD_API_BASE: &str = "https://discord.com/api/v10"; + use exports::near::agent::channel::{ AgentResponse, ChannelConfig, Guest, HttpEndpointConfig, IncomingHttpRequest, OutgoingHttpResponse, PollConfig, StatusUpdate, @@ -427,7 +430,7 @@ impl Guest for DiscordChannel { ( "PATCH", format!( - "https://discord.com/api/v10/webhooks/{}/{}/messages/@original", + "{DISCORD_API_BASE}/webhooks/{}/{}/messages/@original", application_id, token ), ) @@ -438,20 +441,7 @@ impl Guest for DiscordChannel { payload["allowed_mentions"] = serde_json::json!({ "replied_user": true }); - let mention_payload = serde_json::to_vec(&payload) - .map_err(|e| format!("Failed to serialize mention payload: {}", e))?; - let mention_url = format!( - "https://discord.com/api/v10/channels/{}/messages", - metadata.channel_id - ); - let result = channel_host::http_request( - "POST", - &mention_url, - &discord_auth_headers_json(true), - Some(&mention_payload), - None, - ); - return map_discord_response(result); + return send_channel_message(&metadata.channel_id, payload); } else { return Err("Unsupported Discord response metadata".to_string()); }; @@ -469,8 +459,8 @@ impl Guest for DiscordChannel { fn on_status(_update: StatusUpdate) {} - fn on_broadcast(_user_id: String, _response: AgentResponse) -> Result<(), String> { - Err("broadcast not yet implemented for Discord channel".to_string()) + fn on_broadcast(user_id: String, response: AgentResponse) -> Result<(), String> { + broadcast_dm(&user_id, &response.content) } fn on_shutdown() { @@ -501,6 +491,21 @@ fn map_discord_response( } } +/// Post a JSON payload to a Discord channel as a new message. +fn send_channel_message(channel_id: &str, payload: serde_json::Value) -> Result<(), String> { + let payload_bytes = serde_json::to_vec(&payload) + .map_err(|e| format!("Failed to serialize message: {}", e))?; + let url = format!("{DISCORD_API_BASE}/channels/{}/messages", channel_id); + let result = channel_host::http_request( + "POST", + &url, + &discord_auth_headers_json(true), + Some(&payload_bytes), + None, + ); + map_discord_response(result) +} + fn load_runtime_config() -> DiscordRuntimeConfig { channel_host::workspace_read("config.json") .and_then(|raw| serde_json::from_str::(&raw).ok()) @@ -539,7 +544,7 @@ fn get_or_fetch_bot_id() -> Option { let response = channel_host::http_request( "GET", - "https://discord.com/api/v10/users/@me", + &format!("{DISCORD_API_BASE}/users/@me"), &discord_auth_headers_json(false), None, Some(10_000), @@ -659,7 +664,7 @@ fn poll_channel_mentions(channel_id: &str, bot_id: &str) { fn fetch_latest_message_id(channel_id: &str) -> Option { let url = format!( - "https://discord.com/api/v10/channels/{}/messages?limit=1", + "{DISCORD_API_BASE}/channels/{}/messages?limit=1", channel_id ); let response = channel_host::http_request( @@ -697,7 +702,7 @@ fn fetch_messages_after_cursor( for page in 0..MAX_PAGES { let url = format!( - "https://discord.com/api/v10/channels/{}/messages?limit={}&after={}", + "{DISCORD_API_BASE}/channels/{}/messages?limit={}&after={}", channel_id, PAGE_LIMIT, after ); let response = match channel_host::http_request( @@ -986,7 +991,7 @@ fn handle_slash_command(interaction: &DiscordInteraction) -> bool { ); // Attempt to notify user of internal error let url = format!( - "https://discord.com/api/v10/webhooks/{}/{}", + "{DISCORD_API_BASE}/webhooks/{}/{}", interaction.application_id, interaction.token ); let payload = serde_json::json!({ @@ -1106,7 +1111,7 @@ fn check_sender_permission( } let dm_policy = - channel_host::workspace_read(DM_POLICY_PATH).unwrap_or_else(|| default_dm_policy()); + channel_host::workspace_read(DM_POLICY_PATH).unwrap_or_else(default_dm_policy); if dm_policy == "open" { return true; } @@ -1161,7 +1166,7 @@ fn check_sender_permission( /// Send a pairing code as an ephemeral Discord followup message. fn send_pairing_reply(ctx: &PairingReplyCtx, code: &str) -> Result<(), String> { let url = format!( - "https://discord.com/api/v10/webhooks/{}/{}", + "{DISCORD_API_BASE}/webhooks/{}/{}", ctx.application_id, ctx.token ); let payload = serde_json::json!({ @@ -1194,6 +1199,57 @@ fn send_pairing_reply(ctx: &PairingReplyCtx, code: &str) -> Result<(), String> { } } +/// Send a broadcast message to a Discord user via DM. +/// +/// Creates a DM channel with the user (Discord caches this, so repeated calls +/// for the same user reuse the existing channel) and then posts the message. +fn broadcast_dm(user_id: &str, content: &str) -> Result<(), String> { + // Validate user_id is a plausible Discord snowflake (numeric, 17-20 digits) + // to avoid injecting arbitrary strings into API URLs. + if user_id.is_empty() + || !user_id.chars().all(|c| c.is_ascii_digit()) + || user_id.len() < 17 + || user_id.len() > 20 + { + return Err(format!("Invalid Discord user ID: '{}'", user_id)); + } + + // Step 1: Open (or reuse) a DM channel with the target user. + let create_dm_payload = serde_json::json!({ "recipient_id": user_id }); + let create_dm_bytes = serde_json::to_vec(&create_dm_payload) + .map_err(|e| format!("Failed to serialize DM channel request: {}", e))?; + + let dm_response = channel_host::http_request( + "POST", + &format!("{DISCORD_API_BASE}/users/@me/channels"), + &discord_auth_headers_json(true), + Some(&create_dm_bytes), + Some(10_000), + ) + .map_err(|e| format!("Failed to create DM channel: {}", e))?; + + if !(200..300).contains(&dm_response.status) { + let body = String::from_utf8_lossy(&dm_response.body); + return Err(format!( + "Discord create-DM failed: {} - {}", + dm_response.status, body + )); + } + + #[derive(Deserialize)] + struct DmChannelResponse { + id: String, + } + let dm_channel: DmChannelResponse = serde_json::from_slice(&dm_response.body) + .map_err(|e| format!("Failed to parse DM channel response: {}", e))?; + let channel_id = &dm_channel.id; + + // Step 2: Send the message to the DM channel. + let truncated = truncate_message(content); + let payload = serde_json::json!({ "content": truncated }); + send_channel_message(channel_id, payload) +} + fn json_response(status: u16, value: serde_json::Value) -> OutgoingHttpResponse { let body = serde_json::to_vec(&value).unwrap_or_default(); let headers = serde_json::json!({"Content-Type": "application/json"}); @@ -1593,4 +1649,43 @@ mod tests { assert_eq!(interaction.interaction_type, 2); assert!(interaction.data.is_some()); } + + #[test] + fn test_broadcast_dm_payload_format() { + // Verify the DM channel creation payload is well-formed JSON that + // Discord's API expects. + let user_id = "123456789012345678"; + let payload = serde_json::json!({ "recipient_id": user_id }); + let serialized = serde_json::to_vec(&payload).unwrap(); + let parsed: serde_json::Value = serde_json::from_slice(&serialized).unwrap(); + assert_eq!( + parsed.get("recipient_id").and_then(|v| v.as_str()), + Some(user_id) + ); + } + + #[test] + fn test_broadcast_message_truncation() { + // Broadcast uses truncate_message, verify it handles content within + // Discord's 2000-char limit for DMs. + let short = "Hello from broadcast"; + assert_eq!(truncate_message(short), short); + + let long = "x".repeat(2500); + let result = truncate_message(&long); + assert!(result.len() <= 2006); // 1990 content + 16 suffix + assert!(result.ends_with("\n... (truncated)")); + } + + #[test] + fn test_broadcast_dm_validates_snowflake() { + // broadcast_dm rejects invalid Discord snowflake IDs before making + // any API calls. We can call it directly since invalid IDs are + // rejected before any host function is invoked. + assert!(broadcast_dm("", "hi").is_err()); + assert!(broadcast_dm("abc", "hi").is_err()); + assert!(broadcast_dm("12345", "hi").is_err()); // too short + assert!(broadcast_dm("123456789012345678901", "hi").is_err()); // too long + assert!(broadcast_dm("12345678901234567x", "hi").is_err()); // non-digit + } } From de5a1c7b0d0588e1898458870a0796e6dd8a361e Mon Sep 17 00:00:00 2001 From: Joseph Bloggs <252831379+j-bloggs@users.noreply.github.com> Date: Sun, 29 Mar 2026 02:31:49 +1100 Subject: [PATCH 04/23] fix(worker): replace script -qfc with pty-process for injection-safe PTY (#1678) - Add pty-process crate (MIT, tokio async support) for PTY allocation - Spawn claude CLI with pty-process::Command::arg() chaining instead of building a shell string for script -qfc - Eliminates all shell injection surfaces: prompt, model, session_id are passed via execve, never interpreted by a shell - Keep stderr on separate pipe to prevent NDJSON parse breakage (pty-process attaches PTY to all fds by default) - Gate PTY behind #[cfg(unix)] with direct-spawn fallback for Windows CI - Read stdout from PTY master (implements tokio::io::AsyncRead) - Add regression tests: arg vector construction + PTY allocation Addresses review feedback from zmanian and gemini-code-assist. Co-authored-by: j-bloggs Co-authored-by: Claude Opus 4.6 (1M context) --- Cargo.lock | 21 ++++- Cargo.toml | 4 + src/worker/claude_bridge.rs | 176 +++++++++++++++++++++++++++++------- 3 files changed, 162 insertions(+), 39 deletions(-) diff --git a/Cargo.lock b/Cargo.lock index dfea8b45e43..0c524704135 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -3150,7 +3150,7 @@ dependencies = [ "libc", "percent-encoding", "pin-project-lite", - "socket2 0.5.10", + "socket2 0.6.3", "system-configuration", "tokio", "tower-service", @@ -3439,6 +3439,7 @@ dependencies = [ "pgvector", "postgres-types", "pretty_assertions", + "pty-process", "rand 0.8.5", "readabilityrs", "refinery", @@ -3524,7 +3525,7 @@ checksum = "3640c1c38b8e4e43584d8df18be5fc6b0aa314ce6ebf51b53313d4306cca8e46" dependencies = [ "hermit-abi", "libc", - "windows-sys 0.59.0", + "windows-sys 0.61.2", ] [[package]] @@ -4906,6 +4907,16 @@ dependencies = [ "syn 1.0.109", ] +[[package]] +name = "pty-process" +version = "0.5.3" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "71cec9e2670207c5ebb9e477763c74436af3b9091dd550b9fb3c1bec7f3ea266" +dependencies = [ + "rustix 1.1.4", + "tokio", +] + [[package]] name = "pulley-interpreter" version = "28.0.1" @@ -4930,7 +4941,7 @@ dependencies = [ "quinn-udp", "rustc-hash 2.1.1", "rustls 0.23.37", - "socket2 0.5.10", + "socket2 0.6.3", "thiserror 2.0.18", "tokio", "tracing", @@ -4967,9 +4978,9 @@ dependencies = [ "cfg_aliases", "libc", "once_cell", - "socket2 0.5.10", + "socket2 0.6.3", "tracing", - "windows-sys 0.59.0", + "windows-sys 0.60.2", ] [[package]] diff --git a/Cargo.toml b/Cargo.toml index 2d1d5ce6605..0382a2a7a1b 100644 --- a/Cargo.toml +++ b/Cargo.toml @@ -189,6 +189,10 @@ json5 = { version = "0.4", optional = true } [target.'cfg(target_os = "macos")'.dependencies] security-framework = "3" +# PTY allocation for Claude CLI stdout buffering fix (Unix only) +[target.'cfg(unix)'.dependencies] +pty-process = { version = "0.5", features = ["async"] } + # Linux secret-service (GNOME Keyring, KWallet) [target.'cfg(target_os = "linux")'.dependencies] secret-service = { version = "4", features = ["rt-tokio-crypto-rust"] } diff --git a/src/worker/claude_bridge.rs b/src/worker/claude_bridge.rs index b2f674cf029..9b6f475fcd7 100644 --- a/src/worker/claude_bridge.rs +++ b/src/worker/claude_bridge.rs @@ -31,6 +31,7 @@ use std::time::Duration; use serde::{Deserialize, Serialize}; use tokio::io::{AsyncBufReadExt, BufReader}; +#[cfg(not(unix))] use tokio::process::Command; use uuid::Uuid; @@ -340,6 +341,11 @@ impl ClaudeBridgeRuntime { /// Spawn a `claude` CLI process and stream its output. /// + /// Uses a PTY on Unix so Node.js line-buffers stdout instead of + /// full-buffering (which causes the bridge to hang on non-TTY pipes). + /// Arguments are passed via `execve` (no shell) — injection-safe by + /// construction. + /// /// Returns the session_id if captured from the `system` init message. async fn run_claude_session( &self, @@ -347,47 +353,102 @@ impl ClaudeBridgeRuntime { resume_session_id: Option<&str>, extra_env: &std::collections::HashMap, ) -> Result, WorkerError> { - let mut cmd = Command::new("claude"); - cmd.arg("-p") - .arg(prompt) - .arg("--output-format") - .arg("stream-json") - .arg("--verbose") - .arg("--max-turns") - .arg(self.config.max_turns.to_string()) - .arg("--model") - .arg(&self.config.model); - - if let Some(sid) = resume_session_id { - cmd.arg("--resume").arg(sid); - } - - // Inject credentials into the child process environment without - // mutating the global process env (which is unsafe in multi-threaded programs). - cmd.envs(extra_env); + let max_turns_str = self.config.max_turns.to_string(); + + // Spawn with PTY on Unix to fix Node.js stdout buffering. + // All arguments are passed individually via execve — never through + // a shell interpreter. This eliminates shell injection by construction. + #[cfg(unix)] + let (mut child, stdout, stderr) = { + let (pty, pts) = pty_process::open().map_err(|e| WorkerError::ExecutionFailed { + reason: format!("failed to allocate PTY: {}", e), + })?; - cmd.current_dir("/workspace") - .stdout(std::process::Stdio::piped()) - .stderr(std::process::Stdio::piped()); + let mut cmd = pty_process::Command::new("claude"); + cmd = cmd + .arg("-p") + .arg(prompt) + .arg("--output-format") + .arg("stream-json") + .arg("--verbose") + .arg("--max-turns") + .arg(&max_turns_str) + .arg("--model") + .arg(&self.config.model); + + if let Some(sid) = resume_session_id { + cmd = cmd.arg("--resume").arg(sid); + } - let mut child = cmd.spawn().map_err(|e| WorkerError::ExecutionFailed { - reason: format!("failed to spawn claude: {}", e), - })?; + cmd = cmd.envs(extra_env.iter()); + cmd = cmd.current_dir("/workspace"); + // Keep stderr on a separate pipe — pty-process attaches the PTY + // to all fds by default, which would merge stderr into the PTY + // stream and break NDJSON parsing. + cmd = cmd.stderr(std::process::Stdio::piped()); - let stdout = child - .stdout - .take() - .ok_or_else(|| WorkerError::ExecutionFailed { - reason: "failed to capture claude stdout".to_string(), + let mut child = cmd.spawn(pts).map_err(|e| WorkerError::ExecutionFailed { + reason: format!("failed to spawn claude with PTY: {}", e), })?; - let stderr = child - .stderr - .take() - .ok_or_else(|| WorkerError::ExecutionFailed { - reason: "failed to capture claude stderr".to_string(), + let stderr = child + .stderr + .take() + .ok_or_else(|| WorkerError::ExecutionFailed { + reason: "failed to capture claude stderr".to_string(), + })?; + + // stdout comes from the PTY master, which implements AsyncRead + let stdout: Box = Box::new(pty); + (child, stdout, stderr) + }; + + // Non-Unix fallback (Windows CI) — no PTY, direct spawn. + // Claude bridge only runs in Linux Docker containers, so this path + // exists solely for compilation on Windows targets. + #[cfg(not(unix))] + let (mut child, stdout, stderr) = { + let mut cmd = Command::new("claude"); + cmd.arg("-p") + .arg(prompt) + .arg("--output-format") + .arg("stream-json") + .arg("--verbose") + .arg("--max-turns") + .arg(&max_turns_str) + .arg("--model") + .arg(&self.config.model); + + if let Some(sid) = resume_session_id { + cmd.arg("--resume").arg(sid); + } + + cmd.envs(extra_env); + cmd.current_dir("/workspace") + .stdout(std::process::Stdio::piped()) + .stderr(std::process::Stdio::piped()); + + let mut child = cmd.spawn().map_err(|e| WorkerError::ExecutionFailed { + reason: format!("failed to spawn claude: {}", e), })?; + let stdout_pipe = child + .stdout + .take() + .ok_or_else(|| WorkerError::ExecutionFailed { + reason: "failed to capture claude stdout".to_string(), + })?; + let stderr = child + .stderr + .take() + .ok_or_else(|| WorkerError::ExecutionFailed { + reason: "failed to capture claude stderr".to_string(), + })?; + + let stdout: Box = Box::new(stdout_pipe); + (child, stdout, stderr) + }; + // Spawn stderr reader that forwards lines as log events let client_for_stderr = Arc::clone(&self.client); let job_id = self.config.job_id; @@ -1027,4 +1088,51 @@ mod tests { let copied = copy_dir_recursive(nonexistent, dst.path()).unwrap(); assert_eq!(copied, 0); } + + /// Regression test: arguments are passed individually (not via shell string), + /// so shell metacharacters in prompt/model/session_id are harmless. + #[test] + fn command_args_no_shell_interpretation() { + // Prompt, model, and session_id may contain shell metacharacters from + // user-supplied task descriptions or LLM output. Since we use + // Command::arg() (execve), these are passed as literal strings. + let prompt = "Fix the user's bug; echo $HOME && rm -rf /"; + let model = "claude-3-opus-20240229"; + let session_id = "'; DROP TABLE jobs; --"; + + let max_turns = 10u32; + let max_turns_str = max_turns.to_string(); + let args: Vec<&str> = vec![ + "-p", + prompt, + "--output-format", + "stream-json", + "--verbose", + "--max-turns", + &max_turns_str, + "--model", + model, + "--resume", + session_id, + ]; + + // All values present as literal strings — no shell interpretation + // ["-p", prompt, "--output-format", "stream-json", "--verbose", + // "--max-turns", "10", "--model", model, "--resume", session_id] + assert_eq!(args[1], prompt); + assert_eq!(args[8], model); + assert_eq!(args[10], session_id); + // Shell metacharacters preserved, not expanded + assert!(args[1].contains("$HOME")); + assert!(args[1].contains("&&")); + assert!(args[10].contains("'; DROP TABLE")); + } + + /// Verify PTY is available on Unix platforms. + #[cfg(unix)] + #[tokio::test] + async fn pty_opens_successfully() { + let result = pty_process::open(); + assert!(result.is_ok(), "PTY allocation should succeed on Unix"); + } } From fd41bdf4bed3c9b43cf12788b717ac4c0fa8b5b5 Mon Sep 17 00:00:00 2001 From: Joseph Bloggs <252831379+j-bloggs@users.noreply.github.com> Date: Sun, 29 Mar 2026 04:46:08 +1100 Subject: [PATCH 05/23] fix(worker): treat empty LLM response after text output as completion (#1677) * fix(worker): treat empty LLM response after text output as completion When a job's LLM produces a substantive text response (e.g., formatted results from a routine) and the next LLM call returns empty or errors, the worker now treats this as successful completion instead of continuing the loop until failure. Previously, empty responses always triggered TextAction::Continue, causing the loop to re-call the LLM. The LLM had nothing more to say, so the provider returned "Response contained no message or tool call (empty)". This made routine jobs that successfully produced results report as "failed". The fix adds a `has_text_response` flag to JobDelegate: - After any non-empty text response: flag is set - Empty text after flag is set: treated as completion - LLM errors (select_tools/respond_with_tools) after flag: treated as completion instead of propagating - Empty text before any output: still retries (rate-limit backoff) Co-Authored-By: Claude Opus 4.6 (1M context) * fix(worker): restrict error swallowing to EmptyResponse variant only - Add LlmError::EmptyResponse variant for when LLM returns no content - Update nearai_chat and github_copilot providers to emit EmptyResponse instead of InvalidResponse for empty/no-choice responses - try_complete_on_error now only swallows EmptyResponse (not AuthFailed, ContextLengthExceeded, Http, Io, etc.) - Extract is_completion_eligible_error as testable pure function - Log mark_completed errors at warn level instead of silently dropping - Add EmptyResponse to retry and circuit breaker transient classifications - Rewrite test to exercise real classification logic against all variants Addresses review feedback from zmanian and gemini-code-assist. Co-Authored-By: Claude Opus 4.6 (1M context) * refactor(worker): extract mark_completed_or_warn helper to DRY completion logic Extract shared mark-completed + warn-on-failure pattern into a single helper method used by both try_complete_on_error and handle_text_response. Co-Authored-By: Claude Opus 4.6 (1M context) --------- Co-authored-by: j-bloggs Co-authored-by: Claude Opus 4.6 (1M context) --- src/llm/circuit_breaker.rs | 1 + src/llm/error.rs | 3 + src/llm/github_copilot.rs | 6 +- src/llm/nearai_chat.rs | 6 +- src/llm/retry.rs | 1 + src/worker/job.rs | 152 ++++++++++++++++++++++++++++++++++++- 6 files changed, 157 insertions(+), 12 deletions(-) diff --git a/src/llm/circuit_breaker.rs b/src/llm/circuit_breaker.rs index 46f29dedab3..d0f74421292 100644 --- a/src/llm/circuit_breaker.rs +++ b/src/llm/circuit_breaker.rs @@ -234,6 +234,7 @@ fn is_transient(err: &LlmError) -> bool { LlmError::RequestFailed { .. } | LlmError::RateLimited { .. } | LlmError::InvalidResponse { .. } + | LlmError::EmptyResponse { .. } | LlmError::SessionExpired { .. } | LlmError::SessionRenewalFailed { .. } | LlmError::Http(_) diff --git a/src/llm/error.rs b/src/llm/error.rs index 749e7820033..ce516d72022 100644 --- a/src/llm/error.rs +++ b/src/llm/error.rs @@ -17,6 +17,9 @@ pub enum LlmError { #[error("Invalid response from {provider}: {reason}")] InvalidResponse { provider: String, reason: String }, + #[error("Empty response from {provider}: no content returned")] + EmptyResponse { provider: String }, + #[error("Context length exceeded: {used} tokens used, {limit} allowed")] ContextLengthExceeded { used: usize, limit: usize }, diff --git a/src/llm/github_copilot.rs b/src/llm/github_copilot.rs index c7a24b1a32b..6fefe5af656 100644 --- a/src/llm/github_copilot.rs +++ b/src/llm/github_copilot.rs @@ -231,9 +231,8 @@ impl LlmProvider for GithubCopilotProvider { .choices .into_iter() .next() - .ok_or_else(|| LlmError::InvalidResponse { + .ok_or_else(|| LlmError::EmptyResponse { provider: "github_copilot".to_string(), - reason: "No choices in response".to_string(), })?; let (content, _tool_calls) = extract_choice_content(&choice); @@ -309,9 +308,8 @@ impl LlmProvider for GithubCopilotProvider { .choices .into_iter() .next() - .ok_or_else(|| LlmError::InvalidResponse { + .ok_or_else(|| LlmError::EmptyResponse { provider: "github_copilot".to_string(), - reason: "No choices in response".to_string(), })?; let (content, tool_calls) = extract_choice_content(&choice); diff --git a/src/llm/nearai_chat.rs b/src/llm/nearai_chat.rs index 1f6dbb77622..80335d86a5d 100644 --- a/src/llm/nearai_chat.rs +++ b/src/llm/nearai_chat.rs @@ -490,9 +490,8 @@ impl LlmProvider for NearAiChatProvider { .choices .into_iter() .next() - .ok_or_else(|| LlmError::InvalidResponse { + .ok_or_else(|| LlmError::EmptyResponse { provider: "nearai_chat".to_string(), - reason: "No choices in response".to_string(), })?; // Fall back to reasoning_content when content is null (same as @@ -570,9 +569,8 @@ impl LlmProvider for NearAiChatProvider { .choices .into_iter() .next() - .ok_or_else(|| LlmError::InvalidResponse { + .ok_or_else(|| LlmError::EmptyResponse { provider: "nearai_chat".to_string(), - reason: "No choices in response".to_string(), })?; let tool_calls: Vec = choice diff --git a/src/llm/retry.rs b/src/llm/retry.rs index 78a26b27a5b..db76ba8b57e 100644 --- a/src/llm/retry.rs +++ b/src/llm/retry.rs @@ -48,6 +48,7 @@ pub(crate) fn is_retryable(err: &LlmError) -> bool { LlmError::RequestFailed { .. } | LlmError::RateLimited { .. } | LlmError::InvalidResponse { .. } + | LlmError::EmptyResponse { .. } | LlmError::SessionRenewalFailed { .. } | LlmError::Http(_) | LlmError::Io(_) diff --git a/src/worker/job.rs b/src/worker/job.rs index f74d4ec8c6a..94a04290718 100644 --- a/src/worker/job.rs +++ b/src/worker/job.rs @@ -391,6 +391,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), + has_text_response: std::sync::atomic::AtomicBool::new(false), }; let config = AgenticLoopConfig { @@ -1101,6 +1102,15 @@ fn store_fallback_in_metadata( } /// Job delegate: implements `LoopDelegate` for the background job context. +/// Whether an LLM error represents a completion-eligible empty response. +/// +/// Only `EmptyResponse` (provider returned no choices/content) qualifies. +/// Infrastructure errors (`AuthFailed`, `Http`, `Io`, etc.) never qualify — +/// they must propagate even if prior text output was produced. +fn is_completion_eligible_error(error: &crate::error::LlmError) -> bool { + matches!(error, crate::error::LlmError::EmptyResponse { .. }) +} + /// /// Handles: signal channel (stop/ping/user messages), cancellation checks, /// rate-limit retry, parallel tool execution, DB persistence, SSE broadcasting. @@ -1109,6 +1119,10 @@ 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, + /// 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). + has_text_response: std::sync::atomic::AtomicBool, } impl<'a> JobDelegate<'a> { @@ -1161,6 +1175,53 @@ impl<'a> JobDelegate<'a> { finish_reason: crate::llm::FinishReason::Stop, }) } + + /// Mark the job as completed, logging a warning on failure. + async fn mark_completed_or_warn(&self, context: &str) { + if let Err(e) = self.worker.mark_completed().await { + tracing::warn!( + job_id = %self.worker.job_id, + error = %e, + "Failed to mark job completed ({context})" + ); + } + } + + /// If a substantive text response was already produced and the error + /// indicates the LLM simply returned nothing, treat it as successful + /// completion rather than a fatal failure. + /// + /// Only swallows `EmptyResponse` — infrastructure errors (`AuthFailed`, + /// `ContextLengthExceeded`, `Http`, `Io`, etc.) always propagate. + /// + /// Returns `Some(empty RespondOutput)` when the error should be swallowed, + /// `None` when it should propagate normally. + async fn try_complete_on_error( + &self, + context: &str, + error: &crate::error::LlmError, + ) -> Option { + if !is_completion_eligible_error(error) { + return None; + } + if !self + .has_text_response + .load(std::sync::atomic::Ordering::Relaxed) + { + return None; + } + tracing::info!( + job_id = %self.worker.job_id, + error = %error, + "{context} empty response after text output — treating as completion" + ); + self.mark_completed_or_warn(context).await; + Some(crate::llm::RespondOutput { + result: RespondResult::Text(String::new()), + usage: crate::llm::TokenUsage::default(), + finish_reason: crate::llm::FinishReason::Stop, + }) + } } #[async_trait] @@ -1291,7 +1352,12 @@ impl<'a> LoopDelegate for JobDelegate<'a> { Err(crate::error::LlmError::RateLimited { retry_after, .. }) => { return self.handle_rate_limit(retry_after, "tool selection").await; } - Err(e) => return Err(e.into()), + Err(e) => { + if let Some(output) = self.try_complete_on_error("select_tools", &e).await { + return Ok(output); + } + return Err(e.into()); + } }; // Fall back to respond_with_tools @@ -1321,7 +1387,12 @@ impl<'a> LoopDelegate for JobDelegate<'a> { self.handle_rate_limit(retry_after, "respond_with_tools") .await } - Err(e) => Err(e.into()), + Err(e) => { + if let Some(output) = self.try_complete_on_error("respond_with_tools", &e).await { + return Ok(output); + } + Err(e.into()) + } } } @@ -1330,9 +1401,22 @@ impl<'a> LoopDelegate for JobDelegate<'a> { text: &str, reason_ctx: &mut ReasoningContext, ) -> TextAction { - // Empty text from rate-limit backoff retry — skip processing and let the - // loop proceed to the next iteration which will re-call the LLM. + // 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)"). if text.is_empty() { + if self + .has_text_response + .load(std::sync::atomic::Ordering::Relaxed) + { + tracing::debug!( + job_id = %self.worker.job_id, + "Empty response after text output — treating as completion" + ); + self.mark_completed_or_warn("empty text response").await; + return TextAction::Return(LoopOutcome::Response(String::new())); + } + // No prior text response — this is likely a rate-limit backoff retry. return TextAction::Continue; } @@ -1348,6 +1432,10 @@ impl<'a> LoopDelegate for JobDelegate<'a> { return TextAction::Return(LoopOutcome::Response(text.to_string())); } + // Track that a substantive response has been produced. + self.has_text_response + .store(true, std::sync::atomic::Ordering::Relaxed); + // Add assistant response to context reason_ctx.messages.push(ChatMessage::assistant(text)); @@ -2285,4 +2373,60 @@ mod tests { assert_eq!(telegram[0].0, "owner-scope"); assert_eq!(telegram[0].1.content, "hello from routine"); } + + /// Regression test: only `EmptyResponse` errors are eligible for + /// completion-swallowing. Infrastructure errors must always propagate. + #[test] + fn is_completion_eligible_only_matches_empty_response() { + use crate::error::LlmError; + + // EmptyResponse is eligible + assert!(super::is_completion_eligible_error( + &LlmError::EmptyResponse { + provider: "test".to_string(), + } + )); + + // All other variants are NOT eligible + assert!(!super::is_completion_eligible_error( + &LlmError::InvalidResponse { + provider: "test".to_string(), + reason: "parse error".to_string(), + } + )); + assert!(!super::is_completion_eligible_error( + &LlmError::AuthFailed { + provider: "test".to_string(), + } + )); + assert!(!super::is_completion_eligible_error( + &LlmError::ContextLengthExceeded { + used: 100_000, + limit: 50_000, + } + )); + assert!(!super::is_completion_eligible_error( + &LlmError::ModelNotAvailable { + provider: "test".to_string(), + model: "gpt-4".to_string(), + } + )); + assert!(!super::is_completion_eligible_error( + &LlmError::RequestFailed { + provider: "test".to_string(), + reason: "timeout".to_string(), + } + )); + assert!(!super::is_completion_eligible_error( + &LlmError::SessionExpired { + provider: "test".to_string(), + } + )); + assert!(!super::is_completion_eligible_error( + &LlmError::SessionRenewalFailed { + provider: "test".to_string(), + reason: "timeout".to_string(), + } + )); + } } From 8a320ae9db4f7fdada609a30528bee6116cbe71c Mon Sep 17 00:00:00 2001 From: Illia Polosukhin Date: Sat, 28 Mar 2026 12:27:43 -0700 Subject: [PATCH 06/23] fix(routines): complete full_job execution reliability overhaul (#1650) MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit * fix(routines): persist full LLM transcript and remove sandbox gate for full_job Routine execution output was invisible — routine_fire returned a one-liner, routine_history had no actual output, and the conversation thread contained only a summary. Full-job routines also hard-failed without Docker. Three fixes: 1. **Full transcript persistence**: execute_lightweight now persists every message (prompt, LLM responses, tool calls with params, tool results) to the routine's conversation thread as it executes, not just a summary after the fact. 2. **Routine output visibility**: routine_history includes conversation_id and recent_output messages. routine_fire tells the user to check routine_history. Web detail page has a "View Execution Thread" button that navigates to the chat tab. ROUTINE_OK stores "No issues found" instead of None. Full-job summary pulls actual job output instead of generic "Job X finished". 3. **Remove SandboxReadiness gate**: full_job routines dispatch through the scheduler like regular /job commands — no Docker required. The SandboxReadiness enum is removed entirely. Co-Authored-By: Claude Opus 4.6 (1M context) * style: apply cargo fmt Co-Authored-By: Claude Opus 4.6 (1M context) * fix(worker): treat AutonomousUnavailable tool errors as recoverable The job worker crashed the entire job when a tool was denied for autonomous execution (e.g. secret_list). The error was already recorded in reason_ctx for the LLM to see, but process_tool_result_job returned Err which propagated through the agentic loop and terminated the job. Now all tool errors (including AutonomousUnavailable) return Ok, letting the LLM see the denial and try a different approach. Co-Authored-By: Claude Opus 4.6 (1M context) * fix(llm): sanitize tool names for OpenAI Codex Responses API The Codex API requires tool names to match `^[a-zA-Z0-9_-]+$` but MCP/extension tools can have dots in their names (e.g. `mcp.server.tool`). This caused HTTP 400 errors when the job worker sent tool calls back to the LLM. Sanitize tool names in both `convert_tool_definition` and `convert_message` (function_call items) by replacing invalid characters with underscores. Co-Authored-By: Claude Opus 4.6 (1M context) * fix(routines): inject execution context into full_job description [skip-regression-check] When a full_job routine dispatches a job, the LLM had no context that it was already executing inside a routine. It wasted iterations on infrastructure (discovering tools, creating routines, setting up auth) instead of doing the actual work. Prepend a clear directive to the job description telling the LLM that tools and the routine are already configured, and to execute the task directly. Co-Authored-By: Claude Opus 4.6 (1M context) * fix(mcp): auto-refresh expired OAuth tokens on access [skip-regression-check] When IronClaw restarts, MCP servers fail with "Secret has expired" because get_access_token() checks token expiry locally and returns an error before any HTTP request is made — so the existing 401-retry refresh logic never triggers. Now get_access_token() catches SecretError::Expired and automatically calls refresh_access_token() using the stored refresh token. If the refresh succeeds, the new token is returned transparently. If it fails, the error message includes both the expiry and the refresh failure. Co-Authored-By: Claude Opus 4.6 (1M context) * fix(mcp): align refresh token naming and set expiry on stored tokens Two bugs prevented MCP OAuth token auto-refresh on restart: 1. Naming mismatch: the hosted OAuth flow stored the refresh token as `{token_secret_name}_refresh_token` (e.g. `mcp_notion_access_token_refresh_token`) but `McpServerConfig::refresh_token_secret_name()` returned `mcp_notion_refresh_token`. The refresh token was there but unfindable. 2. Missing expiry: `store_tokens` in auth.rs never called `with_expiry()` even though `AccessToken::expires_in` was available. Combined with the fix from the previous commit (auto-refresh on Expired), tokens stored via the MCP auth flow will now also trigger refresh correctly. Co-Authored-By: Claude Opus 4.6 (1M context) * fix(web): show activity and transitions for agent jobs in job detail [skip-regression-check] The job events endpoint only checked sandbox jobs for ownership, returning 404 for agent jobs dispatched from routines. The detail handler also returned empty transitions for agent jobs. - events handler: fall back to agent job ownership check - detail handler: populate transitions from job's state history Co-Authored-By: Claude Opus 4.6 (1M context) * feat(routines): expose max_iterations for full_job routines (default 25) The max_iterations parameter was hardcoded to 10 and not configurable via routine_create or routine_update, causing complex tasks to hit the iteration cap. - Add max_iterations to full_job execution schema (1-200, default 25) - Thread it through parse → build → RoutineAction - Support updating via routine_update - Raise default from 10 to 25 Co-Authored-By: Claude Opus 4.6 (1M context) * fix(routines): break self-dialogue loop after full_job plan execution After plan execution, the completion-check Q&A ("Is the job complete?" / "No, not complete...") was left in the message context, causing the agentic loop to repeat the same analysis instead of calling tools. Replace the stale dialogue with an action-oriented continuation prompt that instructs the LLM to use tools for remaining work. Also strip tags from all job output since they're only meaningful for interactive chat sessions. Co-Authored-By: Claude Opus 4.6 (1M context) * fix(repl): prevent test hang in single-message mode In single-message mode, start() stored a clone of the mpsc sender in self.msg_tx for approval injection. After the thread sent /quit and exited, the stored clone kept the stream alive, so stream.next() blocked forever in the test assertion that the stream ends. Skip storing the sender in single-message mode since interactive approval is not needed. Co-Authored-By: Claude Opus 4.6 (1M context) * fix(jobs): treat text responses as final answer in agentic loop When the LLM produces a non-empty text response with no tool intent (already filtered by the nudge mechanism), it is the job's final answer. Previously, handle_text_response only exited the loop if the text matched rigid completion phrases like "job is complete". Natural summaries like "Weekly review completed and saved to Notion" were added to context and the loop continued, causing the LLM to restate the same summary until max_iterations was hit. Now any non-empty text response marks the job complete and stops the loop, matching the chat dispatcher behavior. Co-Authored-By: Claude Opus 4.6 (1M context) * perf(tests): reduce skills catalog network failure test from 10s to 1s The test_search_returns_error_on_network_failure test connects to an unreachable RFC 5737 TEST-NET IP and waited for the full 10s production REQUEST_TIMEOUT. Add with_url_and_timeout test helper and use a 1s timeout instead. [skip-regression-check] Co-Authored-By: Claude Opus 4.6 (1M context) * fix(tools): accept 'message' as alias for 'content' in message tool LLMs frequently call the message tool with {"message": "..."} instead of {"content": "..."}. Fall back to the 'message' key when 'content' is missing to avoid InvalidParameters errors during autonomous job execution. Co-Authored-By: Claude Opus 4.6 (1M context) * fix(tools): attach thread_id for gateway broadcast in message tool When the message tool broadcasts to all channels (channel=null), it sent an OutgoingResponse without a thread_id. The gateway silently dropped these messages (returned Ok but never sent the SSE event), so they appeared in repl but not in the web UI. The thread_id was only populated when channel was explicitly "gateway". Now it is always populated from notify_thread_id metadata, so broadcast_all delivers to the gateway correctly. Co-Authored-By: Claude Opus 4.6 (1M context) * fix(gateway): return error instead of silently dropping messages Gateway broadcast() and respond() previously returned Ok(()) when thread_id was missing, silently swallowing the message. Callers (message tool, agent loop) believed delivery succeeded when it didn't. Now returns ChannelError::MissingRoutingTarget so callers can detect and report the failure. Four regression tests verify the contract: respond/broadcast with and without thread_id. Co-Authored-By: Claude Opus 4.6 (1M context) * fix: resolve rebase conflicts with staging Restore sandbox_readiness field removed by pre-rebase commits (staging still uses it). Update repl test to match staging's single-message behavior (no longer sends /quit). Add missing reasoning field to ToolCall in codex test. Co-Authored-By: Claude Opus 4.6 (1M context) * fix(tools): log error when routine conversation lookup fails The routine_history tool silently swallowed errors from get_or_create_routine_conversation, returning empty output without any diagnostic logging. Add tracing::warn so failures are visible in logs. [skip-regression-check] Co-Authored-By: Claude Opus 4.6 (1M context) * fix: address PR #1650 review comments - E2E test: accept submitted/accepted as success states in job assertion - TimeTool: remove operation from required schema (defaults to "now") - jobs handler: log DB errors server-side, return generic message to client - routines handler: use read-only find_routine_conversation on GET - codex provider: reverse-map sanitized tool names so MCP tools resolve Co-Authored-By: Claude Opus 4.6 (1M context) * fix: address zmanian review feedback on PR #1650 - MCP refresh token: fall back to legacy secret name (mcp_{name}_refresh_token) so existing users don't need to re-authenticate after the naming fix - Job worker: replace fragile messages.pop() with truncate-to-saved-count to avoid maintenance hazard if message flow changes - Document cost implications of max_iterations 10->25 default bump - Revert Cargo.toml dist profile change (thin LTO comment, codegen-units=16) as it's unrelated to this PR Co-Authored-By: Claude Opus 4.6 (1M context) * fix: resolve rebase conflicts and address new Copilot comments - Fix no_silent_drop tests for updated GatewayConfig (user_id moved to GatewayChannel::new second arg, user_tokens removed) - Fix handle_text_response param name (_reason_ctx -> reason_ctx) - Fix missing has_text_response field in test JobDelegate - Propagate row.get errors in find_routine_conversation instead of unwrap_or_default - Only fall back to legacy refresh token name on NotFound/Expired, propagate real errors (DB, decryption) Co-Authored-By: Claude Opus 4.6 (1M context) --------- Co-authored-by: Claude Opus 4.6 (1M context) --- Cargo.toml | 3 +- src/agent/dispatcher.rs | 20 +++ src/agent/mod.rs | 1 + src/agent/routine.rs | 7 +- src/agent/routine_engine.rs | 14 +- src/channels/repl.rs | 37 ++--- src/channels/web/handlers/jobs.rs | 37 +++-- src/channels/web/handlers/routines.rs | 11 ++ src/channels/web/mod.rs | 16 +- src/channels/web/static/app.js | 12 ++ src/channels/web/tests/mod.rs | 1 + src/channels/web/tests/no_silent_drop.rs | 93 ++++++++++++ src/channels/web/types.rs | 1 + src/db/libsql/conversations.rs | 35 +++++ src/db/mod.rs | 7 + src/db/postgres.rs | 10 ++ src/history/store.rs | 21 +++ src/llm/openai_codex_provider.rs | 136 ++++++++++++++++- src/skills/catalog.rs | 12 +- src/tools/builtin/message.rs | 40 ++++- src/tools/builtin/routine.rs | 72 ++++++++- src/tools/builtin/time.rs | 9 +- src/tools/mcp/auth.rs | 35 ++++- src/tools/mcp/client.rs | 30 ++++ src/tools/mcp/config.rs | 19 +++ src/util.rs | 22 +++ src/worker/job.rs | 146 ++++++++++++++++--- tests/e2e/ironclaw_e2e.egg-info/SOURCES.txt | 3 + tests/e2e/mock_llm.py | 82 +++++++++++ tests/e2e/scenarios/test_routine_full_job.py | 133 +++++++++++++++++ tests/e2e_routine_heartbeat.rs | 22 ++- tests/e2e_telegram_message_routing.rs | 2 +- tests/support/gateway_workflow_harness.rs | 3 +- tests/support/test_rig.rs | 4 +- 34 files changed, 988 insertions(+), 108 deletions(-) create mode 100644 src/channels/web/tests/no_silent_drop.rs create mode 100644 tests/e2e/scenarios/test_routine_full_job.py diff --git a/Cargo.toml b/Cargo.toml index 0382a2a7a1b..fbd3d6eec0d 100644 --- a/Cargo.toml +++ b/Cargo.toml @@ -252,8 +252,7 @@ strip = true # Remove debug symbols from release binaries # The profile that 'cargo dist' will build with [profile.dist] inherits = "release" -lto = "fat" # Full cross-crate LTO (slow build, better codegen) -codegen-units = 1 # Single codegen unit for maximum optimization +lto = "thin" # Config for 'dist' [workspace.metadata.dist] diff --git a/src/agent/dispatcher.rs b/src/agent/dispatcher.rs index 9e639171f52..4420a450470 100644 --- a/src/agent/dispatcher.rs +++ b/src/agent/dispatcher.rs @@ -1269,6 +1269,14 @@ pub(crate) fn extract_suggestions(text: &str) -> (String, Vec) { (cleaned, suggestions) } +/// Remove `` tags from a response, returning only the cleaned text. +/// +/// Convenience wrapper around [`extract_suggestions`] for callers that don't +/// need the parsed suggestion list (e.g. job worker, plan completion check). +pub(crate) fn strip_suggestions(text: &str) -> String { + extract_suggestions(text).0 +} + #[cfg(test)] mod tests { use std::sync::Arc; @@ -2539,6 +2547,18 @@ mod tests { assert_eq!(suggestions, vec!["ok"]); // safety: test } + #[test] + fn test_strip_suggestions_removes_tags() { + let input = "The job is complete.\n[\"Check logs\"]"; + assert_eq!(super::strip_suggestions(input), "The job is complete."); // safety: test + } + + #[test] + fn test_strip_suggestions_no_tag_passthrough() { + let input = "Plain text without tags."; + assert_eq!(super::strip_suggestions(input), input); // safety: test + } + #[test] fn test_tool_error_format_includes_tool_name() { let tool_name = "http"; diff --git a/src/agent/mod.rs b/src/agent/mod.rs index e7242845e99..79616aaed71 100644 --- a/src/agent/mod.rs +++ b/src/agent/mod.rs @@ -36,6 +36,7 @@ pub(crate) use agent_loop::truncate_for_preview; pub use agent_loop::{Agent, AgentDeps}; pub use compaction::{CompactionResult, ContextCompactor}; pub use context_monitor::{CompactionStrategy, ContextBreakdown, ContextMonitor}; +pub(crate) use dispatcher::strip_suggestions; pub use heartbeat::{ HeartbeatConfig, HeartbeatResult, HeartbeatRunner, spawn_heartbeat, spawn_multi_user_heartbeat, }; diff --git a/src/agent/routine.rs b/src/agent/routine.rs index 26e769da7f7..5a57a8a66b7 100644 --- a/src/agent/routine.rs +++ b/src/agent/routine.rs @@ -265,8 +265,13 @@ fn default_max_tokens() -> u32 { 4096 } +/// Default max agentic loop iterations for full_job routines. +/// +/// Raised from 10 to 25 to accommodate multi-step tool chains that +/// stalled at the old cap. Worst-case LLM cost is 2.5x higher per run; +/// callers needing tighter budgets should set `max_iterations` explicitly. fn default_max_iterations() -> u32 { - 10 + 25 } fn default_max_tool_rounds() -> u32 { diff --git a/src/agent/routine_engine.rs b/src/agent/routine_engine.rs index 64c3b94c1be..3687ebd4f8e 100644 --- a/src/agent/routine_engine.rs +++ b/src/agent/routine_engine.rs @@ -1292,11 +1292,23 @@ async fn execute_full_job( } metadata["notify_user"] = serde_json::json!(&routine.notify.user); + // Prepend execution context so the LLM knows it's already inside a + // routine and should execute the task directly — not set up infrastructure. + let contextualized_description = format!( + "IMPORTANT: You are executing inside routine \"{routine_name}\". \ + The routine and its schedule are already configured. \ + Tools and credentials are already set up. \ + Do NOT create routines, jobs, or try to discover/install/authenticate tools. \ + Execute the task directly.\n\n{desc}", + routine_name = routine.name, + desc = execution.description, + ); + let job_id = scheduler .dispatch_job( &routine.user_id, execution.title, - execution.description, + &contextualized_description, Some(metadata), ) .await diff --git a/src/channels/repl.rs b/src/channels/repl.rs index 41d73a8c094..27b6ea40259 100644 --- a/src/channels/repl.rs +++ b/src/channels/repl.rs @@ -492,10 +492,12 @@ impl Channel for ReplChannel { async fn start(&self) -> Result { let (tx, rx) = mpsc::channel(32); - // Approval prompts inject responses back through this sender. - // In single-message mode we keep it until the turn finishes, then - // drop it after enqueuing /quit so the receiver stream can close. - if let Ok(mut guard) = self.msg_tx.lock() { + // Store tx so send_status can inject approval responses directly. + // Skip for single-message mode — no interactive approval is needed + // and the extra sender would keep the stream open after /quit. + if self.single_message.is_none() + && let Ok(mut guard) = self.msg_tx.lock() + { *guard = Some(tx.clone()); } let single_message = self.single_message.clone(); @@ -914,8 +916,10 @@ mod tests { use super::*; + /// Regression: single-message mode must close the stream after the one + /// message so callers (and tests) don't hang forever. #[tokio::test] - async fn single_message_mode_sends_message_then_quit() { + async fn single_message_mode_sends_message_and_closes_stream() { let repl = ReplChannel::with_message("hi".to_string()); let mut stream = repl.start().await.expect("repl start should succeed"); @@ -926,30 +930,15 @@ mod tests { assert_eq!(first.channel, "repl"); assert_eq!(first.content, "hi"); - assert!( - timeout(Duration::from_millis(100), stream.next()) - .await - .is_err(), - "single-message mode should wait for the turn to finish before quitting" - ); - - repl.respond(&first, OutgoingResponse::text("done")) - .await - .expect("respond should succeed"); - - let second = timeout(Duration::from_secs(1), stream.next()) - .await - .expect("timed out waiting for quit message") - .expect("quit message missing"); - assert_eq!(second.channel, "repl"); - assert_eq!(second.content, "/quit"); - + // The spawned thread sent the message and returned, dropping its + // sender. Because we skip storing a clone in msg_tx for single- + // message mode, the stream should close immediately. assert!( timeout(Duration::from_secs(1), stream.next()) .await .expect("timed out waiting for stream to close") .is_none(), - "stream should end after /quit" + "stream should end after the single message" ); } } diff --git a/src/channels/web/handlers/jobs.rs b/src/channels/web/handlers/jobs.rs index b171561c767..aca3e97cdb8 100644 --- a/src/channels/web/handlers/jobs.rs +++ b/src/channels/web/handlers/jobs.rs @@ -236,6 +236,18 @@ pub async fn jobs_detail_handler( (end - start).num_seconds().max(0) as u64 }); + // Build transitions from the job's state transition history. + let transitions: Vec = ctx + .transitions + .iter() + .map(|t| TransitionInfo { + from: t.from.to_string(), + to: t.to.to_string(), + timestamp: t.timestamp.to_rfc3339(), + reason: t.reason.clone(), + }) + .collect(); + // Only show prompt bar for jobs that have a running worker (Pending/InProgress). // Stuck jobs have no active worker loop, so messages would be silently dropped. let is_promptable = matches!( @@ -255,7 +267,7 @@ pub async fn jobs_detail_handler( project_dir: None, browse_url: None, job_mode: None, - transitions: Vec::new(), + transitions, can_restart: state.scheduler.is_some(), can_prompt: is_promptable && state.scheduler.is_some(), job_kind: Some("agent".to_string()), @@ -643,25 +655,28 @@ pub async fn jobs_events_handler( .parse() .map_err(|_| (StatusCode::BAD_REQUEST, "Invalid job ID".to_string()))?; - // Verify ownership before returning events. - match store.get_sandbox_job(job_id).await { - Ok(Some(job)) => { - if job.user_id != user.user_id { - return Err((StatusCode::NOT_FOUND, "Job not found".to_string())); - } - } + // Verify ownership before returning events (check both sandbox and agent jobs). + let is_owner = match store.get_sandbox_job(job_id).await { + Ok(Some(job)) => job.user_id == user.user_id, Ok(None) => { - return Err((StatusCode::NOT_FOUND, "Job not found".to_string())); + // Fall back to agent job ownership check. + match store.get_job(job_id).await { + Ok(Some(ctx)) => ctx.user_id == user.user_id, + _ => false, + } } Err(e) => { - return Err(db_error("jobs_handler", e)); + return Err(db_error("jobs_events_handler", e)); } + }; + if !is_owner { + return Err((StatusCode::NOT_FOUND, "Job not found".to_string())); } let events = store .list_job_events(job_id, None) .await - .map_err(|e| (StatusCode::INTERNAL_SERVER_ERROR, e.to_string()))?; + .map_err(|e| db_error("jobs_events_handler", e))?; let events_json: Vec = events .into_iter() diff --git a/src/channels/web/handlers/routines.rs b/src/channels/web/handlers/routines.rs index fc56b187fdd..5597a47c92c 100644 --- a/src/channels/web/handlers/routines.rs +++ b/src/channels/web/handlers/routines.rs @@ -122,6 +122,16 @@ pub async fn routines_detail_handler( .collect(); let routine_info = RoutineInfo::from_routine(&routine); + // Read-only lookup — do not create a conversation on a GET request. + // The conversation is created lazily when the routine first executes. + let conversation_id = store + .find_routine_conversation(routine.id, &routine.user_id) + .await + .unwrap_or_else(|e| { + tracing::warn!(routine_id = %routine.id, error = %e, "Failed to look up routine conversation"); + None + }); + Ok(Json(RoutineDetailResponse { id: routine.id, name: routine.name.clone(), @@ -139,6 +149,7 @@ pub async fn routines_detail_handler( run_count: routine.run_count, consecutive_failures: routine.consecutive_failures, created_at: routine.created_at.to_rfc3339(), + conversation_id, recent_runs, })) } diff --git a/src/channels/web/mod.rs b/src/channels/web/mod.rs index 2cadece607e..77968223c96 100644 --- a/src/channels/web/mod.rs +++ b/src/channels/web/mod.rs @@ -347,10 +347,10 @@ impl Channel for GatewayChannel { let thread_id = match &msg.thread_id { Some(tid) => tid.clone(), None => { - tracing::warn!( - "Gateway respond with no thread_id — skipping (clients would drop it)" - ); - return Ok(()); + return Err(ChannelError::MissingRoutingTarget { + name: "gateway".to_string(), + reason: "respond() requires a thread_id on the incoming message".to_string(), + }); } }; @@ -507,10 +507,10 @@ impl Channel for GatewayChannel { let thread_id = match response.thread_id { Some(tid) => tid, None => { - tracing::warn!( - "Gateway broadcast with no thread_id — skipping (clients would drop it)" - ); - return Ok(()); + return Err(ChannelError::MissingRoutingTarget { + name: "gateway".to_string(), + reason: "broadcast() requires a thread_id on the response".to_string(), + }); } }; self.state.sse.broadcast_for_user( diff --git a/src/channels/web/static/app.js b/src/channels/web/static/app.js index 9cfd35df420..c0c15acff9a 100644 --- a/src/channels/web/static/app.js +++ b/src/channels/web/static/app.js @@ -4265,6 +4265,13 @@ function renderRoutineDetail(routine) { html += '

Action

' + '
' + escapeHtml(JSON.stringify(routine.action, null, 2)) + '
'; + // Conversation thread link + if (routine.conversation_id) { + html += ''; + } + // Recent runs if (routine.recent_runs && routine.recent_runs.length > 0) { html += '

Recent Runs

' @@ -6190,6 +6197,11 @@ document.addEventListener('click', function(e) { switchTab('jobs'); openJobDetail(el.dataset.id); break; + case 'view-routine-thread': + e.preventDefault(); + switchTab('chat'); + switchThread(el.dataset.id); + break; case 'copy-tee-report': copyTeeReport(); break; diff --git a/src/channels/web/tests/mod.rs b/src/channels/web/tests/mod.rs index fa6db197135..daeee1bf982 100644 --- a/src/channels/web/tests/mod.rs +++ b/src/channels/web/tests/mod.rs @@ -1,3 +1,4 @@ //! Integration tests for the web gateway module. mod multi_tenant; +mod no_silent_drop; diff --git a/src/channels/web/tests/no_silent_drop.rs b/src/channels/web/tests/no_silent_drop.rs new file mode 100644 index 00000000000..5ffd9f04cf2 --- /dev/null +++ b/src/channels/web/tests/no_silent_drop.rs @@ -0,0 +1,93 @@ +//! Regression tests: the gateway channel must never silently drop messages. +//! +//! Previously, `respond()` and `broadcast()` returned `Ok(())` when thread_id +//! was missing, making callers believe the message was delivered when it wasn't. +//! These tests ensure that missing routing info produces an explicit error. + +use crate::channels::channel::{Channel, IncomingMessage, OutgoingResponse}; +use crate::channels::web::GatewayChannel; +use crate::config::GatewayConfig; +use crate::error::ChannelError; + +fn test_gateway() -> GatewayChannel { + GatewayChannel::new( + GatewayConfig { + host: "127.0.0.1".to_string(), + port: 0, + auth_token: Some("test-token".to_string()), + workspace_read_scopes: vec![], + memory_layers: vec![], + }, + "test-user".to_string(), + ) +} + +#[tokio::test] +async fn gateway_respond_without_thread_id_returns_error() { + let gw = test_gateway(); + let msg = IncomingMessage::new("gateway", "test-user", "hello"); + // msg has no thread_id by default + assert!(msg.thread_id.is_none()); + + let response = OutgoingResponse::text("reply"); + let result = gw.respond(&msg, response).await; + + assert!( + result.is_err(), + "respond() must not silently succeed without thread_id" + ); + assert!( + matches!(result, Err(ChannelError::MissingRoutingTarget { .. })), + "Expected MissingRoutingTarget, got: {:?}", + result + ); +} + +#[tokio::test] +async fn gateway_respond_with_thread_id_succeeds() { + let gw = test_gateway(); + let mut msg = IncomingMessage::new("gateway", "test-user", "hello"); + msg.thread_id = Some("thread-123".to_string()); + + let response = OutgoingResponse::text("reply"); + let result = gw.respond(&msg, response).await; + + assert!( + result.is_ok(), + "respond() should succeed with thread_id: {:?}", + result + ); +} + +#[tokio::test] +async fn gateway_broadcast_without_thread_id_returns_error() { + let gw = test_gateway(); + let response = OutgoingResponse::text("notification"); + // response has no thread_id by default + + let result = gw.broadcast("test-user", response).await; + + assert!( + result.is_err(), + "broadcast() must not silently succeed without thread_id" + ); + assert!( + matches!(result, Err(ChannelError::MissingRoutingTarget { .. })), + "Expected MissingRoutingTarget, got: {:?}", + result + ); +} + +#[tokio::test] +async fn gateway_broadcast_with_thread_id_succeeds() { + let gw = test_gateway(); + let response = OutgoingResponse::text("notification").in_thread("thread-456".to_string()); + + let result = gw.broadcast("test-user", response).await; + + assert!( + result.is_ok(), + "broadcast() should succeed with thread_id: {:?}", + result + ); +} diff --git a/src/channels/web/types.rs b/src/channels/web/types.rs index 8698c030792..9ecece5725e 100644 --- a/src/channels/web/types.rs +++ b/src/channels/web/types.rs @@ -768,6 +768,7 @@ pub struct RoutineDetailResponse { pub run_count: u64, pub consecutive_failures: u32, pub created_at: String, + pub conversation_id: Option, pub recent_runs: Vec, } diff --git a/src/db/libsql/conversations.rs b/src/db/libsql/conversations.rs index 911ee8631f1..4f9f1079309 100644 --- a/src/db/libsql/conversations.rs +++ b/src/db/libsql/conversations.rs @@ -290,6 +290,41 @@ impl ConversationStore for LibSqlBackend { result } + async fn find_routine_conversation( + &self, + routine_id: Uuid, + user_id: &str, + ) -> Result, DatabaseError> { + let conn = self.connect().await?; + let rid = routine_id.to_string(); + let mut rows = conn + .query( + r#" + SELECT id FROM conversations + WHERE user_id = ?1 AND json_extract(metadata, '$.routine_id') = ?2 + LIMIT 1 + "#, + params![user_id, rid], + ) + .await + .map_err(|e| DatabaseError::Query(e.to_string()))?; + + if let Some(row) = rows + .next() + .await + .map_err(|e| DatabaseError::Query(e.to_string()))? + { + let id_str: String = row.get(0).map_err(|e| { + DatabaseError::Query(format!("Failed to read conversation id: {e}")) + })?; + let id = id_str + .parse() + .map_err(|_| DatabaseError::Serialization("Invalid UUID".to_string()))?; + return Ok(Some(id)); + } + Ok(None) + } + /// Uses BEGIN IMMEDIATE to serialize concurrent writers and prevent /// duplicate heartbeat conversations (TOCTOU race). async fn get_or_create_heartbeat_conversation( diff --git a/src/db/mod.rs b/src/db/mod.rs index 14cad543c5d..e2a81412c18 100644 --- a/src/db/mod.rs +++ b/src/db/mod.rs @@ -391,6 +391,13 @@ pub trait ConversationStore: Send + Sync { routine_name: &str, user_id: &str, ) -> Result; + /// Read-only lookup for an existing routine conversation. Returns `None` + /// if the routine has never executed (no conversation created yet). + async fn find_routine_conversation( + &self, + routine_id: Uuid, + user_id: &str, + ) -> Result, DatabaseError>; async fn get_or_create_heartbeat_conversation( &self, user_id: &str, diff --git a/src/db/postgres.rs b/src/db/postgres.rs index 2fba0b5320f..462c3c46020 100644 --- a/src/db/postgres.rs +++ b/src/db/postgres.rs @@ -137,6 +137,16 @@ impl ConversationStore for PgBackend { .await } + async fn find_routine_conversation( + &self, + routine_id: Uuid, + user_id: &str, + ) -> Result, DatabaseError> { + self.store + .find_routine_conversation(routine_id, user_id) + .await + } + async fn get_or_create_heartbeat_conversation( &self, user_id: &str, diff --git a/src/history/store.rs b/src/history/store.rs index e6e869b6f7c..c2b75a88c3f 100644 --- a/src/history/store.rs +++ b/src/history/store.rs @@ -1771,6 +1771,27 @@ impl Store { Ok(row.get("id")) } + /// Read-only lookup for an existing routine conversation. + pub async fn find_routine_conversation( + &self, + routine_id: Uuid, + user_id: &str, + ) -> Result, DatabaseError> { + let conn = self.conn().await?; + let rid = routine_id.to_string(); + let row = conn + .query_opt( + r#" + SELECT id FROM conversations + WHERE user_id = $1 AND metadata->>'routine_id' = $2 + LIMIT 1 + "#, + &[&user_id, &rid], + ) + .await?; + Ok(row.map(|r| r.get("id"))) + } + /// Get or create the singleton heartbeat conversation for a user. /// /// Looks for a conversation where `metadata->>'thread_type' = 'heartbeat'`. diff --git a/src/llm/openai_codex_provider.rs b/src/llm/openai_codex_provider.rs index 3449a08a390..f1f688ac8ad 100644 --- a/src/llm/openai_codex_provider.rs +++ b/src/llm/openai_codex_provider.rs @@ -276,8 +276,33 @@ impl LlmProvider for OpenAiCodexProvider { &self, request: ToolCompletionRequest, ) -> Result { + // Build a reverse map so we can translate sanitized names back to originals. + // Only needed when sanitization actually changes a name (e.g. MCP tools with dots). + let name_map: std::collections::HashMap = request + .tools + .iter() + .filter_map(|t| { + let sanitized = sanitize_tool_name(&t.name); + if sanitized != t.name { + Some((sanitized, t.name.clone())) + } else { + None + } + }) + .collect(); + let body = self.build_request_body(&request.messages, Some(&request.tools)); - let parsed = self.send_request(body).await?; + let mut parsed = self.send_request(body).await?; + + // Reverse-map sanitized tool names back to originals so the caller + // can look them up in the tool registry. + if !name_map.is_empty() { + for tc in &mut parsed.tool_calls { + if let Some(original) = name_map.get(&tc.name) { + tc.name = original.clone(); + } + } + } let finish_reason = if !parsed.tool_calls.is_empty() { FinishReason::ToolUse @@ -421,7 +446,7 @@ fn convert_message(msg: &ChatMessage, index: usize) -> Vec { serde_json::json!({ "type": "function_call", "call_id": tc.id, - "name": tc.name, + "name": sanitize_tool_name(&tc.name), "arguments": args_str, }) }) @@ -452,6 +477,20 @@ fn convert_message(msg: &ChatMessage, index: usize) -> Vec { } } +/// Sanitize a tool name to match the OpenAI Responses API pattern `^[a-zA-Z0-9_-]+$`. +/// Replaces any invalid character (e.g. dots in MCP tool names) with underscores. +fn sanitize_tool_name(name: &str) -> String { + name.chars() + .map(|c| { + if c.is_ascii_alphanumeric() || c == '_' || c == '-' { + c + } else { + '_' + } + }) + .collect() +} + /// Convert a `ToolDefinition` to Responses API tool format. /// /// Applies strict-mode schema normalization (same as OpenAI Chat Completions): @@ -461,7 +500,7 @@ fn convert_tool_definition(tool: &ToolDefinition) -> serde_json::Value { serde_json::json!({ "type": "function", - "name": tool.name, + "name": sanitize_tool_name(&tool.name), "description": tool.description, "parameters": normalize_schema_strict(&tool.parameters), }) @@ -1093,4 +1132,95 @@ data: {"type":"response.completed","response":{"status":"completed","usage":{"in assert_eq!(parsed.tool_calls[1].name, "read_file"); assert_eq!(parsed.finish_reason, FinishReason::ToolUse); } + + /// Regression test: tool names with dots (e.g. MCP tools) must be sanitized + /// to match OpenAI's `^[a-zA-Z0-9_-]+$` pattern. + #[test] + fn test_sanitize_tool_name_replaces_dots() { + assert_eq!(super::sanitize_tool_name("memory_search"), "memory_search"); + assert_eq!( + super::sanitize_tool_name("mcp.server.tool"), + "mcp_server_tool" + ); + assert_eq!(super::sanitize_tool_name("tool@v2"), "tool_v2"); + assert_eq!(super::sanitize_tool_name("my-tool"), "my-tool"); + } + + /// Regression test: convert_tool_definition sanitizes the name. + #[test] + fn test_convert_tool_definition_sanitizes_name() { + let tool = ToolDefinition { + name: "mcp.server.search".to_string(), + description: "Search".to_string(), + parameters: serde_json::json!({"type": "object", "properties": {}}), + }; + let json = super::convert_tool_definition(&tool); + assert_eq!(json["name"], "mcp_server_search"); + } + + /// Regression test: function_call items sanitize tool names. + #[test] + fn test_convert_message_sanitizes_tool_call_name() { + let tool_calls = vec![ToolCall { + id: "call_1".to_string(), + name: "mcp.server.search".to_string(), + arguments: serde_json::json!({"q": "test"}), + reasoning: None, + }]; + let msg = ChatMessage::assistant_with_tool_calls(None, tool_calls); + let items = super::convert_message(&msg, 0); + assert_eq!(items[0]["name"], "mcp_server_search"); + } + + /// Regression: sanitized tool names in API responses must be reverse-mapped + /// back to original names so the tool registry can look them up. + #[test] + fn test_sanitized_name_reverse_mapping() { + use std::collections::HashMap; + + let tools = [ + ToolDefinition { + name: "mcp.server.search".to_string(), + description: "Search".to_string(), + parameters: serde_json::json!({"type": "object", "properties": {}}), + }, + ToolDefinition { + name: "memory_search".to_string(), + description: "Memory".to_string(), + parameters: serde_json::json!({"type": "object", "properties": {}}), + }, + ]; + + // Build name map (same logic as complete_with_tools) + let name_map: HashMap = tools + .iter() + .filter_map(|t| { + let sanitized = super::sanitize_tool_name(&t.name); + if sanitized != t.name { + Some((sanitized, t.name.clone())) + } else { + None + } + }) + .collect(); + + // Only the MCP tool should appear (its name changed) + assert_eq!(name_map.len(), 1); + assert_eq!( + name_map.get("mcp_server_search"), + Some(&"mcp.server.search".to_string()) + ); + + // Simulate a tool call coming back with the sanitized name + let mut tc = ToolCall { + id: "call_1".to_string(), + name: "mcp_server_search".to_string(), + arguments: serde_json::json!({}), + reasoning: None, + }; + if let Some(original) = name_map.get(&tc.name) { + tc.name = original.clone(); + } + assert_eq!(tc.name, "mcp.server.search"); + } } diff --git a/src/skills/catalog.rs b/src/skills/catalog.rs index 93584f5f47f..30759546977 100644 --- a/src/skills/catalog.rs +++ b/src/skills/catalog.rs @@ -182,8 +182,14 @@ impl SkillCatalog { /// Create a catalog with a custom registry URL (for testing). #[cfg(test)] pub fn with_url(url: &str) -> Self { + Self::with_url_and_timeout(url, REQUEST_TIMEOUT) + } + + /// Create a catalog with a custom registry URL and timeout (for testing). + #[cfg(test)] + pub fn with_url_and_timeout(url: &str, timeout: Duration) -> Self { let client = reqwest::Client::builder() - .timeout(REQUEST_TIMEOUT) + .timeout(timeout) .user_agent(concat!("ironclaw/", env!("CARGO_PKG_VERSION"))) .build() .unwrap_or_default(); @@ -458,7 +464,9 @@ mod tests { #[tokio::test] async fn test_search_returns_error_on_network_failure() { // Use RFC 5737 TEST-NET-1 (192.0.2.0/24) for reliable failure even behind proxies. - let catalog = SkillCatalog::with_url("http://192.0.2.1:9999"); + // Short timeout so the test doesn't block for the full 10s REQUEST_TIMEOUT. + let catalog = + SkillCatalog::with_url_and_timeout("http://192.0.2.1:9999", Duration::from_secs(1)); let outcome = catalog.search("test").await; assert!(outcome.results.is_empty()); assert!(outcome.error.is_some()); diff --git a/src/tools/builtin/message.rs b/src/tools/builtin/message.rs index 08029d6fbe1..d42de4cebae 100644 --- a/src/tools/builtin/message.rs +++ b/src/tools/builtin/message.rs @@ -224,7 +224,13 @@ impl Tool for MessageTool { ) -> Result { let start = std::time::Instant::now(); - let content = require_str(¶ms, "content")?; + // Accept "message" as an alias for "content" — LLMs frequently use + // the wrong parameter name in autonomous job execution. + let content = require_str(¶ms, "content").or_else(|_| { + require_str(¶ms, "message").map_err(|_| { + ToolError::InvalidParameters("missing 'content' parameter".to_string()) + }) + })?; let explicit_channel = params .get("channel") @@ -323,8 +329,11 @@ impl Tool for MessageTool { if !attachments.is_empty() { response = response.with_attachments(attachments); } - if channel.as_deref() == Some("gateway") - && response.thread_id.is_none() + // Attach thread_id so the gateway can route the message into the + // correct conversation. Previously this only fired when channel was + // explicitly "gateway", which meant broadcast_all (channel=null) sent + // a response without a thread_id and the gateway silently dropped it. + if response.thread_id.is_none() && let Some(thread_id) = metadata_string(&ctx.metadata, "notify_thread_id") { response = response.in_thread(thread_id); @@ -480,6 +489,31 @@ mod tests { assert!(params.get("attachments").is_some()); } + /// Regression: LLMs frequently pass {"message": "..."} instead of + /// {"content": "..."}. The tool should accept both. + #[tokio::test] + async fn message_param_alias_accepted() { + let tool = MessageTool::new(Arc::new(ChannelManager::new())); + tool.set_context(Some("gateway".to_string()), Some("user".to_string())) + .await; + + let ctx = crate::context::JobContext::new("test", "test"); + + // "message" alias should not produce InvalidParameters + let result = tool + .execute(serde_json::json!({"message": "hello from alias"}), &ctx) + .await; + // Execution may fail for other reasons (no real channel), but + // the error must NOT be about a missing 'content' parameter. + if let Err(ref e) = result { + let msg = e.to_string(); + assert!( + !msg.contains("missing 'content'"), + "Should accept 'message' as alias for 'content', got: {msg}" + ); + } + } + #[tokio::test] async fn message_tool_set_context_updates_defaults() { let tool = MessageTool::new(Arc::new(ChannelManager::new())); diff --git a/src/tools/builtin/routine.rs b/src/tools/builtin/routine.rs index 76f6e38be1f..20257f16eee 100644 --- a/src/tools/builtin/routine.rs +++ b/src/tools/builtin/routine.rs @@ -65,6 +65,7 @@ struct NormalizedExecutionRequest { context_paths: Vec, use_tools: bool, max_tool_rounds: u32, + max_iterations: u32, } #[derive(Debug, Clone, PartialEq, Eq)] @@ -328,6 +329,13 @@ fn full_job_execution_variant() -> Value { "type": "string", "enum": ["full_job"], "description": "Full-job execution mode." + }, + "max_iterations": { + "type": "integer", + "description": "Maximum LLM iterations for the job (default: 25). Increase for complex multi-step tasks.", + "default": 25, + "minimum": 1, + "maximum": 200 } }, "required": ["mode"] @@ -644,6 +652,12 @@ pub(crate) fn routine_update_parameters_schema() -> Value { "description": { "type": "string", "description": "New description" + }, + "max_iterations": { + "type": "integer", + "description": "Maximum LLM iterations for full_job routines (1-200).", + "minimum": 1, + "maximum": 200 } }, "required": ["name"] @@ -887,11 +901,16 @@ fn parse_routine_execution( .clamp(1, crate::agent::routine::MAX_TOOL_ROUNDS_LIMIT as u64) as u32; + let max_iterations = u64_field(params, "execution", "max_iterations", &["max_iterations"]) + .unwrap_or(25) + .clamp(1, 200) as u32; + Ok(NormalizedExecutionRequest { mode, context_paths, use_tools, max_tool_rounds, + max_iterations, }) } @@ -972,7 +991,7 @@ fn build_routine_action( NormalizedExecutionMode::FullJob => RoutineAction::FullJob { title: name.to_string(), description: prompt.to_string(), - max_iterations: 10, + max_iterations: execution.max_iterations, }, } } @@ -1317,6 +1336,12 @@ impl Tool for RoutineUpdateTool { } } + if let Some(iters) = params.get("max_iterations").and_then(|v| v.as_u64()) + && let RoutineAction::FullJob { max_iterations, .. } = &mut routine.action + { + *max_iterations = (iters.clamp(1, 200)) as u32; + } + // Validate timezone param if provided let new_timezone = params .get("timezone") @@ -1544,6 +1569,7 @@ impl Tool for RoutineFireTool { "name": name, "run_id": run_id.to_string(), "status": "fired", + "note": "Routine is executing asynchronously. Use routine_history to check the result.", }); Ok(ToolOutput::success(result, start.elapsed())) @@ -1642,10 +1668,47 @@ impl Tool for RoutineHistoryTool { }) .collect(); + // Look up the routine's conversation thread and fetch recent messages + // so the user can see the full output of routine runs. + let (conversation_id, recent_output) = match self + .store + .get_or_create_routine_conversation(routine.id, name, &ctx.user_id) + .await + { + Ok(conv_id) => { + let messages = self + .store + .list_conversation_messages_paginated(conv_id, None, limit) + .await + .map(|(msgs, _)| msgs) + .unwrap_or_default(); + let msg_list: Vec = messages + .iter() + .map(|m| { + serde_json::json!({ + "role": m.role, + "content": m.content, + "timestamp": m.created_at.to_rfc3339(), + }) + }) + .collect(); + (Some(conv_id.to_string()), msg_list) + } + Err(e) => { + tracing::warn!( + routine = %name, + "Failed to fetch routine conversation thread: {e}" + ); + (None, Vec::new()) + } + }; + let result = serde_json::json!({ "routine": name, "total_runs": routine.run_count, + "conversation_id": conversation_id, "runs": run_list, + "recent_output": recent_output, }); Ok(ToolOutput::success(result, start.elapsed())) @@ -2282,8 +2345,8 @@ mod tests { .and_then(Value::as_object) .expect("full_job properties"); assert!( - full_job_props.len() == 1 && full_job_props.contains_key("mode"), - "full_job variant should only expose the execution mode", + full_job_props.contains_key("mode") && full_job_props.contains_key("max_iterations"), + "full_job variant should expose mode and max_iterations", ); } @@ -2491,6 +2554,7 @@ mod tests { context_paths: Vec::new(), use_tools: false, max_tool_rounds: 3, + max_iterations: 25, }; let action = build_routine_action("issue-1316", "Run it", &execution); @@ -2503,7 +2567,7 @@ mod tests { max_iterations, } if title == "issue-1316" && description == "Run it" - && max_iterations == 10 + && max_iterations == 25 )); } } diff --git a/src/tools/builtin/time.rs b/src/tools/builtin/time.rs index 5f0379640fb..f5c944a6bef 100644 --- a/src/tools/builtin/time.rs +++ b/src/tools/builtin/time.rs @@ -5,7 +5,7 @@ use chrono::{DateTime, LocalResult, NaiveDate, NaiveDateTime, TimeZone, Utc}; use chrono_tz::Tz; use crate::context::JobContext; -use crate::tools::tool::{Tool, ToolError, ToolOutput, require_str}; +use crate::tools::tool::{Tool, ToolError, ToolOutput}; /// Tool for getting current time and date operations. pub struct TimeTool; @@ -62,7 +62,7 @@ impl Tool for TimeTool { "description": "Second timestamp for diff." } }, - "required": ["operation"] + "required": [] }) } @@ -73,7 +73,10 @@ impl Tool for TimeTool { ) -> Result { let start = std::time::Instant::now(); - let operation = require_str(¶ms, "operation")?; + let operation = params + .get("operation") + .and_then(|v| v.as_str()) + .unwrap_or("now"); let result = match operation { "now" => execute_now(¶ms, ctx)?, diff --git a/src/tools/mcp/auth.rs b/src/tools/mcp/auth.rs index 1926e78db85..9e8792ad1cc 100644 --- a/src/tools/mcp/auth.rs +++ b/src/tools/mcp/auth.rs @@ -954,16 +954,22 @@ pub async fn store_tokens( server_config: &McpServerConfig, token: &AccessToken, ) -> Result<(), AuthError> { - // Store access token - let params = CreateSecretParams::new(server_config.token_secret_name(), &token.access_token) - .with_provider(format!("mcp:{}", server_config.name)); + // Store access token (with expiry if provided) + let mut params = + CreateSecretParams::new(server_config.token_secret_name(), &token.access_token) + .with_provider(format!("mcp:{}", server_config.name)); + + if let Some(secs) = token.expires_in { + let expires_at = chrono::Utc::now() + chrono::Duration::seconds(secs as i64); + params = params.with_expiry(expires_at); + } secrets .create(user_id, params) .await .map_err(|e| AuthError::Secrets(e.to_string()))?; - // Store refresh token if present + // Store refresh token if present (no expiry — long-lived) if let Some(ref refresh_token) = token.refresh_token { let params = CreateSecretParams::new(server_config.refresh_token_secret_name(), refresh_token) @@ -1064,11 +1070,26 @@ pub async fn refresh_access_token( // Get client_id (from config or stored DCR) let client_id = get_client_id(server_config, secrets, user_id).await?; - // Get the refresh token - let refresh_token = secrets + // Get the refresh token (try current name, fall back to legacy name for + // users who authenticated before the naming convention was fixed). + // Only fall back on NotFound/Expired — propagate real errors (DB, decryption). + let refresh_token = match secrets .get_decrypted(user_id, &server_config.refresh_token_secret_name()) .await - .map_err(|e| AuthError::RefreshFailed(format!("No refresh token: {}", e)))?; + { + Ok(token) => token, + Err(crate::secrets::SecretError::NotFound(_) | crate::secrets::SecretError::Expired) => { + secrets + .get_decrypted(user_id, &server_config.legacy_refresh_token_secret_name()) + .await + .map_err(|e| AuthError::RefreshFailed(format!("No refresh token: {}", e)))? + } + Err(e) => { + return Err(AuthError::RefreshFailed(format!( + "Failed to read refresh token: {e}" + ))); + } + }; // Discover the token endpoint let token_url = if let Some(ref oauth) = server_config.oauth { diff --git a/src/tools/mcp/client.rs b/src/tools/mcp/client.rs index 32c5767d1f7..4125e1f4490 100644 --- a/src/tools/mcp/client.rs +++ b/src/tools/mcp/client.rs @@ -259,6 +259,9 @@ impl McpClient { } /// Get the access token for this server (if authenticated). + /// + /// If the stored token has expired, automatically attempts a refresh using + /// the stored refresh token before failing. async fn get_access_token(&self) -> Result, ToolError> { let Some(ref secrets) = self.secrets else { return Ok(None); @@ -272,6 +275,33 @@ impl McpClient { { Ok(token) => Ok(Some(token.expose().to_string())), Err(crate::secrets::SecretError::NotFound(_)) => Ok(None), + Err(crate::secrets::SecretError::Expired) => { + // Token expired — attempt refresh before failing. + tracing::info!( + server = %self.server_name, + "Access token expired, attempting refresh" + ); + match refresh_access_token(config, secrets, &self.user_id).await { + Ok(new_token) => { + tracing::info!( + server = %self.server_name, + "Access token refreshed successfully" + ); + Ok(Some(new_token.access_token)) + } + Err(e) => { + tracing::warn!( + server = %self.server_name, + "Token refresh failed: {}", e + ); + Err(ToolError::ExternalService(format!( + "Failed to get access token: Secret has expired \ + and refresh failed: {}", + e + ))) + } + } + } Err(e) => Err(ToolError::ExternalService(format!( "Failed to get access token: {}", e diff --git a/src/tools/mcp/config.rs b/src/tools/mcp/config.rs index 06adbd3dc56..c7eb62f649c 100644 --- a/src/tools/mcp/config.rs +++ b/src/tools/mcp/config.rs @@ -250,7 +250,19 @@ impl McpServerConfig { } /// Get the secret name used to store the refresh token. + /// + /// Matches the convention used by the hosted OAuth flow in + /// `store_oauth_tokens`: `{token_secret_name}_refresh_token`. pub fn refresh_token_secret_name(&self) -> String { + format!("{}_refresh_token", self.token_secret_name()) + } + + /// Legacy secret name for refresh tokens (pre-v0.22). + /// + /// Earlier versions stored refresh tokens as `mcp_{name}_refresh_token` + /// instead of `{token_secret_name}_refresh_token`. Used as a fallback + /// during lookup to avoid forcing re-auth on existing users. + pub fn legacy_refresh_token_secret_name(&self) -> String { format!("mcp_{}_refresh_token", self.name) } @@ -750,8 +762,15 @@ mod tests { fn test_token_secret_names() { let config = McpServerConfig::new("notion", "https://mcp.notion.com"); assert_eq!(config.token_secret_name(), "mcp_notion_access_token"); + // Refresh token name follows the hosted OAuth convention: + // {token_secret_name}_refresh_token assert_eq!( config.refresh_token_secret_name(), + "mcp_notion_access_token_refresh_token" + ); + // Legacy name used before v0.22 — fallback lookup prevents forced re-auth + assert_eq!( + config.legacy_refresh_token_secret_name(), "mcp_notion_refresh_token" ); } diff --git a/src/util.rs b/src/util.rs index a76f3b27b5c..568e943d7a6 100644 --- a/src/util.rs +++ b/src/util.rs @@ -225,4 +225,26 @@ mod tests { "The tool returned: TASK_COMPLETE signal" )); } + + #[test] + fn signals_completion_after_suggestions_stripped() { + // Regression: after stripping tags, the completion + // signal should still be detected in the cleaned text. + assert!(llm_signals_completion( + "The job is complete. All requested work has been finished." + )); + } + + #[test] + fn signals_completion_self_dialogue_pattern() { + // Regression: the "not complete" pattern that caused the self-dialogue + // loop when left in job context after plan completion. + assert!(!llm_signals_completion( + "No — the job is **not complete**.\n\n\ + What still needs to be done:\n\ + 1. Fetch actual meeting note contents\n\ + 2. Create the Notion page\n\ + 3. Send the completion message" + )); + } } diff --git a/src/worker/job.rs b/src/worker/job.rs index 94a04290718..edf87bf8265 100644 --- a/src/worker/job.rs +++ b/src/worker/job.rs @@ -828,14 +828,11 @@ Report when the job is complete or if you encounter issues you cannot resolve."# }), ); - if matches!( - &e, - Error::Tool(crate::error::ToolError::AutonomousUnavailable { .. }) - ) { - Err(e) - } else { - Ok(()) - } + // All tool errors (including AutonomousUnavailable) are + // recoverable — the error message is already recorded in + // reason_ctx so the LLM can see it and try a different + // approach. Returning Err here would kill the entire job. + Ok(()) } } } @@ -930,17 +927,31 @@ Report when the job is complete or if you encounter issues you cannot resolve."# tokio::time::sleep(Duration::from_millis(100)).await; } - // Plan completed, check with LLM if job is done + // Plan completed — ask the LLM whether the job is done. + let msg_count_before = reason_ctx.messages.len(); reason_ctx.messages.push(ChatMessage::user( - "All planned actions have been executed. Is the job complete? If not, what else needs to be done?", + "All planned actions have been executed. Assess the results: \ + if the job is fully complete, state that the job is complete. \ + Otherwise, briefly list what remains.", )); let response = reasoning.respond(reason_ctx).await?; - reason_ctx.messages.push(ChatMessage::assistant(&response)); + let response = crate::agent::strip_suggestions(&response); if crate::util::llm_signals_completion(&response) { + reason_ctx.messages.push(ChatMessage::assistant(&response)); self.mark_completed().await?; } else { + // Replace the completion-check exchange with an action-oriented + // continuation prompt. Leaving the "Is the job complete?" / "No" + // dialogue in context causes the agentic loop to repeat the same + // analysis instead of calling tools (self-dialogue loop). + reason_ctx.messages.truncate(msg_count_before); + reason_ctx.messages.push(ChatMessage::user(format!( + "The planned actions are done but the job is not yet complete. \ + Remaining work:\n\n{response}\n\n\ + Continue executing now — use tools to finish the job." + ))); tracing::info!( "Job {} plan completed but work remains, falling back to direct selection", self.job_id @@ -1420,16 +1431,20 @@ impl<'a> LoopDelegate for JobDelegate<'a> { return TextAction::Continue; } - // Check for explicit completion - if crate::util::llm_signals_completion(text) { - if let Err(e) = self.worker.mark_completed().await { - tracing::warn!( - "Failed to mark job {} as completed: {}", - self.worker.job_id, - e - ); - } - return TextAction::Return(LoopOutcome::Response(text.to_string())); + // Jobs run autonomously — strip tags that are only + // meaningful for interactive chat sessions. + let text = crate::agent::strip_suggestions(text); + + // A non-empty text response with no tool intent (already filtered + // by the agentic loop's nudge mechanism) is the LLM's final answer. + // Mark the job complete and stop the loop. Without this, the LLM + // restates its summary every iteration until the cap is hit. + if let Err(e) = self.worker.mark_completed().await { + tracing::warn!( + "Failed to mark job {} as completed: {}", + self.worker.job_id, + e + ); } // Track that a substantive response has been produced. @@ -1437,7 +1452,7 @@ impl<'a> LoopDelegate for JobDelegate<'a> { .store(true, std::sync::atomic::Ordering::Relaxed); // Add assistant response to context - reason_ctx.messages.push(ChatMessage::assistant(text)); + reason_ctx.messages.push(ChatMessage::assistant(&text)); self.worker.log_event( "message", @@ -1447,7 +1462,7 @@ impl<'a> LoopDelegate for JobDelegate<'a> { }), ); - TextAction::Continue + TextAction::Return(LoopOutcome::Response(text)) } async fn execute_tool_calls( @@ -1456,6 +1471,9 @@ impl<'a> LoopDelegate for JobDelegate<'a> { content: Option, reason_ctx: &mut ReasoningContext, ) -> Result, crate::error::Error> { + // Strip suggestions from accompanying text (not useful in job context). + let content = content.map(|c| crate::agent::strip_suggestions(&c)); + if let Some(ref text) = content { self.worker.log_event( "message", @@ -2156,6 +2174,53 @@ mod tests { ); } + /// Regression: a text response without rigid completion phrases (e.g. + /// "Weekly review completed and saved to Notion") must still terminate the + /// agentic loop and mark the job complete, rather than continuing until + /// max_iterations. + #[tokio::test] + async fn test_text_response_terminates_loop_without_explicit_completion_phrase() { + let worker = make_worker(vec![]).await; + worker + .context_manager() + .update_context(worker.job_id, |ctx| { + ctx.transition_to(JobState::InProgress, None) + }) + .await + .unwrap() // safety: test + .unwrap(); // safety: test + + let (_, mut rx) = tokio::sync::mpsc::channel(1); + let delegate = JobDelegate { + worker: &worker, + rx: tokio::sync::Mutex::new(&mut rx), + consecutive_rate_limits: std::sync::atomic::AtomicUsize::new(0), + has_text_response: std::sync::atomic::AtomicBool::new(false), + }; + + let mut reason_ctx = ReasoningContext::new(); + + // Text that a real LLM would produce but doesn't match llm_signals_completion + let action = delegate + .handle_text_response( + "Weekly review created in Notion and notification sent.", + &mut reason_ctx, + ) + .await; + + assert!( + matches!(action, TextAction::Return(_)), + "Text response should terminate the loop, got Continue" + ); // safety: test + + let ctx = worker + .context_manager() + .get_context(worker.job_id) + .await + .unwrap(); // safety: test + assert_eq!(ctx.state, JobState::Completed); // safety: test + } + /// Regression test: selections_to_tool_calls must preserve tool_call_id /// so that tool_result messages match the assistant_with_tool_calls message /// and are not treated as orphaned by sanitize_tool_messages. @@ -2429,4 +2494,39 @@ mod tests { } )); } + + /// Regression test: AutonomousUnavailable errors must be recoverable. + /// Previously the job worker treated them as fatal, killing the entire + /// job instead of feeding the error back to the LLM. + #[tokio::test] + async fn test_autonomous_unavailable_is_recoverable() { + let worker = make_worker(vec![]).await; + let mut reason_ctx = ReasoningContext::new(); + let selection = ToolSelection { + tool_name: "secret_list".to_string(), + parameters: serde_json::json!({}), + reasoning: "list secrets".to_string(), + alternatives: vec![], + tool_call_id: "call_123".to_string(), + }; + let err = Error::Tool(crate::error::ToolError::AutonomousUnavailable { + name: "secret_list".to_string(), + reason: "not available in autonomous jobs".to_string(), + }); + + let result = worker + .process_tool_result_job(&mut reason_ctx, &selection, Err(err)) + .await; + + assert!( + result.is_ok(), + "AutonomousUnavailable must be recoverable, not fatal: {:?}", + result + ); + // The error should be fed back to the LLM as a message. + assert!( + !reason_ctx.messages.is_empty(), + "Error message should be added to reason_ctx for the LLM" + ); + } } diff --git a/tests/e2e/ironclaw_e2e.egg-info/SOURCES.txt b/tests/e2e/ironclaw_e2e.egg-info/SOURCES.txt index c2784f643b8..aa4ea46c3d3 100644 --- a/tests/e2e/ironclaw_e2e.egg-info/SOURCES.txt +++ b/tests/e2e/ironclaw_e2e.egg-info/SOURCES.txt @@ -14,9 +14,12 @@ scenarios/test_extensions.py scenarios/test_html_injection.py scenarios/test_mcp_auth_flow.py scenarios/test_oauth_credential_fallback.py +scenarios/test_oauth_refresh.py +scenarios/test_oauth_url_parameters.py scenarios/test_owner_scope.py scenarios/test_pairing.py scenarios/test_routine_event_batch.py +scenarios/test_routine_full_job.py scenarios/test_routine_oauth_credential_injection.py scenarios/test_skills.py scenarios/test_sse_reconnect.py diff --git a/tests/e2e/mock_llm.py b/tests/e2e/mock_llm.py index 1147662cced..f049e5049e3 100644 --- a/tests/e2e/mock_llm.py +++ b/tests/e2e/mock_llm.py @@ -121,6 +121,75 @@ def _last_user_content(messages: list[dict]) -> str: return "" +def _is_job_mode(messages: list[dict]) -> bool: + """Detect if this conversation is a background job (not chat).""" + for msg in messages: + if msg.get("role") == "system": + content = msg.get("content", "") + if "autonomous agent working on a job" in content: + return True + return False + + +def _count_tool_results(messages: list[dict]) -> int: + """Count how many tool result messages are in the conversation.""" + return sum(1 for m in messages if m.get("role") == "tool") + + +def match_job_response(messages: list[dict], has_tools: bool) -> dict | None: + """Handle background job conversations. + + Returns a dict with either {"text": ...} or {"tool_call": ...}, + or None if this isn't a job conversation. + """ + if not _is_job_mode(messages): + return None + + last_user = _last_user_content(messages) + tool_result_count = _count_tool_results(messages) + + # Planning call (no tools available = complete() not complete_with_tools()) + if "create a plan" in last_user.lower(): + return {"text": json.dumps({ + "goal": "Complete the requested routine job", + "actions": [ + { + "tool_name": "echo", + "parameters": {"message": "job-step-1"}, + "reasoning": "First step: echo a test message", + "expected_outcome": "Echo returns the message", + }, + { + "tool_name": "time", + "parameters": {"operation": "now"}, + "reasoning": "Second step: get the current time", + "expected_outcome": "Returns current timestamp", + }, + ], + "estimated_cost": 0.001, + "estimated_time_secs": 5, + "confidence": 0.95, + })} + + # Post-plan completion check: after tool results, say complete + if "planned actions" in last_user.lower() and tool_result_count >= 2: + return {"text": "The job is complete. All tasks are done."} + + # Continuation prompt (from our fix): the plan didn't fully complete, + # now the agentic loop should call tools + if "continue executing now" in last_user.lower() and has_tools: + return {"tool_call": { + "tool_name": "echo", + "arguments": {"message": "continuation-step"}, + }} + + # After a tool result in the agentic loop, signal completion + if tool_result_count > 0 and has_tools: + return {"text": "The job is complete. All requested work has been finished."} + + return None + + def match_response(messages: list[dict]) -> str: content = _last_user_content(messages) for pattern, response in CANNED_RESPONSES: @@ -193,6 +262,19 @@ async def chat_completions(request: web.Request) -> web.StreamResponse: has_tools = bool(body.get("tools")) cid = f"mock-{uuid.uuid4().hex[:8]}" + # Job-mode conversations (background routine/job execution) + job_resp = match_job_response(messages, has_tools) + if job_resp: + if "tool_call" in job_resp: + tc = job_resp["tool_call"] + if not stream: + return _tool_call_response(cid, tc) + return await _stream_tool_call(request, cid, tc) + text = job_resp["text"] + if not stream: + return _text_response(cid, text) + return await _stream_text(request, cid, text) + # Tool result in messages -> text summary tr = _find_tool_result(messages) if tr: diff --git a/tests/e2e/scenarios/test_routine_full_job.py b/tests/e2e/scenarios/test_routine_full_job.py new file mode 100644 index 00000000000..64472dc99dc --- /dev/null +++ b/tests/e2e/scenarios/test_routine_full_job.py @@ -0,0 +1,133 @@ +"""E2E tests for full_job routine execution. + +Exercises the complete lifecycle: create a full_job routine via the +web UI, trigger it via the API, and verify the job runs tools and +completes without hitting the iteration cap. + +Requires Playwright (browser-based tests). +""" + +import asyncio +import uuid + +from helpers import SEL, api_get, api_post + + +# -- Helpers ------------------------------------------------------------------ + +async def _send_chat_message(page, message: str) -> None: + """Send a chat message and wait for the assistant turn to appear.""" + chat_input = page.locator(SEL["chat_input"]) + await chat_input.wait_for(state="visible", timeout=5000) + assistant_messages = page.locator(SEL["message_assistant"]) + before_count = await assistant_messages.count() + + await chat_input.fill(message) + await chat_input.press("Enter") + + await page.wait_for_function( + """({ selector, expectedCount }) => { + return document.querySelectorAll(selector).length >= expectedCount; + }""", + arg={ + "selector": SEL["message_assistant"], + "expectedCount": before_count + 1, + }, + timeout=30000, + ) + + +async def _wait_for_routine(base_url: str, name: str, timeout: float = 20.0) -> dict: + """Poll until the named routine exists.""" + for _ in range(int(timeout * 2)): + resp = await api_get(base_url, "/api/routines") + resp.raise_for_status() + for routine in resp.json()["routines"]: + if routine["name"] == name: + return routine + await asyncio.sleep(0.5) + raise AssertionError(f"Routine '{name}' not created within {timeout}s") + + +async def _get_routine_runs(base_url: str, routine_id: str) -> list[dict]: + """Fetch routine runs.""" + resp = await api_get(base_url, f"/api/routines/{routine_id}/runs") + resp.raise_for_status() + return resp.json()["runs"] + + +async def _wait_for_completed_run( + base_url: str, + routine_id: str, + *, + timeout: float = 60.0, +) -> dict: + """Poll until the newest run reaches a terminal state.""" + for _ in range(int(timeout * 2)): + runs = await _get_routine_runs(base_url, routine_id) + if runs and runs[0]["status"].lower() not in ("running", "pending"): + return runs[0] + await asyncio.sleep(0.5) + raise AssertionError( + f"Routine '{routine_id}' did not complete within {timeout}s" + ) + + +async def _wait_for_job_terminal( + base_url: str, + job_id: str, + *, + timeout: float = 60.0, +) -> dict: + """Poll until a job reaches a terminal state.""" + terminal = {"completed", "failed", "cancelled", "submitted", "accepted"} + for _ in range(int(timeout * 2)): + resp = await api_get(base_url, f"/api/jobs/{job_id}") + resp.raise_for_status() + detail = resp.json() + if detail.get("state", "").lower() in terminal: + return detail + await asyncio.sleep(0.5) + raise AssertionError(f"Job '{job_id}' did not reach terminal state within {timeout}s") + + +# -- Tests -------------------------------------------------------------------- + +async def test_full_job_routine_completes_with_tools(page, ironclaw_server): + """A full_job routine should plan, execute tools, and complete.""" + name = f"fjob-{uuid.uuid4().hex[:8]}" + + # Step 1: Create full_job routine via chat + await _send_chat_message(page, f"create full-job owner routine {name}") + routine = await _wait_for_routine(ironclaw_server, name) + + assert routine["id"] + assert routine["action_type"] == "full_job" + + # Step 2: Trigger the routine + resp = await api_post(ironclaw_server, f"/api/routines/{routine['id']}/trigger") + resp.raise_for_status() + trigger_data = resp.json() + assert trigger_data["status"] == "triggered" + + # Step 3: Wait for the run to complete + completed_run = await _wait_for_completed_run( + ironclaw_server, routine["id"], timeout=60 + ) + + # The run should have succeeded (not failed) + assert completed_run["status"].lower() != "failed", ( + f"Full job routine run failed: {completed_run}" + ) + + # Step 4: Verify the job reached a success state. + # Jobs may advance past "completed" to "submitted" or "accepted", + # so treat all post-completion states as success. + success_states = {"completed", "submitted", "accepted"} + if completed_run.get("job_id"): + job = await _wait_for_job_terminal( + ironclaw_server, completed_run["job_id"], timeout=30 + ) + assert job["state"].lower() in success_states, ( + f"Expected job state in {success_states}, got '{job['state']}'" + ) diff --git a/tests/e2e_routine_heartbeat.rs b/tests/e2e_routine_heartbeat.rs index 6849ee055eb..36d87a07eb9 100644 --- a/tests/e2e_routine_heartbeat.rs +++ b/tests/e2e_routine_heartbeat.rs @@ -21,9 +21,7 @@ mod tests { NotifyConfig, Routine, RoutineAction, RoutineGuardrails, RoutineRun, RunStatus, Trigger, }; use ironclaw::agent::routine_engine::RoutineEngine; - use ironclaw::agent::{ - HeartbeatConfig, HeartbeatRunner, SandboxReadiness, Scheduler, SchedulerDeps, - }; + use ironclaw::agent::{HeartbeatConfig, HeartbeatRunner, Scheduler, SchedulerDeps}; use ironclaw::channels::IncomingMessage; use ironclaw::config::{AgentConfig, RoutineConfig, SafetyConfig}; use ironclaw::context::{ContextManager, JobContext}; @@ -352,7 +350,7 @@ mod tests { extension_manager, registry, safety, - SandboxReadiness::Available, + ironclaw::agent::routine_engine::SandboxReadiness::DisabledByConfig, )) } @@ -456,7 +454,7 @@ mod tests { None, tools, safety, - SandboxReadiness::DisabledByConfig, + ironclaw::agent::routine_engine::SandboxReadiness::DisabledByConfig, )); // Insert a cron routine with next_fire_at in the past. @@ -535,7 +533,7 @@ mod tests { None, tools, safety, - SandboxReadiness::DisabledByConfig, + ironclaw::agent::routine_engine::SandboxReadiness::DisabledByConfig, )); // Insert an event routine matching "deploy.*production". @@ -622,7 +620,7 @@ mod tests { None, tools, safety, - SandboxReadiness::DisabledByConfig, + ironclaw::agent::routine_engine::SandboxReadiness::DisabledByConfig, )); let routine = make_routine( @@ -731,7 +729,7 @@ mod tests { None, tools, safety, - SandboxReadiness::DisabledByConfig, + ironclaw::agent::routine_engine::SandboxReadiness::DisabledByConfig, )); let mut filters = std::collections::HashMap::new(); @@ -874,7 +872,7 @@ mod tests { None, tools, safety, - SandboxReadiness::DisabledByConfig, + ironclaw::agent::routine_engine::SandboxReadiness::DisabledByConfig, )); // Insert an event routine with 1-hour cooldown. @@ -1057,7 +1055,7 @@ mod tests { None, tools, safety, - SandboxReadiness::DisabledByConfig, + ironclaw::agent::routine_engine::SandboxReadiness::DisabledByConfig, )); (engine, db, dir) @@ -1179,7 +1177,7 @@ mod tests { None, tools, safety, - SandboxReadiness::DisabledByConfig, + ironclaw::agent::routine_engine::SandboxReadiness::DisabledByConfig, )); // Create a full_job routine with max_concurrent = 1 @@ -1287,7 +1285,7 @@ mod tests { None, tools, safety, - SandboxReadiness::DisabledByConfig, + ironclaw::agent::routine_engine::SandboxReadiness::DisabledByConfig, )); // Insert a due cron routine diff --git a/tests/e2e_telegram_message_routing.rs b/tests/e2e_telegram_message_routing.rs index 810fc218f00..e91aae8e948 100644 --- a/tests/e2e_telegram_message_routing.rs +++ b/tests/e2e_telegram_message_routing.rs @@ -198,7 +198,7 @@ mod tests { http_interceptor: None, transcription: None, document_extraction: None, - sandbox_readiness: ironclaw::agent::SandboxReadiness::DisabledByConfig, + sandbox_readiness: ironclaw::agent::routine_engine::SandboxReadiness::DisabledByConfig, builder: None, llm_backend: "nearai".to_string(), tenant_rates: std::sync::Arc::new(ironclaw::tenant::TenantRateRegistry::new(4, 3)), diff --git a/tests/support/gateway_workflow_harness.rs b/tests/support/gateway_workflow_harness.rs index d6341704ca7..efb5cf96815 100644 --- a/tests/support/gateway_workflow_harness.rs +++ b/tests/support/gateway_workflow_harness.rs @@ -264,7 +264,8 @@ impl GatewayWorkflowHarness { http_interceptor: None, transcription: None, document_extraction: None, - sandbox_readiness: ironclaw::agent::SandboxReadiness::DisabledByConfig, + sandbox_readiness: + ironclaw::agent::routine_engine::SandboxReadiness::DisabledByConfig, builder: None, llm_backend: "nearai".to_string(), tenant_rates: std::sync::Arc::new(ironclaw::tenant::TenantRateRegistry::new(4, 3)), diff --git a/tests/support/test_rig.rs b/tests/support/test_rig.rs index 5775b86da6d..0e2883901de 100644 --- a/tests/support/test_rig.rs +++ b/tests/support/test_rig.rs @@ -650,7 +650,7 @@ impl TestRigBuilder { None, components.tools.clone(), components.safety.clone(), - ironclaw::agent::SandboxReadiness::Available, // tests don't use real Docker + ironclaw::agent::routine_engine::SandboxReadiness::DisabledByConfig, )); components .tools @@ -759,7 +759,7 @@ impl TestRigBuilder { http_interceptor, transcription: None, document_extraction: None, - sandbox_readiness: ironclaw::agent::SandboxReadiness::Available, // tests don't use real Docker + sandbox_readiness: ironclaw::agent::routine_engine::SandboxReadiness::DisabledByConfig, builder: None, llm_backend: "nearai".to_string(), tenant_rates: std::sync::Arc::new(ironclaw::tenant::TenantRateRegistry::new(4, 3)), From a8e83210ff01e7317f7af96bb3dc5b705ab7a63b Mon Sep 17 00:00:00 2001 From: Feri Muhammad <83993897+hdward-dev@users.noreply.github.com> Date: Sun, 29 Mar 2026 13:20:40 +0800 Subject: [PATCH 07/23] feat(discord): add gateway channel flow in wasm (#944) * feat(discord): restore gateway channel flow in wasm * chore(discord): bump channel version to 0.2.1 * fix(discord): address review feedback on gateway channel PR - Add #[serde(default)] to DiscordMessageMetadata for backward compat with old Option serialized metadata - Restore mention polling alongside Gateway (on_poll processes gateway events first, then runs poll_for_mentions if configured) - Update on_respond to handle source_message_id with message_reference for mention-poll reply threading - Implement Gateway presence status: dnd before pairing, online after - Implement Gateway resume (OP 6) with session_id tracking, falling back to fresh identify on Invalid Session (OP 9) - Extract WebsocketSessionState and spawn_websocket_poll to reduce nesting in start_websocket_runtime - Simplify should_apply_dm_pairing tautology - Remove completed plan docs - Fix clippy items_after_test_module in extensions handler Co-Authored-By: Claude Opus 4.6 (1M context) * fix(discord): address review findings in gateway channel PR - Fix gateway presence always showing "online" by filtering empty owner_id strings from workspace store reads - Fix interaction followup using POST instead of PATCH to /messages/@original, which left deferred "thinking" state unresolved - Restore mention-poll pagination (up to 5 pages of 100 messages) - Remove dead ed25519-dalek and hex dependencies from WASM crate - Remove unused _channel_id parameter from remember_processed_id - Clean up redundant let binding in send_pairing_reply Co-Authored-By: Claude Opus 4.6 (1M context) * fix(discord): address second-round review findings - Log warning when gateway event queue JSON fails to deserialize instead of silently returning empty (zmanian review item 1) - Defer presence update from OP 10 Hello to after OP 0 READY, per Discord gateway protocol which requires READY before non-Identify commands (zmanian review item 2) - Add 0-25% random jitter to websocket reconnect backoff per Discord's reconnection recommendations (zmanian suggestion) - Extract WebsocketPollContext struct to replace 19-parameter spawn_websocket_poll function (zmanian suggestion) - Document intent bitmask 4609 = GUILDS + GUILD_MESSAGES + DIRECT_MESSAGES in capabilities JSON (zmanian suggestion) Co-Authored-By: Claude Opus 4.6 (1M context) --------- Co-authored-by: zhyaoyu Co-authored-by: ilblackdragon@gmail.com Co-authored-by: Claude Opus 4.6 (1M context) --- Cargo.lock | 6 + Cargo.toml | 2 +- FEATURE_PARITY.md | 2 +- channels-src/discord/Cargo.lock | 227 +- channels-src/discord/Cargo.toml | 4 +- channels-src/discord/README.md | 19 + .../discord/discord.capabilities.json | 19 +- channels-src/discord/src/lib.rs | 2543 ++++++++++------- src/channels/wasm/host.rs | 104 + src/channels/wasm/wrapper.rs | 1109 ++++++- src/channels/web/handlers/extensions.rs | 117 +- src/channels/web/server.rs | 23 +- src/pairing/store.rs | 6 + src/tools/wasm/capabilities.rs | 3 + src/tools/wasm/capabilities_schema.rs | 51 + 15 files changed, 2979 insertions(+), 1256 deletions(-) diff --git a/Cargo.lock b/Cargo.lock index 0c524704135..0e7d6521092 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -6789,7 +6789,11 @@ checksum = "7a9daff607c6d2bf6c16fd681ccb7eecc83e4e2cdc1ca067ffaadfca5de7f084" dependencies = [ "futures-util", "log", + "rustls 0.23.37", + "rustls-native-certs 0.8.3", + "rustls-pki-types", "tokio", + "tokio-rustls 0.26.4", "tungstenite 0.26.2", ] @@ -7150,6 +7154,8 @@ dependencies = [ "httparse", "log", "rand 0.9.2", + "rustls 0.23.37", + "rustls-pki-types", "sha1", "thiserror 2.0.18", "utf-8", diff --git a/Cargo.toml b/Cargo.toml index fbd3d6eec0d..b62f102696b 100644 --- a/Cargo.toml +++ b/Cargo.toml @@ -40,6 +40,7 @@ eula = false tokio = { version = "1", features = ["full"] } tokio-stream = { version = "0.1", features = ["sync"] } futures = "0.3" +tokio-tungstenite = { version = "0.26", features = ["rustls-tls-native-roots"] } eventsource-stream = "0.2" # HTTP client @@ -201,7 +202,6 @@ zbus = "4" [dev-dependencies] tokio-test = "0.4" tracing-test = "0.2" -tokio-tungstenite = "0.26" testcontainers-modules = { version = "0.11", features = ["postgres"] } pretty_assertions = "1" tempfile = "3" diff --git a/FEATURE_PARITY.md b/FEATURE_PARITY.md index ad2db551177..1946dce6ee5 100644 --- a/FEATURE_PARITY.md +++ b/FEATURE_PARITY.md @@ -70,7 +70,7 @@ This document tracks feature parity between IronClaw (Rust implementation) and O | WASM channels | ❌ | ✅ | - | IronClaw innovation; host resolves owner scope vs sender identity | | WhatsApp | ✅ | ❌ | P1 | Baileys (Web), same-phone mode with echo detection | | Telegram | ✅ | ✅ | - | WASM channel(MTProto), DM pairing, caption, /start, bot_username, DM topics, setup-time owner auto-verification, owner-scoped persistence | -| Discord | ✅ | ❌ | P2 | discord.js, thread parent binding inheritance | +| Discord | ✅ | 🚧 | P2 | Gateway `MESSAGE_CREATE` intake restored via websocket queue + WASM poll; Gateway DMs now respect pairing; thread parent binding inheritance and reply/thread parity still incomplete | | Signal | ✅ | ✅ | P2 | signal-cli daemonPC, SSE listener HTTP/JSON-R, user/group allowlists, DM pairing | | Slack | ✅ | ✅ | - | WASM tool | | iMessage | ✅ | ❌ | P3 | BlueBubbles or Linq recommended | diff --git a/channels-src/discord/Cargo.lock b/channels-src/discord/Cargo.lock index f25ce5511b5..f6e4a814278 100644 --- a/channels-src/discord/Cargo.lock +++ b/channels-src/discord/Cargo.lock @@ -20,162 +20,33 @@ version = "1.0.102" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "7f202df86484c868dbad7eaa557ef785d5c66295e41b460ef922eca0723b842c" -[[package]] -name = "base64ct" -version = "1.8.3" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "2af50177e190e07a26ab74f8b1efbfe2ef87da2116221318cb1c2e82baf7de06" - [[package]] name = "bitflags" version = "2.11.0" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "843867be96c8daad0d758b57df9392b6d8d271134fce549de6ce169ff98a92af" -[[package]] -name = "block-buffer" -version = "0.10.4" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "3078c7629b62d3f0439517fa394996acacc5cbc91c5a20d8c658e77abd503a71" -dependencies = [ - "generic-array", -] - [[package]] name = "cfg-if" version = "1.0.4" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "9330f8b2ff13f34540b44e946ef35111825727b38d33286ef986142615121801" -[[package]] -name = "const-oid" -version = "0.9.6" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "c2459377285ad874054d797f3ccebf984978aa39129f6eafde5cdc8315b612f8" - -[[package]] -name = "cpufeatures" -version = "0.2.17" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "59ed5838eebb26a2bb2e58f6d5b5316989ae9d08bab10e0e6d103e656d1b0280" -dependencies = [ - "libc", -] - -[[package]] -name = "crypto-common" -version = "0.1.7" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "78c8292055d1c1df0cce5d180393dc8cce0abec0a7102adb6c7b1eef6016d60a" -dependencies = [ - "generic-array", - "typenum", -] - -[[package]] -name = "curve25519-dalek" -version = "4.1.3" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "97fb8b7c4503de7d6ae7b42ab72a5a59857b4c937ec27a3d4539dba95b5ab2be" -dependencies = [ - "cfg-if", - "cpufeatures", - "curve25519-dalek-derive", - "digest", - "fiat-crypto", - "rustc_version", - "subtle", - "zeroize", -] - -[[package]] -name = "curve25519-dalek-derive" -version = "0.1.1" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "f46882e17999c6cc590af592290432be3bce0428cb0d5f8b6715e4dc7b383eb3" -dependencies = [ - "proc-macro2", - "quote", - "syn", -] - -[[package]] -name = "der" -version = "0.7.10" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "e7c1832837b905bbfb5101e07cc24c8deddf52f93225eee6ead5f4d63d53ddcb" -dependencies = [ - "const-oid", - "zeroize", -] - -[[package]] -name = "digest" -version = "0.10.7" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "9ed9a281f7bc9b7576e61468ba615a66a5c8cfdff42420a70aa82701a3b1e292" -dependencies = [ - "block-buffer", - "crypto-common", -] - [[package]] name = "discord-channel" -version = "0.2.0" +version = "0.2.1" dependencies = [ - "ed25519-dalek", - "hex", "serde", "serde_json", "wit-bindgen", ] -[[package]] -name = "ed25519" -version = "2.2.3" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "115531babc129696a58c64a4fef0a8bf9e9698629fb97e9e40767d235cfbcd53" -dependencies = [ - "pkcs8", - "signature", -] - -[[package]] -name = "ed25519-dalek" -version = "2.2.0" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "70e796c081cee67dc755e1a36a0a172b897fab85fc3f6bc48307991f64e4eca9" -dependencies = [ - "curve25519-dalek", - "ed25519", - "serde", - "sha2", - "subtle", - "zeroize", -] - [[package]] name = "equivalent" version = "1.0.2" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "877a4ace8713b0bcf2a4e7eec82529c029f1d0619886d18145fea96c3ffe5c0f" -[[package]] -name = "fiat-crypto" -version = "0.2.9" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "28dea519a9695b9977216879a3ebfddf92f1c08c05d984f8996aecd6ecdc811d" - -[[package]] -name = "generic-array" -version = "0.14.7" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "85649ca51fd72272d7821adaf274ad91c288277713d9c18820d8499a7ff69e9a" -dependencies = [ - "typenum", - "version_check", -] - [[package]] name = "hashbrown" version = "0.14.5" @@ -197,12 +68,6 @@ version = "0.5.0" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "2304e00983f87ffb38b55b444b5e3b60a884b5d30c0fca7d82fe33449bbe55ea" -[[package]] -name = "hex" -version = "0.4.3" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "7f24254aa9a54b5c858eaee2f5bccdb46aaf0e486a595ed5fd8f86ba55232a70" - [[package]] name = "id-arena" version = "2.3.0" @@ -223,9 +88,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 = "leb128" @@ -233,12 +98,6 @@ version = "0.2.5" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "884e2677b40cc8c339eaefcb701c32ef1fd2493d71118dc0ca4b6a736c93bd67" -[[package]] -name = "libc" -version = "0.2.182" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "6800badb6cb2082ffd7b6a67e6125bb39f18782f793520caee8cb8846be06112" - [[package]] name = "log" version = "0.4.29" @@ -253,19 +112,9 @@ checksum = "f8ca58f447f06ed17d5fc4043ce1b10dd205e060fb3ce5b979b8ed8e59ff3f79" [[package]] name = "once_cell" -version = "1.21.3" +version = "1.21.4" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "42f5e15c9953c5e4ccceeb2e7382a716482c34515315f7b03532b8b4e8393d2d" - -[[package]] -name = "pkcs8" -version = "0.10.2" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "f950b2377845cebe5cf8b5165cb3cc1a5e0fa5cfa3e1f7f55707d8fd82e0a7b7" -dependencies = [ - "der", - "spki", -] +checksum = "9f7c3e4beb33f85d45ae3e3a1792185706c8e16d043238c593331cc7cd313b50" [[package]] name = "prettyplease" @@ -288,22 +137,13 @@ dependencies = [ [[package]] name = "quote" -version = "1.0.44" +version = "1.0.45" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "21b2ebcf727b7760c461f091f9f0f539b77b8e87f2fd88131e7f1b433b3cece4" +checksum = "41f2619966050689382d2b44f664f4bc593e129785a36d6ee376ddf37259b924" dependencies = [ "proc-macro2", ] -[[package]] -name = "rustc_version" -version = "0.4.1" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "cfcb3a22ef46e85b45de6ee7e79d063319ebb6594faafcf1c225ea92ab6e9b92" -dependencies = [ - "semver", -] - [[package]] name = "semver" version = "1.0.27" @@ -353,23 +193,6 @@ dependencies = [ "zmij", ] -[[package]] -name = "sha2" -version = "0.10.9" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "a7507d819769d01a365ab707794a4084392c824f54a7a6a7862f8c3d0892b283" -dependencies = [ - "cfg-if", - "cpufeatures", - "digest", -] - -[[package]] -name = "signature" -version = "2.2.0" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "77549399552de45a898a580c1b41d445bf730df867cc44e6c0233bbc4b8329de" - [[package]] name = "smallvec" version = "1.15.1" @@ -385,22 +208,6 @@ dependencies = [ "smallvec", ] -[[package]] -name = "spki" -version = "0.7.3" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "d91ed6c858b01f942cd56b37a94b3e0a1798290327d1236e4d9cf4eaca44d29d" -dependencies = [ - "base64ct", - "der", -] - -[[package]] -name = "subtle" -version = "2.6.1" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "13c2bddecc57b384dee18652358fb23172facb8a2c51ccc10d74c157bdea3292" - [[package]] name = "syn" version = "2.0.117" @@ -412,12 +219,6 @@ dependencies = [ "unicode-ident", ] -[[package]] -name = "typenum" -version = "1.19.0" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "562d481066bde0658276a35467c4af00bdc6ee726305698a55b86e61d7ad82bb" - [[package]] name = "unicode-ident" version = "1.0.24" @@ -575,30 +376,24 @@ dependencies = [ [[package]] name = "zerocopy" -version = "0.8.39" +version = "0.8.47" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "db6d35d663eadb6c932438e763b262fe1a70987f9ae936e60158176d710cae4a" +checksum = "efbb2a062be311f2ba113ce66f697a4dc589f85e78a4aea276200804cea0ed87" dependencies = [ "zerocopy-derive", ] [[package]] name = "zerocopy-derive" -version = "0.8.39" +version = "0.8.47" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "4122cd3169e94605190e77839c9a40d40ed048d305bfdc146e7df40ab0f3e517" +checksum = "0e8bc7269b54418e7aeeef514aa68f8690b8c0489a06b0136e5f57c4c5ccab89" dependencies = [ "proc-macro2", "quote", "syn", ] -[[package]] -name = "zeroize" -version = "1.8.2" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "b97154e67e32c85465826e8bcc1c59429aaaf107c1e4a9e53c8d8ccd5eff88d0" - [[package]] name = "zmij" version = "1.0.21" diff --git a/channels-src/discord/Cargo.toml b/channels-src/discord/Cargo.toml index a2892494a84..6388178c0a1 100644 --- a/channels-src/discord/Cargo.toml +++ b/channels-src/discord/Cargo.toml @@ -1,6 +1,6 @@ [package] name = "discord-channel" -version = "0.2.0" +version = "0.2.1" edition = "2021" description = "Discord channel for IronClaw" license = "MIT OR Apache-2.0" @@ -10,8 +10,6 @@ publish = false serde = { version = "1.0", features = ["derive"] } serde_json = "1.0" wit-bindgen = "0.36" -ed25519-dalek = { version = "2", default-features = false, features = ["alloc", "fast", "zeroize"] } -hex = "0.4" [lib] crate-type = ["cdylib"] diff --git a/channels-src/discord/README.md b/channels-src/discord/README.md index 333e7670db0..927d2fae9b4 100644 --- a/channels-src/discord/README.md +++ b/channels-src/discord/README.md @@ -86,6 +86,24 @@ If an internal error occurs (e.g., metadata serialization failure), the tool att Check the host logs for detailed error information. ## Advanced Usage +### Gateway Mode + +The Discord channel now defaults to Discord Gateway transport for inbound message intake. +The bundled identify payload requests intents `4609`, which expands to: + +- `GUILDS` (`1`) +- `GUILD_MESSAGES` (`512`) +- `DIRECT_MESSAGES` (`4096`) + +Gateway DMs now follow the same pairing policy as webhook DMs. Unpaired users receive a pairing +instruction reply in the DM channel before the message is allowed through to the agent. If you +want stricter access control than pairing, set `owner_id`; that lock still applies to both +webhook and Gateway traffic. + +Gateway presence simply reflects a successful authenticated Gateway connection and advertises +`online`. Pairing still controls whether DMs are allowed through to the agent, but it no longer +changes the visible Discord status. + ### Mention Polling The Discord channel can also poll configured channels for `@bot` mentions. @@ -110,6 +128,7 @@ Example channel config: - `owner_id`: when set, only that Discord user can interact with the bot. - `dm_policy`: `open` allows all DMs; `pairing` requires approval. - `allow_from`: allowlist entries for DM pairing checks (`*`, user id, or username). +- Gateway DMs respect `dm_policy` and pairing just like webhook DMs. ### Embeds diff --git a/channels-src/discord/discord.capabilities.json b/channels-src/discord/discord.capabilities.json index 9ff7a8905d6..00ee25858eb 100644 --- a/channels-src/discord/discord.capabilities.json +++ b/channels-src/discord/discord.capabilities.json @@ -1,5 +1,5 @@ { - "version": "0.2.0", + "version": "0.2.1", "wit_version": "0.3.0", "type": "channel", "name": "discord", @@ -22,7 +22,8 @@ "capabilities": { "http": { "allowlist": [ - { "host": "discord.com", "path_prefix": "/api/v10" } + { "host": "discord.com", "path_prefix": "/api/v10" }, + { "host": "gateway.discord.gg", "path_prefix": "/", "methods": ["GET"] } ], "credentials": { "discord_bot_token": { @@ -36,6 +37,20 @@ "requests_per_hour": 3600 } }, + "websocket": { + "url": "wss://gateway.discord.gg/?v=10&encoding=json", + "connect_on_start": true, + "identify_secret_name": "discord_bot_token", + "identify": { + "_intents_doc": "GUILDS(1) + GUILD_MESSAGES(512) + DIRECT_MESSAGES(4096)", + "intents": 4609, + "properties": { + "os": "linux", + "browser": "ironclaw", + "device": "ironclaw" + } + } + }, "secrets": { "allowed_names": ["discord_bot_token", "discord_*"] }, diff --git a/channels-src/discord/src/lib.rs b/channels-src/discord/src/lib.rs index 249d1e5b341..fb0e1fc8d5e 100644 --- a/channels-src/discord/src/lib.rs +++ b/channels-src/discord/src/lib.rs @@ -10,11 +10,11 @@ //! - Message event parsing (@mentions, DMs) //! - Thread support for conversations //! - Response posting via Discord Web API -//! - Automatic message truncation (> 2000 chars) +//! - Markdown attachment fallback for oversized replies //! //! # Security //! -//! - Signature validation is handled in-channel using Discord's Ed25519 headers +//! - Signature validation is handled by the host (webhook secrets) //! - Bot token is injected by host during HTTP requests //! - WASM never sees raw credentials @@ -23,20 +23,18 @@ wit_bindgen::generate!({ path: "../../wit/channel.wit", }); -use std::{cmp::Ordering, collections::HashMap}; - -use ed25519_dalek::{Signature, Verifier, VerifyingKey}; use serde::{Deserialize, Serialize}; - -/// Discord REST API v10 base URL. -const DISCORD_API_BASE: &str = "https://discord.com/api/v10"; +use std::sync::atomic::{AtomicU64, Ordering}; +use std::time::{SystemTime, UNIX_EPOCH}; use exports::near::agent::channel::{ AgentResponse, ChannelConfig, Guest, HttpEndpointConfig, IncomingHttpRequest, - OutgoingHttpResponse, PollConfig, StatusUpdate, + OutgoingHttpResponse, PollConfig, StatusType, StatusUpdate, }; use near::agent::channel_host::{self, EmittedMessage}; +const DISCORD_API_BASE: &str = "https://discord.com/api/v10"; + /// Discord interaction wrapper. #[derive(Debug, Deserialize)] struct DiscordInteraction { @@ -111,146 +109,453 @@ struct DiscordMessage { author: DiscordUser, } -#[derive(Debug, Deserialize)] -struct DiscordChannelMessage { - id: String, - content: String, +/// Deserialize a String that may be null or missing (backward compat with old Option fields). +fn deserialize_nullable_string<'de, D>(deserializer: D) -> Result +where + D: serde::Deserializer<'de>, +{ + Option::::deserialize(deserializer).map(|opt| opt.unwrap_or_default()) +} + +/// Metadata stored with emitted messages for response routing. +#[derive(Debug, Serialize, Deserialize)] +struct DiscordMessageMetadata { + /// Discord channel ID channel_id: String, - author: DiscordChannelAuthor, - #[serde(default)] - mentions: Vec, + + /// Interaction ID for followups + #[serde(default, deserialize_with = "deserialize_nullable_string")] + interaction_id: String, + + /// Interaction token for responding + #[serde(default, deserialize_with = "deserialize_nullable_string")] + token: String, + + /// Application ID + #[serde(default, deserialize_with = "deserialize_nullable_string")] + application_id: String, + + /// Source message ID when handling mention-poll events. #[serde(default)] - webhook_id: Option, + source_message_id: Option, + + /// Thread ID (for forum threads) + thread_id: Option, } -#[derive(Debug, Deserialize)] -struct DiscordChannelAuthor { - id: String, - username: String, - global_name: Option, - #[serde(default)] - bot: bool, +#[derive(Debug, PartialEq, Eq)] +enum DiscordResponseRoute { + InteractionWebhook(String), + ChannelMessage(String), } -#[derive(Debug, Clone, Serialize, Deserialize)] -struct DiscordRuntimeConfig { - #[serde(default = "default_require_signature_verification")] - require_signature_verification: bool, - #[serde(default)] - webhook_secret: Option, - #[serde(default)] - polling_enabled: bool, - #[serde(default = "default_poll_interval_ms")] - poll_interval_ms: u32, - #[serde(default)] - mention_channel_ids: Vec, - #[serde(default)] - owner_id: Option, - #[serde(default = "default_dm_policy")] - dm_policy: String, - #[serde(default)] - allow_from: Vec, +fn response_route_for_metadata(metadata: &DiscordMessageMetadata) -> DiscordResponseRoute { + if !metadata.application_id.is_empty() && !metadata.token.is_empty() { + DiscordResponseRoute::InteractionWebhook(format!( + "{DISCORD_API_BASE}/webhooks/{}/{}/messages/@original", + metadata.application_id, metadata.token + )) + } else { + DiscordResponseRoute::ChannelMessage(format!( + "{DISCORD_API_BASE}/channels/{}/messages", + metadata.channel_id + )) + } } -fn default_poll_interval_ms() -> u32 { - 30_000 +fn typing_request_url_for_update(update: &StatusUpdate) -> Option { + if update.status != StatusType::Thinking { + return None; + } + + let metadata: DiscordMessageMetadata = serde_json::from_str(&update.metadata_json).ok()?; + if metadata.channel_id.is_empty() { + return None; + } + + Some(format!( + "{DISCORD_API_BASE}/channels/{}/typing", + metadata.channel_id + )) } -fn default_require_signature_verification() -> bool { - true +const DISCORD_MESSAGE_CHAR_LIMIT: usize = 2000; +const DISCORD_MULTIPART_BOUNDARY: &str = "ironclaw-discord-response-boundary"; +const DISCORD_ATTACHMENT_FILENAME: &str = "response.md"; +const DISCORD_ATTACHMENT_NOTICE: &str = "Response too long for Discord; attached as response.md."; +static MULTIPART_BOUNDARY_COUNTER: AtomicU64 = AtomicU64::new(0); + +#[derive(Debug, PartialEq, Eq)] +struct DiscordHttpRequest { + headers_json: String, + body: Vec, +} + +#[derive(Debug, PartialEq, Eq)] +enum DiscordReplyPlan { + Inline(DiscordHttpRequest), + Attachment { + upload: DiscordHttpRequest, + fallback: DiscordHttpRequest, + }, +} + +fn embeds_from_metadata_json(metadata_json: &str) -> Option { + serde_json::from_str::(metadata_json) + .ok()? + .get("embeds") + .cloned() +} + +fn build_discord_json_request( + content: &str, + embeds: Option<&serde_json::Value>, +) -> Result { + let mut payload = serde_json::json!({ + "content": content, + }); + + if let Some(embeds) = embeds { + payload["embeds"] = embeds.clone(); + } + + Ok(DiscordHttpRequest { + headers_json: serde_json::json!({ + "Content-Type": "application/json" + }) + .to_string(), + body: serde_json::to_vec(&payload).map_err(|e| format!("Failed to serialize: {}", e))?, + }) +} + +fn build_discord_attachment_request( + content: &str, + embeds: Option<&serde_json::Value>, +) -> Result { + let boundary = next_multipart_boundary(); + let mut payload = serde_json::json!({ + "content": DISCORD_ATTACHMENT_NOTICE, + }); + + if let Some(embeds) = embeds { + payload["embeds"] = embeds.clone(); + } + + let payload_json = + serde_json::to_string(&payload).map_err(|e| format!("Failed to serialize: {}", e))?; + + let mut body = Vec::new(); + body.extend_from_slice( + format!( + "--{boundary}\r\nContent-Disposition: form-data; name=\"payload_json\"\r\nContent-Type: application/json\r\n\r\n{payload_json}\r\n", + boundary = boundary, + ) + .as_bytes(), + ); + body.extend_from_slice( + format!( + "--{boundary}\r\nContent-Disposition: form-data; name=\"files[0]\"; filename=\"{filename}\"\r\nContent-Type: text/markdown\r\n\r\n", + boundary = boundary, + filename = DISCORD_ATTACHMENT_FILENAME, + ) + .as_bytes(), + ); + body.extend_from_slice(content.as_bytes()); + body.extend_from_slice(format!("\r\n--{}--\r\n", boundary).as_bytes()); + + Ok(DiscordHttpRequest { + headers_json: serde_json::json!({ + "Content-Type": format!( + "multipart/form-data; boundary={}", + boundary + ) + }) + .to_string(), + body, + }) +} + +fn next_multipart_boundary() -> String { + let counter = MULTIPART_BOUNDARY_COUNTER.fetch_add(1, Ordering::Relaxed); + let nanos = SystemTime::now() + .duration_since(UNIX_EPOCH) + .map(|duration| duration.as_nanos()) + .unwrap_or(0); + format!("{}-{:x}-{:x}", DISCORD_MULTIPART_BOUNDARY, nanos, counter) } -fn default_dm_policy() -> String { - "pairing".to_string() +fn build_discord_reply_plan(response: &AgentResponse) -> Result { + let embeds = embeds_from_metadata_json(&response.metadata_json); + + if response.content.chars().count() <= DISCORD_MESSAGE_CHAR_LIMIT { + return build_discord_json_request(&response.content, embeds.as_ref()) + .map(DiscordReplyPlan::Inline); + } + + Ok(DiscordReplyPlan::Attachment { + upload: build_discord_attachment_request(&response.content, embeds.as_ref())?, + fallback: build_discord_json_request( + &truncate_message(&response.content), + embeds.as_ref(), + )?, + }) } -fn default_runtime_config() -> DiscordRuntimeConfig { - DiscordRuntimeConfig { - require_signature_verification: default_require_signature_verification(), - webhook_secret: None, - polling_enabled: false, - poll_interval_ms: default_poll_interval_ms(), - mention_channel_ids: Vec::new(), - owner_id: None, - dm_policy: default_dm_policy(), - allow_from: Vec::new(), +fn send_discord_request( + method: &str, + url: &str, + request: &DiscordHttpRequest, +) -> Result<(), String> { + match channel_host::http_request( + method, + url, + &request.headers_json, + Some(&request.body), + None, + ) { + Ok(http_response) => { + if http_response.status >= 200 && http_response.status < 300 { + channel_host::log(channel_host::LogLevel::Debug, "Posted followup to Discord"); + Ok(()) + } else { + let body_str = String::from_utf8_lossy(&http_response.body); + Err(format!( + "Discord API error: {} - {}", + http_response.status, body_str + )) + } + } + Err(e) => Err(format!("HTTP request failed: {}", e)), } } /// Workspace path for persisting owner_id across WASM callbacks. const OWNER_ID_PATH: &str = "state/owner_id"; +/// Workspace path for persisting polling_enabled flag. +const POLLING_ENABLED_PATH: &str = "state/polling_enabled"; +/// Workspace path for persisting mention channel IDs (JSON array). +const MENTION_CHANNEL_IDS_PATH: &str = "state/mention_channel_ids"; /// Workspace path for persisting dm_policy across WASM callbacks. 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 the current gateway text-frame batch prepared by the host runtime. +const GATEWAY_EVENT_QUEUE_PATH: &str = "state/gateway_event_queue_processing"; +/// Workspace path for persisting the bot user id learned from READY dispatches. +const BOT_USER_ID_PATH: &str = "state/bot_user_id"; /// Channel name for pairing store (used by pairing host APIs). const CHANNEL_NAME: &str = "discord"; -/// Metadata stored with emitted messages for response routing. -#[derive(Debug, Serialize, Deserialize)] -struct DiscordMessageMetadata { - /// Discord channel ID - channel_id: String, +#[derive(Debug, Deserialize)] +struct DiscordGatewayEvent { + op: u64, + #[serde(default)] + t: Option, + #[serde(default)] + d: serde_json::Value, +} - /// Interaction ID for followups +#[derive(Debug, Deserialize)] +struct DiscordGatewayReady { + user: DiscordGatewayAuthor, +} + +#[derive(Debug, Deserialize, Clone)] +struct DiscordGatewayAuthor { + id: String, + username: String, + global_name: Option, #[serde(default)] - interaction_id: Option, + bot: bool, +} - /// Interaction token for responding +#[derive(Debug, Deserialize)] +struct DiscordGatewayMessageCreate { + channel_id: String, #[serde(default)] - token: Option, + guild_id: Option, + content: String, + author: DiscordGatewayAuthor, +} - /// Application ID +/// A message returned by the Discord REST channel-messages endpoint. +#[derive(Debug, Deserialize)] +struct DiscordChannelMessage { + id: String, + content: String, + channel_id: String, + author: DiscordChannelAuthor, + #[serde(default)] + mentions: Vec, #[serde(default)] - application_id: Option, + webhook_id: Option, +} - /// Source message ID when handling mention-poll events. +/// Author sub-object for REST channel messages. +#[derive(Debug, Deserialize)] +struct DiscordChannelAuthor { + id: String, + username: String, + global_name: Option, #[serde(default)] - source_message_id: Option, + bot: bool, +} - /// Thread ID (for forum threads) - thread_id: Option, +#[derive(Debug, PartialEq, Eq)] +struct ParsedGatewayMessage { + user_id: String, + user_name: String, + channel_id: String, + content: String, + is_dm: bool, } -struct DiscordChannel; +#[derive(Debug, Default, PartialEq, Eq)] +struct GatewayPollResult { + bot_user_id: Option, + messages: Vec, +} -impl Guest for DiscordChannel { - fn on_start(config_json: String) -> Result { - channel_host::log(channel_host::LogLevel::Info, "Discord channel starting"); +fn parse_gateway_event_queue( + queue_json: &str, + known_bot_user_id: Option<&str>, +) -> GatewayPollResult { + let frames: Vec = match serde_json::from_str(queue_json) { + Ok(v) => v, + Err(e) => { + channel_host::log( + channel_host::LogLevel::Warn, + &format!("Failed to deserialize gateway event queue: {}", e), + ); + return GatewayPollResult::default(); + } + }; + let mut result = GatewayPollResult::default(); + let mut bot_user_id = known_bot_user_id.map(ToOwned::to_owned); - let config = - serde_json::from_str::(&config_json).unwrap_or_else(|e| { - channel_host::log( - channel_host::LogLevel::Warn, - &format!("Invalid config JSON, using defaults: {}", e), - ); - default_runtime_config() - }); + for frame in frames { + let event: DiscordGatewayEvent = match serde_json::from_str(&frame) { + Ok(value) => value, + Err(_) => continue, + }; - if let Ok(serialized) = serde_json::to_string(&config) { - let _ = channel_host::workspace_write("config.json", &serialized); + if event.op != 0 { + continue; } - if config.require_signature_verification - && config - .webhook_secret - .as_deref() - .map(str::trim) - .filter(|s| !s.is_empty()) - .is_none() - { - channel_host::log( - channel_host::LogLevel::Error, - "Discord channel misconfigured: require_signature_verification=true but webhook_secret is empty", - ); - } else if !config.require_signature_verification { - channel_host::log( - channel_host::LogLevel::Warn, - "Discord signature verification is disabled; webhook endpoint is unprotected", - ); + match event.t.as_deref() { + Some("READY") => { + if let Ok(ready) = serde_json::from_value::(event.d) { + if !ready.user.id.is_empty() { + bot_user_id = Some(ready.user.id); + } + } + } + Some("MESSAGE_CREATE") => { + let message = match serde_json::from_value::(event.d) { + Ok(value) => value, + Err(_) => continue, + }; + + let active_bot_user_id = bot_user_id.as_deref().or(known_bot_user_id); + if message.author.bot + || active_bot_user_id.is_some_and(|bot_id| message.author.id == bot_id) + { + continue; + } + + let is_dm = message.guild_id.is_none(); + let content = + match gateway_content_for_agent(&message.content, active_bot_user_id, is_dm) { + Some(value) => value, + None => continue, + }; + + result.messages.push(ParsedGatewayMessage { + user_id: message.author.id, + user_name: message + .author + .global_name + .unwrap_or(message.author.username), + channel_id: message.channel_id, + content, + is_dm, + }); + } + _ => {} + } + } + + result.bot_user_id = bot_user_id; + result +} + +fn gateway_content_for_agent( + content: &str, + bot_user_id: Option<&str>, + is_dm: bool, +) -> Option { + let trimmed = content.trim(); + if trimmed.is_empty() { + return None; + } + + if is_dm { + return Some(trimmed.to_string()); + } + + let bot_user_id = bot_user_id?; + for mention in [ + format!("<@{}>", bot_user_id), + format!("<@!{}>", bot_user_id), + ] { + if let Some(stripped) = trimmed.strip_prefix(&mention) { + let cleaned = stripped.trim(); + return if cleaned.is_empty() { + None + } else { + Some(cleaned.to_string()) + }; } + } + + None +} + +fn default_poll_interval_ms() -> u32 { + 30_000 +} + +/// Channel configuration from capabilities file. +#[derive(Debug, Deserialize)] +struct DiscordConfig { + #[serde(default)] + #[allow(dead_code)] + require_signature_verification: bool, + #[serde(default)] + owner_id: Option, + #[serde(default)] + dm_policy: Option, + #[serde(default)] + allow_from: Option>, + #[serde(default)] + polling_enabled: bool, + #[serde(default = "default_poll_interval_ms")] + poll_interval_ms: u32, + #[serde(default)] + mention_channel_ids: Vec, +} + +struct DiscordChannel; + +impl Guest for DiscordChannel { + fn on_start(config_json: String) -> Result { + let config: DiscordConfig = serde_json::from_str(&config_json) + .map_err(|e| format!("Failed to parse config: {}", e))?; - // Persist owner_id so subsequent callbacks can read it. + channel_host::log(channel_host::LogLevel::Info, "Discord channel starting"); + + // Persist owner_id so subsequent callbacks can read it if let Some(ref owner_id) = config.owner_id { let _ = channel_host::workspace_write(OWNER_ID_PATH, owner_id); channel_host::log( @@ -261,18 +566,29 @@ impl Guest for DiscordChannel { let _ = channel_host::workspace_write(OWNER_ID_PATH, ""); } - // Persist dm_policy and allow_from for DM pairing. - let _ = channel_host::workspace_write(DM_POLICY_PATH, &config.dm_policy); - let allow_from_json = - serde_json::to_string(&config.allow_from).unwrap_or_else(|_| "[]".to_string()); + // Persist dm_policy and allow_from for DM pairing + let dm_policy = config.dm_policy.as_deref().unwrap_or("pairing"); + let _ = channel_host::workspace_write(DM_POLICY_PATH, dm_policy); + + let allow_from_json = serde_json::to_string(&config.allow_from.unwrap_or_default()) + .unwrap_or_else(|_| "[]".to_string()); let _ = channel_host::workspace_write(ALLOW_FROM_PATH, &allow_from_json); + // Persist polling config + let _ = channel_host::workspace_write( + POLLING_ENABLED_PATH, + &config.polling_enabled.to_string(), + ); + let mention_ids_json = + serde_json::to_string(&config.mention_channel_ids).unwrap_or_else(|_| "[]".to_string()); + let _ = channel_host::workspace_write(MENTION_CHANNEL_IDS_PATH, &mention_ids_json); + Ok(ChannelConfig { display_name: "Discord".to_string(), http_endpoints: vec![HttpEndpointConfig { path: "/webhook/discord".to_string(), methods: vec!["POST".to_string()], - require_secret: false, + require_secret: true, }], poll: if config.polling_enabled { Some(PollConfig { @@ -286,45 +602,6 @@ impl Guest for DiscordChannel { } fn on_http_request(req: IncomingHttpRequest) -> OutgoingHttpResponse { - let config = load_runtime_config(); - let headers: HashMap = - serde_json::from_str(&req.headers_json).unwrap_or_default(); - if config.require_signature_verification { - if config - .webhook_secret - .as_deref() - .map(str::trim) - .filter(|s| !s.is_empty()) - .is_none() - { - channel_host::log( - channel_host::LogLevel::Error, - "Discord channel misconfigured: webhook_secret not set while verification is required", - ); - return json_response( - 500, - serde_json::json!({"error": "Channel misconfigured: webhook_secret not set"}), - ); - } - - if !verify_discord_request_signature( - headers, - &req.body, - config.webhook_secret.as_deref(), - ) { - channel_host::log( - channel_host::LogLevel::Warn, - "Discord signature verification failed", - ); - return json_response(401, serde_json::json!({"error": "Invalid signature"})); - } - } else { - channel_host::log( - channel_host::LogLevel::Warn, - "Discord signature verification is disabled; accepting unverified webhook request", - ); - } - let body_str = match std::str::from_utf8(&req.body) { Ok(s) => s, Err(_) => { @@ -353,23 +630,16 @@ impl Guest for DiscordChannel { // Application Command (slash command) 2 => { if handle_slash_command(&interaction) { + json_response(200, serde_json::json!({"type": 5})) + } else { + // Permission denied — ephemeral response json_response( 200, serde_json::json!({ - "type": 5, + "type": 4, "data": { - "content": "🤔 Thinking..." - } - }), - ) - } else { - json_response( - 200, - serde_json::json!({ - "type": 4, - "data": { - "content": "You are not authorized to use this bot.", - "flags": 64 + "content": "You are not authorized to use this bot.", + "flags": 64 } }), ) @@ -398,530 +668,200 @@ impl Guest for DiscordChannel { } fn on_poll() { - poll_for_mentions(); - } - - fn on_respond(response: AgentResponse) -> Result<(), String> { - let metadata: DiscordMessageMetadata = serde_json::from_str(&response.metadata_json) - .map_err(|e| format!("Failed to parse metadata: {}", e))?; - - // Truncate content to 2000 characters to comply with Discord limits - let content = truncate_message(&response.content); - - let mut payload = serde_json::json!({ "content": content }); - - // Check for embeds in metadata - if let Ok(meta_json) = serde_json::from_str::(&response.metadata_json) { - if let Some(embeds) = meta_json.get("embeds") { - payload["embeds"] = embeds.clone(); - } - } - - let payload_bytes = - serde_json::to_vec(&payload).map_err(|e| format!("Failed to serialize: {}", e))?; - - let headers = serde_json::json!({ - "Content-Type": "application/json" - }); - - let (method, url) = if let (Some(application_id), Some(token)) = - (metadata.application_id.as_ref(), metadata.token.as_ref()) - { - ( - "PATCH", - format!( - "{DISCORD_API_BASE}/webhooks/{}/{}/messages/@original", - application_id, token - ), - ) - } else if let Some(source_message_id) = metadata.source_message_id.as_ref() { - payload["message_reference"] = serde_json::json!({ - "message_id": source_message_id - }); - payload["allowed_mentions"] = serde_json::json!({ - "replied_user": true - }); - return send_channel_message(&metadata.channel_id, payload); - } else { - return Err("Unsupported Discord response metadata".to_string()); - }; - - let result = channel_host::http_request( - method, - &url, - &headers.to_string(), - Some(&payload_bytes), - None, - ); - - map_discord_response(result) - } - - fn on_status(_update: StatusUpdate) {} - - fn on_broadcast(user_id: String, response: AgentResponse) -> Result<(), String> { - broadcast_dm(&user_id, &response.content) - } - - fn on_shutdown() { - channel_host::log( - channel_host::LogLevel::Info, - "Discord channel shutting down", - ); - } -} - -fn map_discord_response( - result: Result, -) -> Result<(), String> { - match result { - Ok(http_response) => { - if http_response.status >= 200 && http_response.status < 300 { - channel_host::log(channel_host::LogLevel::Debug, "Posted response to Discord"); - Ok(()) - } else { - let body_str = String::from_utf8_lossy(&http_response.body); - Err(format!( - "Discord API error: {} - {}", - http_response.status, body_str - )) - } - } - Err(e) => Err(format!("HTTP request failed: {}", e)), - } -} - -/// Post a JSON payload to a Discord channel as a new message. -fn send_channel_message(channel_id: &str, payload: serde_json::Value) -> Result<(), String> { - let payload_bytes = serde_json::to_vec(&payload) - .map_err(|e| format!("Failed to serialize message: {}", e))?; - let url = format!("{DISCORD_API_BASE}/channels/{}/messages", channel_id); - let result = channel_host::http_request( - "POST", - &url, - &discord_auth_headers_json(true), - Some(&payload_bytes), - None, - ); - map_discord_response(result) -} - -fn load_runtime_config() -> DiscordRuntimeConfig { - channel_host::workspace_read("config.json") - .and_then(|raw| serde_json::from_str::(&raw).ok()) - .unwrap_or_else(default_runtime_config) -} - -fn poll_for_mentions() { - let config = load_runtime_config(); - if !config.polling_enabled || config.mention_channel_ids.is_empty() { - return; - } - - let bot_id = match get_or_fetch_bot_id() { - Some(id) => id, - None => { - channel_host::log( - channel_host::LogLevel::Warn, - "Skipping mention polling: failed to resolve bot user id", - ); - return; - } - }; - - for channel_id in &config.mention_channel_ids { - poll_channel_mentions(channel_id, &bot_id); - } -} - -fn get_or_fetch_bot_id() -> Option { - if let Some(id) = channel_host::workspace_read("bot_user_id.txt") { - let trimmed = id.trim(); - if !trimmed.is_empty() { - return Some(trimmed.to_string()); - } - } - - let response = channel_host::http_request( - "GET", - &format!("{DISCORD_API_BASE}/users/@me"), - &discord_auth_headers_json(false), - None, - Some(10_000), - ) - .ok()?; - - if !(200..300).contains(&response.status) { - return None; - } - - let value: serde_json::Value = serde_json::from_slice(&response.body).ok()?; - let id = value.get("id")?.as_str()?.to_string(); - let _ = channel_host::workspace_write("bot_user_id.txt", &id); - Some(id) -} - -fn poll_channel_mentions(channel_id: &str, bot_id: &str) { - let cursor_path = format!("cursor_{}.txt", channel_id); - let last_seen = channel_host::workspace_read(&cursor_path).map(|s| s.trim().to_string()); - - // On first run for a channel, initialize the cursor to "latest seen" and - // skip back-processing historical messages. - if last_seen.is_none() { - if let Some(latest) = fetch_latest_message_id(channel_id) { - let _ = channel_host::workspace_write(&cursor_path, &latest); - } - return; - } - - let Some(mut messages) = - fetch_messages_after_cursor(channel_id, last_seen.as_deref().unwrap_or("")) - else { - return; - }; - if messages.is_empty() { - return; - } - - messages.sort_by(|a, b| compare_message_ids(&a.id, &b.id)); - let mut max_seen = last_seen.clone(); - let mut recent_ids = load_recent_processed_ids(channel_id); - let mut dedup_updated = false; - - for msg in messages { - if is_new_message(max_seen.as_deref(), &msg.id) { - max_seen = Some(msg.id.clone()); - } - - if msg.webhook_id.is_some() || msg.author.bot || msg.author.id == bot_id { - continue; - } - - if !message_mentions_bot(&msg, bot_id) { - continue; - } - - if recent_ids.iter().any(|id| id == &msg.id) { - continue; - } - - let user_name = msg - .author - .global_name - .as_ref() - .filter(|s| !s.is_empty()) - .unwrap_or(&msg.author.username) - .clone(); - if !check_sender_permission(&msg.author.id, Some(&user_name), false, None) { - continue; - } - - let content = strip_bot_mention(&msg.content, bot_id); - let metadata = DiscordMessageMetadata { - channel_id: msg.channel_id.clone(), - interaction_id: None, - token: None, - application_id: None, - source_message_id: Some(msg.id.clone()), - thread_id: None, - }; - - let metadata_json = match serde_json::to_string(&metadata) { - Ok(v) => v, - Err(e) => { - channel_host::log( - channel_host::LogLevel::Warn, - &format!("Failed to serialize mention metadata: {}", e), - ); - continue; - } - }; - - channel_host::emit_message(&EmittedMessage { - user_id: msg.author.id.clone(), - user_name: Some(user_name.clone()), - content: if content.is_empty() { - "mention".to_string() - } else { - content - }, - thread_id: None, - metadata_json, - attachments: vec![], - }); + // 1. Process Gateway event queue + let queue_json = channel_host::workspace_read(GATEWAY_EVENT_QUEUE_PATH).unwrap_or_default(); + let has_gateway_events = !queue_json.trim().is_empty() && queue_json.trim() != "[]"; - remember_processed_id(&mut recent_ids, &msg.id); - dedup_updated = true; - } - - if let Some(cursor) = max_seen { - let _ = channel_host::workspace_write(&cursor_path, &cursor); - } - if dedup_updated { - let _ = save_recent_processed_ids(channel_id, &recent_ids); - } -} + if has_gateway_events { + let known_bot_user_id = channel_host::workspace_read(BOT_USER_ID_PATH); + let parsed = parse_gateway_event_queue(&queue_json, known_bot_user_id.as_deref()); -fn fetch_latest_message_id(channel_id: &str) -> Option { - let url = format!( - "{DISCORD_API_BASE}/channels/{}/messages?limit=1", - channel_id - ); - let response = channel_host::http_request( - "GET", - &url, - &discord_auth_headers_json(false), - None, - Some(10_000), - ) - .ok()?; - if !(200..300).contains(&response.status) { - let body = String::from_utf8_lossy(&response.body); - channel_host::log( - channel_host::LogLevel::Warn, - &format!( - "Discord initial poll failed for channel {}: status={} body={}", - channel_id, response.status, body - ), - ); - return None; - } - let messages: Vec = serde_json::from_slice(&response.body).ok()?; - messages.first().map(|m| m.id.clone()) -} - -fn fetch_messages_after_cursor( - channel_id: &str, - last_seen: &str, -) -> Option> { - const PAGE_LIMIT: usize = 100; - const MAX_PAGES: usize = 50; - - let mut all_messages = Vec::new(); - let mut after = last_seen.to_string(); - - for page in 0..MAX_PAGES { - let url = format!( - "{DISCORD_API_BASE}/channels/{}/messages?limit={}&after={}", - channel_id, PAGE_LIMIT, after - ); - let response = match channel_host::http_request( - "GET", - &url, - &discord_auth_headers_json(false), - None, - Some(10_000), - ) { - Ok(r) => r, - Err(e) => { + if let Err(error) = channel_host::workspace_write(GATEWAY_EVENT_QUEUE_PATH, "[]") { channel_host::log( channel_host::LogLevel::Warn, - &format!( - "Discord poll request failed for channel {}: {}", - channel_id, e - ), + &format!("Failed to clear Discord gateway queue: {}", error), ); - return None; } - }; - if !(200..300).contains(&response.status) { - let body = String::from_utf8_lossy(&response.body); - channel_host::log( - channel_host::LogLevel::Warn, - &format!( - "Discord poll failed for channel {}: status={} body={}", - channel_id, response.status, body - ), - ); - return None; - } - - let messages: Vec = match serde_json::from_slice(&response.body) { - Ok(v) => v, - Err(e) => { - channel_host::log( - channel_host::LogLevel::Warn, - &format!("Failed to parse polled Discord messages: {}", e), - ); - return None; + if let Some(bot_user_id) = parsed.bot_user_id.as_deref() { + if let Err(error) = channel_host::workspace_write(BOT_USER_ID_PATH, bot_user_id) { + channel_host::log( + channel_host::LogLevel::Warn, + &format!("Failed to persist Discord bot user id: {}", error), + ); + } } - }; - let page_len = messages.len(); - if messages.is_empty() { - break; - } - - let page_max_id = messages - .iter() - .map(|m| m.id.as_str()) - .max_by(|a, b| compare_message_ids(a, b)) - .map(str::to_string); - all_messages.extend(messages.into_iter()); - - if page_len < PAGE_LIMIT { - break; - } + for message in parsed.messages { + if !check_sender_permission( + &message.user_id, + Some(&message.user_name), + message.is_dm, + PermissionSource::Gateway, + Some(&PairingReplyCtx { + channel_id: message.channel_id.clone(), + application_id: String::new(), + token: String::new(), + }), + ) { + continue; + } - if let Some(max_id) = page_max_id { - if max_id == after { - break; + let metadata = DiscordMessageMetadata { + channel_id: message.channel_id, + interaction_id: String::new(), + token: String::new(), + application_id: String::new(), + source_message_id: None, + thread_id: None, + }; + + let metadata_json = match serde_json::to_string(&metadata) { + Ok(json) => json, + Err(error) => { + channel_host::log( + channel_host::LogLevel::Error, + &format!("Failed to serialize gateway metadata: {}", error), + ); + continue; + } + }; + + channel_host::emit_message(&EmittedMessage { + user_id: message.user_id, + user_name: Some(message.user_name), + content: message.content, + thread_id: None, + metadata_json, + attachments: vec![], + }); } - after = max_id; - } else { - break; - } - - if page + 1 == MAX_PAGES { - channel_host::log( - channel_host::LogLevel::Warn, - &format!( - "Discord poll pagination limit reached for channel {}; processing partial batch", - channel_id - ), - ); } - } - - Some(all_messages) -} -fn compare_message_ids(a: &str, b: &str) -> Ordering { - match (a.parse::(), b.parse::()) { - (Ok(left), Ok(right)) => left.cmp(&right), - _ => a.cmp(b), + // 2. Run mention polling if configured + poll_for_mentions(); } -} -fn dedup_ids_path(channel_id: &str) -> String { - format!("dedup_{}.json", channel_id) -} - -fn load_recent_processed_ids(channel_id: &str) -> Vec { - let path = dedup_ids_path(channel_id); - channel_host::workspace_read(&path) - .and_then(|raw| serde_json::from_str::>(&raw).ok()) - .unwrap_or_default() -} + fn on_respond(response: AgentResponse) -> Result<(), String> { + let metadata: DiscordMessageMetadata = serde_json::from_str(&response.metadata_json) + .map_err(|e| format!("Failed to parse metadata: {}", e))?; -fn save_recent_processed_ids(channel_id: &str, ids: &[String]) -> Result<(), String> { - let path = dedup_ids_path(channel_id); - let raw = - serde_json::to_string(ids).map_err(|e| format!("Failed to serialize dedup ids: {}", e))?; - channel_host::workspace_write(&path, &raw) -} + // Mention-poll replies: include message_reference so Discord renders as a reply + if let Some(ref source_id) = metadata.source_message_id { + if let DiscordResponseRoute::ChannelMessage(ref url) = + response_route_for_metadata(&metadata) + { + let embeds = embeds_from_metadata_json(&response.metadata_json); + let content = if response.content.chars().count() > DISCORD_MESSAGE_CHAR_LIMIT { + truncate_message(&response.content) + } else { + response.content.clone() + }; + + let mut payload = serde_json::json!({ + "content": content, + "message_reference": { + "message_id": source_id + }, + "allowed_mentions": { + "replied_user": true + } + }); -fn remember_processed_id(ids: &mut Vec, message_id: &str) { - const MAX_RECENT_IDS: usize = 200; - if ids.iter().any(|id| id == message_id) { - return; - } - ids.push(message_id.to_string()); - if ids.len() > MAX_RECENT_IDS { - let drop_count = ids.len() - MAX_RECENT_IDS; - ids.drain(0..drop_count); - } -} + if let Some(ref e) = embeds { + payload["embeds"] = e.clone(); + } -fn is_new_message(last_seen: Option<&str>, current: &str) -> bool { - match last_seen { - None => true, - Some(prev) => { - let prev_num = prev.parse::().ok(); - let cur_num = current.parse::().ok(); - match (prev_num, cur_num) { - (Some(p), Some(c)) => c > p, - _ => current > prev, + let headers = discord_auth_headers_json(true); + let body = serde_json::to_vec(&payload) + .map_err(|e| format!("Failed to serialize: {}", e))?; + + return send_discord_request( + "POST", + url, + &DiscordHttpRequest { + headers_json: headers, + body, + }, + ); } } - } -} -fn message_mentions_bot(msg: &DiscordChannelMessage, bot_id: &str) -> bool { - msg.mentions.iter().any(|u| u.id == bot_id) - || msg.content.contains(&format!("<@{}>", bot_id)) - || msg.content.contains(&format!("<@!{}>", bot_id)) -} + let route = response_route_for_metadata(&metadata); + let plan = build_discord_reply_plan(&response)?; -fn strip_bot_mention(content: &str, bot_id: &str) -> String { - content - .replace(&format!("<@{}>", bot_id), "") - .replace(&format!("<@!{}>", bot_id), "") - .trim() - .to_string() -} + let (method, url) = match &route { + DiscordResponseRoute::InteractionWebhook(url) => ("PATCH", url.as_str()), + DiscordResponseRoute::ChannelMessage(url) => ("POST", url.as_str()), + }; -fn discord_auth_headers_json(include_content_type: bool) -> String { - if include_content_type { - serde_json::json!({ - "Content-Type": "application/json", - "Authorization": "Bot {DISCORD_BOT_TOKEN}" - }) - .to_string() - } else { - serde_json::json!({ - "Authorization": "Bot {DISCORD_BOT_TOKEN}" - }) - .to_string() + match plan { + DiscordReplyPlan::Inline(request) => send_discord_request(method, url, &request), + DiscordReplyPlan::Attachment { upload, fallback } => { + match send_discord_request(method, url, &upload) { + Ok(()) => Ok(()), + Err(upload_error) => { + channel_host::log( + channel_host::LogLevel::Warn, + &format!( + "Discord attachment upload failed, falling back to truncated text: {}", + upload_error + ), + ); + send_discord_request(method, url, &fallback).map_err(|fallback_error| { + format!( + "Discord attachment upload failed: {}; fallback also failed: {}", + upload_error, fallback_error + ) + }) + } + } + } + } } -} - -fn verify_discord_request_signature( - headers: HashMap, - body: &[u8], - public_key_hex: Option<&str>, -) -> bool { - let Some(public_key_hex) = public_key_hex.map(str::trim).filter(|s| !s.is_empty()) else { - return false; - }; - let Some(signature_hex) = header_case_insensitive(&headers, "x-signature-ed25519") else { - return false; - }; - let Some(timestamp) = header_case_insensitive(&headers, "x-signature-timestamp") else { - return false; - }; - let public_key_bytes = match hex::decode(public_key_hex) { - Ok(v) => v, - Err(_) => return false, - }; - let public_key_arr: [u8; 32] = match public_key_bytes.try_into() { - Ok(v) => v, - Err(_) => return false, - }; - let verifying_key = match VerifyingKey::from_bytes(&public_key_arr) { - Ok(v) => v, - Err(_) => return false, - }; + fn on_status(update: StatusUpdate) { + let Some(url) = typing_request_url_for_update(&update) else { + return; + }; - let sig_bytes = match hex::decode(signature_hex.trim()) { - Ok(v) => v, - Err(_) => return false, - }; - let sig_arr: [u8; 64] = match sig_bytes.try_into() { - Ok(v) => v, - Err(_) => return false, - }; - let signature = Signature::from_bytes(&sig_arr); + let headers = serde_json::json!({ + "Content-Type": "application/json" + }); - let mut signed_message = Vec::with_capacity(timestamp.len() + body.len()); - signed_message.extend_from_slice(timestamp.as_bytes()); - signed_message.extend_from_slice(body); + match channel_host::http_request("POST", &url, &headers.to_string(), None, None) { + Ok(response) if (200..300).contains(&response.status) => {} + Ok(response) => { + channel_host::log( + channel_host::LogLevel::Warn, + &format!( + "Discord typing indicator failed with status {}", + response.status + ), + ); + } + Err(error) => { + channel_host::log( + channel_host::LogLevel::Warn, + &format!("Discord typing indicator request failed: {}", error), + ); + } + } + } - verifying_key.verify(&signed_message, &signature).is_ok() -} + fn on_broadcast(user_id: String, response: AgentResponse) -> Result<(), String> { + broadcast_dm(&user_id, &response.content) + } -fn header_case_insensitive<'a>( - headers: &'a HashMap, - name: &str, -) -> Option<&'a str> { - headers - .iter() - .find(|(k, _)| k.eq_ignore_ascii_case(name)) - .map(|(_, v)| v.as_str()) + fn on_shutdown() { + channel_host::log( + channel_host::LogLevel::Info, + "Discord channel shutting down", + ); + } } +/// Returns true if the message was emitted, false if permission denied. fn handle_slash_command(interaction: &DiscordInteraction) -> bool { let user = interaction .member @@ -939,13 +879,17 @@ fn handle_slash_command(interaction: &DiscordInteraction) -> bool { }) .unwrap_or_default(); - // DM if no guild member context (only direct user field set). + // DM if no guild member context (only direct user field set) let is_dm = interaction.member.is_none(); + + // Permission check if !check_sender_permission( &user_id, Some(&user_name), is_dm, + PermissionSource::Webhook, Some(&PairingReplyCtx { + channel_id: interaction.channel_id.clone().unwrap_or_default(), application_id: interaction.application_id.clone(), token: interaction.token.clone(), }), @@ -975,9 +919,9 @@ fn handle_slash_command(interaction: &DiscordInteraction) -> bool { let metadata = DiscordMessageMetadata { channel_id: channel_id.clone(), - interaction_id: Some(interaction.id.clone()), - token: Some(interaction.token.clone()), - application_id: Some(interaction.application_id.clone()), + interaction_id: interaction.id.clone(), + token: interaction.token.clone(), + application_id: interaction.application_id.clone(), source_message_id: None, thread_id: None, }; @@ -989,14 +933,13 @@ fn handle_slash_command(interaction: &DiscordInteraction) -> bool { channel_host::LogLevel::Error, &format!("Failed to serialize metadata: {}", e), ); - // Attempt to notify user of internal error let url = format!( "{DISCORD_API_BASE}/webhooks/{}/{}", interaction.application_id, interaction.token ); let payload = serde_json::json!({ "content": "❌ Internal Error: Failed to process command metadata.", - "flags": 64 // Ephemeral + "flags": 64 }); let _ = channel_host::http_request( "POST", @@ -1005,7 +948,7 @@ fn handle_slash_command(interaction: &DiscordInteraction) -> bool { Some(&serde_json::to_vec(&payload).unwrap_or_default()), None, ); - return true; + return true; // Error, but not a permission denial } }; @@ -1021,7 +964,6 @@ fn handle_slash_command(interaction: &DiscordInteraction) -> bool { } fn handle_message_component(interaction: &DiscordInteraction, message: &DiscordMessage) { - // Check member first (for server contexts), then user (for DMs) let user = interaction .member .as_ref() @@ -1039,7 +981,13 @@ fn handle_message_component(interaction: &DiscordInteraction, message: &DiscordM .unwrap_or_default(); let is_dm = interaction.member.is_none(); - if !check_sender_permission(&user_id, Some(&user_name), is_dm, None) { + if !check_sender_permission( + &user_id, + Some(&user_name), + is_dm, + PermissionSource::Webhook, + None, + ) { return; } @@ -1047,9 +995,9 @@ fn handle_message_component(interaction: &DiscordInteraction, message: &DiscordM let metadata = DiscordMessageMetadata { channel_id: channel_id.clone(), - interaction_id: Some(interaction.id.clone()), - token: Some(interaction.token.clone()), - application_id: Some(interaction.application_id.clone()), + interaction_id: interaction.id.clone(), + token: interaction.token.clone(), + application_id: interaction.application_id.clone(), source_message_id: None, thread_id: None, }; @@ -1075,21 +1023,39 @@ fn handle_message_component(interaction: &DiscordInteraction, message: &DiscordM }); } +// ============================================================================ +// Permission & Pairing +// ============================================================================ + /// Context needed to send a pairing reply via Discord webhook followup. struct PairingReplyCtx { + channel_id: String, application_id: String, token: String, } +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +enum PermissionSource { + Webhook, + Gateway, +} + +fn should_apply_dm_pairing(_source: PermissionSource, is_dm: bool) -> bool { + // All current permission sources (Webhook, Gateway) apply DM pairing equally. + // Kept as a function for future sources that may bypass pairing (e.g. internal). + is_dm +} + /// Check if a sender is permitted to interact with the bot. /// Returns true if allowed, false if denied (pairing reply sent if applicable). fn check_sender_permission( user_id: &str, username: Option<&str>, is_dm: bool, + source: PermissionSource, reply_ctx: Option<&PairingReplyCtx>, ) -> bool { - // 1. Owner check (highest priority, applies to all contexts). + // 1. Owner check (highest priority, applies to all contexts) let owner_id = channel_host::workspace_read(OWNER_ID_PATH).filter(|s| !s.is_empty()); if let Some(ref owner) = owner_id { if user_id != owner { @@ -1105,26 +1071,28 @@ fn check_sender_permission( return true; } - // 2. DM policy (only for DMs when no owner_id). - if !is_dm { + // 2. DM policy (only for DMs when no owner_id) + if !should_apply_dm_pairing(source, is_dm) { return true; } let dm_policy = - channel_host::workspace_read(DM_POLICY_PATH).unwrap_or_else(default_dm_policy); + channel_host::workspace_read(DM_POLICY_PATH).unwrap_or_else(|| "pairing".to_string()); + if dm_policy == "open" { return true; } - // 3. Build merged allow list: config allow_from + pairing store. + // 3. Build merged allow list: config allow_from + pairing store let mut allowed: Vec = channel_host::workspace_read(ALLOW_FROM_PATH) .and_then(|s| serde_json::from_str(&s).ok()) .unwrap_or_default(); + if let Ok(store_allowed) = channel_host::pairing_read_allow_from(CHANNEL_NAME) { allowed.extend(store_allowed); } - // 4. Check sender against allow list. + // 4. Check sender against allow list let is_allowed = allowed.contains(&"*".to_string()) || allowed.contains(&user_id.to_string()) || username.is_some_and(|u| allowed.contains(&u.to_string())); @@ -1133,13 +1101,14 @@ fn check_sender_permission( return true; } - // 5. Not allowed - handle by policy. + // 5. Not allowed — handle by policy if dm_policy == "pairing" { let meta = serde_json::json!({ "user_id": user_id, "username": username, }) .to_string(); + match channel_host::pairing_upsert_request(CHANNEL_NAME, user_id, &meta) { Ok(result) => { channel_host::log( @@ -1163,29 +1132,52 @@ fn check_sender_permission( false } -/// Send a pairing code as an ephemeral Discord followup message. +fn pairing_reply_route(ctx: &PairingReplyCtx) -> DiscordResponseRoute { + if !ctx.application_id.is_empty() && !ctx.token.is_empty() { + DiscordResponseRoute::InteractionWebhook(format!( + "{DISCORD_API_BASE}/webhooks/{}/{}", + ctx.application_id, ctx.token + )) + } else { + DiscordResponseRoute::ChannelMessage(format!( + "{DISCORD_API_BASE}/channels/{}/messages", + ctx.channel_id + )) + } +} + +/// Send a pairing code reply via webhook followup or channel message. fn send_pairing_reply(ctx: &PairingReplyCtx, code: &str) -> Result<(), String> { - let url = format!( - "{DISCORD_API_BASE}/webhooks/{}/{}", - ctx.application_id, ctx.token - ); - let payload = serde_json::json!({ + let route = pairing_reply_route(ctx); + + let mut payload = serde_json::json!({ "content": format!( "To pair with this bot, run: `ironclaw pairing approve discord {}`", code - ), - "flags": 64 + ) }); + + if matches!(route, DiscordResponseRoute::InteractionWebhook(_)) { + payload["flags"] = serde_json::json!(64); + } + let payload_bytes = serde_json::to_vec(&payload).map_err(|e| format!("Failed to serialize: {}", e))?; + let headers = serde_json::json!({"Content-Type": "application/json"}); + let url = match &route { + DiscordResponseRoute::InteractionWebhook(url) => url, + DiscordResponseRoute::ChannelMessage(url) => url, + }; + let result = channel_host::http_request( "POST", - &url, + url, &headers.to_string(), Some(&payload_bytes), None, ); + match result { Ok(response) if response.status >= 200 && response.status < 300 => Ok(()), Ok(response) => { @@ -1195,14 +1187,336 @@ fn send_pairing_reply(ctx: &PairingReplyCtx, code: &str) -> Result<(), String> { response.status, body_str )) } - Err(e) => Err(format!("HTTP request failed: {}", e)), + Err(e) => Err(format!("HTTP request failed: {}", e)), + } +} + +// ============================================================================ +// Mention Polling +// ============================================================================ + +/// Maximum number of processed message IDs to keep per channel for dedup. +const DEDUP_CAP: usize = 200; + +/// Poll configured channels for new messages that mention the bot. +fn poll_for_mentions() { + let enabled = channel_host::workspace_read(POLLING_ENABLED_PATH) + .map(|v| v.trim() == "true") + .unwrap_or(false); + + if !enabled { + return; + } + + let bot_id = match get_or_fetch_bot_id() { + Some(id) => id, + None => { + channel_host::log( + channel_host::LogLevel::Warn, + "Mention polling: unable to determine bot user id", + ); + return; + } + }; + + let channel_ids: Vec = channel_host::workspace_read(MENTION_CHANNEL_IDS_PATH) + .and_then(|s| serde_json::from_str(&s).ok()) + .unwrap_or_default(); + + for channel_id in &channel_ids { + poll_channel_mentions(channel_id, &bot_id); + } +} + +/// Read the bot user ID from workspace or fetch it from the Discord API. +fn get_or_fetch_bot_id() -> Option { + if let Some(id) = channel_host::workspace_read(BOT_USER_ID_PATH).filter(|s| !s.is_empty()) { + return Some(id); + } + + let headers = discord_auth_headers_json(false); + let resp = channel_host::http_request( + "GET", + "{DISCORD_API_BASE}/users/@me", + &headers, + None, + None, + ) + .ok()?; + + if resp.status < 200 || resp.status >= 300 { + return None; + } + + let body: serde_json::Value = serde_json::from_slice(&resp.body).ok()?; + let id = body["id"].as_str()?.to_string(); + + let _ = channel_host::workspace_write(BOT_USER_ID_PATH, &id); + Some(id) +} + +/// Poll a single channel for new mention messages. +fn poll_channel_mentions(channel_id: &str, bot_id: &str) { + let cursor_path = format!("state/mention_cursor/{}", channel_id); + let last_seen = channel_host::workspace_read(&cursor_path).unwrap_or_default(); + + let messages = if last_seen.is_empty() { + // First poll: initialise cursor without emitting any messages. + if let Some(latest_id) = fetch_latest_message_id(channel_id) { + let _ = channel_host::workspace_write(&cursor_path, &latest_id); + } + return; + } else { + match fetch_messages_after_cursor(channel_id, &last_seen) { + Some(msgs) => msgs, + None => return, + } + }; + + let mut processed_ids = load_recent_processed_ids(channel_id); + let mut new_cursor = last_seen.clone(); + + for msg in &messages { + if !is_new_message(&last_seen, &msg.id) { + continue; + } + if processed_ids.contains(&msg.id) { + continue; + } + if msg.author.bot || msg.author.id == bot_id { + remember_processed_id(&msg.id, &mut processed_ids); + continue; + } + if msg.webhook_id.is_some() { + remember_processed_id(&msg.id, &mut processed_ids); + continue; + } + if !message_mentions_bot(msg, bot_id) { + remember_processed_id(&msg.id, &mut processed_ids); + continue; + } + + // Permission check (API-based poll uses Webhook source) + if !check_sender_permission( + &msg.author.id, + Some(&msg.author.username), + false, + PermissionSource::Webhook, + None, + ) { + remember_processed_id(&msg.id, &mut processed_ids); + continue; + } + + let content = strip_bot_mention(&msg.content, bot_id); + if content.is_empty() { + remember_processed_id(&msg.id, &mut processed_ids); + continue; + } + + let user_name = msg + .author + .global_name + .clone() + .unwrap_or_else(|| msg.author.username.clone()); + + let metadata = DiscordMessageMetadata { + channel_id: msg.channel_id.clone(), + interaction_id: String::new(), + token: String::new(), + application_id: String::new(), + source_message_id: Some(msg.id.clone()), + thread_id: None, + }; + + let metadata_json = match serde_json::to_string(&metadata) { + Ok(json) => json, + Err(error) => { + channel_host::log( + channel_host::LogLevel::Error, + &format!("Failed to serialize mention-poll metadata: {}", error), + ); + continue; + } + }; + + channel_host::emit_message(&EmittedMessage { + user_id: msg.author.id.clone(), + user_name: Some(user_name), + content, + thread_id: None, + metadata_json, + attachments: vec![], + }); + + remember_processed_id(&msg.id, &mut processed_ids); + + if compare_message_ids(&msg.id, &new_cursor) == std::cmp::Ordering::Greater { + new_cursor = msg.id.clone(); + } + } + + if new_cursor != last_seen { + let _ = channel_host::workspace_write(&cursor_path, &new_cursor); + } + + save_recent_processed_ids(channel_id, &processed_ids); +} + +/// Fetch the latest message ID in a channel (used for cursor initialisation). +fn fetch_latest_message_id(channel_id: &str) -> Option { + let url = format!( + "{DISCORD_API_BASE}/channels/{}/messages?limit=1", + channel_id + ); + let headers = discord_auth_headers_json(false); + let resp = channel_host::http_request("GET", &url, &headers, None, None).ok()?; + + if resp.status < 200 || resp.status >= 300 { + return None; + } + + let messages: Vec = serde_json::from_slice(&resp.body).ok()?; + messages + .first() + .and_then(|m| m["id"].as_str().map(String::from)) +} + +/// Maximum number of pages to fetch when catching up on missed messages. +const MENTION_POLL_MAX_PAGES: usize = 5; + +/// Fetch messages after `last_seen` using the `after` parameter, paginating up +/// to [`MENTION_POLL_MAX_PAGES`] pages of 100 messages each. +fn fetch_messages_after_cursor( + channel_id: &str, + last_seen: &str, +) -> Option> { + let headers = discord_auth_headers_json(false); + let mut all_messages: Vec = Vec::new(); + let mut after = last_seen.to_string(); + + for _ in 0..MENTION_POLL_MAX_PAGES { + let url = format!( + "{DISCORD_API_BASE}/channels/{}/messages?after={}&limit=100", + channel_id, after + ); + let resp = channel_host::http_request("GET", &url, &headers, None, None).ok()?; + + if resp.status < 200 || resp.status >= 300 { + let body_str = String::from_utf8_lossy(&resp.body); + channel_host::log( + channel_host::LogLevel::Warn, + &format!( + "Mention poll: failed to fetch messages for channel {}: {} - {}", + channel_id, resp.status, body_str + ), + ); + return None; + } + + let page: Vec = serde_json::from_slice(&resp.body).ok()?; + let page_len = page.len(); + + if page.is_empty() { + break; + } + + // Discord returns newest-first; find the max ID for the next page cursor + let page_max_id = page + .iter() + .map(|m| m.id.as_str()) + .max_by(|a, b| compare_message_ids(a, b)) + .map(str::to_string); + + all_messages.extend(page); + + if page_len < 100 { + break; + } + + match page_max_id { + Some(max_id) if max_id != after => after = max_id, + _ => break, + } + } + + Some(all_messages) +} + +/// Compare two Discord snowflake IDs. Falls back to lexical comparison. +fn compare_message_ids(a: &str, b: &str) -> std::cmp::Ordering { + match (a.parse::(), b.parse::()) { + (Ok(a_num), Ok(b_num)) => a_num.cmp(&b_num), + _ => a.cmp(b), + } +} + +fn dedup_ids_path(channel_id: &str) -> String { + format!("state/mention_dedup/{}", channel_id) +} + +fn load_recent_processed_ids(channel_id: &str) -> Vec { + channel_host::workspace_read(&dedup_ids_path(channel_id)) + .and_then(|s| serde_json::from_str(&s).ok()) + .unwrap_or_default() +} + +fn save_recent_processed_ids(channel_id: &str, ids: &[String]) { + let json = serde_json::to_string(ids).unwrap_or_else(|_| "[]".to_string()); + let _ = channel_host::workspace_write(&dedup_ids_path(channel_id), &json); +} + +fn remember_processed_id(msg_id: &str, ids: &mut Vec) { + if ids.contains(&msg_id.to_string()) { + return; + } + ids.push(msg_id.to_string()); + if ids.len() > DEDUP_CAP { + let excess = ids.len() - DEDUP_CAP; + ids.drain(0..excess); + } +} + +/// Returns true when `current` is strictly newer than `last_seen`. +fn is_new_message(last_seen: &str, current: &str) -> bool { + compare_message_ids(current, last_seen) == std::cmp::Ordering::Greater +} + +/// Returns true if the message mentions the bot (by mention objects or content). +fn message_mentions_bot(msg: &DiscordChannelMessage, bot_id: &str) -> bool { + if msg.mentions.iter().any(|u| u.id == bot_id) { + return true; + } + let mention = format!("<@{}>", bot_id); + let mention_nick = format!("<@!{}>", bot_id); + msg.content.contains(&mention) || msg.content.contains(&mention_nick) +} + +/// Strip the bot mention prefix from content. +fn strip_bot_mention(content: &str, bot_id: &str) -> String { + let trimmed = content.trim(); + for mention in [format!("<@{}>", bot_id), format!("<@!{}>", bot_id)] { + if let Some(rest) = trimmed.strip_prefix(&mention) { + return rest.trim().to_string(); + } + } + trimmed.to_string() +} + +/// Build JSON headers string with Discord bot authorization. +/// When `include_content_type` is true, includes `Content-Type: application/json`. +fn discord_auth_headers_json(include_content_type: bool) -> String { + if include_content_type { + serde_json::json!({ + "Content-Type": "application/json" + }) + .to_string() + } else { + serde_json::json!({}).to_string() } } -/// Send a broadcast message to a Discord user via DM. -/// -/// Creates a DM channel with the user (Discord caches this, so repeated calls -/// for the same user reuse the existing channel) and then posts the message. +/// Send a DM to a Discord user by opening (or reusing) a DM channel. fn broadcast_dm(user_id: &str, content: &str) -> Result<(), String> { // Validate user_id is a plausible Discord snowflake (numeric, 17-20 digits) // to avoid injecting arbitrary strings into API URLs. @@ -1242,12 +1556,20 @@ fn broadcast_dm(user_id: &str, content: &str) -> Result<(), String> { } let dm_channel: DmChannelResponse = serde_json::from_slice(&dm_response.body) .map_err(|e| format!("Failed to parse DM channel response: {}", e))?; - let channel_id = &dm_channel.id; // Step 2: Send the message to the DM channel. let truncated = truncate_message(content); let payload = serde_json::json!({ "content": truncated }); - send_channel_message(channel_id, payload) + let body = + serde_json::to_vec(&payload).map_err(|e| format!("Failed to serialize: {}", e))?; + send_discord_request( + "POST", + &format!("{DISCORD_API_BASE}/channels/{}/messages", dm_channel.id), + &DiscordHttpRequest { + headers_json: discord_auth_headers_json(true), + body, + }, + ) } fn json_response(status: u16, value: serde_json::Value) -> OutgoingHttpResponse { @@ -1264,17 +1586,12 @@ fn json_response(status: u16, value: serde_json::Value) -> OutgoingHttpResponse export!(DiscordChannel); fn truncate_message(content: &str) -> String { - if content.len() <= 2000 { + if content.chars().count() <= DISCORD_MESSAGE_CHAR_LIMIT { content.to_string() } else { - let max_bytes = 1990; - let cutoff = content - .char_indices() - .map(|(i, c)| i + c.len_utf8()) - .take_while(|&end| end <= max_bytes) - .last() - .unwrap_or(0); - let mut truncated = content[..cutoff].to_string(); + let suffix = "\n... (truncated)"; + let allowed_chars = DISCORD_MESSAGE_CHAR_LIMIT.saturating_sub(suffix.chars().count()); + let mut truncated = content.chars().take(allowed_chars).collect::(); truncated.push_str("\n... (truncated)"); truncated } @@ -1283,7 +1600,8 @@ fn truncate_message(content: &str) -> String { #[cfg(test)] mod tests { use super::*; - use ed25519_dalek::{Signer, SigningKey}; + + const DISCORD_CAPABILITIES_JSON: &str = include_str!("../discord.capabilities.json"); #[test] fn test_truncate_message() { @@ -1292,19 +1610,14 @@ mod tests { let long = "a".repeat(2005); let truncated = truncate_message(&long); - assert_eq!(truncated.len(), 2006); // 1990 + 16 chars suffix + assert_eq!(truncated.chars().count(), 2000); assert!(truncated.ends_with("\n... (truncated)")); // Test with multibyte characters (Euro sign is 3 bytes) - // 1000 chars * 3 bytes = 3000 bytes - let multi = "€".repeat(1000); + let multi = "€".repeat(2005); let truncated_multi = truncate_message(&multi); - // 1990 bytes limit. 1990 / 3 = 663 with remainder 1. - // Should truncate at 663 chars (1989 bytes). - // Suffix is 16 bytes. Total: 1989 + 16 = 2005 bytes. - assert!(truncated_multi.len() <= 2006); - assert!(truncated_multi.len() >= 2006 - 4); // Allow for max utf8 char width variance + assert_eq!(truncated_multi.chars().count(), 2000); assert!(truncated_multi.ends_with("\n... (truncated)")); let content_part = &truncated_multi[..truncated_multi.len() - 16]; @@ -1312,312 +1625,284 @@ mod tests { } #[test] - fn test_metadata_serialization() { - let metadata = DiscordMessageMetadata { - channel_id: "123".into(), - interaction_id: Some("456".into()), - token: Some("abc".into()), - application_id: Some("789".into()), - source_message_id: None, - thread_id: None, - }; - let json = serde_json::to_string(&metadata).unwrap(); - let parsed: DiscordMessageMetadata = serde_json::from_str(&json).unwrap(); - assert_eq!(parsed.channel_id, "123"); - assert_eq!(parsed.interaction_id.as_deref(), Some("456")); - } - - #[test] - fn test_is_new_message() { - assert!(is_new_message(None, "100")); - assert!(is_new_message(Some("100"), "200")); - assert!(!is_new_message(Some("200"), "100")); - assert!(!is_new_message(Some("100"), "100")); - assert!(is_new_message(Some("abc"), "abd")); - assert!(!is_new_message(Some("abd"), "abc")); - } + fn test_reply_plan_uses_character_count_for_attachment_threshold() { + let inline = + build_discord_reply_plan(&test_response(test_metadata_json(), "€".repeat(2000))) + .unwrap(); - #[test] - fn test_strip_bot_mention() { - assert_eq!(strip_bot_mention("<@123> hello", "123"), "hello"); - assert_eq!(strip_bot_mention("<@!123> hello", "123"), "hello"); - assert_eq!(strip_bot_mention("<@123>", "123"), ""); - assert_eq!( - strip_bot_mention("hello <@123> world <@!123>", "123"), - "hello world" - ); + assert!(matches!(inline, DiscordReplyPlan::Inline(_))); } - #[test] - fn test_message_mentions_bot() { - let msg = DiscordChannelMessage { - id: "1".to_string(), - content: "hello <@123>".to_string(), - channel_id: "10".to_string(), - author: DiscordChannelAuthor { - id: "u1".to_string(), - username: "alice".to_string(), - global_name: None, - bot: false, - }, - mentions: vec![], - webhook_id: None, - }; - assert!(message_mentions_bot(&msg, "123")); - assert!(!message_mentions_bot(&msg, "999")); + fn test_response(metadata_json: String, content: String) -> AgentResponse { + AgentResponse { + message_id: "msg-1".to_string(), + content, + thread_id: None, + metadata_json, + attachments: vec![], + } } - #[test] - fn test_message_mentions_bot_via_mentions_array() { - let msg = DiscordChannelMessage { - id: "2".to_string(), - content: "hello".to_string(), - channel_id: "10".to_string(), - author: DiscordChannelAuthor { - id: "u1".to_string(), - username: "alice".to_string(), - global_name: None, - bot: false, - }, - mentions: vec![DiscordUser { - id: "777".to_string(), - username: "bot".to_string(), - global_name: None, - }], - webhook_id: None, - }; - assert!(message_mentions_bot(&msg, "777")); + fn test_metadata_json() -> String { + serde_json::json!({ + "channel_id": "chan-1", + "interaction_id": "int-1", + "token": "tok-1", + "application_id": "app-1", + "thread_id": null, + "embeds": [{"title": "embed title"}] + }) + .to_string() } #[test] - fn test_compare_message_ids_numeric_and_lexical_fallback() { - assert_eq!(compare_message_ids("100", "20"), Ordering::Greater); - assert_eq!(compare_message_ids("20", "100"), Ordering::Less); - assert_eq!(compare_message_ids("abc", "abd"), Ordering::Less); - assert_eq!(compare_message_ids("abd", "abc"), Ordering::Greater); + fn test_reply_plan_threshold_uses_attachment_only_above_2000_chars() { + let inline = + build_discord_reply_plan(&test_response(test_metadata_json(), "a".repeat(2000))) + .unwrap(); + assert!(matches!(inline, DiscordReplyPlan::Inline(_))); + + let attachment = + build_discord_reply_plan(&test_response(test_metadata_json(), "a".repeat(2001))) + .unwrap(); + assert!(matches!(attachment, DiscordReplyPlan::Attachment { .. })); } #[test] - fn test_remember_processed_id_dedup_and_cap() { - let mut ids = Vec::new(); - for i in 0..220 { - remember_processed_id(&mut ids, &format!("{}", i)); - } - assert_eq!(ids.len(), 200); - assert_eq!(ids.first().map(String::as_str), Some("20")); - assert_eq!(ids.last().map(String::as_str), Some("219")); - - remember_processed_id(&mut ids, "219"); - assert_eq!(ids.len(), 200); - assert_eq!(ids.last().map(String::as_str), Some("219")); - } + fn test_reply_plan_preserves_short_message_content_and_embeds() { + let plan = build_discord_reply_plan(&test_response( + test_metadata_json(), + "short reply".to_string(), + )) + .unwrap(); + + let DiscordReplyPlan::Inline(request) = plan else { + panic!("expected inline plan"); + }; - #[test] - fn test_header_case_insensitive() { - let mut headers = HashMap::new(); - headers.insert("X-Signature-Timestamp".to_string(), "123".to_string()); assert_eq!( - header_case_insensitive(&headers, "x-signature-timestamp"), - Some("123") + request.headers_json, + r#"{"Content-Type":"application/json"}"# ); - assert_eq!(header_case_insensitive(&headers, "missing"), None); + + let payload: serde_json::Value = serde_json::from_slice(&request.body).unwrap(); + assert_eq!(payload["content"], "short reply"); + assert_eq!(payload["embeds"][0]["title"], "embed title"); } #[test] - fn test_discord_auth_headers_json_shape() { - let with_ct: serde_json::Value = - serde_json::from_str(&discord_auth_headers_json(true)).unwrap(); - assert_eq!( - with_ct.get("Content-Type").and_then(|v| v.as_str()), - Some("application/json") - ); - assert_eq!( - with_ct.get("Authorization").and_then(|v| v.as_str()), - Some("Bot {DISCORD_BOT_TOKEN}") - ); + fn test_reply_plan_builds_markdown_attachment_multipart_payload() { + let content = "# Heading\n\nA long markdown reply".repeat(80); + let plan = build_discord_reply_plan(&test_response(test_metadata_json(), content.clone())) + .unwrap(); - let no_ct: serde_json::Value = - serde_json::from_str(&discord_auth_headers_json(false)).unwrap(); - assert!(no_ct.get("Content-Type").is_none()); - assert_eq!( - no_ct.get("Authorization").and_then(|v| v.as_str()), - Some("Bot {DISCORD_BOT_TOKEN}") - ); + let DiscordReplyPlan::Attachment { upload, .. } = plan else { + panic!("expected attachment plan"); + }; + + assert!(upload + .headers_json + .contains("multipart/form-data; boundary=")); + + let body = String::from_utf8(upload.body).unwrap(); + assert!(body.contains("name=\"payload_json\"")); + assert!(body.contains("filename=\"response.md\"")); + assert!(body.contains("Content-Type: text/markdown")); + assert!(body.contains(DISCORD_ATTACHMENT_NOTICE)); + assert!(body.contains("embed title")); + assert!(body.contains(&content)); } #[test] - fn test_verify_discord_request_signature_valid() { - let signing_key = SigningKey::from_bytes(&[7u8; 32]); - let public_key_hex = hex::encode(signing_key.verifying_key().to_bytes()); - let timestamp = "1234567890"; - let body = br#"{"type":1}"#; - - let mut signed = Vec::new(); - signed.extend_from_slice(timestamp.as_bytes()); - signed.extend_from_slice(body); - let signature = signing_key.sign(&signed); - - let mut headers = HashMap::new(); - headers.insert( - "x-signature-ed25519".to_string(), - hex::encode(signature.to_bytes()), - ); - headers.insert("x-signature-timestamp".to_string(), timestamp.to_string()); + fn test_reply_plan_uses_dynamic_multipart_boundary() { + let content = "# Heading\n\nA long markdown reply".repeat(80); + + let first = build_discord_reply_plan(&test_response(test_metadata_json(), content.clone())) + .unwrap(); + let second = + build_discord_reply_plan(&test_response(test_metadata_json(), content)).unwrap(); + + let DiscordReplyPlan::Attachment { + upload: first_upload, + .. + } = first + else { + panic!("expected attachment plan"); + }; + let DiscordReplyPlan::Attachment { + upload: second_upload, + .. + } = second + else { + panic!("expected attachment plan"); + }; - assert!(verify_discord_request_signature( - headers, - body, - Some(&public_key_hex) - )); + let first_headers: serde_json::Value = + serde_json::from_str(&first_upload.headers_json).unwrap(); + let second_headers: serde_json::Value = + serde_json::from_str(&second_upload.headers_json).unwrap(); + + let first_boundary = first_headers["Content-Type"] + .as_str() + .unwrap() + .strip_prefix("multipart/form-data; boundary=") + .unwrap(); + let second_boundary = second_headers["Content-Type"] + .as_str() + .unwrap() + .strip_prefix("multipart/form-data; boundary=") + .unwrap(); + + assert!(first_boundary.starts_with(DISCORD_MULTIPART_BOUNDARY)); + assert!(second_boundary.starts_with(DISCORD_MULTIPART_BOUNDARY)); + assert_ne!(first_boundary, second_boundary); + + let first_body = String::from_utf8(first_upload.body).unwrap(); + let second_body = String::from_utf8(second_upload.body).unwrap(); + assert!(first_body.contains(&format!("--{first_boundary}\r\n"))); + assert!(second_body.contains(&format!("--{second_boundary}\r\n"))); } #[test] - fn test_verify_discord_request_signature_tampered_body() { - let signing_key = SigningKey::from_bytes(&[9u8; 32]); - let public_key_hex = hex::encode(signing_key.verifying_key().to_bytes()); - let timestamp = "1234567890"; - let body = b"hello"; - - let mut signed = Vec::new(); - signed.extend_from_slice(timestamp.as_bytes()); - signed.extend_from_slice(body); - let signature = signing_key.sign(&signed); - - let mut headers = HashMap::new(); - headers.insert( - "x-signature-ed25519".to_string(), - hex::encode(signature.to_bytes()), - ); - headers.insert("x-signature-timestamp".to_string(), timestamp.to_string()); + fn test_reply_plan_includes_truncated_text_fallback_for_attachment_failures() { + let content = "a".repeat(2400); + let plan = build_discord_reply_plan(&test_response(test_metadata_json(), content.clone())) + .unwrap(); - assert!(!verify_discord_request_signature( - headers, - b"hello-modified", - Some(&public_key_hex) - )); + let DiscordReplyPlan::Attachment { fallback, .. } = plan else { + panic!("expected attachment plan"); + }; + + let payload: serde_json::Value = serde_json::from_slice(&fallback.body).unwrap(); + assert_eq!(payload["content"], truncate_message(&content)); + assert_eq!(payload["embeds"][0]["title"], "embed title"); } #[test] - fn test_verify_discord_request_signature_wrong_public_key() { - let signing_key = SigningKey::from_bytes(&[11u8; 32]); - let wrong_key = SigningKey::from_bytes(&[12u8; 32]); - let timestamp = "1234567890"; - let body = b"payload"; - - let mut signed = Vec::new(); - signed.extend_from_slice(timestamp.as_bytes()); - signed.extend_from_slice(body); - let signature = signing_key.sign(&signed); - - let mut headers = HashMap::new(); - headers.insert( - "x-signature-ed25519".to_string(), - hex::encode(signature.to_bytes()), - ); - headers.insert("x-signature-timestamp".to_string(), timestamp.to_string()); - - assert!(!verify_discord_request_signature( - headers, - body, - Some(&hex::encode(wrong_key.verifying_key().to_bytes())) - )); + fn test_metadata_serialization() { + let metadata = DiscordMessageMetadata { + channel_id: "123".into(), + interaction_id: "456".into(), + token: "abc".into(), + application_id: "789".into(), + source_message_id: None, + thread_id: None, + }; + let json = serde_json::to_string(&metadata).unwrap(); + let parsed: DiscordMessageMetadata = serde_json::from_str(&json).unwrap(); + assert_eq!(parsed.channel_id, "123"); + assert_eq!(parsed.interaction_id, "456"); } #[test] - fn test_verify_discord_request_signature_missing_headers() { - let headers = HashMap::new(); - assert!(!verify_discord_request_signature( - headers, - b"abc", - Some("00112233445566778899aabbccddeeff00112233445566778899aabbccddeeff") - )); + fn test_metadata_backward_compat_with_old_option_format() { + // Old metadata format used Option for these fields + let old_json = r#"{ + "channel_id": "123", + "interaction_id": null, + "token": null, + "application_id": null, + "thread_id": null + }"#; + let parsed: DiscordMessageMetadata = serde_json::from_str(old_json).unwrap(); + assert_eq!(parsed.channel_id, "123"); + assert!(parsed.interaction_id.is_empty()); + + // Old format without the fields at all + let minimal_json = r#"{"channel_id": "456"}"#; + let parsed: DiscordMessageMetadata = serde_json::from_str(minimal_json).unwrap(); + assert_eq!(parsed.channel_id, "456"); + assert!(parsed.interaction_id.is_empty()); + assert!(parsed.token.is_empty()); + assert!(parsed.application_id.is_empty()); } #[test] - fn test_verify_discord_request_signature_invalid_signature_hex() { - let mut headers = HashMap::new(); - headers.insert("x-signature-ed25519".to_string(), "not-hex".to_string()); - headers.insert( - "x-signature-timestamp".to_string(), - "1234567890".to_string(), + fn test_response_route_uses_webhook_for_interactions() { + let metadata = DiscordMessageMetadata { + channel_id: "123".into(), + interaction_id: "456".into(), + token: "tok".into(), + application_id: "app".into(), + source_message_id: None, + thread_id: None, + }; + + assert_eq!( + response_route_for_metadata(&metadata), + DiscordResponseRoute::InteractionWebhook( + format!("{DISCORD_API_BASE}/webhooks/app/tok/messages/@original") + ) ); - assert!(!verify_discord_request_signature( - headers, - b"abc", - Some("00112233445566778899aabbccddeeff00112233445566778899aabbccddeeff") - )); } #[test] - fn test_verify_discord_request_signature_invalid_public_key_hex() { - let mut headers = HashMap::new(); - headers.insert("x-signature-ed25519".to_string(), "00".repeat(64)); - headers.insert( - "x-signature-timestamp".to_string(), - "1234567890".to_string(), + fn test_response_route_uses_channel_messages_for_gateway_metadata() { + let metadata = DiscordMessageMetadata { + channel_id: "chan-1".into(), + interaction_id: String::new(), + token: String::new(), + application_id: String::new(), + source_message_id: None, + thread_id: None, + }; + + assert_eq!( + response_route_for_metadata(&metadata), + DiscordResponseRoute::ChannelMessage( + format!("{DISCORD_API_BASE}/channels/chan-1/messages") + ) ); - assert!(!verify_discord_request_signature( - headers, - b"abc", - Some("not-hex") - )); } #[test] - fn test_verify_discord_request_signature_invalid_lengths() { - let mut headers = HashMap::new(); - headers.insert("x-signature-ed25519".to_string(), "00".repeat(10)); - headers.insert( - "x-signature-timestamp".to_string(), - "1234567890".to_string(), + fn test_typing_request_url_uses_channel_id_for_thinking_status() { + let update = StatusUpdate { + status: StatusType::Thinking, + message: "Thinking...".to_string(), + metadata_json: serde_json::json!({ + "channel_id": "chan-42", + "interaction_id": "", + "token": "", + "application_id": "", + "thread_id": null + }) + .to_string(), + }; + + assert_eq!( + typing_request_url_for_update(&update), + Some(format!("{DISCORD_API_BASE}/channels/chan-42/typing")) ); - assert!(!verify_discord_request_signature( - headers.clone(), - b"abc", - Some("00".repeat(31).as_str()) - )); - assert!(!verify_discord_request_signature( - headers, - b"abc", - Some("00".repeat(32).as_str()) - )); } #[test] - fn test_verify_discord_request_signature_case_insensitive_headers() { - let signing_key = SigningKey::from_bytes(&[13u8; 32]); - let public_key_hex = hex::encode(signing_key.verifying_key().to_bytes()); - let timestamp = "1234567890"; - let body = b"case-header"; - - let mut signed = Vec::new(); - signed.extend_from_slice(timestamp.as_bytes()); - signed.extend_from_slice(body); - let signature = signing_key.sign(&signed); - - let mut headers = HashMap::new(); - headers.insert( - "X-Signature-Ed25519".to_string(), - hex::encode(signature.to_bytes()), - ); - headers.insert("X-Signature-Timestamp".to_string(), timestamp.to_string()); + fn test_typing_request_url_ignores_non_thinking_status() { + let update = StatusUpdate { + status: StatusType::Done, + message: "Done".to_string(), + metadata_json: serde_json::json!({ + "channel_id": "chan-42", + "interaction_id": "", + "token": "", + "application_id": "", + "thread_id": null + }) + .to_string(), + }; - assert!(verify_discord_request_signature( - headers, - body, - Some(&public_key_hex) - )); + assert_eq!(typing_request_url_for_update(&update), None); } #[test] - fn test_verify_discord_request_signature_empty_public_key() { - let mut headers = HashMap::new(); - headers.insert("x-signature-ed25519".to_string(), "00".repeat(64)); - headers.insert( - "x-signature-timestamp".to_string(), - "1234567890".to_string(), - ); - assert!(!verify_discord_request_signature(headers, b"abc", Some(""))); + fn test_typing_request_url_ignores_invalid_metadata() { + let update = StatusUpdate { + status: StatusType::Thinking, + message: "Thinking...".to_string(), + metadata_json: "not-json".to_string(), + }; + + assert_eq!(typing_request_url_for_update(&update), None); } #[test] @@ -1651,41 +1936,319 @@ mod tests { } #[test] - fn test_broadcast_dm_payload_format() { - // Verify the DM channel creation payload is well-formed JSON that - // Discord's API expects. - let user_id = "123456789012345678"; - let payload = serde_json::json!({ "recipient_id": user_id }); - let serialized = serde_json::to_vec(&payload).unwrap(); - let parsed: serde_json::Value = serde_json::from_slice(&serialized).unwrap(); + fn test_capabilities_default_to_gateway_mode() { + let caps: serde_json::Value = + serde_json::from_str(DISCORD_CAPABILITIES_JSON).expect("capabilities parse"); + let allowlist = caps["capabilities"]["http"]["allowlist"] + .as_array() + .expect("http allowlist array"); + assert_eq!( - parsed.get("recipient_id").and_then(|v| v.as_str()), - Some(user_id) + caps["capabilities"]["channel"]["allow_polling"], + serde_json::Value::Bool(true) + ); + assert!(allowlist.iter().any(|entry| { + entry["host"] == serde_json::Value::String("gateway.discord.gg".to_string()) + && entry["methods"] == serde_json::json!(["GET"]) + })); + assert_eq!( + caps["capabilities"]["websocket"]["url"], + serde_json::Value::String("wss://gateway.discord.gg/?v=10&encoding=json".to_string()) + ); + assert_eq!( + caps["capabilities"]["websocket"]["connect_on_start"], + serde_json::Value::Bool(true) + ); + assert_eq!( + caps["capabilities"]["websocket"]["identify_secret_name"], + serde_json::Value::String("discord_bot_token".to_string()) + ); + assert_eq!( + caps["capabilities"]["websocket"]["identify"]["intents"], + serde_json::Value::Number(4609u64.into()) ); } #[test] - fn test_broadcast_message_truncation() { - // Broadcast uses truncate_message, verify it handles content within - // Discord's 2000-char limit for DMs. - let short = "Hello from broadcast"; - assert_eq!(truncate_message(short), short); + fn test_parse_gateway_event_queue_emits_message_create_after_ready() { + let queue_json = serde_json::json!([ + serde_json::json!({ + "op": 0, + "t": "READY", + "d": { + "user": { + "id": "bot-1", + "username": "ironclaw", + "global_name": "IronClaw", + "bot": true + } + } + }) + .to_string(), + serde_json::json!({ + "op": 0, + "t": "MESSAGE_CREATE", + "d": { + "channel_id": "chan-1", + "guild_id": "guild-1", + "content": "<@bot-1> hello from discord", + "author": { + "id": "user-1", + "username": "alice", + "global_name": "Alice", + "bot": false + } + } + }) + .to_string() + ]) + .to_string(); + + let result = parse_gateway_event_queue(&queue_json, None); + + assert_eq!(result.bot_user_id.as_deref(), Some("bot-1")); + assert_eq!( + result.messages, + vec![ParsedGatewayMessage { + user_id: "user-1".to_string(), + user_name: "Alice".to_string(), + channel_id: "chan-1".to_string(), + content: "hello from discord".to_string(), + is_dm: false, + }] + ); + } + + #[test] + fn test_parse_gateway_event_queue_ignores_bot_and_unmentioned_guild_messages() { + let queue_json = serde_json::json!([ + serde_json::json!({ + "op": 0, + "t": "MESSAGE_CREATE", + "d": { + "channel_id": "chan-1", + "guild_id": "guild-1", + "content": "this should not trigger", + "author": { + "id": "user-1", + "username": "alice", + "global_name": "Alice", + "bot": false + } + } + }) + .to_string(), + serde_json::json!({ + "op": 0, + "t": "MESSAGE_CREATE", + "d": { + "channel_id": "dm-1", + "content": "bot echo", + "author": { + "id": "bot-1", + "username": "ironclaw", + "global_name": "IronClaw", + "bot": true + } + } + }) + .to_string(), + serde_json::json!({ + "op": 0, + "t": "MESSAGE_CREATE", + "d": { + "channel_id": "dm-2", + "content": "direct message", + "author": { + "id": "user-2", + "username": "bob", + "global_name": null, + "bot": false + } + } + }) + .to_string() + ]) + .to_string(); + + let result = parse_gateway_event_queue(&queue_json, Some("bot-1")); + + assert_eq!(result.bot_user_id.as_deref(), Some("bot-1")); + assert_eq!( + result.messages, + vec![ParsedGatewayMessage { + user_id: "user-2".to_string(), + user_name: "bob".to_string(), + channel_id: "dm-2".to_string(), + content: "direct message".to_string(), + is_dm: true, + }] + ); + } + + #[test] + fn test_non_gateway_dm_pairing_behavior_is_unchanged() { + assert!(should_apply_dm_pairing(PermissionSource::Webhook, true)); + assert!(!should_apply_dm_pairing(PermissionSource::Webhook, false)); + } + + #[test] + fn test_gateway_dm_pairing_behavior_matches_webhook_dm() { + assert!(should_apply_dm_pairing(PermissionSource::Gateway, true)); + assert!(!should_apply_dm_pairing(PermissionSource::Gateway, false)); + } + + #[test] + fn test_pairing_reply_route_uses_channel_messages_for_gateway_metadata() { + let route = pairing_reply_route(&PairingReplyCtx { + channel_id: "chan-1".to_string(), + application_id: String::new(), + token: String::new(), + }); + + assert_eq!( + route, + DiscordResponseRoute::ChannelMessage( + format!("{DISCORD_API_BASE}/channels/chan-1/messages") + ) + ); + } + + #[test] + fn test_pairing_reply_route_uses_webhook_for_interactions() { + let route = pairing_reply_route(&PairingReplyCtx { + channel_id: "chan-1".to_string(), + application_id: "app-1".to_string(), + token: "tok-1".to_string(), + }); + + assert_eq!( + route, + DiscordResponseRoute::InteractionWebhook( + format!("{DISCORD_API_BASE}/webhooks/app-1/tok-1") + ) + ); + } + + // ====================================================================== + // Mention polling tests + // ====================================================================== + + #[test] + fn test_is_new_message() { + assert!(is_new_message("100", "200")); + assert!(!is_new_message("200", "100")); + assert!(!is_new_message("100", "100")); + // Large snowflake-like IDs + assert!(is_new_message("1234567890123456789", "1234567890123456790")); + assert!(!is_new_message( + "1234567890123456790", + "1234567890123456789" + )); + } + + #[test] + fn test_strip_bot_mention() { + assert_eq!( + strip_bot_mention("<@bot-123> hello world", "bot-123"), + "hello world" + ); + assert_eq!( + strip_bot_mention("<@!bot-123> hi there", "bot-123"), + "hi there" + ); + // No mention prefix — return content as-is + assert_eq!( + strip_bot_mention("no mention here", "bot-123"), + "no mention here" + ); + // Only mention, no content after stripping + assert_eq!(strip_bot_mention("<@bot-123>", "bot-123"), ""); + assert_eq!(strip_bot_mention("<@bot-123> ", "bot-123"), ""); + } + + #[test] + fn test_message_mentions_bot() { + // Via mentions array + let msg = DiscordChannelMessage { + id: "1".to_string(), + content: "hello".to_string(), + channel_id: "ch-1".to_string(), + author: DiscordChannelAuthor { + id: "user-1".to_string(), + username: "alice".to_string(), + global_name: None, + bot: false, + }, + mentions: vec![DiscordUser { + id: "bot-1".to_string(), + username: "ironclaw".to_string(), + global_name: None, + }], + webhook_id: None, + }; + assert!(message_mentions_bot(&msg, "bot-1")); + assert!(!message_mentions_bot(&msg, "other-bot")); + + // Via content + let msg2 = DiscordChannelMessage { + id: "2".to_string(), + content: "<@bot-2> do something".to_string(), + channel_id: "ch-1".to_string(), + author: DiscordChannelAuthor { + id: "user-1".to_string(), + username: "alice".to_string(), + global_name: None, + bot: false, + }, + mentions: vec![], + webhook_id: None, + }; + assert!(message_mentions_bot(&msg2, "bot-2")); + assert!(!message_mentions_bot(&msg2, "other-bot")); + } + + #[test] + fn test_compare_message_ids() { + use std::cmp::Ordering; + assert_eq!(compare_message_ids("100", "200"), Ordering::Less); + assert_eq!(compare_message_ids("200", "100"), Ordering::Greater); + assert_eq!(compare_message_ids("100", "100"), Ordering::Equal); + // Non-numeric fallback + assert_eq!(compare_message_ids("abc", "abd"), Ordering::Less); + assert_eq!(compare_message_ids("abd", "abc"), Ordering::Greater); + } + + #[test] + fn test_remember_processed_id_dedup_and_cap() { + let mut ids = Vec::new(); + + // Basic add + remember_processed_id("msg-1", &mut ids); + assert_eq!(ids, vec!["msg-1".to_string()]); - let long = "x".repeat(2500); - let result = truncate_message(&long); - assert!(result.len() <= 2006); // 1990 content + 16 suffix - assert!(result.ends_with("\n... (truncated)")); + // Duplicate is ignored + remember_processed_id("msg-1", &mut ids); + assert_eq!(ids.len(), 1); + + // Fill beyond DEDUP_CAP + for i in 2..=(DEDUP_CAP + 5) { + remember_processed_id(&format!("msg-{}", i), &mut ids); + } + assert_eq!(ids.len(), DEDUP_CAP); + // Oldest entries should have been drained + assert!(!ids.contains(&"msg-1".to_string())); + assert!(ids.contains(&format!("msg-{}", DEDUP_CAP + 5))); } #[test] - fn test_broadcast_dm_validates_snowflake() { - // broadcast_dm rejects invalid Discord snowflake IDs before making - // any API calls. We can call it directly since invalid IDs are - // rejected before any host function is invoked. - assert!(broadcast_dm("", "hi").is_err()); - assert!(broadcast_dm("abc", "hi").is_err()); - assert!(broadcast_dm("12345", "hi").is_err()); // too short - assert!(broadcast_dm("123456789012345678901", "hi").is_err()); // too long - assert!(broadcast_dm("12345678901234567x", "hi").is_err()); // non-digit + fn test_discord_auth_headers_json_shape() { + let with_ct = discord_auth_headers_json(true); + let parsed: serde_json::Value = serde_json::from_str(&with_ct).unwrap(); + assert_eq!(parsed["Content-Type"], "application/json"); + + let without_ct = discord_auth_headers_json(false); + let parsed: serde_json::Value = serde_json::from_str(&without_ct).unwrap(); + assert!(parsed.get("Content-Type").is_none()); } } diff --git a/src/channels/wasm/host.rs b/src/channels/wasm/host.rs index eeaccb20561..e8a59a88c1d 100644 --- a/src/channels/wasm/host.rs +++ b/src/channels/wasm/host.rs @@ -539,6 +539,65 @@ impl ChannelWorkspaceStore { } } } + + /// Append a text frame to a JSON queue stored at `path`. + /// + /// The queue is stored as a JSON array of strings and bounded to the most + /// recent `max_items` entries so websocket runtimes cannot grow it without + /// limit. + pub fn append_json_text_queue( + &self, + path: &str, + text: &str, + max_items: usize, + ) -> Result<(), String> { + let mut data = self + .data + .write() + .map_err(|_| "workspace store lock poisoned".to_string())?; + + let mut queue: Vec = data + .get(path) + .and_then(|raw| serde_json::from_str(raw).ok()) + .unwrap_or_default(); + + queue.push(text.to_string()); + + if queue.len() > max_items { + let overflow = queue.len() - max_items; + queue.drain(0..overflow); + } + + let serialized = serde_json::to_string(&queue) + .map_err(|error| format!("failed to serialize websocket queue: {error}"))?; + + data.insert(path.to_string(), serialized); + Ok(()) + } + + /// Atomically move a queued JSON array of text frames from `source_path` to `dest_path`. + pub fn move_json_text_queue(&self, source_path: &str, dest_path: &str) -> Result { + let mut data = self + .data + .write() + .map_err(|_| "workspace store lock poisoned".to_string())?; + + let Some(raw_queue) = data.remove(source_path) else { + data.remove(dest_path); + return Ok(false); + }; + + let queue: Vec = serde_json::from_str(&raw_queue) + .map_err(|error| format!("failed to deserialize websocket queue: {error}"))?; + + if queue.is_empty() { + data.remove(dest_path); + return Ok(false); + } + + data.insert(dest_path.to_string(), raw_queue); + Ok(true) + } } impl crate::tools::wasm::WorkspaceReader for ChannelWorkspaceStore { @@ -818,6 +877,51 @@ mod tests { ); } + #[test] + fn test_channel_workspace_store_append_json_text_queue_is_bounded() { + use crate::channels::wasm::host::ChannelWorkspaceStore; + use crate::tools::wasm::WorkspaceReader; + + let store = ChannelWorkspaceStore::new(); + let path = "channels/discord/state/gateway_event_queue"; + + store.append_json_text_queue(path, "frame-1", 2).unwrap(); + store.append_json_text_queue(path, "frame-2", 2).unwrap(); + store.append_json_text_queue(path, "frame-3", 2).unwrap(); + + let queue: Vec = serde_json::from_str(&store.read(path).unwrap()).unwrap(); + assert_eq!(queue, vec!["frame-2".to_string(), "frame-3".to_string()]); + } + + #[test] + fn test_channel_workspace_store_move_json_text_queue_is_atomic() { + use crate::channels::wasm::host::ChannelWorkspaceStore; + use crate::tools::wasm::WorkspaceReader; + + let store = ChannelWorkspaceStore::new(); + let live_path = "channels/discord/state/gateway_event_queue"; + let drain_path = "channels/discord/state/gateway_event_queue_processing"; + + store + .append_json_text_queue(live_path, "frame-1", 4) + .unwrap(); + store + .append_json_text_queue(live_path, "frame-2", 4) + .unwrap(); + + assert!(store.move_json_text_queue(live_path, drain_path).unwrap()); + assert_eq!(store.read(live_path), None); + + let drained: Vec = serde_json::from_str(&store.read(drain_path).unwrap()).unwrap(); + assert_eq!(drained, vec!["frame-1".to_string(), "frame-2".to_string()]); + + store + .append_json_text_queue(live_path, "frame-3", 4) + .unwrap(); + let live: Vec = serde_json::from_str(&store.read(live_path).unwrap()).unwrap(); + assert_eq!(live, vec!["frame-3".to_string()]); + } + // === QA Plan P2 - 2.3: WASM channel lifecycle tests === #[test] diff --git a/src/channels/wasm/wrapper.rs b/src/channels/wasm/wrapper.rs index a0f9689f0a7..d3feb2318c8 100644 --- a/src/channels/wasm/wrapper.rs +++ b/src/channels/wasm/wrapper.rs @@ -33,8 +33,10 @@ use std::sync::Arc; use std::time::Duration; use async_trait::async_trait; -use tokio::sync::{RwLock, mpsc, oneshot}; +use futures::{SinkExt, StreamExt}; +use tokio::sync::{Mutex, RwLock, mpsc, oneshot}; use tokio_stream::wrappers::ReceiverStream; +use tokio_tungstenite::tungstenite::protocol::Message as WebsocketMessage; use uuid::Uuid; use wasmtime::Store; use wasmtime::component::Linker; @@ -59,6 +61,10 @@ use crate::tools::wasm::credential_injector::{ InjectedCredentials, host_matches_pattern, inject_credential, }; +const WEBSOCKET_EVENT_QUEUE_RELATIVE_PATH: &str = "state/gateway_event_queue"; +const WEBSOCKET_EVENT_PROCESSING_QUEUE_RELATIVE_PATH: &str = "state/gateway_event_queue_processing"; +const WEBSOCKET_EVENT_QUEUE_MAX_ITEMS: usize = 100; + // Generate component model bindings from the WIT file wasmtime::component::bindgen!({ path: "wit/channel.wit", @@ -691,6 +697,12 @@ pub struct WasmChannel { /// Polling shutdown signal sender (keeps polling alive while held). poll_shutdown_tx: RwLock>>, + /// Websocket runtime shutdown signal sender. + websocket_shutdown_tx: RwLock>>, + + /// Serializes websocket-triggered poll executions. + websocket_poll_lock: Arc>, + /// Registered HTTP endpoints. endpoints: RwLock>, @@ -839,6 +851,8 @@ impl WasmChannel { rate_limiter: Arc::new(RwLock::new(rate_limiter)), shutdown_tx: RwLock::new(None), poll_shutdown_tx: RwLock::new(None), + websocket_shutdown_tx: RwLock::new(None), + websocket_poll_lock: Arc::new(Mutex::new(())), endpoints: RwLock::new(Vec::new()), credentials: Arc::new(RwLock::new(HashMap::new())), typing_task: RwLock::new(None), @@ -1070,6 +1084,228 @@ impl WasmChannel { Ok(()) } + fn start_websocket_runtime( + &self, + config: WebsocketRuntimeConfig, + shutdown_rx: oneshot::Receiver<()>, + ) { + let channel_name = self.name.clone(); + let runtime = Arc::clone(&self.runtime); + let prepared = Arc::clone(&self.prepared); + let capabilities = Self::inject_workspace_reader(&self.capabilities, &self.workspace_store); + let poll_capabilities = self.capabilities.clone(); + let message_tx = self.message_tx.clone(); + let rate_limiter = self.rate_limiter.clone(); + let credentials = self.credentials.clone(); + let pairing_store = self.pairing_store.clone(); + let callback_timeout = self.runtime.config().callback_timeout; + let workspace_store = self.workspace_store.clone(); + let last_broadcast_metadata = self.last_broadcast_metadata.clone(); + let settings_store = self.settings_store.clone(); + let owner_scope_id = self.owner_scope_id.clone(); + let owner_actor_id = self.owner_actor_id.clone(); + let websocket_secrets_store = self.secrets_store.clone(); + let websocket_poll_lock = Arc::clone(&self.websocket_poll_lock); + + tokio::spawn(async move { + let mut shutdown = std::pin::pin!(shutdown_rx); + let mut reconnect_attempt = 0u32; + let (outbound_tx, mut outbound_rx) = mpsc::unbounded_channel::(); + + tracing::info!( + channel = %channel_name, + url = %config.url, + "Starting websocket runtime" + ); + let queue_path = websocket_queue_path(&channel_name); + let processing_queue_path = websocket_processing_queue_path(&channel_name); + let identify_payload = + resolve_websocket_identify_message(&config, websocket_secrets_store.as_deref()) + .await; + let mut session_state = WebsocketSessionState::new(identify_payload.as_deref()); + + 'reconnect: loop { + let connect_url = session_state.connect_url(&config.url); + let connect_result = tokio_tungstenite::connect_async(connect_url).await; + let (stream, _) = match connect_result { + Ok(parts) => { + reconnect_attempt = 0; + tracing::info!(channel = %channel_name, "Websocket runtime connected"); + parts + } + Err(error) => { + let backoff = websocket_reconnect_backoff(reconnect_attempt); + reconnect_attempt = reconnect_attempt.saturating_add(1); + tracing::warn!( + channel = %channel_name, + url = %config.url, + error = %error, + backoff_secs = backoff.as_secs(), + "Websocket runtime connection failed; retrying" + ); + tokio::select! { + _ = tokio::time::sleep(backoff) => continue 'reconnect, + _ = &mut shutdown => { + tracing::info!(channel = %channel_name, "Stopping websocket runtime"); + break 'reconnect; + } + } + } + }; + + let (mut write, mut read) = stream.split(); + let mut next_heartbeat: Option>> = None; + session_state.reset_connection(); + + loop { + tokio::select! { + _ = async { + if let Some(sleep) = next_heartbeat.as_mut() { + sleep.as_mut().await; + } else { + std::future::pending::<()>().await; + } + } => { + if let Some(payload) = build_websocket_heartbeat_message(session_state.last_sequence.clone()) + && let Err(error) = write.send(WebsocketMessage::Text(payload.into())).await + { + tracing::warn!(channel = %channel_name, error = %error, "Websocket heartbeat send failed"); + break; + } + + next_heartbeat = session_state.heartbeat_interval_ms + .map(|interval_ms| Box::pin(tokio::time::sleep(websocket_heartbeat_sleep_duration(interval_ms)))); + } + outbound = outbound_rx.recv() => { + if let Some(payload) = outbound + && let Err(error) = write.send(WebsocketMessage::Text(payload.into())).await + { + tracing::warn!(channel = %channel_name, error = %error, "Websocket outbound control send failed"); + break; + } + } + _ = &mut shutdown => { + tracing::info!(channel = %channel_name, "Stopping websocket runtime"); + break 'reconnect; + } + message = read.next() => { + match message { + Some(Ok(WebsocketMessage::Text(text))) => { + log_websocket_diagnostic(&channel_name, &WebsocketMessage::Text(text.clone())); + let text = text.to_string(); + + let actions = session_state.process_text_frame( + &text, + &channel_name, + identify_payload.as_deref(), + workspace_store.as_ref(), + pairing_store.as_ref(), + ); + + let mut should_break = false; + let mut should_reconnect = false; + for action in actions { + match action { + WebsocketFrameAction::SetHeartbeat { interval_ms } => { + next_heartbeat = Some(Box::pin(tokio::time::sleep( + websocket_heartbeat_sleep_duration(interval_ms), + ))); + } + WebsocketFrameAction::Send(payload) => { + if let Err(error) = write.send(WebsocketMessage::Text(payload.into())).await { + tracing::warn!(channel = %channel_name, error = %error, "Websocket send failed"); + should_break = true; + break; + } + } + WebsocketFrameAction::Enqueue(raw_text) => { + if let Err(error) = workspace_store.append_json_text_queue( + &queue_path, + &raw_text, + WEBSOCKET_EVENT_QUEUE_MAX_ITEMS, + ) { + tracing::warn!(channel = %channel_name, error = %error, "Failed to enqueue websocket text frame"); + continue; + } + + if let Ok(poll_guard) = Arc::clone(&websocket_poll_lock).try_lock_owned() { + spawn_websocket_poll( + poll_guard, + WebsocketPollContext { + channel_name: channel_name.clone(), + runtime: Arc::clone(&runtime), + prepared: Arc::clone(&prepared), + capabilities: capabilities.clone(), + poll_capabilities: poll_capabilities.clone(), + credentials: Arc::clone(&credentials), + pairing_store: pairing_store.clone(), + workspace_store: workspace_store.clone(), + message_tx: message_tx.clone(), + rate_limiter: Arc::clone(&rate_limiter), + last_broadcast_metadata: Arc::clone(&last_broadcast_metadata), + settings_store: settings_store.clone(), + owner_scope_id: owner_scope_id.clone(), + owner_actor_id: owner_actor_id.clone(), + secrets_store: websocket_secrets_store.clone(), + outbound_tx: outbound_tx.clone(), + queue_path: queue_path.clone(), + processing_queue_path: processing_queue_path.clone(), + callback_timeout, + }, + ); + } + } + WebsocketFrameAction::InvalidateAndReconnect => { + should_reconnect = true; + break; + } + } + } + if should_reconnect { + break; + } + if should_break { + break; + } + } + Some(Ok(other)) => { + log_websocket_diagnostic(&channel_name, &other); + } + Some(Err(error)) => { + tracing::warn!( + channel = %channel_name, + error = %error, + "Websocket runtime receive error" + ); + break; + } + None => { + tracing::info!(channel = %channel_name, "Websocket runtime closed by peer"); + break; + } + } + } + } + } + + let backoff = websocket_reconnect_backoff(reconnect_attempt); + reconnect_attempt = reconnect_attempt.saturating_add(1); + tracing::info!( + channel = %channel_name, + backoff_secs = backoff.as_secs(), + "Websocket runtime disconnected; reconnect scheduled" + ); + tokio::select! { + _ = tokio::time::sleep(backoff) => {} + _ = &mut shutdown => { + tracing::info!(channel = %channel_name, "Stopping websocket runtime"); + break 'reconnect; + } + } + } + }); + } + /// Create a fresh store configured for WASM execution. fn create_store( runtime: &WasmChannelRuntime, @@ -1504,6 +1740,8 @@ impl WasmChannel { let channel_name = self.name.clone(); match result { Ok(Ok(((), mut host_state))) => { + let _ = drain_guest_logs(&channel_name, "on_poll", &mut host_state); + // Process emitted messages let emitted = host_state.take_emitted_messages(); self.process_emitted_messages(emitted).await?; @@ -2439,6 +2677,7 @@ impl WasmChannel { match result { Ok(Ok(mut host_state)) => { + let _ = drain_guest_logs(channel_name, "on_poll", &mut host_state); let emitted = host_state.take_emitted_messages(); tracing::debug!( channel = %channel_name, @@ -2660,6 +2899,15 @@ impl Channel for WasmChannel { self.start_polling(Duration::from_millis(interval as u64), poll_shutdown_rx); } + if let Some(websocket_config) = + WebsocketRuntimeConfig::from_capabilities(&self.capabilities) + && websocket_config.connect_on_start + { + let (websocket_shutdown_tx, websocket_shutdown_rx) = oneshot::channel(); + *self.websocket_shutdown_tx.write().await = Some(websocket_shutdown_tx); + self.start_websocket_runtime(websocket_config, websocket_shutdown_rx); + } + tracing::info!( channel = %self.name, display_name = %config.display_name, @@ -2776,6 +3024,9 @@ impl Channel for WasmChannel { // Stop polling by dropping the sender (receiver will complete) let _ = self.poll_shutdown_tx.write().await.take(); + // Stop websocket runtime by dropping the sender (receiver will complete) + let _ = self.websocket_shutdown_tx.write().await.take(); + // Clear the message sender *self.message_tx.write().await = None; @@ -2788,6 +3039,564 @@ impl Channel for WasmChannel { } } +#[derive(Debug, Clone, PartialEq)] +pub(crate) struct WebsocketRuntimeConfig { + pub(crate) url: String, + pub(crate) connect_on_start: bool, + pub(crate) identify: Option, + pub(crate) identify_secret_name: Option, +} + +impl WebsocketRuntimeConfig { + pub(crate) fn from_capabilities(capabilities: &ChannelCapabilities) -> Option { + let raw = capabilities.tool_capabilities.websocket.as_ref()?; + let url = raw.get("url")?.as_str()?.trim(); + if url.is_empty() { + return None; + } + + let parsed = url::Url::parse(url).ok()?; + let scheme = parsed.scheme(); + if scheme != "ws" && scheme != "wss" { + return None; + } + + let host = parsed.host_str()?; + let path = parsed.path(); + let http = capabilities.tool_capabilities.http.as_ref()?; + if !http + .allowlist + .iter() + .any(|pattern| pattern.matches(host, path, "GET")) + { + return None; + } + + Some(Self { + url: url.to_string(), + connect_on_start: raw + .get("connect_on_start") + .and_then(serde_json::Value::as_bool) + .unwrap_or(false), + identify: raw.get("identify").cloned(), + identify_secret_name: raw + .get("identify_secret_name") + .and_then(serde_json::Value::as_str) + .map(ToOwned::to_owned), + }) + } +} + +fn websocket_queue_path(channel_name: &str) -> String { + format!("channels/{channel_name}/{WEBSOCKET_EVENT_QUEUE_RELATIVE_PATH}") +} + +fn websocket_processing_queue_path(channel_name: &str) -> String { + format!("channels/{channel_name}/{WEBSOCKET_EVENT_PROCESSING_QUEUE_RELATIVE_PATH}") +} + +async fn resolve_websocket_identify_message( + config: &WebsocketRuntimeConfig, + store: Option<&(dyn SecretsStore + Send + Sync)>, +) -> Option { + let identify = config.identify.clone()?; + let secret_name = config.identify_secret_name.as_ref()?; + let store = store?; + let secret = store.get_decrypted("default", secret_name).await.ok()?; + build_websocket_identify_message(&identify, secret.expose()) +} + +fn build_websocket_identify_message(identify: &serde_json::Value, token: &str) -> Option { + let mut payload = identify.as_object()?.clone(); + payload.insert( + "token".to_string(), + serde_json::Value::String(token.to_string()), + ); + + serde_json::to_string(&serde_json::json!({ + "op": 2, + "d": serde_json::Value::Object(payload), + })) + .ok() +} + +fn build_websocket_heartbeat_message(sequence: Option) -> Option { + serde_json::to_string(&serde_json::json!({ + "op": 1, + "d": sequence.unwrap_or(serde_json::Value::Null), + })) + .ok() +} + +fn build_discord_gateway_presence_update(status: &str) -> Option { + serde_json::to_string(&serde_json::json!({ + "op": 3, + "d": { + "since": serde_json::Value::Null, + "activities": [], + "status": status, + "afk": false + } + })) + .ok() +} + +fn build_gateway_presence_update( + channel_name: &str, + workspace_store: &crate::channels::wasm::host::ChannelWorkspaceStore, + pairing_store: &PairingStore, +) -> Option { + if channel_name != "discord" { + return None; + } + + build_discord_gateway_presence_update(discord_gateway_presence_status( + channel_name, + workspace_store, + pairing_store, + )) +} + +fn discord_gateway_presence_status( + channel_name: &str, + workspace_store: &crate::channels::wasm::host::ChannelWorkspaceStore, + pairing_store: &PairingStore, +) -> &'static str { + use crate::tools::wasm::WorkspaceReader; + + let owner_key = format!("channels/{}/state/owner_id", channel_name); + if workspace_store + .read(&owner_key) + .filter(|s| !s.is_empty()) + .is_some() + { + return "online"; + } + + if pairing_store + .read_allow_from(channel_name) + .ok() + .is_some_and(|v| !v.is_empty()) + { + return "online"; + } + + "dnd" +} + +fn parse_websocket_hello_heartbeat_interval_ms(text: &str) -> Option { + let payload: serde_json::Value = serde_json::from_str(text).ok()?; + if payload.get("op")?.as_u64()? != 10 { + return None; + } + + payload.get("d")?.get("heartbeat_interval")?.as_u64() +} + +fn websocket_reconnect_backoff(attempt: u32) -> Duration { + use rand::Rng; + + let exponent = attempt.min(6); + let base_ms = (1u64 << exponent) * 1_000; + // Add 0-25% jitter per Discord's reconnection recommendations to avoid + // thundering-herd when many bots reconnect after a Discord deploy. + let jitter_ms = rand::thread_rng().gen_range(0..=base_ms / 4); + Duration::from_millis(base_ms + jitter_ms) +} + +fn websocket_heartbeat_sleep_duration(interval_ms: u64) -> Duration { + Duration::from_millis(interval_ms.max(1)) +} + +fn should_warn_on_heartbeat_interval(interval_ms: u64) -> bool { + interval_ms < 1_000 +} + +fn parse_websocket_sequence(text: &str) -> Option { + let payload: serde_json::Value = serde_json::from_str(text).ok()?; + payload.get("s")?.as_u64() +} + +fn parse_websocket_ready_session(text: &str) -> Option<(String, Option)> { + let payload: serde_json::Value = serde_json::from_str(text).ok()?; + if payload.get("op")?.as_u64()? != 0 { + return None; + } + if payload.get("t")?.as_str()? != "READY" { + return None; + } + let d = payload.get("d")?; + let sid = d.get("session_id")?.as_str()?.to_string(); + let resume_url = d + .get("resume_gateway_url") + .and_then(|v| v.as_str()) + .map(ToOwned::to_owned); + Some((sid, resume_url)) +} + +fn build_websocket_resume_message( + token: &str, + session_id: &str, + sequence: Option<&serde_json::Value>, +) -> Option { + serde_json::to_string(&serde_json::json!({ + "op": 6, + "d": { + "token": token, + "session_id": session_id, + "seq": sequence.cloned().unwrap_or(serde_json::Value::Null), + } + })) + .ok() +} + +fn parse_websocket_invalid_session(text: &str) -> Option { + let payload: serde_json::Value = serde_json::from_str(text).ok()?; + if payload.get("op")?.as_u64()? != 9 { + return None; + } + Some(payload.get("d")?.as_bool().unwrap_or(false)) +} + +fn extract_token_from_identify_payload(identify_payload: &str) -> Option { + let payload: serde_json::Value = serde_json::from_str(identify_payload).ok()?; + payload + .get("d")? + .get("token")? + .as_str() + .map(ToOwned::to_owned) +} + +fn drain_guest_logs( + channel_name: &str, + callback: &str, + host_state: &mut ChannelHostState, +) -> Vec { + let entries = host_state.take_logs(); + + for entry in &entries { + match entry.level { + crate::tools::wasm::LogLevel::Error => { + tracing::error!(channel = %channel_name, callback = callback, "{}", entry.message); + } + crate::tools::wasm::LogLevel::Warn => { + tracing::warn!(channel = %channel_name, callback = callback, "{}", entry.message); + } + crate::tools::wasm::LogLevel::Info => { + tracing::info!(channel = %channel_name, callback = callback, "{}", entry.message); + } + crate::tools::wasm::LogLevel::Debug => { + tracing::debug!(channel = %channel_name, callback = callback, "{}", entry.message); + } + crate::tools::wasm::LogLevel::Trace => { + tracing::trace!(channel = %channel_name, callback = callback, "{}", entry.message); + } + } + } + + entries +} + +/// Shared state for websocket-triggered poll tasks. +/// +/// Groups the many `Arc` handles needed by [`spawn_websocket_poll`] into a +/// single cloneable context so the call site stays readable. +struct WebsocketPollContext { + channel_name: String, + runtime: Arc, + prepared: Arc, + capabilities: ChannelCapabilities, + poll_capabilities: ChannelCapabilities, + credentials: Arc>>, + pairing_store: Arc, + workspace_store: Arc, + message_tx: Arc>>>, + rate_limiter: Arc>, + last_broadcast_metadata: Arc>>, + settings_store: Option>, + owner_scope_id: String, + owner_actor_id: Option, + secrets_store: Option>, + outbound_tx: mpsc::UnboundedSender, + queue_path: String, + processing_queue_path: String, + callback_timeout: Duration, +} + +/// Spawn the websocket-triggered poll task. +/// +/// Extracted from the select loop to reduce nesting. Moves items from the +/// event queue to a processing queue and runs the WASM `on_poll` callback. +fn spawn_websocket_poll(poll_guard: tokio::sync::OwnedMutexGuard<()>, ctx: WebsocketPollContext) { + tokio::spawn(async move { + let _poll_guard = poll_guard; + + loop { + let moved = match ctx + .workspace_store + .move_json_text_queue(&ctx.queue_path, &ctx.processing_queue_path) + { + Ok(value) => value, + Err(error) => { + tracing::warn!(channel = %ctx.channel_name, error = %error, "Failed to snapshot websocket queue for polling"); + break; + } + }; + + if !moved { + break; + } + + let host_credentials = resolve_channel_host_credentials( + &ctx.poll_capabilities, + ctx.secrets_store.as_deref(), + &ctx.owner_scope_id, + ) + .await; + + match WasmChannel::execute_poll( + &ctx.channel_name, + &ctx.runtime, + &ctx.prepared, + &ctx.capabilities, + &ctx.credentials, + host_credentials, + ctx.pairing_store.clone(), + ctx.callback_timeout, + &ctx.workspace_store, + ) + .await + { + Ok(emitted_messages) => { + if !emitted_messages.is_empty() + && let Err(error) = WasmChannel::dispatch_emitted_messages( + EmitDispatchContext { + channel_name: &ctx.channel_name, + owner_scope_id: &ctx.owner_scope_id, + owner_actor_id: ctx.owner_actor_id.as_deref(), + message_tx: &ctx.message_tx, + rate_limiter: &ctx.rate_limiter, + last_broadcast_metadata: &ctx.last_broadcast_metadata, + settings_store: ctx.settings_store.as_ref(), + }, + emitted_messages, + ) + .await + { + tracing::warn!(channel = %ctx.channel_name, error = %error, "Failed to dispatch emitted websocket poll messages"); + } + } + Err(error) => { + tracing::warn!(channel = %ctx.channel_name, error = %error, "Websocket-triggered poll failed"); + } + } + + if let Some(payload) = build_gateway_presence_update( + &ctx.channel_name, + ctx.workspace_store.as_ref(), + ctx.pairing_store.as_ref(), + ) { + let _ = ctx.outbound_tx.send(payload); + } + } + }); +} + +/// Actions produced by websocket text frame processing. +/// +/// Returned from [`WebsocketSessionState::process_text_frame`] so the caller +/// can perform the actual I/O (send messages, break loops) while keeping the +/// parsing logic synchronous and testable. +enum WebsocketFrameAction { + /// Update the heartbeat timer to fire after `interval_ms` milliseconds. + SetHeartbeat { interval_ms: u64 }, + /// Send a text payload over the websocket. + Send(String), + /// Enqueue the raw text into the workspace event queue. + Enqueue(String), + /// Clear session state and reconnect with a fresh identify. + InvalidateAndReconnect, +} + +/// Tracks websocket session state across reconnects. +/// +/// Keeps heartbeat interval, sequence counter, and Discord Gateway session +/// resumption fields. The [`process_text_frame`] method parses incoming frames +/// and returns a list of [`WebsocketFrameAction`]s the caller should execute. +struct WebsocketSessionState { + heartbeat_interval_ms: Option, + last_sequence: Option, + session_id: Option, + resume_gateway_url: Option, + /// Raw bot token extracted from the identify payload. + token: Option, + /// Whether we attempted a resume on this connection. + attempted_resume: bool, +} + +impl WebsocketSessionState { + fn new(identify_payload: Option<&str>) -> Self { + let token = identify_payload.and_then(extract_token_from_identify_payload); + Self { + heartbeat_interval_ms: None, + last_sequence: None, + session_id: None, + resume_gateway_url: None, + token, + attempted_resume: false, + } + } + + /// Determine the URL to use for the next connection attempt. + fn connect_url<'a>(&'a self, default_url: &'a str) -> &'a str { + if self.session_id.is_some() + && let Some(ref url) = self.resume_gateway_url + { + return url.as_str(); + } + default_url + } + + /// Reset per-connection state when starting a fresh connection. + fn reset_connection(&mut self) { + self.heartbeat_interval_ms = None; + self.attempted_resume = false; + } + + /// Clear all session state so the next reconnect performs a fresh identify. + fn invalidate_session(&mut self) { + self.session_id = None; + self.resume_gateway_url = None; + self.last_sequence = None; + } + + /// Process a text frame and return a list of actions for the caller to + /// execute. This keeps the select loop thin and the parsing logic testable. + fn process_text_frame( + &mut self, + text: &str, + channel_name: &str, + identify_payload: Option<&str>, + workspace_store: &crate::channels::wasm::host::ChannelWorkspaceStore, + pairing_store: &PairingStore, + ) -> Vec { + let mut actions = Vec::new(); + + // OP 10 Hello: extract heartbeat interval, send identify or resume + if let Some(interval_ms) = parse_websocket_hello_heartbeat_interval_ms(text) { + if should_warn_on_heartbeat_interval(interval_ms) { + tracing::warn!( + channel = %channel_name, + heartbeat_interval_ms = interval_ms, + "Websocket hello provided unexpectedly low heartbeat interval" + ); + } + + self.heartbeat_interval_ms = Some(interval_ms); + actions.push(WebsocketFrameAction::SetHeartbeat { interval_ms }); + + // Try resume if we have a session, otherwise fresh identify + let sent_resume = if let (Some(token), Some(sid)) = (&self.token, &self.session_id) { + if let Some(payload) = + build_websocket_resume_message(token, sid, self.last_sequence.as_ref()) + { + self.attempted_resume = true; + actions.push(WebsocketFrameAction::Send(payload)); + true + } else { + false + } + } else { + false + }; + + if !sent_resume && let Some(payload) = identify_payload { + actions.push(WebsocketFrameAction::Send(payload.to_string())); + } + } + + // OP 0 Dispatch READY: capture session_id and resume_gateway_url. + // Presence update is sent here (after READY) rather than on Hello, + // because Discord's gateway protocol requires waiting for READY/RESUMED + // before sending non-Identify commands. + if let Some((sid, resume_url)) = parse_websocket_ready_session(text) { + self.session_id = Some(sid); + self.resume_gateway_url = resume_url; + + if let Some(payload) = + build_gateway_presence_update(channel_name, workspace_store, pairing_store) + { + actions.push(WebsocketFrameAction::Send(payload)); + } + } + + // Track sequence number from any dispatch + if let Some(sequence) = parse_websocket_sequence(text) { + self.last_sequence = Some(serde_json::Value::Number(sequence.into())); + } + + // OP 9 Invalid Session: if not resumable, clear state and reconnect + if let Some(resumable) = parse_websocket_invalid_session(text) + && !resumable + { + tracing::info!( + channel = %channel_name, + "Received non-resumable invalid session; will reconnect with fresh identify" + ); + self.invalidate_session(); + actions.push(WebsocketFrameAction::InvalidateAndReconnect); + return actions; + } + + // Always enqueue the raw frame for the poll callback + actions.push(WebsocketFrameAction::Enqueue(text.to_string())); + + actions + } +} + +fn log_websocket_diagnostic(channel_name: &str, message: &WebsocketMessage) { + match message { + WebsocketMessage::Text(text) => { + tracing::trace!( + channel = %channel_name, + bytes = text.len(), + "Websocket runtime received text frame" + ); + } + WebsocketMessage::Binary(bytes) => { + tracing::debug!( + channel = %channel_name, + bytes = bytes.len(), + "Websocket runtime received binary frame" + ); + } + WebsocketMessage::Close(frame) => { + tracing::info!( + channel = %channel_name, + code = ?frame.as_ref().map(|f| f.code), + reason = ?frame.as_ref().map(|f| f.reason.to_string()), + "Websocket runtime received close frame" + ); + } + WebsocketMessage::Ping(payload) => { + tracing::trace!( + channel = %channel_name, + bytes = payload.len(), + "Websocket runtime received ping" + ); + } + WebsocketMessage::Pong(payload) => { + tracing::trace!( + channel = %channel_name, + bytes = payload.len(), + "Websocket runtime received pong" + ); + } + WebsocketMessage::Frame(_) => {} + } +} + impl std::fmt::Debug for WasmChannel { fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { f.debug_struct("WasmChannel") @@ -3326,19 +4135,29 @@ fn read_attachments(paths: &[String]) -> Result, St #[cfg(test)] mod tests { use std::sync::Arc; + use std::time::Duration; use crate::channels::Channel; use crate::channels::OutgoingResponse; use crate::channels::wasm::capabilities::ChannelCapabilities; + use crate::channels::wasm::host::{ChannelHostState, PendingWorkspaceWrite}; use crate::channels::wasm::runtime::{ PreparedChannelModule, WasmChannelRuntime, WasmChannelRuntimeConfig, }; use crate::channels::wasm::wrapper::{ - EmitDispatchContext, HttpResponse, WasmChannel, uses_owner_broadcast_target, + EmitDispatchContext, HttpResponse, WasmChannel, WebsocketRuntimeConfig, + build_discord_gateway_presence_update, build_websocket_identify_message, + build_websocket_resume_message, discord_gateway_presence_status, drain_guest_logs, + parse_websocket_invalid_session, parse_websocket_ready_session, + should_warn_on_heartbeat_interval, uses_owner_broadcast_target, + websocket_heartbeat_sleep_duration, websocket_reconnect_backoff, }; use crate::pairing::PairingStore; use crate::testing::credentials::TEST_TELEGRAM_BOT_TOKEN; - use crate::tools::wasm::ResourceLimits; + use crate::tools::wasm::{ + Capabilities as ToolCapabilities, EndpointPattern, HttpCapability, LogLevel, ResourceLimits, + }; + use tempfile::tempdir; fn create_test_channel() -> WasmChannel { create_test_channel_with_owner_scope("default") @@ -3368,6 +4187,219 @@ mod tests { ) } + #[test] + fn test_websocket_runtime_config_reads_capability_payload() { + let mut tool_capabilities = ToolCapabilities::default(); + let mut http = HttpCapability::new(vec![EndpointPattern::host("gateway.discord.gg")]); + http.credentials.insert( + "discord_bot_token".to_string(), + crate::secrets::CredentialMapping { + secret_name: "discord_bot_token".to_string(), + location: crate::secrets::CredentialLocation::Header { + name: "Authorization".to_string(), + prefix: Some("Bot ".to_string()), + }, + host_patterns: vec!["discord.com".to_string()], + }, + ); + tool_capabilities.http = Some(http); + tool_capabilities.websocket = Some(serde_json::json!({ + "url": "wss://gateway.discord.gg/?v=10&encoding=json", + "connect_on_start": true, + "identify_secret_name": "discord_bot_token", + "identify": { + "intents": 513, + "properties": { + "os": "linux", + "browser": "ironclaw", + "device": "ironclaw" + } + } + })); + + let capabilities = + ChannelCapabilities::for_channel("discord").with_tool_capabilities(tool_capabilities); + + let config = WebsocketRuntimeConfig::from_capabilities(&capabilities) + .expect("websocket config should be parsed"); + + assert_eq!(config.url, "wss://gateway.discord.gg/?v=10&encoding=json"); + assert!(config.connect_on_start); + assert_eq!( + config.identify_secret_name.as_deref(), + Some("discord_bot_token") + ); + assert_eq!( + config.identify, + Some(serde_json::json!({ + "intents": 513, + "properties": { + "os": "linux", + "browser": "ironclaw", + "device": "ironclaw" + } + })) + ); + } + + #[test] + fn test_build_websocket_identify_message_includes_token() { + let identify = serde_json::json!({ + "intents": 513, + "properties": { + "os": "linux", + "browser": "ironclaw", + "device": "ironclaw" + } + }); + + let payload = build_websocket_identify_message(&identify, "bot-token").unwrap(); + let json: serde_json::Value = serde_json::from_str(&payload).unwrap(); + + assert_eq!(json["op"], serde_json::json!(2)); + assert_eq!(json["d"]["token"], serde_json::json!("bot-token")); + assert_eq!(json["d"]["intents"], serde_json::json!(513)); + } + + #[test] + fn test_websocket_runtime_config_requires_allowlisted_host() { + let tool_capabilities = ToolCapabilities { + http: Some(HttpCapability::new(vec![EndpointPattern::host( + "discord.com", + )])), + websocket: Some(serde_json::json!({ + "url": "wss://gateway.discord.gg/?v=10&encoding=json", + "connect_on_start": true + })), + ..Default::default() + }; + + let capabilities = + ChannelCapabilities::for_channel("discord").with_tool_capabilities(tool_capabilities); + + assert!(WebsocketRuntimeConfig::from_capabilities(&capabilities).is_none()); + } + + #[test] + fn test_drain_guest_logs_collects_poll_entries() { + let mut host_state = ChannelHostState::new("poll-test", ChannelCapabilities::default()); + host_state + .log(LogLevel::Warn, "poll warning".to_string()) + .expect("log entry should be stored"); + + let logs = drain_guest_logs("poll-test", "on_poll", &mut host_state); + + assert_eq!(logs.len(), 1); + assert_eq!(logs[0].message, "poll warning"); + assert_eq!(logs[0].level, LogLevel::Warn); + assert!(host_state.take_logs().is_empty(), "logs should be drained"); + } + + #[test] + fn test_websocket_reconnect_backoff_caps_at_sixty_four_seconds_with_jitter() { + // Backoff = base + 0-25% jitter, so check range [base, base * 1.25]. + let check = |attempt: u32, base_secs: u64| { + let d = websocket_reconnect_backoff(attempt); + let base = Duration::from_secs(base_secs); + let max = base + base / 4; + assert!( + d >= base && d <= max, + "attempt {attempt}: {d:?} not in [{base:?}, {max:?}]" + ); + }; + check(0, 1); + check(1, 2); + check(5, 32); + check(6, 64); + check(10, 64); // capped at 2^6 + } + + #[test] + fn test_websocket_heartbeat_helpers_guard_low_intervals() { + assert!(should_warn_on_heartbeat_interval(0)); + assert!(should_warn_on_heartbeat_interval(999)); + assert!(!should_warn_on_heartbeat_interval(1_000)); + assert_eq!( + websocket_heartbeat_sleep_duration(0), + Duration::from_millis(1) + ); + assert_eq!( + websocket_heartbeat_sleep_duration(42), + Duration::from_millis(42) + ); + } + + #[test] + fn test_discord_gateway_presence_defaults_to_dnd() { + let store = crate::channels::wasm::host::ChannelWorkspaceStore::new(); + let pairing_dir = tempdir().unwrap(); + let pairing_store = PairingStore::with_base_dir(pairing_dir.path().to_path_buf()); + + assert_eq!( + discord_gateway_presence_status("discord", &store, &pairing_store), + "dnd" + ); + } + + #[test] + fn test_discord_gateway_presence_empty_owner_id_is_dnd() { + let store = crate::channels::wasm::host::ChannelWorkspaceStore::new(); + let pairing_dir = tempdir().unwrap(); + let pairing_store = PairingStore::with_base_dir(pairing_dir.path().to_path_buf()); + // Simulate on_start writing empty string when no owner_id is configured + store.commit_writes(&[PendingWorkspaceWrite { + path: "channels/discord/state/owner_id".to_string(), + content: String::new(), + }]); + + assert_eq!( + discord_gateway_presence_status("discord", &store, &pairing_store), + "dnd" + ); + } + + #[test] + fn test_discord_gateway_presence_pairing_approved_is_online() { + let store = crate::channels::wasm::host::ChannelWorkspaceStore::new(); + let pairing_dir = tempdir().unwrap(); + let pairing_store = PairingStore::with_base_dir(pairing_dir.path().to_path_buf()); + let request = pairing_store + .upsert_request("discord", "user-1", None) + .unwrap(); + pairing_store.approve("discord", &request.code).unwrap(); + + assert_eq!( + discord_gateway_presence_status("discord", &store, &pairing_store), + "online" + ); + } + + #[test] + fn test_discord_gateway_presence_owner_id_is_online() { + let store = crate::channels::wasm::host::ChannelWorkspaceStore::new(); + let pairing_dir = tempdir().unwrap(); + let pairing_store = PairingStore::with_base_dir(pairing_dir.path().to_path_buf()); + store.commit_writes(&[PendingWorkspaceWrite { + path: "channels/discord/state/owner_id".to_string(), + content: "owner-1".to_string(), + }]); + + assert_eq!( + discord_gateway_presence_status("discord", &store, &pairing_store), + "online" + ); + } + + #[test] + fn test_build_discord_gateway_presence_update_uses_status() { + let payload = build_discord_gateway_presence_update("dnd").unwrap(); + let json: serde_json::Value = serde_json::from_str(&payload).unwrap(); + + assert_eq!(json["op"], serde_json::json!(3)); + assert_eq!(json["d"]["status"], serde_json::json!("dnd")); + assert_eq!(json["d"]["afk"], serde_json::json!(false)); + } + #[test] fn test_channel_name() { let channel = create_test_channel(); @@ -4762,6 +5794,77 @@ mod tests { assert!(msg.attachments.is_empty()); // safety: test-only assertion } + #[test] + fn test_parse_websocket_ready_session() { + let ready = serde_json::json!({ + "op": 0, + "s": 1, + "t": "READY", + "d": { + "session_id": "abc123", + "resume_gateway_url": "wss://gateway-resume.discord.gg", + "user": {"id": "12345"} + } + }); + let (sid, resume_url) = parse_websocket_ready_session(&ready.to_string()).unwrap(); + assert_eq!(sid, "abc123"); + assert_eq!( + resume_url.as_deref(), + Some("wss://gateway-resume.discord.gg") + ); + + // Non-READY dispatch returns None + let message_create = serde_json::json!({ + "op": 0, + "s": 2, + "t": "MESSAGE_CREATE", + "d": {"content": "hello"} + }); + assert!(parse_websocket_ready_session(&message_create.to_string()).is_none()); + + // Non-dispatch opcode returns None + let hello = serde_json::json!({"op": 10, "d": {"heartbeat_interval": 41250}}); + assert!(parse_websocket_ready_session(&hello.to_string()).is_none()); + } + + #[test] + fn test_build_websocket_resume_message() { + let seq = serde_json::Value::Number(42.into()); + let payload = build_websocket_resume_message("bot-token", "session-1", Some(&seq)).unwrap(); + let json: serde_json::Value = serde_json::from_str(&payload).unwrap(); + + assert_eq!(json["op"], serde_json::json!(6)); + assert_eq!(json["d"]["token"], serde_json::json!("bot-token")); + assert_eq!(json["d"]["session_id"], serde_json::json!("session-1")); + assert_eq!(json["d"]["seq"], serde_json::json!(42)); + + // With no sequence, seq should be null + let payload_null = build_websocket_resume_message("bot-token", "session-1", None).unwrap(); + let json_null: serde_json::Value = serde_json::from_str(&payload_null).unwrap(); + assert!(json_null["d"]["seq"].is_null()); + } + + #[test] + fn test_parse_websocket_invalid_session() { + // Non-resumable invalid session (d: false) + let not_resumable = serde_json::json!({"op": 9, "d": false}); + assert_eq!( + parse_websocket_invalid_session(¬_resumable.to_string()), + Some(false) + ); + + // Resumable invalid session (d: true) + let resumable = serde_json::json!({"op": 9, "d": true}); + assert_eq!( + parse_websocket_invalid_session(&resumable.to_string()), + Some(true) + ); + + // Different opcode returns None + let hello = serde_json::json!({"op": 10, "d": {"heartbeat_interval": 41250}}); + assert!(parse_websocket_invalid_session(&hello.to_string()).is_none()); + } + #[test] fn test_mime_from_extension() { use super::mime_from_extension; diff --git a/src/channels/web/handlers/extensions.rs b/src/channels/web/handlers/extensions.rs index d705591e7a1..34982f7435a 100644 --- a/src/channels/web/handlers/extensions.rs +++ b/src/channels/web/handlers/extensions.rs @@ -12,6 +12,37 @@ use crate::channels::web::auth::AuthenticatedUser; use crate::channels::web::server::GatewayState; use crate::channels::web::types::*; +pub(crate) fn derive_activation_status( + ext: &crate::extensions::InstalledExtension, + pairing_store: &crate::pairing::PairingStore, + has_owner_binding: bool, +) -> Option { + if ext.kind == crate::extensions::ExtensionKind::WasmChannel { + let allowlist_exists = pairing_store + .has_allow_from_file(&ext.name) + .unwrap_or(false); + let has_paired = pairing_store + .read_allow_from(&ext.name) + .map(|list| !list.is_empty()) + .unwrap_or(false); + classify_wasm_channel_activation( + ext, + has_paired, + has_owner_binding || (ext.active && !allowlist_exists), + ) + } else if ext.kind == crate::extensions::ExtensionKind::ChannelRelay { + Some(if ext.active { + ExtensionActivationStatus::Active + } else if ext.authenticated { + ExtensionActivationStatus::Configured + } else { + ExtensionActivationStatus::Installed + }) + } else { + None + } +} + pub async fn extensions_list_handler( State(state): State>, AuthenticatedUser(user): AuthenticatedUser, @@ -38,27 +69,11 @@ pub async fn extensions_list_handler( let extensions = installed .into_iter() .map(|ext| { - let activation_status = if ext.kind == crate::extensions::ExtensionKind::WasmChannel { - let has_paired = pairing_store - .read_allow_from(&ext.name) - .map(|list| !list.is_empty()) - .unwrap_or(false); - crate::channels::web::types::classify_wasm_channel_activation( - &ext, - has_paired, - owner_bound_channels.contains(&ext.name), - ) - } else if ext.kind == crate::extensions::ExtensionKind::ChannelRelay { - Some(if ext.active { - crate::channels::web::types::ExtensionActivationStatus::Active - } else if ext.authenticated { - crate::channels::web::types::ExtensionActivationStatus::Configured - } else { - crate::channels::web::types::ExtensionActivationStatus::Installed - }) - } else { - None - }; + let activation_status = derive_activation_status( + &ext, + &pairing_store, + owner_bound_channels.contains(&ext.name), + ); ExtensionInfo { name: ext.name, display_name: ext.display_name, @@ -143,3 +158,63 @@ pub async fn extensions_remove_handler( Err(e) => Ok(Json(ActionResponse::fail(e.to_string()))), } } + +#[cfg(test)] +mod tests { + use std::fs; + + use tempfile::TempDir; + + use super::derive_activation_status; + use crate::channels::web::types::ExtensionActivationStatus; + use crate::extensions::{ExtensionKind, InstalledExtension}; + use crate::pairing::PairingStore; + + fn active_authenticated_wasm_channel(name: &str) -> InstalledExtension { + InstalledExtension { + name: name.to_string(), + kind: ExtensionKind::WasmChannel, + display_name: None, + description: None, + url: None, + authenticated: true, + active: true, + tools: Vec::new(), + needs_setup: false, + has_auth: false, + installed: true, + activation_error: None, + version: None, + } + } + + #[test] + fn active_authenticated_wasm_channel_without_allowlist_file_is_active() { + let temp_dir = TempDir::new().expect("temp dir"); + let pairing_store = PairingStore::with_base_dir(temp_dir.path().to_path_buf()); + let ext = active_authenticated_wasm_channel("discord"); + + assert_eq!( + derive_activation_status(&ext, &pairing_store, false), + Some(ExtensionActivationStatus::Active) + ); + } + + #[test] + fn active_authenticated_wasm_channel_with_empty_allowlist_file_is_pairing() { + let temp_dir = TempDir::new().expect("temp dir"); + let pairing_store = PairingStore::with_base_dir(temp_dir.path().to_path_buf()); + let ext = active_authenticated_wasm_channel("discord"); + + fs::write( + temp_dir.path().join("discord-allowFrom.json"), + r#"{"version":1,"allowFrom":[]}"#, + ) + .expect("write empty allowlist"); + + assert_eq!( + derive_activation_status(&ext, &pairing_store, false), + Some(ExtensionActivationStatus::Pairing) + ); + } +} diff --git a/src/channels/web/server.rs b/src/channels/web/server.rs index 8c9ecddec6d..d403e93c399 100644 --- a/src/channels/web/server.rs +++ b/src/channels/web/server.rs @@ -2159,27 +2159,12 @@ async fn extensions_list_handler( let extensions = installed .into_iter() .map(|ext| { - let activation_status = if ext.kind == crate::extensions::ExtensionKind::WasmChannel { - let has_paired = pairing_store - .read_allow_from(&ext.name) - .map(|list| !list.is_empty()) - .unwrap_or(false); - crate::channels::web::types::classify_wasm_channel_activation( + let activation_status = + crate::channels::web::handlers::extensions::derive_activation_status( &ext, - has_paired, + &pairing_store, owner_bound_channels.contains(&ext.name), - ) - } else if ext.kind == crate::extensions::ExtensionKind::ChannelRelay { - Some(if ext.active { - ExtensionActivationStatus::Active - } else if ext.authenticated { - ExtensionActivationStatus::Configured - } else { - ExtensionActivationStatus::Installed - }) - } else { - None - }; + ); ExtensionInfo { name: ext.name, display_name: ext.display_name, diff --git a/src/pairing/store.rs b/src/pairing/store.rs index 6c0882fd994..d82a7f1409f 100644 --- a/src/pairing/store.rs +++ b/src/pairing/store.rs @@ -420,6 +420,12 @@ impl PairingStore { Ok(Some(entry)) } + /// Read the allowFrom list for a channel. + pub fn has_allow_from_file(&self, channel: &str) -> Result { + let path = allow_from_path(&self.base_dir, channel)?; + Ok(path.exists()) + } + /// Read the allowFrom list for a channel. pub fn read_allow_from(&self, channel: &str) -> Result, PairingStoreError> { let path = allow_from_path(&self.base_dir, channel)?; diff --git a/src/tools/wasm/capabilities.rs b/src/tools/wasm/capabilities.rs index ff98ae03491..608a53071b9 100644 --- a/src/tools/wasm/capabilities.rs +++ b/src/tools/wasm/capabilities.rs @@ -34,6 +34,8 @@ pub struct Capabilities { pub secrets: Option, /// Webhook authentication and signature verification. pub webhook: Option, + /// Arbitrary websocket configuration preserved from capabilities JSON. + pub websocket: Option, } impl Capabilities { @@ -341,6 +343,7 @@ mod tests { assert!(caps.tool_invoke.is_none()); assert!(caps.secrets.is_none()); assert!(caps.webhook.is_none()); + assert!(caps.websocket.is_none()); } #[test] diff --git a/src/tools/wasm/capabilities_schema.rs b/src/tools/wasm/capabilities_schema.rs index b275832957b..8ac7806ca58 100644 --- a/src/tools/wasm/capabilities_schema.rs +++ b/src/tools/wasm/capabilities_schema.rs @@ -75,6 +75,10 @@ pub struct CapabilitiesFile { #[serde(default)] pub webhook: Option, + /// Arbitrary websocket configuration preserved for runtime consumers. + #[serde(default)] + pub websocket: Option, + /// Authentication setup instructions. /// Used by `ironclaw config` to guide users through auth setup. #[serde(default)] @@ -155,6 +159,7 @@ impl CapabilitiesFile { self.tool_invoke = self.tool_invoke.or(inner.tool_invoke); self.workspace = self.workspace.or(inner.workspace); self.webhook = self.webhook.or(inner.webhook); + self.websocket = self.websocket.or(inner.websocket); self.auth = self.auth.or(inner.auth); self.setup = self.setup.or(inner.setup); } @@ -250,6 +255,8 @@ impl CapabilitiesFile { caps.webhook = Some(webhook.to_webhook_capability()); } + caps.websocket = self.websocket.clone(); + caps } } @@ -745,6 +752,8 @@ fn default_tool_setup_field_input_type() -> ToolSetupFieldInputType { #[cfg(test)] mod tests { + use serde_json::json; + use crate::tools::wasm::capabilities_schema::{CapabilitiesFile, CredentialLocationSchema}; #[test] @@ -1402,6 +1411,48 @@ mod tests { ); } + #[test] + fn test_discord_websocket_config_preserved_in_runtime_capabilities() { + let json = r#"{ + "capabilities": { + "http": { + "allowlist": [{ "host": "discord.com", "path_prefix": "/api/v10" }] + }, + "websocket": { + "url": "wss://gateway.discord.gg/?v=10&encoding=json", + "connect_on_start": true, + "identify": { + "intents": 513, + "properties": { + "os": "linux", + "browser": "ironclaw", + "device": "ironclaw" + } + } + } + } + }"#; + + let file = CapabilitiesFile::from_json(json).unwrap(); + let caps = file.to_capabilities(); + + assert_eq!( + caps.websocket, + Some(json!({ + "url": "wss://gateway.discord.gg/?v=10&encoding=json", + "connect_on_start": true, + "identify": { + "intents": 513, + "properties": { + "os": "linux", + "browser": "ironclaw", + "device": "ironclaw" + } + } + })) + ); + } + // ── Tool description ──────────────────────────────────────────────── #[test] From e0e530e646c4cdb05a2ea0cac8534d6e5fbb6afc Mon Sep 17 00:00:00 2001 From: "firat.sertgoz" Date: Sun, 29 Mar 2026 09:11:58 +0300 Subject: [PATCH 08/23] docs: tighten contribution and PR guidance (#1704) --- .github/pull_request_template.md | 17 ++++--- CONTRIBUTING.md | 77 +++++++++++++++++++++++++++++++- 2 files changed, 88 insertions(+), 6 deletions(-) diff --git a/.github/pull_request_template.md b/.github/pull_request_template.md index 4fc7cbf233b..e6fe6128f82 100644 --- a/.github/pull_request_template.md +++ b/.github/pull_request_template.md @@ -6,7 +6,7 @@ ## Change Type - + - [ ] Bug fix - [ ] New feature @@ -18,16 +18,19 @@ ## Linked Issue - + ## Validation -- [ ] `cargo fmt` -- [ ] `cargo clippy --all --benches --tests --examples --all-features` +- [ ] `cargo fmt --all -- --check` +- [ ] `cargo clippy --all --benches --tests --examples --all-features -- -D warnings` +- [ ] `cargo build` - [ ] Relevant tests pass: +- [ ] `cargo test --features integration` if database-backed or integration behavior changed - [ ] Manual testing: +- [ ] If a coding agent was used and supports it, `review-pr` or `pr-shepherd --fix` was run before requesting review ## Security Impact @@ -45,6 +48,10 @@ +## Review Follow-Through + + + --- -**Review track**: +**Review track**: diff --git a/CONTRIBUTING.md b/CONTRIBUTING.md index 1c5c6d88194..c7a2b2dfc6a 100644 --- a/CONTRIBUTING.md +++ b/CONTRIBUTING.md @@ -10,6 +10,42 @@ cd ironclaw This installs the Rust toolchain, WASM targets, git hooks, and runs initial checks. +## How to Contribute + +- Bug fixes, docs improvements, and focused cleanup tied to a concrete problem are welcome. +- Search existing issues and PRs before opening a new one to avoid duplicates. +- Keep changes scoped. One bug, one feature, or one documentation improvement per PR. + +### Creating Issues + +Open an issue when you are reporting a bug, proposing a feature, or documenting a gap in behavior. + +For bug reports, include: + +- What you expected to happen +- What actually happened +- Clear reproduction steps +- Relevant logs, screenshots, or error output +- Environment details when they matter (OS, database backend, feature flags, commit/branch) + +For feature requests: + +- Open an issue first before writing code +- Explain the problem being solved, not just the implementation idea +- Wait for maintainer feedback before investing in a large PR + +We require an issue for new features so maintainers can prioritize the work and confirm it fits the roadmap before anyone spends time implementing it. + +### Fixing Bugs + +- Small, targeted bug-fix PRs are welcome +- If there is already an issue, link it in your PR +- If the bug is non-trivial, security-sensitive, or changes behavior across subsystems, open or confirm an issue first so the approach can be aligned before implementation + +### Refactor-Only PRs + +Refactor-only PRs are not accepted from contributors outside the core team. If a refactor is necessary to land a bug fix or approved feature, keep it minimal and clearly tied to that change. + ## Development Workflow ```bash @@ -19,6 +55,45 @@ cargo test # unit tests cargo test --features integration # + PostgreSQL tests ``` +These commands are for day-to-day iteration while you are developing locally. The pre-submission checks below are intentionally stricter and use CI-style flags so you can catch formatting drift and clippy warnings before requesting review. + +## Before You Open a PR + +Run the local validation checks required before requesting a review. These are stricter than the commands for iterative development: + +```bash +cargo fmt --all -- --check +cargo clippy --all --benches --tests --examples --all-features -- -D warnings +cargo build +cargo test +``` + +Also run this when your change touches database-backed or integration behavior: + +```bash +cargo test --features integration +``` + +Before asking for review: + +- Build and exercise the changed path locally, not just the narrowest unit test +- Keep the PR focused and avoid mixing unrelated concerns +- Fill out the PR template with a clear summary, validation notes, and impact assessment +- If your change affects tracked behavior, update `FEATURE_PARITY.md` in the same branch +- If onboarding or setup behavior changes, update the relevant setup docs in the same branch +- If you are using a coding agent and it supports them, run `review-pr` or `pr-shepherd --fix` before opening or updating the PR +- `codex review --base origin/main` is also encouraged before requesting review + +## Review Follow-Through + +Review conversations are author-owned. + +- Address each review comment with a code change or a clear explanation +- Resolve conversations you have handled; leave them open only when reviewer judgment is still needed +- Do not leave review cleanup for maintainers when the follow-through belongs to the author + +If a PR is stale for more than 48 hours after review feedback is posted, maintainers may take over the follow-up work and land the changes needed to accomplish the original PR or issue intent. + ## Code Style - Zero clippy warnings policy @@ -46,7 +121,7 @@ All PRs follow a risk-based review process: | Track | Scope | Requirements | |-------|-------|-------------| | **A** | Docs, tests, chore, dependency bumps | 1 approval + CI green | -| **B** | Features, refactors, new tools/channels | 1 approval + CI green + test evidence | +| **B** | Features, maintainer-requested refactors, new tools/channels | 1 approval + CI green + test evidence | | **C** | Security (`src/safety/`, `src/secrets/`), runtime (`src/agent/`, `src/worker/`), database schema, CI workflows | 2 approvals + rollback plan documented | Select the appropriate track in the PR template based on what your changes touch. From 6f6a5f1dbdea2b291c9922696487180d09953bef Mon Sep 17 00:00:00 2001 From: rajulbhatnagar Date: Sun, 29 Mar 2026 02:38:20 -0700 Subject: [PATCH 09/23] feat(skills): recursive bundle directory scanning for skill discovery (#1667) Add support for bundle layouts where directories without SKILL.md are recursed into to find nested skills (e.g., skills/my-org/skill-a/SKILL.md). - Add configurable max_scan_depth (SKILLS_MAX_SCAN_DEPTH env, default 3) - Recurse into subdirectories lacking SKILL.md up to depth limit - Share remaining discovery cap across recursive levels - Replace try_exists + read_dir with single read_dir (eliminates TOCTOU) - Box::pin recursive async calls for correct future sizing Closes #1664 Co-authored-by: Rajul Bhatnagar --- src/app.rs | 3 +- src/cli/skills.rs | 3 +- src/config/skills.rs | 5 + src/skills/registry.rs | 395 +++++++++++++++++++++++++++++++++++------ 4 files changed, 348 insertions(+), 58 deletions(-) diff --git a/src/app.rs b/src/app.rs index 94262c3ac7f..0ca9cd9ec4d 100644 --- a/src/app.rs +++ b/src/app.rs @@ -857,7 +857,8 @@ impl AppBuilder { // Skills system let (skill_registry, skill_catalog) = if self.config.skills.enabled { let mut registry = SkillRegistry::new(self.config.skills.local_dir.clone()) - .with_installed_dir(self.config.skills.installed_dir.clone()); + .with_installed_dir(self.config.skills.installed_dir.clone()) + .with_max_scan_depth(self.config.skills.max_scan_depth); let loaded = registry.discover_all().await; if !loaded.is_empty() { tracing::debug!("Loaded {} skill(s): {}", loaded.len(), loaded.join(", ")); diff --git a/src/cli/skills.rs b/src/cli/skills.rs index 1f3cc46b761..62d81e0ce09 100644 --- a/src/cli/skills.rs +++ b/src/cli/skills.rs @@ -69,7 +69,8 @@ pub async fn run_skills_command( /// Discover skills from all configured directories. async fn discover_skills(config: &SkillsConfig) -> SkillRegistry { let mut registry = SkillRegistry::new(config.local_dir.clone()) - .with_installed_dir(config.installed_dir.clone()); + .with_installed_dir(config.installed_dir.clone()) + .with_max_scan_depth(config.max_scan_depth); registry.discover_all().await; registry } diff --git a/src/config/skills.rs b/src/config/skills.rs index 97970784181..e893d2316e7 100644 --- a/src/config/skills.rs +++ b/src/config/skills.rs @@ -19,6 +19,9 @@ pub struct SkillsConfig { pub max_active_skills: usize, /// Maximum total context tokens allocated to skill prompts. pub max_context_tokens: usize, + /// Maximum recursion depth when scanning skill directories for bundle layouts. + /// Subdirectories without `SKILL.md` are recursed into up to this depth. + pub max_scan_depth: usize, } impl Default for SkillsConfig { @@ -29,6 +32,7 @@ impl Default for SkillsConfig { installed_dir: default_installed_skills_dir(), max_active_skills: 3, max_context_tokens: 4000, + max_scan_depth: 3, } } } @@ -55,6 +59,7 @@ impl SkillsConfig { .unwrap_or_else(default_installed_skills_dir), max_active_skills: parse_optional_env("SKILLS_MAX_ACTIVE", 3)?, max_context_tokens: parse_optional_env("SKILLS_MAX_CONTEXT_TOKENS", 4000)?, + max_scan_depth: parse_optional_env("SKILLS_MAX_SCAN_DEPTH", 3)?, }) } } diff --git a/src/skills/registry.rs b/src/skills/registry.rs index 6f881f77849..0eccb027054 100644 --- a/src/skills/registry.rs +++ b/src/skills/registry.rs @@ -1,12 +1,15 @@ //! Skill registry for discovering, loading, and managing available skills. //! -//! Skills are discovered from two filesystem locations: +//! Skills are discovered from three filesystem locations: //! 1. Workspace skills directory (`/skills/`) -- Trusted //! 2. User skills directory (`~/.ironclaw/skills/`) -- Trusted +//! 3. Installed skills directory (`~/.ironclaw/installed_skills/`) -- Installed //! //! Both flat (`skills/SKILL.md`) and subdirectory (`skills//SKILL.md`) -//! layouts are supported. Earlier locations win on name collision (workspace -//! overrides user). Uses async I/O throughout to avoid blocking the tokio runtime. +//! layouts are supported. Subdirectories without `SKILL.md` are treated as +//! bundle directories and recursed into (up to `SKILLS_MAX_SCAN_DEPTH`, +//! default 3). Earlier locations win on name collision (workspace overrides +//! user). Uses async I/O throughout to avoid blocking the tokio runtime. use std::collections::HashSet; use std::path::{Path, PathBuf}; @@ -20,10 +23,14 @@ use crate::skills::{ normalize_line_endings, }; -/// Maximum number of skills that can be discovered from a single directory. -/// Prevents resource exhaustion from a directory with thousands of entries. +/// Maximum total number of skills that can be discovered across all sources. +/// Shared across workspace, user, and installed directories. +/// Prevents resource exhaustion from directories with thousands of entries. const MAX_DISCOVERED_SKILLS: usize = 100; +/// Default recursion depth for bundle directory scanning. +const DEFAULT_MAX_SCAN_DEPTH: usize = 3; + fn to_lowercase_vec(items: &[String]) -> Vec { items.iter().map(|s| s.to_lowercase()).collect() } @@ -78,6 +85,8 @@ pub struct SkillRegistry { installed_dir: Option, /// Optional workspace skills directory. workspace_dir: Option, + /// Maximum recursion depth for bundle directory scanning (default: 3). + max_scan_depth: usize, } impl SkillRegistry { @@ -88,6 +97,7 @@ impl SkillRegistry { user_dir, installed_dir: None, workspace_dir: None, + max_scan_depth: DEFAULT_MAX_SCAN_DEPTH, } } @@ -108,6 +118,12 @@ impl SkillRegistry { self } + /// Set the maximum recursion depth for bundle directory scanning. + pub fn with_max_scan_depth(mut self, depth: usize) -> Self { + self.max_scan_depth = depth; + self + } + /// Discover and load skills from all configured directories. /// /// Discovery order (earlier wins on name collision): @@ -120,91 +136,110 @@ impl SkillRegistry { // 1. Workspace skills (highest priority) if let Some(ws_dir) = self.workspace_dir.clone() { - let ws_skills = self - .discover_from_dir(&ws_dir, SkillTrust::Trusted, SkillSource::Workspace) + let cap = MAX_DISCOVERED_SKILLS.saturating_sub(loaded_names.len()); + let skills = self + .discover_from_dir( + &ws_dir, + SkillTrust::Trusted, + &SkillSource::Workspace, + cap, + 0, + ) .await; - for (name, skill) in ws_skills { - if seen.contains(&name) { - continue; - } - seen.insert(name.clone()); - loaded_names.push(name); - self.skills.push(skill); - } + self.absorb(skills, &mut seen, &mut loaded_names, "workspace"); } // 2. User skills - let user_dir = self.user_dir.clone(); - let user_skills = self - .discover_from_dir(&user_dir, SkillTrust::Trusted, SkillSource::User) - .await; - for (name, skill) in user_skills { - if seen.contains(&name) { - tracing::debug!("Skipping user skill '{}' (overridden by workspace)", name); - continue; - } - seen.insert(name.clone()); - loaded_names.push(name); - self.skills.push(skill); + if loaded_names.len() < MAX_DISCOVERED_SKILLS { + let cap = MAX_DISCOVERED_SKILLS.saturating_sub(loaded_names.len()); + let user_dir = self.user_dir.clone(); + let skills = self + .discover_from_dir(&user_dir, SkillTrust::Trusted, &SkillSource::User, cap, 0) + .await; + self.absorb(skills, &mut seen, &mut loaded_names, "user/workspace"); } // 3. Installed skills (registry-installed, lowest priority) - if let Some(inst_dir) = self.installed_dir.clone() { - let inst_skills = self - .discover_from_dir(&inst_dir, SkillTrust::Installed, SkillSource::User) + if loaded_names.len() < MAX_DISCOVERED_SKILLS + && let Some(inst_dir) = self.installed_dir.clone() + { + let cap = MAX_DISCOVERED_SKILLS.saturating_sub(loaded_names.len()); + let skills = self + .discover_from_dir(&inst_dir, SkillTrust::Installed, &SkillSource::User, cap, 0) .await; - for (name, skill) in inst_skills { - if seen.contains(&name) { - tracing::debug!( - "Skipping installed skill '{}' (overridden by user/workspace)", - name - ); - continue; - } - seen.insert(name.clone()); - loaded_names.push(name); - self.skills.push(skill); - } + self.absorb(skills, &mut seen, &mut loaded_names, "user/workspace"); + } + + if loaded_names.len() >= MAX_DISCOVERED_SKILLS { + tracing::warn!( + "Global skill discovery cap reached ({} skills)", + MAX_DISCOVERED_SKILLS + ); } loaded_names } - /// Discover skills from a single directory. + /// Dedup and absorb discovered skills into the registry. + fn absorb( + &mut self, + skills: Vec<(String, LoadedSkill)>, + seen: &mut HashSet, + loaded_names: &mut Vec, + override_source: &str, + ) { + for (name, skill) in skills { + if seen.contains(&name) { + tracing::debug!( + "Skipping skill '{}' (overridden by {})", + name, + override_source + ); + continue; + } + seen.insert(name.clone()); + loaded_names.push(name); + self.skills.push(skill); + } + } + + /// Discover skills from a single directory, recursing into bundle directories. /// - /// Supports both layouts: + /// Supports three layouts: /// - Flat: `dir/SKILL.md` (skill name derived from parent dir or file stem) /// - Subdirectory: `dir//SKILL.md` + /// - Bundle: `dir///SKILL.md` (bundle has no `SKILL.md`, recursed into) async fn discover_from_dir( &self, dir: &Path, trust: SkillTrust, - make_source: F, + make_source: &F, + remaining_cap: usize, + current_depth: usize, ) -> Vec<(String, LoadedSkill)> where - F: Fn(PathBuf) -> SkillSource, + F: Fn(PathBuf) -> SkillSource + Send + Sync, { let mut results = Vec::new(); - if !tokio::fs::try_exists(dir).await.unwrap_or(false) { - tracing::debug!("Skills directory does not exist: {:?}", dir); - return results; - } - let mut entries = match tokio::fs::read_dir(dir).await { Ok(entries) => entries, Err(e) => { - tracing::warn!("Failed to read skills directory {:?}: {}", dir, e); + if e.kind() == std::io::ErrorKind::NotFound { + tracing::debug!("Skills directory does not exist: {:?}", dir); + } else { + tracing::warn!("Failed to read skills directory {:?}: {}", dir, e); + } return results; } }; let mut count = 0usize; while let Ok(Some(entry)) = entries.next_entry().await { - if count >= MAX_DISCOVERED_SKILLS { + if count >= remaining_cap { tracing::warn!( - "Skill discovery cap reached ({} skills), skipping remaining", - MAX_DISCOVERED_SKILLS + "Skill discovery cap reached ({} skills in this scan), skipping remaining", + count ); break; } @@ -218,7 +253,6 @@ impl SkillRegistry { } }; - // Reject symlinks if meta.is_symlink() { tracing::warn!( "Skipping symlink in skills directory: {:?}", @@ -246,6 +280,22 @@ impl SkillRegistry { ); } } + } else if current_depth < self.max_scan_depth { + tracing::debug!( + "Recursing into bundle directory {:?} (depth {})", + path.file_name().unwrap_or_default(), + current_depth + 1 + ); + let nested = Box::pin(self.discover_from_dir( + &path, + trust, + make_source, + remaining_cap.saturating_sub(count), + current_depth + 1, + )) + .await; + count += nested.len(); + results.extend(nested); } continue; } @@ -1091,4 +1141,237 @@ mod tests { let skill = registry.find_by_name("my-skill").unwrap(); assert_eq!(skill.trust, SkillTrust::Trusted); } + + #[tokio::test] + async fn test_discover_nested_bundle_directory() { + let dir = tempfile::tempdir().unwrap(); + + // Bundle directory (no SKILL.md) + let bundle = dir.path().join("my-org"); + fs::create_dir(&bundle).unwrap(); + + // Two skills inside the bundle + let skill_a = bundle.join("skill-a"); + fs::create_dir(&skill_a).unwrap(); + fs::write( + skill_a.join("SKILL.md"), + "---\nname: skill-a\n---\n\nSkill A prompt.\n", + ) + .unwrap(); + + let skill_b = bundle.join("skill-b"); + fs::create_dir(&skill_b).unwrap(); + fs::write( + skill_b.join("SKILL.md"), + "---\nname: skill-b\n---\n\nSkill B prompt.\n", + ) + .unwrap(); + + let mut registry = SkillRegistry::new(dir.path().to_path_buf()); + let loaded = registry.discover_all().await; + + assert_eq!(registry.count(), 2); + assert!(loaded.contains(&"skill-a".to_string())); + assert!(loaded.contains(&"skill-b".to_string())); + } + + #[tokio::test] + async fn test_discover_respects_depth_limit() { + let dir = tempfile::tempdir().unwrap(); + + // Create skill nested 3 levels deep (a/b/c/deep-skill/SKILL.md) + let nested = dir.path().join("a").join("b").join("c").join("deep-skill"); + fs::create_dir_all(&nested).unwrap(); + fs::write( + nested.join("SKILL.md"), + "---\nname: deep-skill\n---\n\nDeep prompt.\n", + ) + .unwrap(); + + // Depth 2 should NOT find it (3 intermediate dirs: a, b, c) + let mut registry = SkillRegistry::new(dir.path().to_path_buf()).with_max_scan_depth(2); + let loaded = registry.discover_all().await; + assert!(loaded.is_empty(), "depth=2 should not reach 3 levels deep"); + + // Depth 3 SHOULD find it + let mut registry = SkillRegistry::new(dir.path().to_path_buf()).with_max_scan_depth(3); + let loaded = registry.discover_all().await; + assert_eq!(loaded, vec!["deep-skill"]); + } + + #[tokio::test] + async fn test_discover_cap_spans_recursive_levels() { + let dir = tempfile::tempdir().unwrap(); + + // Spread skills across two bundle directories so the cap must be + // shared across separate recursive calls (not just within one). + // Each bundle has 60 skills; with a global cap of 100, the second + // bundle should be cut short. + for bundle_name in &["bundle-a", "bundle-b"] { + let bundle = dir.path().join(bundle_name); + fs::create_dir(&bundle).unwrap(); + for i in 0..60 { + let skill_dir = bundle.join(format!("{}-skill-{:02}", bundle_name, i)); + fs::create_dir(&skill_dir).unwrap(); + fs::write( + skill_dir.join("SKILL.md"), + format!( + "---\nname: {}-skill-{:02}\n---\n\nPrompt.\n", + bundle_name, i + ), + ) + .unwrap(); + } + } + + let mut registry = SkillRegistry::new(dir.path().to_path_buf()); + registry.discover_all().await; + + assert!( + registry.count() <= MAX_DISCOVERED_SKILLS, + "global cap should limit total to {} but got {}", + MAX_DISCOVERED_SKILLS, + registry.count() + ); + } + + #[cfg(unix)] + #[tokio::test] + async fn test_symlink_rejected_in_nested_directory() { + let dir = tempfile::tempdir().unwrap(); + + // Real skill outside the bundle + let real_dir = dir.path().join("real-skill"); + fs::create_dir(&real_dir).unwrap(); + fs::write( + real_dir.join("SKILL.md"), + "---\nname: real-skill\n---\n\nReal prompt.\n", + ) + .unwrap(); + + // Bundle directory with a symlink inside + let bundle = dir.path().join("bundle"); + fs::create_dir(&bundle).unwrap(); + std::os::unix::fs::symlink(&real_dir, bundle.join("linked-skill")).unwrap(); + + let mut registry = SkillRegistry::new(dir.path().to_path_buf()); + let loaded = registry.discover_all().await; + + // The real skill at top level is found, but the symlinked one inside bundle is rejected + assert_eq!(loaded, vec!["real-skill"]); + assert_eq!(registry.count(), 1); + } + + #[tokio::test] + async fn test_discover_nested_plus_direct() { + let dir = tempfile::tempdir().unwrap(); + + // Direct skill at depth 1 + let direct = dir.path().join("direct-skill"); + fs::create_dir(&direct).unwrap(); + fs::write( + direct.join("SKILL.md"), + "---\nname: direct-skill\n---\n\nDirect prompt.\n", + ) + .unwrap(); + + // Bundle with nested skill + let bundle = dir.path().join("bundle"); + fs::create_dir(&bundle).unwrap(); + let nested = bundle.join("nested-skill"); + fs::create_dir(&nested).unwrap(); + fs::write( + nested.join("SKILL.md"), + "---\nname: nested-skill\n---\n\nNested prompt.\n", + ) + .unwrap(); + + let mut registry = SkillRegistry::new(dir.path().to_path_buf()); + let loaded = registry.discover_all().await; + + assert_eq!(registry.count(), 2); + assert!(loaded.contains(&"direct-skill".to_string())); + assert!(loaded.contains(&"nested-skill".to_string())); + } + + #[tokio::test] + async fn test_discover_dedup_direct_vs_bundle_same_name() { + let dir = tempfile::tempdir().unwrap(); + + // Direct skill at depth 1 + let direct = dir.path().join("my-skill"); + fs::create_dir(&direct).unwrap(); + fs::write( + direct.join("SKILL.md"), + "---\nname: my-skill\n---\n\nDirect version.\n", + ) + .unwrap(); + + // Bundle directory containing a skill with the same name + let bundle = dir.path().join("org-bundle"); + fs::create_dir(&bundle).unwrap(); + let nested = bundle.join("my-skill"); + fs::create_dir(&nested).unwrap(); + fs::write( + nested.join("SKILL.md"), + "---\nname: my-skill\n---\n\nBundle version.\n", + ) + .unwrap(); + + let mut registry = SkillRegistry::new(dir.path().to_path_buf()); + let loaded = registry.discover_all().await; + + // Only one instance should survive dedup + assert_eq!(registry.count(), 1); + assert_eq!(loaded, vec!["my-skill"]); + } + + #[tokio::test] + async fn test_global_cap_shared_across_sources() { + // The global cap (MAX_DISCOVERED_SKILLS=100) is shared across all + // sources. Workspace skills are discovered first, consuming part of + // the budget, leaving less for user skills. + let user_dir = tempfile::tempdir().unwrap(); + let ws_dir = tempfile::tempdir().unwrap(); + + // 10 workspace skills (discovered first, highest priority) + for i in 0..10 { + let skill_dir = ws_dir.path().join(format!("ws-skill-{:02}", i)); + fs::create_dir(&skill_dir).unwrap(); + fs::write( + skill_dir.join("SKILL.md"), + format!("---\nname: ws-skill-{:02}\n---\n\nPrompt.\n", i), + ) + .unwrap(); + } + + // 120 user skills (more than the remaining budget of 90) + for i in 0..120 { + let skill_dir = user_dir.path().join(format!("user-skill-{:03}", i)); + fs::create_dir(&skill_dir).unwrap(); + fs::write( + skill_dir.join("SKILL.md"), + format!("---\nname: user-skill-{:03}\n---\n\nPrompt.\n", i), + ) + .unwrap(); + } + + let mut registry = SkillRegistry::new(user_dir.path().to_path_buf()) + .with_workspace_dir(ws_dir.path().to_path_buf()); + registry.discover_all().await; + + // Total capped at 100 globally + assert_eq!(registry.count(), MAX_DISCOVERED_SKILLS); + + // All 10 workspace skills must be present (discovered first) + for i in 0..10 { + assert!( + registry + .find_by_name(&format!("ws-skill-{:02}", i)) + .is_some(), + "workspace skill ws-skill-{:02} should be discoverable", + i + ); + } + } } From d97c0145cf8b67b1b2be4ddbb49f798d5fa90d33 Mon Sep 17 00:00:00 2001 From: Henry Park Date: Sun, 29 Mar 2026 02:55:34 -0700 Subject: [PATCH 10/23] Clarify message tool vs channel setup guidance (#1715) * Clarify message tool and channel setup guidance * Add target format hints to proactive messaging prompt * Clarify search and message tool edge cases * Fix stale tool_search e2e assertion * Update src/llm/reasoning.rs Co-authored-by: Copilot <175728472+Copilot@users.noreply.github.com> * Tighten prompt guidance for message replies * Format prompt guidance assertions --------- Co-authored-by: Copilot <175728472+Copilot@users.noreply.github.com> --- src/extensions/mod.rs | 3 +- src/llm/reasoning.rs | 70 +++++++++++++++-- src/tools/builtin/extension_tools.rs | 18 ++++- src/tools/builtin/message.rs | 18 +++-- tests/e2e_builtin_tool_coverage.rs | 110 ++++++++++++++++++++++++++- tests/multi_tenant_system_prompt.rs | 43 +++++++++++ 6 files changed, 247 insertions(+), 15 deletions(-) diff --git a/src/extensions/mod.rs b/src/extensions/mod.rs index 4c32767b48b..2602884d457 100644 --- a/src/extensions/mod.rs +++ b/src/extensions/mod.rs @@ -2,7 +2,8 @@ //! and activation of channels, tools, and MCP servers. //! //! Extensions are the user-facing abstraction that unifies three runtime kinds: -//! - **Channels** (Telegram, Slack, Discord) — messaging integrations (WASM) +//! - **Channels** (Telegram, Slack, Discord) — messaging platform connections +//! and conversation transports (WASM) //! - **Tools** — sandboxed capabilities (WASM) //! - **MCP servers** — external API integrations via Model Context Protocol //! diff --git a/src/llm/reasoning.rs b/src/llm/reasoning.rs index a0852cefff9..f5fc6b8a132 100644 --- a/src/llm/reasoning.rs +++ b/src/llm/reasoning.rs @@ -1017,8 +1017,11 @@ Example: "\n\n## Extensions\n\ You can search, install, and activate extensions to add new capabilities:\n\ - - **Channels** (Telegram, Slack, Discord) — messaging integrations. \ - When users ask about connecting a messaging platform, search for it as a channel.\n\ + - **Channels** (Telegram, Slack, Discord) — connect messaging platforms so users can \ + talk to you there. When users ask about connecting a messaging platform, search for it \ + as a channel. Channels are not separate send-message tools; use normal assistant output \ + to reply in the current conversation, and use the `message` tool only for proactive, \ + background, or cross-channel outbound sends.\n\ - **Tools** — sandboxed functions that extend your abilities.\n\ - **MCP servers** — external API integrations via the Model Context Protocol.\n\n\ Use `tool_search` to find extensions by name. Refer to them by their kind \ @@ -1059,15 +1062,20 @@ Example: let message_tool_hint = "\ \n\n## Proactive Messaging\n\ +For ordinary replies in the current conversation, respond normally without calling `message`.\n\ Send messages via Signal, Telegram, Slack, or other connected channels:\n\ - `content` (required): the message text\n\ - `attachments` (optional): array of file paths to send\n\ - `channel` (optional): which channel to use (signal, telegram, slack, etc.)\n\ - `target` (optional): who to send to (phone number, group ID, etc.)\n\ -\nOmit both `channel` and `target` to send to the current conversation.\n\ +\nOmit both `channel` and `target` for a proactive follow-up in the current conversation.\n\ +Target formats:\n\ +- Signal: E.164 phone number (`+1234567890`) or group ID\n\ +- Telegram: username or chat ID\n\ +- Slack: channel name (`#general`) or user ID\n\ Examples (tool calls use JSON format):\n\ -- Reply here: {\"content\": \"Hi!\"}\n\ -- Send file here: {\"content\": \"Here's the file\", \"attachments\": [\"/path/to/file.txt\"]}\n\ +- Proactive follow-up here: {\"content\": \"Hi again!\"}\n\ +- Send file here proactively: {\"content\": \"Here's the file\", \"attachments\": [\"/path/to/file.txt\"]}\n\ - Message a different user: {\"channel\": \"signal\", \"target\": \"+1234567890\", \"content\": \"Hi!\"}\n\ - Message a different group: {\"channel\": \"signal\", \"target\": \"group:abc123\", \"content\": \"Hi!\"}"; @@ -1105,7 +1113,9 @@ Examples (tool calls use JSON format):\n\ format!( "\n\n## Current Conversation\n\ - This is who you're talking to (omit 'target' to send here):\n{}", + This is who you're talking to in the active conversation. Use normal assistant \ + output to reply here; only use the `message` tool for proactive, background, or \ + cross-channel outbound sends:\n{}", lines.join("\n") ) } @@ -2506,6 +2516,54 @@ That's my plan."#; ); } + #[test] + fn test_extensions_section_clarifies_channels_are_not_send_tools() { + let reasoning = make_test_reasoning(); + let tool_defs = vec![ToolDefinition { + name: "tool_search".to_string(), + description: "Search extensions".to_string(), + parameters: serde_json::json!({}), + }]; + + let section = reasoning.build_extensions_section_for_tools(&tool_defs); + assert!(section.contains("connect messaging platforms so users can talk to you there")); + assert!(section.contains("Channels are not separate send-message tools")); + assert!( + section.contains("use normal assistant output to reply in the current conversation") + ); + assert!(section.contains( + "`message` tool only for proactive, background, or cross-channel outbound sends" + )); + } + + #[test] + fn test_channel_section_separates_normal_replies_from_message_tool() { + let reasoning = make_test_reasoning().with_channel("telegram"); + + let section = reasoning.build_channel_section(); + assert!(section.contains("respond normally without calling `message`")); + assert!(section.contains("proactive follow-up in the current conversation")); + assert!(section.contains("Target formats:")); + assert!(section.contains("Signal: E.164 phone number")); + assert!(section.contains("Telegram: username or chat ID")); + assert!(section.contains("Slack: channel name")); + assert!(section.contains("Proactive follow-up here")); + } + + #[test] + fn test_current_conversation_section_does_not_imply_message_tool_for_replies() { + let reasoning = make_test_reasoning() + .with_channel("telegram") + .with_conversation_data("User", "telegram-user"); + + let section = reasoning.build_conversation_section(); + assert!(section.contains("Use normal assistant output to reply here")); + assert!(section.contains( + "only use the `message` tool for proactive, background, or cross-channel outbound sends" + )); + assert!(!section.contains("omit 'target' to send here")); + } + // ---- plan/evaluate bypass clean_response (Bug #564-2) ---- #[test] diff --git a/src/tools/builtin/extension_tools.rs b/src/tools/builtin/extension_tools.rs index fba61613bda..6d6efa4b620 100644 --- a/src/tools/builtin/extension_tools.rs +++ b/src/tools/builtin/extension_tools.rs @@ -31,8 +31,11 @@ impl Tool for ToolSearchTool { fn description(&self) -> &str { "Search for available extensions to add new capabilities. Extensions include \ - channels (Telegram, Slack, Discord — for messaging), tools, and MCP servers. \ - Use discover:true to search online if the built-in registry has no results." + channels (Telegram, Slack, Discord — connect messaging platforms so IronClaw can \ + receive and reply there), tools, and MCP servers. Use `tool_install` and \ + `tool_activate` to install and enable channels; use the `message` tool for proactive \ + outbound sends. Use discover:true to search online if the built-in registry has no \ + results." } fn parameters_schema(&self) -> serde_json::Value { @@ -634,6 +637,17 @@ mod tests { assert!(schema["properties"].get("query").is_some()); } + #[test] + fn test_tool_search_description_clarifies_channel_setup_vs_sending() { + let tool = ToolSearchTool { + manager: test_manager_stub(), + }; + + let description = tool.description(); + assert!(description.contains("Use `tool_install` and `tool_activate`")); + assert!(description.contains("use the `message` tool for proactive outbound sends")); + } + #[test] fn test_tool_install_schema() { use crate::tools::tool::ApprovalRequirement; diff --git a/src/tools/builtin/message.rs b/src/tools/builtin/message.rs index d42de4cebae..33f1b74563b 100644 --- a/src/tools/builtin/message.rs +++ b/src/tools/builtin/message.rs @@ -181,11 +181,15 @@ impl Tool for MessageTool { } fn description(&self) -> &str { - "Send a message to a channel. If channel/target omitted, uses the current conversation's \ - channel and sender/group. Use to proactively message users on any connected channel. \ + "Send a proactive message to a channel. Use normal assistant output to reply in the \ + active conversation; use this tool for proactive notifications, routine/background \ + follow-ups, attachments, or sending to a different channel/recipient. If channel/target \ + are omitted, reuses the current conversation's channel and sender/group when available. \ + If you provide `target` without `channel` and no scoped channel can be resolved, the \ + message may be broadcast across connected channels instead of sent to just one. \ Supports file attachments: first download the file with the http tool using save_to \ - (e.g., http GET https://picsum.photos/800/600 save_to=/tmp/photo.jpg), then pass \ - the file path in the attachments array. Images are sent as photos on Telegram. \ + (e.g., http GET https://picsum.photos/800/600 save_to=/tmp/photo.jpg), then pass the \ + file path in the attachments array. Images are sent as photos on Telegram. \ - Signal: target accepts E.164 (+1234567890) or group ID \ - Telegram: target accepts username or chat ID \ - Slack: target accepts channel (#general) or user ID" @@ -460,7 +464,11 @@ mod tests { #[test] fn message_tool_description() { let tool = MessageTool::new(Arc::new(ChannelManager::new())); - assert!(!tool.description().is_empty()); + let description = tool.description(); + assert!(!description.is_empty()); + assert!(description.contains("Use normal assistant output to reply")); + assert!(description.contains("proactive notifications")); + assert!(description.contains("provide `target` without `channel`")); } #[test] diff --git a/tests/e2e_builtin_tool_coverage.rs b/tests/e2e_builtin_tool_coverage.rs index 7c0c7bc7855..a6781d1ee3c 100644 --- a/tests/e2e_builtin_tool_coverage.rs +++ b/tests/e2e_builtin_tool_coverage.rs @@ -13,7 +13,7 @@ mod tests { use ironclaw::agent::routine::{RoutineAction, Trigger}; use crate::support::test_rig::TestRigBuilder; - use crate::support::trace_llm::LlmTrace; + use crate::support::trace_llm::{LlmTrace, TraceResponse, TraceStep, TraceToolCall, TraceTurn}; // ----------------------------------------------------------------------- // Test 1: time_parse_and_diff @@ -875,4 +875,112 @@ mod tests { rig.shutdown(); } + + #[tokio::test] + async fn tool_info_clarifies_message_and_channel_setup_roles() { + let trace = LlmTrace::new( + "test-tool-info-channel-message-clarity", + vec![TraceTurn { + user_input: "How do message and channels differ?".to_string(), + steps: vec![ + TraceStep { + request_hint: None, + response: TraceResponse::ToolCalls { + tool_calls: vec![TraceToolCall { + id: "call_tool_info_message".to_string(), + name: "tool_info".to_string(), + arguments: serde_json::json!({"name": "message"}), + }], + input_tokens: 100, + output_tokens: 20, + }, + expected_tool_results: Vec::new(), + }, + TraceStep { + request_hint: None, + response: TraceResponse::ToolCalls { + tool_calls: vec![TraceToolCall { + id: "call_tool_info_tool_search".to_string(), + name: "tool_info".to_string(), + arguments: serde_json::json!({"name": "tool_search"}), + }], + input_tokens: 140, + output_tokens: 20, + }, + expected_tool_results: Vec::new(), + }, + TraceStep { + request_hint: None, + response: TraceResponse::Text { + content: "I checked both tool descriptions.".to_string(), + input_tokens: 220, + output_tokens: 30, + }, + expected_tool_results: Vec::new(), + }, + ], + expects: Default::default(), + }], + ); + + let rig = TestRigBuilder::new() + .with_trace(trace.clone()) + .with_auto_approve_tools(true) + .build() + .await; + + rig.send_message("How do message and channels differ?") + .await; + let responses = rig.wait_for_responses(1, Duration::from_secs(15)).await; + + rig.verify_trace_expects(&trace, &responses); + + let results = rig.tool_results(); + let info_results: Vec<_> = results.iter().filter(|(n, _)| n == "tool_info").collect(); + assert_eq!(info_results.len(), 2, "Expected two tool_info results"); + + let info_json: Vec = info_results + .iter() + .map(|(_, preview)| { + serde_json::from_str(preview) + .expect("tool_info result preview should be valid JSON") + }) + .collect(); + + let message_json = info_json + .iter() + .find(|info| info["name"] == "message") + .expect("tool_info result should contain 'message'"); + let message_description = message_json["description"] + .as_str() + .expect("message description should be a string"); + assert!( + message_description.contains("Use normal assistant output to reply"), + "message description should distinguish normal replies: {message_description}" + ); + assert!( + message_description.contains("proactive notifications"), + "message description should describe proactive sends: {message_description}" + ); + + let tool_search_json = info_json + .iter() + .find(|info| info["name"] == "tool_search") + .expect("tool_info result should contain 'tool_search'"); + let tool_search_description = tool_search_json["description"] + .as_str() + .expect("tool_search description should be a string"); + assert!( + tool_search_description.contains("`tool_install`") + && tool_search_description.contains("`tool_activate`"), + "tool_search description should describe setup/activation via tool_install and \ + tool_activate: {tool_search_description}" + ); + assert!( + tool_search_description.contains("use the `message` tool for proactive outbound sends"), + "tool_search description should point outbound sends to message: {tool_search_description}" + ); + + rig.shutdown(); + } } diff --git a/tests/multi_tenant_system_prompt.rs b/tests/multi_tenant_system_prompt.rs index b89e6cb5ca9..0f6aa697141 100644 --- a/tests/multi_tenant_system_prompt.rs +++ b/tests/multi_tenant_system_prompt.rs @@ -237,4 +237,47 @@ mod tests { rig.shutdown(); } + + #[tokio::test] + async fn telegram_system_prompt_clarifies_reply_vs_proactive_message_tool() { + let trace = simple_trace(1); + let rig = TestRigBuilder::new().with_trace(trace).build().await; + + let msg = IncomingMessage::new("telegram", "telegram-user", "Hello there"); + rig.send_incoming(msg).await; + let _responses = rig.wait_for_responses(1, TIMEOUT).await; + + let requests = rig.captured_llm_requests(); + let system_prompt = + extract_system_prompt(&requests).expect("Expected a system prompt in the LLM request"); + + assert!( + system_prompt.contains("Channels are not separate send-message tools"), + "System prompt should describe channels as setup/integration surfaces.\n\ + Actual system prompt:\n{system_prompt}" + ); + assert!( + system_prompt + .contains("use normal assistant output to reply in the current conversation"), + "System prompt should route ordinary replies through normal assistant output.\n\ + Actual system prompt:\n{system_prompt}" + ); + assert!( + system_prompt.contains("respond normally without calling `message`"), + "System prompt should say normal replies do not use the message tool.\n\ + Actual system prompt:\n{system_prompt}" + ); + assert!( + system_prompt.contains("proactive follow-up in the current conversation"), + "System prompt should reserve omitted channel/target for proactive follow-ups.\n\ + Actual system prompt:\n{system_prompt}" + ); + assert!( + !system_prompt.contains("omit 'target' to send here"), + "System prompt should not imply the message tool is the default way to reply \ + in-thread.\nActual system prompt:\n{system_prompt}" + ); + + rig.shutdown(); + } } From fcab4f0adaa94feb533ace75afe53838aecb8390 Mon Sep 17 00:00:00 2001 From: Henry Park Date: Sun, 29 Mar 2026 02:58:49 -0700 Subject: [PATCH 11/23] fix: pin staging ci jobs to a single tested sha (#1628) * fix: pin staging ci jobs to a single tested sha * chore(ci): retrigger regression gate [skip-regression-check] --------- Co-authored-by: Firat Sertgoz --- .github/workflows/e2e.yml | 9 +++++++++ .github/workflows/staging-ci.yml | 10 +++++++--- .github/workflows/test.yml | 20 ++++++++++++++++++++ 3 files changed, 36 insertions(+), 3 deletions(-) diff --git a/.github/workflows/e2e.yml b/.github/workflows/e2e.yml index bc705df7280..a008979bb80 100644 --- a/.github/workflows/e2e.yml +++ b/.github/workflows/e2e.yml @@ -1,6 +1,11 @@ name: E2E Tests on: workflow_call: + inputs: + ref: + description: Commit SHA or ref to test + required: false + type: string schedule: - cron: "0 6 * * 1" # Weekly Monday 6 AM UTC workflow_dispatch: @@ -19,6 +24,8 @@ jobs: timeout-minutes: 30 steps: - uses: actions/checkout@v6 + with: + ref: ${{ inputs.ref || github.sha }} - uses: dtolnay/rust-toolchain@stable @@ -59,6 +66,8 @@ jobs: files: "tests/e2e/scenarios/test_owner_scope.py tests/e2e/scenarios/test_routine_event_batch.py" steps: - uses: actions/checkout@v6 + with: + ref: ${{ inputs.ref || github.sha }} - name: Download binary uses: actions/download-artifact@v4 diff --git a/.github/workflows/staging-ci.yml b/.github/workflows/staging-ci.yml index 2df7bf6f70d..99ee2bfc63a 100644 --- a/.github/workflows/staging-ci.yml +++ b/.github/workflows/staging-ci.yml @@ -62,7 +62,7 @@ jobs: steps: - uses: actions/checkout@v6 with: - ref: staging + ref: ${{ github.sha }} fetch-depth: 0 fetch-tags: true @@ -117,6 +117,8 @@ jobs: needs: check-changes if: needs.check-changes.outputs.has_changes == 'true' uses: ./.github/workflows/test.yml + with: + ref: ${{ needs.check-changes.outputs.current_head }} # ── Run E2E browser tests ──────────────────────────────────────── e2e: @@ -124,6 +126,8 @@ jobs: needs: check-changes if: needs.check-changes.outputs.has_changes == 'true' uses: ./.github/workflows/e2e.yml + with: + ref: ${{ needs.check-changes.outputs.current_head }} # ── Create promotion PR (triggers claude-review.yml on the PR) ── create-promotion-pr: @@ -137,7 +141,7 @@ jobs: steps: - uses: actions/checkout@v6 with: - ref: staging + ref: ${{ needs.check-changes.outputs.current_head }} fetch-depth: 0 - name: Generate GitHub App token @@ -163,7 +167,7 @@ jobs: PROMOTION_BASE: ${{ needs.resolve-promotion-base.outputs.promotion_base }} run: | git fetch origin "${PROMOTION_BASE}" - AHEAD=$(git rev-list --count "origin/${PROMOTION_BASE}..origin/staging") + AHEAD=$(git rev-list --count "origin/${PROMOTION_BASE}..HEAD") echo "commits_ahead=${AHEAD}" >> "$GITHUB_OUTPUT" if [ "$AHEAD" -eq 0 ]; then echo "Staging is not ahead of ${PROMOTION_BASE}. Nothing to promote." diff --git a/.github/workflows/test.yml b/.github/workflows/test.yml index 5d4eabc0e8c..27a32a502c0 100644 --- a/.github/workflows/test.yml +++ b/.github/workflows/test.yml @@ -1,6 +1,11 @@ name: Run Tests on: workflow_call: + inputs: + ref: + description: Commit SHA or ref to test + required: false + type: string pull_request: branches: - main @@ -29,6 +34,8 @@ jobs: steps: - name: Checkout repository uses: actions/checkout@v6 + with: + ref: ${{ inputs.ref || github.sha }} - name: Install Rust uses: dtolnay/rust-toolchain@stable with: @@ -52,6 +59,8 @@ jobs: steps: - name: Checkout repository uses: actions/checkout@v6 + with: + ref: ${{ inputs.ref || github.sha }} - name: Install Rust uses: dtolnay/rust-toolchain@stable with: @@ -80,6 +89,8 @@ jobs: steps: - name: Checkout repository uses: actions/checkout@v6 + with: + ref: ${{ inputs.ref || github.sha }} - name: Install Rust uses: dtolnay/rust-toolchain@stable - uses: Swatinem/rust-cache@v2 @@ -107,6 +118,8 @@ jobs: steps: - name: Checkout repository uses: actions/checkout@v6 + with: + ref: ${{ inputs.ref || github.sha }} - name: Install Rust uses: dtolnay/rust-toolchain@stable - uses: Swatinem/rust-cache@v2 @@ -125,6 +138,8 @@ jobs: steps: - name: Checkout repository uses: actions/checkout@v6 + with: + ref: ${{ inputs.ref || github.sha }} - name: Install Rust uses: dtolnay/rust-toolchain@stable with: @@ -147,6 +162,8 @@ jobs: steps: - name: Checkout repository uses: actions/checkout@v6 + with: + ref: ${{ inputs.ref || github.sha }} - name: Install Rust uses: dtolnay/rust-toolchain@stable - uses: Swatinem/rust-cache@v2 @@ -164,6 +181,8 @@ jobs: steps: - name: Checkout repository uses: actions/checkout@v6 + with: + ref: ${{ inputs.ref || github.sha }} - name: Build Docker image run: docker build -t ironclaw-test:ci . @@ -175,6 +194,7 @@ jobs: - name: Checkout repository uses: actions/checkout@v6 with: + ref: ${{ inputs.ref || github.sha }} fetch-depth: 0 - name: Check version bumps for changed extensions env: From 86389dab23ef4c49c56caf8eb0e9da451d916798 Mon Sep 17 00:00:00 2001 From: Nige Date: Sun, 29 Mar 2026 11:01:34 +0100 Subject: [PATCH 12/23] fix(gemini): preserve thought signatures on all tool calls (#1565) --- src/llm/gemini_oauth.rs | 35 +++++++++++++++++++++++++++++++++-- 1 file changed, 33 insertions(+), 2 deletions(-) diff --git a/src/llm/gemini_oauth.rs b/src/llm/gemini_oauth.rs index a19eec1291f..a12e5f82fba 100644 --- a/src/llm/gemini_oauth.rs +++ b/src/llm/gemini_oauth.rs @@ -972,7 +972,7 @@ impl GeminiOauthProvider { None => return, }; - // For each model turn in the active loop, ensure the first functionCall has a thoughtSignature. + // For each model turn in the active loop, ensure functionCall parts have a thoughtSignature. for item in contents.iter_mut().skip(start) { let is_model = item.get("role").and_then(|r| r.as_str()) == Some("model"); if !is_model { @@ -992,7 +992,6 @@ impl GeminiOauthProvider { ); } modified = true; - break; // Only the first functionCall } } if modified { @@ -2583,4 +2582,36 @@ mod tests { assert_eq!(curated[0]["parts"][0]["text"], "hello"); assert_eq!(curated[1]["parts"][0]["text"], "again"); } + + #[test] + fn test_ensure_thought_signatures_adds_signatures_to_all_function_calls() { + let mut contents = vec![ + serde_json::json!({ + "role": "user", + "parts": [{ "text": "call tools" }] + }), + serde_json::json!({ + "role": "model", + "parts": [ + { "functionCall": { "name": "memory_write", "args": { "key": "a" } } }, + { "functionCall": { "name": "memory_write", "args": { "key": "b" } } } + ] + }), + ]; + + GeminiOauthProvider::ensure_thought_signatures(&mut contents); + + let parts = contents[1] + .get("parts") + .and_then(|p| p.as_array()) + .expect("model turn should have parts"); + + let signed_calls = parts + .iter() + .filter(|part| part.get("functionCall").is_some()) + .filter(|part| part.get("thoughtSignature").is_some()) + .count(); + + assert_eq!(signed_calls, 2); // safety: test-only assertion + } } From 64fe9ba6077e3478e2684daf97d43556aff6996d Mon Sep 17 00:00:00 2001 From: Zaki Manian Date: Sun, 29 Mar 2026 13:49:21 -0700 Subject: [PATCH 13/23] 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 14/23] 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 15/23] 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 16/23] 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 17/23] 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 18/23] 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 19/23] 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 20/23] 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 21/23] 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 22/23] 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 23/23] 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}"