diff --git a/.github/workflows/e2e.yml b/.github/workflows/e2e.yml index a008979bb80..65ba6c3d12f 100644 --- a/.github/workflows/e2e.yml +++ b/.github/workflows/e2e.yml @@ -61,7 +61,7 @@ jobs: - group: features files: "tests/e2e/scenarios/test_skills.py tests/e2e/scenarios/test_tool_approval.py tests/e2e/scenarios/test_webhook.py" - group: extensions - files: "tests/e2e/scenarios/test_extensions.py tests/e2e/scenarios/test_extension_oauth.py tests/e2e/scenarios/test_oauth_url_parameters.py tests/e2e/scenarios/test_telegram_token_validation.py tests/e2e/scenarios/test_telegram_hot_activation.py tests/e2e/scenarios/test_wasm_lifecycle.py tests/e2e/scenarios/test_tool_execution.py tests/e2e/scenarios/test_pairing.py tests/e2e/scenarios/test_mcp_auth_flow.py tests/e2e/scenarios/test_oauth_credential_fallback.py tests/e2e/scenarios/test_routine_oauth_credential_injection.py" + files: "tests/e2e/scenarios/test_extensions.py tests/e2e/scenarios/test_extension_oauth.py tests/e2e/scenarios/test_oauth_url_parameters.py tests/e2e/scenarios/test_telegram_token_validation.py tests/e2e/scenarios/test_telegram_hot_activation.py tests/e2e/scenarios/test_wasm_lifecycle.py tests/e2e/scenarios/test_tool_execution.py tests/e2e/scenarios/test_agent_loop_recovery.py tests/e2e/scenarios/test_pairing.py tests/e2e/scenarios/test_mcp_auth_flow.py tests/e2e/scenarios/test_oauth_credential_fallback.py tests/e2e/scenarios/test_routine_oauth_credential_injection.py" - group: routines files: "tests/e2e/scenarios/test_owner_scope.py tests/e2e/scenarios/test_routine_event_batch.py" steps: diff --git a/FEATURE_PARITY.md b/FEATURE_PARITY.md index 915f529e4ea..3296228c7ba 100644 --- a/FEATURE_PARITY.md +++ b/FEATURE_PARITY.md @@ -97,6 +97,7 @@ This document tracks feature parity between IronClaw (Rust implementation) and O | Cron/heartbeat topic targeting | ✅ | ❌ | Messages land in correct topic | | DM topics support | ✅ | ❌ | Agent/topic bindings in DMs and agent-scoped SessionKeys | | Persistent ACP topic binding | ✅ | ❌ | ACP harness sessions can pin to Telegram forum or DM topics | +| sendVoice (voice note replies) | ✅ | ✅ | audio/ogg attachments sent as voice notes; prerequisite for TTS (#90) | ### Discord-Specific Features (since Feb 2025) diff --git a/channels-src/telegram/Cargo.lock b/channels-src/telegram/Cargo.lock index 8d40f01e0ff..7ef6912c7bd 100644 --- a/channels-src/telegram/Cargo.lock +++ b/channels-src/telegram/Cargo.lock @@ -212,7 +212,7 @@ dependencies = [ [[package]] name = "telegram-channel" -version = "0.2.1" +version = "0.2.6" dependencies = [ "serde", "serde_json", diff --git a/channels-src/telegram/Cargo.toml b/channels-src/telegram/Cargo.toml index 182e5f5de5d..982329246f2 100644 --- a/channels-src/telegram/Cargo.toml +++ b/channels-src/telegram/Cargo.toml @@ -1,6 +1,6 @@ [package] name = "telegram-channel" -version = "0.2.1" +version = "0.2.6" edition = "2021" description = "Telegram Bot API channel for IronClaw" license = "MIT OR Apache-2.0" diff --git a/channels-src/telegram/src/lib.rs b/channels-src/telegram/src/lib.rs index f34ed68aa7a..4b5c4ea6062 100644 --- a/channels-src/telegram/src/lib.rs +++ b/channels-src/telegram/src/lib.rs @@ -407,7 +407,8 @@ fn split_message(text: &str) -> Vec { let window = &remaining[..window_bytes]; // 1. Double newline — best paragraph boundary - let split_at = window.rfind("\n\n") + let split_at = window + .rfind("\n\n") // 2. Single newline .or_else(|| window.rfind('\n')) // 3. Sentence-ending punctuation followed by space. @@ -417,9 +418,9 @@ fn split_message(text: &str) -> Vec { .or_else(|| { let bytes = window.as_bytes(); // Search backwards for '. ', '! ', '? ' - (1..bytes.len()).rev().find(|&i| { - matches!(bytes[i - 1], b'.' | b'!' | b'?') && bytes[i] == b' ' - }) + (1..bytes.len()) + .rev() + .find(|&i| matches!(bytes[i - 1], b'.' | b'!' | b'?') && bytes[i] == b' ') }) // 4. Word boundary (last space) .or_else(|| window.rfind(' ')) @@ -427,7 +428,11 @@ fn split_message(text: &str) -> Vec { .unwrap_or(window_bytes); // Avoid empty chunks (e.g. text starting with \n\n). - let split_at = if split_at == 0 { window_bytes } else { split_at }; + let split_at = if split_at == 0 { + window_bytes + } else { + split_at + }; // Trim whitespace at chunk boundaries for clean Telegram display. // Note: this drops leading/trailing spaces at split points, which is @@ -1090,12 +1095,9 @@ fn download_telegram_file(file_id: &str) -> Result, String> { } // ============================================================================ -// Attachment Sending (Photo / Document) +// Attachment Sending (Photo / Voice / Document) // ============================================================================ -/// Maximum photo size for Telegram sendPhoto (10 MB). -const MAX_PHOTO_SIZE: usize = 10 * 1024 * 1024; - /// Write a multipart/form-data text field. fn write_multipart_field(body: &mut Vec, boundary: &str, name: &str, value: &str) { body.extend_from_slice(format!("--{}\r\n", boundary).as_bytes()); @@ -1138,10 +1140,27 @@ fn write_multipart_file( body.extend_from_slice(b"\r\n"); } -/// Send a photo via the Telegram Bot API (multipart upload). +/// Image MIME types that Telegram's sendPhoto API supports. +const PHOTO_MIME_TYPES: &[&str] = &["image/jpeg", "image/png", "image/gif", "image/webp"]; + +/// Audio MIME types that Telegram's sendVoice API supports (ogg/opus container). +const VOICE_MIME_TYPES: &[&str] = &["audio/ogg", "audio/opus"]; + +/// Maximum photo size for Telegram sendPhoto (10 MB). +const MAX_PHOTO_SIZE: usize = 10 * 1024 * 1024; + +/// Maximum voice note size for Telegram sendVoice (50 MB). +const MAX_VOICE_SIZE: usize = 50 * 1024 * 1024; + +/// Send a multipart file upload to a Telegram Bot API endpoint. /// -/// Falls back to `send_document()` if the photo exceeds 10 MB. -fn send_photo( +/// Shared implementation for sendPhoto, sendVoice, and sendDocument. +/// `api_method` is the Telegram method name (e.g. "sendPhoto"), +/// `field_name` is the multipart field (e.g. "photo", "voice", "document"). +#[allow(clippy::too_many_arguments)] +fn send_multipart_upload( + api_method: &str, + field_name: &str, chat_id: i64, filename: &str, mime_type: &str, @@ -1151,25 +1170,6 @@ fn send_photo( ) -> Result<(), String> { let message_thread_id = normalize_thread_id(message_thread_id); - if data.len() > MAX_PHOTO_SIZE { - channel_host::log( - channel_host::LogLevel::Info, - &format!( - "Photo {} exceeds 10MB ({}), sending as document", - filename, - data.len() - ), - ); - return send_document( - chat_id, - filename, - mime_type, - data, - reply_to_message_id, - message_thread_id, - ); - } - let boundary = format!("ironclaw-{}", channel_host::now_millis()); let mut body = Vec::new(); @@ -1190,16 +1190,21 @@ fn send_photo( &thread_id.to_string(), ); } - write_multipart_file(&mut body, &boundary, "photo", filename, mime_type, data); + write_multipart_file(&mut body, &boundary, field_name, filename, mime_type, data); body.extend_from_slice(format!("--{}--\r\n", boundary).as_bytes()); let headers = serde_json::json!({ "Content-Type": format!("multipart/form-data; boundary={}", boundary) }); + let url = format!( + "https://api.telegram.org/bot{{TELEGRAM_BOT_TOKEN}}/{}", + api_method + ); + let result = channel_host::http_request( "POST", - "https://api.telegram.org/bot{TELEGRAM_BOT_TOKEN}/sendPhoto", + &url, &headers.to_string(), Some(&body), Some(60_000), // 60s timeout for file uploads @@ -1209,23 +1214,25 @@ fn send_photo( Ok(resp) if resp.status == 200 => { channel_host::log( channel_host::LogLevel::Debug, - &format!("Sent photo '{}' to chat {}", filename, chat_id), + &format!("Sent {} '{}' to chat {}", field_name, filename, chat_id), ); Ok(()) } Ok(resp) => { let body_str = String::from_utf8_lossy(&resp.body); Err(format!( - "sendPhoto failed (HTTP {}): {}", - resp.status, body_str + "{} failed (HTTP {}): {}", + api_method, resp.status, body_str )) } - Err(e) => Err(format!("sendPhoto HTTP request failed: {}", e)), + Err(e) => Err(format!("{} HTTP request failed: {}", api_method, e)), } } -/// Send a document via the Telegram Bot API (multipart upload). -fn send_document( +/// Send a photo via the Telegram Bot API (multipart upload). +/// +/// Falls back to `send_document()` if the photo exceeds 10 MB. +fn send_photo( chat_id: i64, filename: &str, mime_type: &str, @@ -1233,65 +1240,100 @@ fn send_document( reply_to_message_id: Option, message_thread_id: Option, ) -> Result<(), String> { - let message_thread_id = normalize_thread_id(message_thread_id); - - let boundary = format!("ironclaw-{}", channel_host::now_millis()); - let mut body = Vec::new(); - - write_multipart_field(&mut body, &boundary, "chat_id", &chat_id.to_string()); - if let Some(msg_id) = reply_to_message_id { - write_multipart_field( - &mut body, - &boundary, - "reply_to_message_id", - &msg_id.to_string(), + if data.len() > MAX_PHOTO_SIZE { + channel_host::log( + channel_host::LogLevel::Info, + &format!( + "Photo {} exceeds 10MB ({}), sending as document", + filename, + data.len() + ), ); - } - if let Some(thread_id) = message_thread_id { - write_multipart_field( - &mut body, - &boundary, - "message_thread_id", - &thread_id.to_string(), + return send_document( + chat_id, + filename, + mime_type, + data, + reply_to_message_id, + message_thread_id, ); } - write_multipart_file(&mut body, &boundary, "document", filename, mime_type, data); - body.extend_from_slice(format!("--{}--\r\n", boundary).as_bytes()); - - let headers = serde_json::json!({ - "Content-Type": format!("multipart/form-data; boundary={}", boundary) - }); + send_multipart_upload( + "sendPhoto", + "photo", + chat_id, + filename, + mime_type, + data, + reply_to_message_id, + message_thread_id, + ) +} - let result = channel_host::http_request( - "POST", - "https://api.telegram.org/bot{TELEGRAM_BOT_TOKEN}/sendDocument", - &headers.to_string(), - Some(&body), - Some(60_000), // 60s timeout for file uploads - ); +/// Send a document via the Telegram Bot API (multipart upload). +fn send_document( + chat_id: i64, + filename: &str, + mime_type: &str, + data: &[u8], + reply_to_message_id: Option, + message_thread_id: Option, +) -> Result<(), String> { + send_multipart_upload( + "sendDocument", + "document", + chat_id, + filename, + mime_type, + data, + reply_to_message_id, + message_thread_id, + ) +} - match result { - Ok(resp) if resp.status == 200 => { - channel_host::log( - channel_host::LogLevel::Debug, - &format!("Sent document '{}' to chat {}", filename, chat_id), - ); - Ok(()) - } - Ok(resp) => { - let body_str = String::from_utf8_lossy(&resp.body); - Err(format!( - "sendDocument failed (HTTP {}): {}", - resp.status, body_str - )) - } - Err(e) => Err(format!("sendDocument HTTP request failed: {}", e)), +/// Send a voice note via the Telegram Bot API (multipart upload). +/// +/// Telegram's `sendVoice` requires ogg/opus audio and displays it as an +/// in-chat voice note with waveform and playback controls. +/// Falls back to `send_document()` if the voice note exceeds 50 MB. +fn send_voice( + chat_id: i64, + filename: &str, + mime_type: &str, + data: &[u8], + reply_to_message_id: Option, + message_thread_id: Option, +) -> Result<(), String> { + if data.len() > MAX_VOICE_SIZE { + channel_host::log( + channel_host::LogLevel::Info, + &format!( + "Voice note {} exceeds 50MB ({}), sending as document", + filename, + data.len() + ), + ); + return send_document( + chat_id, + filename, + mime_type, + data, + reply_to_message_id, + message_thread_id, + ); } + send_multipart_upload( + "sendVoice", + "voice", + chat_id, + filename, + mime_type, + data, + reply_to_message_id, + message_thread_id, + ) } -/// Image MIME types that Telegram's sendPhoto API supports. -const PHOTO_MIME_TYPES: &[&str] = &["image/jpeg", "image/png", "image/gif", "image/webp"]; - /// Send a full agent response (attachments + text) to a chat. /// /// Shared implementation for both `on_respond` and `on_broadcast`. @@ -1321,7 +1363,13 @@ fn send_response( for (i, chunk) in chunks.into_iter().enumerate() { // Try Markdown, fall back to plain text on parse errors - let result = send_message(chat_id, &chunk, reply_to, Some("Markdown"), message_thread_id); + let result = send_message( + chat_id, + &chunk, + reply_to, + Some("Markdown"), + message_thread_id, + ); let msg_id = match result { Ok(id) => { @@ -1371,31 +1419,65 @@ fn send_response( Ok(()) } -/// Send a single attachment, choosing sendPhoto or sendDocument based on MIME type. +/// Extract the base MIME type, stripping any parameters after `;`. +/// +/// e.g. `"audio/ogg; codecs=opus"` → `"audio/ogg"` +fn base_mime_type(mime: &str) -> &str { + mime.split(';').next().unwrap_or(mime).trim() +} + +/// Attachment routing category. +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +enum AttachmentKind { + Photo, + Voice, + Document, +} + +/// Classify an attachment's send method based on its MIME type. +fn classify_attachment(mime_type: &str) -> AttachmentKind { + let base = base_mime_type(mime_type); + if PHOTO_MIME_TYPES.contains(&base) { + AttachmentKind::Photo + } else if VOICE_MIME_TYPES.contains(&base) { + AttachmentKind::Voice + } else { + AttachmentKind::Document + } +} + +/// Send a single attachment, choosing sendPhoto, sendVoice, or sendDocument based on MIME type. fn send_attachment( chat_id: i64, attachment: &Attachment, reply_to_message_id: Option, message_thread_id: Option, ) -> Result<(), String> { - if PHOTO_MIME_TYPES.contains(&attachment.mime_type.as_str()) { - send_photo( + match classify_attachment(&attachment.mime_type) { + AttachmentKind::Photo => send_photo( chat_id, &attachment.filename, &attachment.mime_type, &attachment.data, reply_to_message_id, message_thread_id, - ) - } else { - send_document( + ), + AttachmentKind::Voice => send_voice( chat_id, &attachment.filename, &attachment.mime_type, &attachment.data, reply_to_message_id, message_thread_id, - ) + ), + AttachmentKind::Document => send_document( + chat_id, + &attachment.filename, + &attachment.mime_type, + &attachment.data, + reply_to_message_id, + message_thread_id, + ), } } @@ -2969,4 +3051,38 @@ mod tests { // Verify the constant is 20 MB, matching the Slack channel limit assert_eq!(MAX_DOWNLOAD_SIZE_BYTES, 20 * 1024 * 1024); } + + #[test] + fn test_base_mime_type() { + assert_eq!(base_mime_type("audio/ogg"), "audio/ogg"); + assert_eq!(base_mime_type("audio/ogg; codecs=opus"), "audio/ogg"); + assert_eq!(base_mime_type("image/jpeg"), "image/jpeg"); + assert_eq!(base_mime_type("text/plain; charset=utf-8"), "text/plain"); + assert_eq!(base_mime_type(""), ""); + } + + #[test] + fn test_classify_attachment_routing() { + // Photos + assert_eq!(classify_attachment("image/jpeg"), AttachmentKind::Photo); + assert_eq!(classify_attachment("image/png"), AttachmentKind::Photo); + assert_eq!(classify_attachment("image/gif"), AttachmentKind::Photo); + assert_eq!(classify_attachment("image/webp"), AttachmentKind::Photo); + + // Voice notes — exact and parameterized + assert_eq!(classify_attachment("audio/ogg"), AttachmentKind::Voice); + assert_eq!(classify_attachment("audio/opus"), AttachmentKind::Voice); + assert_eq!( + classify_attachment("audio/ogg; codecs=opus"), + AttachmentKind::Voice + ); + + // Everything else falls through to document + assert_eq!( + classify_attachment("application/pdf"), + AttachmentKind::Document + ); + assert_eq!(classify_attachment("audio/mpeg"), AttachmentKind::Document); + assert_eq!(classify_attachment("video/mp4"), AttachmentKind::Document); + } } diff --git a/channels-src/telegram/telegram.capabilities.json b/channels-src/telegram/telegram.capabilities.json index 1526762dedf..13e177d2da7 100644 --- a/channels-src/telegram/telegram.capabilities.json +++ b/channels-src/telegram/telegram.capabilities.json @@ -18,6 +18,14 @@ "name": "telegram_bot_token", "prompt": "Enter your Telegram Bot API token (from @BotFather)", "optional": false + }, + { + "name": "telegram_webhook_secret", + "prompt": "Webhook secret (leave empty to auto-generate)", + "optional": true, + "auto_generate": { + "length": 64 + } } ], "setup_url": "https://t.me/BotFather", diff --git a/channels-src/whatsapp/Cargo.lock b/channels-src/whatsapp/Cargo.lock index 0e55d1e5324..adefa9aa3b2 100644 --- a/channels-src/whatsapp/Cargo.lock +++ b/channels-src/whatsapp/Cargo.lock @@ -269,7 +269,7 @@ dependencies = [ [[package]] name = "whatsapp-channel" -version = "0.1.0" +version = "0.2.0" dependencies = [ "serde", "serde_json", diff --git a/crates/ironclaw_common/src/lib.rs b/crates/ironclaw_common/src/lib.rs index f52dc0aaa6d..75c0843c288 100644 --- a/crates/ironclaw_common/src/lib.rs +++ b/crates/ironclaw_common/src/lib.rs @@ -5,3 +5,8 @@ mod util; pub use event::{AppEvent, ToolDecisionDto}; pub use util::truncate_preview; + +/// Maximum worker agent loop iterations. Used by the orchestrator (server-side +/// clamp in `create_job_inner`) and the worker runtime (`worker/job.rs`). +/// A single source of truth prevents the two from drifting. +pub const MAX_WORKER_ITERATIONS: u32 = 500; diff --git a/docs/TELEGRAM_SETUP.md b/docs/TELEGRAM_SETUP.md index f9ec24eb235..3882334e58c 100644 --- a/docs/TELEGRAM_SETUP.md +++ b/docs/TELEGRAM_SETUP.md @@ -33,7 +33,7 @@ ironclaw onboard When prompted, enable the Telegram channel and paste your bot token. The wizard will: - Validate the token -- Optionally configure a webhook secret +- Auto-generate a webhook secret for webhook mode - Set up tunnel (if you want webhook mode) ### 3. (Optional) Configure Tunnel for Webhooks diff --git a/migrations/V15__conversation_source_channel.sql b/migrations/V15__conversation_source_channel.sql new file mode 100644 index 00000000000..340b48a65d4 --- /dev/null +++ b/migrations/V15__conversation_source_channel.sql @@ -0,0 +1,4 @@ +-- Add source_channel to conversations for cross-channel approval authorization. +-- Tracks which channel originally created a conversation so that approval +-- messages from other channels can be validated. +ALTER TABLE conversations ADD COLUMN source_channel TEXT; diff --git a/registry/channels/telegram.json b/registry/channels/telegram.json index 52f66ce306d..267d9c18c34 100644 --- a/registry/channels/telegram.json +++ b/registry/channels/telegram.json @@ -2,7 +2,7 @@ "name": "telegram", "display_name": "Telegram Channel", "kind": "channel", - "version": "0.2.5", + "version": "0.2.6", "wit_version": "0.3.0", "description": "Talk to your agent through a Telegram bot", "keywords": [ diff --git a/registry/tools/github.json b/registry/tools/github.json index bb351259603..5af24523601 100644 --- a/registry/tools/github.json +++ b/registry/tools/github.json @@ -4,13 +4,15 @@ "kind": "tool", "version": "0.2.2", "wit_version": "0.3.0", - "description": "GitHub integration for issues, PRs, repos, and code search", + "description": "GitHub integration for repositories, issues, pull requests, search, branches, file writes, releases, and workflows", "keywords": [ "git", "code", "issues", "pull-requests", - "repositories" + "repositories", + "search", + "releases" ], "source": { "dir": "tools-src/github", diff --git a/src/agent/agent_loop.rs b/src/agent/agent_loop.rs index 59f1f87d342..6a4a7a2660a 100644 --- a/src/agent/agent_loop.rs +++ b/src/agent/agent_loop.rs @@ -843,7 +843,10 @@ impl Agent { { use crate::agent::session::Thread; let mut sess = session.lock().await; - let thread = Thread::with_id(id, sess.id); + // Bootstrap thread has no incoming message -- use the + // "__bootstrap__" sentinel so approvals from any channel are + // permitted. None means "deny by default" (fail-closed). + let thread = Thread::with_id(id, sess.id, Some("__bootstrap__")); sess.active_thread = Some(id); sess.threads.entry(id).or_insert(thread); } @@ -1148,7 +1151,37 @@ impl Agent { .get_or_create_session(&message.user_id) .await; let mut sess = session.lock().await; - if sess.threads.contains_key(&target_thread_id) { + if let Some(thread) = sess.threads.get(&target_thread_id) { + // Verify the thread actually has a pending approval before + // allowing approval-shaped messages to target it. Without this + // check, an attacker could use approval messages to hijack any + // thread by UUID. + if thread.pending_approval.is_none() { + tracing::warn!( + %target_thread_id, + approval_channel = %message.channel, + "Blocked approval for thread with no pending approval" + ); + drop(sess); + return Ok(Some("Error: no pending approval on this thread".into())); + } + + let authorized = crate::agent::session::is_approval_authorized( + thread.source_channel.as_deref(), + &message.channel, + ); + if !authorized { + tracing::warn!( + %target_thread_id, + source_channel = ?thread.source_channel, + approval_channel = %message.channel, + "Blocked cross-channel approval attempt" + ); + drop(sess); + return Ok(Some( + "Error: approval not authorized for this channel".into(), + )); + } sess.active_thread = Some(target_thread_id); sess.last_active_at = chrono::Utc::now(); drop(sess); diff --git a/src/agent/compaction.rs b/src/agent/compaction.rs index 30bb2b6c64e..c69f9608b54 100644 --- a/src/agent/compaction.rs +++ b/src/agent/compaction.rs @@ -319,7 +319,7 @@ mod tests { #[test] fn test_format_turns() { - let mut thread = Thread::new(Uuid::new_v4()); + let mut thread = Thread::new(Uuid::new_v4(), None); thread.start_turn("Hello"); thread.complete_turn("Hi there"); thread.start_turn("How are you?"); @@ -351,7 +351,7 @@ mod tests { /// Helper: build a thread with `n` completed turns. /// Turn `i` has user_input "msg-{i}" and response "resp-{i}". fn make_thread(n: usize) -> Thread { - let mut thread = Thread::new(Uuid::new_v4()); + let mut thread = Thread::new(Uuid::new_v4(), None); for i in 0..n { thread.start_turn(format!("msg-{}", i)); thread.complete_turn(format!("resp-{}", i)); @@ -457,7 +457,7 @@ mod tests { async fn test_compact_truncate_empty_turns() { let llm = Arc::new(StubLlm::new("unused")); let compactor = make_compactor(llm); - let mut thread = Thread::new(Uuid::new_v4()); + let mut thread = Thread::new(Uuid::new_v4(), None); assert!(thread.turns.is_empty()); let result = compactor @@ -698,7 +698,7 @@ mod tests { #[test] fn test_format_turns_for_storage_with_tool_calls() { - let mut thread = Thread::new(Uuid::new_v4()); + let mut thread = Thread::new(Uuid::new_v4(), None); thread.start_turn("Search for X"); // Record a tool call on the current turn if let Some(turn) = thread.turns.last_mut() { @@ -719,7 +719,7 @@ mod tests { #[test] fn test_format_turns_for_storage_incomplete_turn() { - let mut thread = Thread::new(Uuid::new_v4()); + let mut thread = Thread::new(Uuid::new_v4(), None); thread.start_turn("In progress message"); // Don't complete the turn diff --git a/src/agent/dispatcher.rs b/src/agent/dispatcher.rs index 99dd294ad90..71d9115f87a 100644 --- a/src/agent/dispatcher.rs +++ b/src/agent/dispatcher.rs @@ -2313,7 +2313,7 @@ mod tests { // Initialize a thread in the session so the loop can record tool calls. let thread_id = { let mut sess = session.lock().await; - sess.create_thread().id + sess.create_thread(Some("test")).id }; let message = IncomingMessage::new("test", "test-user", "do something"); @@ -2426,7 +2426,7 @@ mod tests { let session = Arc::new(Mutex::new(Session::new("test-user"))); let thread_id = { let mut sess = session.lock().await; - sess.create_thread().id + sess.create_thread(Some("test")).id }; let message = IncomingMessage::new("test", "test-user", "keep calling tools"); diff --git a/src/agent/routine_engine.rs b/src/agent/routine_engine.rs index 410cc69d56d..2018228fc9b 100644 --- a/src/agent/routine_engine.rs +++ b/src/agent/routine_engine.rs @@ -97,6 +97,10 @@ pub(crate) fn routine_matches_message(routine: &Routine, message: &IncomingMessa true } +fn trigger_uses_event_cache(trigger: &Trigger) -> bool { + matches!(trigger, Trigger::Event { .. } | Trigger::SystemEvent { .. }) +} + /// The routine execution engine. pub struct RoutineEngine { config: RoutineConfig, @@ -666,7 +670,7 @@ impl RoutineEngine { None }; - if let Err(e) = self + let runtime_updated = match self .store .update_routine_runtime( routine.id, @@ -678,10 +682,25 @@ impl RoutineEngine { ) .await { - tracing::error!( - routine = %routine.name, - "Failed to update routine runtime after dispatched run: {}", e - ); + Ok(()) => true, + Err(e) => { + tracing::error!( + routine = %routine.name, + "Failed to update routine runtime after dispatched run: {}", e + ); + false + } + }; + + if runtime_updated && trigger_uses_event_cache(&routine.trigger) { + update_cached_event_runtime( + self.event_cache.as_ref(), + routine.id, + now, + routine.run_count + 1, + new_failures, + ) + .await; } // Persist result to the routine's conversation thread @@ -812,6 +831,7 @@ impl RoutineEngine { tools: self.tools.clone(), safety: self.safety.clone(), sandbox_readiness: self.sandbox_readiness, + event_cache: Arc::clone(&self.event_cache), }; tokio::spawn(async move { @@ -897,6 +917,7 @@ impl RoutineEngine { tools: self.tools.clone(), safety: self.safety.clone(), sandbox_readiness: self.sandbox_readiness, + event_cache: Arc::clone(&self.event_cache), }; tokio::spawn(async move { @@ -951,6 +972,7 @@ impl RoutineEngine { tools: self.tools.clone(), safety: self.safety.clone(), sandbox_readiness: self.sandbox_readiness, + event_cache: Arc::clone(&self.event_cache), }; // Record the run in DB, then spawn execution @@ -1089,6 +1111,7 @@ struct EngineContext { tools: Arc, safety: Arc, sandbox_readiness: SandboxReadiness, + event_cache: Arc>>, } /// Execute a routine run. Handles both lightweight and full_job modes. @@ -1175,7 +1198,7 @@ async fn execute_routine(ctx: EngineContext, mut routine: Routine, run: RoutineR 0 }; - if let Err(e) = ctx + let runtime_updated = match ctx .store .update_routine_runtime( routine.id, @@ -1187,7 +1210,22 @@ async fn execute_routine(ctx: EngineContext, mut routine: Routine, run: RoutineR ) .await { - tracing::error!(routine = %routine.name, "Failed to update runtime state: {}", e); + Ok(()) => true, + Err(e) => { + tracing::error!(routine = %routine.name, "Failed to update runtime state: {}", e); + false + } + }; + + if runtime_updated && trigger_uses_event_cache(&routine.trigger) { + update_cached_event_runtime( + ctx.event_cache.as_ref(), + routine.id, + now, + routine.run_count + 1, + new_failures, + ) + .await; } // Persist routine result to its dedicated conversation thread @@ -1236,6 +1274,27 @@ async fn execute_routine(ctx: EngineContext, mut routine: Routine, run: RoutineR .await; } +async fn update_cached_event_runtime( + event_cache: &RwLock>, + routine_id: Uuid, + last_run_at: chrono::DateTime, + run_count: u64, + consecutive_failures: u32, +) { + let mut cache = event_cache.write().await; + for matcher in cache.iter_mut() { + let routine = match matcher { + EventMatcher::Message { routine, .. } | EventMatcher::System { routine } => routine, + }; + if routine.id == routine_id { + routine.last_run_at = Some(last_run_at); + routine.run_count = run_count; + routine.consecutive_failures = consecutive_failures; + break; + } + } +} + /// Sanitize a routine name for use in workspace paths. /// Only keeps alphanumeric, dash, and underscore characters; replaces everything else. fn sanitize_routine_name(name: &str) -> String { @@ -1613,13 +1672,23 @@ async fn execute_lightweight_with_tools( let mut total_input_tokens = 0; let mut total_output_tokens = 0; - // Create a minimal job context for tool execution with unique run ID + // Create a minimal job context for tool execution with unique run ID. + // Carry the routine's notify config in metadata so the message tool can + // resolve channel/target — mirrors the full-job path in execute_full_job(). let run_id = Uuid::new_v4(); + let mut lw_metadata = serde_json::json!({ + "owner_id": routine.user_id + }); + if let Some(channel) = &routine.notify.channel { + lw_metadata["notify_channel"] = serde_json::json!(channel); + } + lw_metadata["notify_user"] = serde_json::json!(&routine.notify.user); let job_ctx = JobContext { job_id: run_id, user_id: routine.user_id.clone(), title: "Lightweight Routine".to_string(), description: routine.name.clone(), + metadata: lw_metadata, ..Default::default() }; let allowed_tools = @@ -2575,4 +2644,31 @@ mod tests { assert!(result.len() <= 503); assert!(result.ends_with("...")); } + + /// Regression: lightweight routines must carry notify metadata in JobContext + /// so the message tool can route to the correct channel. Previously, + /// `..Default::default()` left metadata as null, causing messages to land + /// in the user's DM instead of the originating Slack channel. + #[test] + fn test_build_lightweight_prompt_preserves_notify_config() { + let notify = NotifyConfig { + channel: Some("slack-relay".to_string()), + user: Some("C088K6C3SQZ".to_string()), + on_attention: true, + on_failure: true, + on_success: false, + }; + + let prompt = + super::build_lightweight_prompt("Send Ping in this channel.", &[], None, ¬ify, true); + + assert!( + prompt.contains("slack-relay"), + "prompt should mention configured delivery channel: {prompt}", + ); + assert!( + prompt.contains("C088K6C3SQZ"), + "prompt should mention configured delivery target: {prompt}", + ); + } } diff --git a/src/agent/session.rs b/src/agent/session.rs index 6c873e46535..8d598c49a88 100644 --- a/src/agent/session.rs +++ b/src/agent/session.rs @@ -68,8 +68,8 @@ impl Session { } /// Create a new thread in this session. - pub fn create_thread(&mut self) -> &mut Thread { - let thread = Thread::new(self.id); + pub fn create_thread(&mut self, channel: Option<&str>) -> &mut Thread { + let thread = Thread::new(self.id, channel); let thread_id = thread.id; self.active_thread = Some(thread_id); self.last_active_at = Utc::now(); @@ -87,9 +87,9 @@ impl Session { } /// Get or create the active thread. - pub fn get_or_create_thread(&mut self) -> &mut Thread { + pub fn get_or_create_thread(&mut self, channel: Option<&str>) -> &mut Thread { match self.active_thread { - None => self.create_thread(), + None => self.create_thread(channel), Some(id) => { if self.threads.contains_key(&id) { // Entry existence confirmed by contains_key above. @@ -100,7 +100,7 @@ impl Session { } else { // Stale active_thread ID: create a new thread, which // updates self.active_thread to the new thread's ID. - self.create_thread() + self.create_thread(channel) } } } @@ -225,6 +225,9 @@ pub struct Thread { /// Messages queued while the thread was processing a turn. #[serde(default, skip_serializing_if = "VecDeque::is_empty")] pub pending_messages: VecDeque, + /// Channel that created this thread (for approval authorization). + #[serde(default)] + pub source_channel: Option, } /// Maximum number of messages that can be queued while a thread is processing. @@ -233,9 +236,34 @@ pub struct Thread { /// rapid follow-ups. The drain loop processes them as one newline-delimited turn. pub const MAX_PENDING_MESSAGES: usize = 10; +/// Sentinel value for bootstrap threads that accept approvals from any channel. +pub const BOOTSTRAP_SOURCE_CHANNEL: &str = "__bootstrap__"; + +/// Channels that are always authorized to approve tool calls on any thread, +/// regardless of which channel originally created the thread. These are +/// trusted UI surfaces (the web dashboard and its gateway). +pub const TRUSTED_APPROVAL_CHANNELS: &[&str] = &["web", "gateway"]; + +/// Check whether an approval from `requesting_channel` is authorized for a +/// thread whose `source_channel` is `source`. +/// +/// Rules: +/// - `None` (unknown origin) -> denied (fail-closed) +/// - `Some("__bootstrap__")` -> authorized from any channel +/// - `Some(src) == requesting` -> same channel, authorized +/// - requesting is in `TRUSTED_APPROVAL_CHANNELS` -> always authorized +/// - Otherwise -> denied +pub fn is_approval_authorized(source: Option<&str>, requesting: &str) -> bool { + match source { + None => false, + Some(src) if src == BOOTSTRAP_SOURCE_CHANNEL => true, + Some(src) => src == requesting || TRUSTED_APPROVAL_CHANNELS.contains(&requesting), + } +} + impl Thread { /// Create a new thread. - pub fn new(session_id: Uuid) -> Self { + pub fn new(session_id: Uuid, source_channel: Option<&str>) -> Self { let now = Utc::now(); Self { id: Uuid::new_v4(), @@ -248,11 +276,12 @@ impl Thread { pending_approval: None, pending_auth: None, pending_messages: VecDeque::new(), + source_channel: source_channel.map(String::from), } } /// Create a thread with a specific ID (for DB hydration). - pub fn with_id(id: Uuid, session_id: Uuid) -> Self { + pub fn with_id(id: Uuid, session_id: Uuid, source_channel: Option<&str>) -> Self { let now = Utc::now(); Self { id, @@ -265,6 +294,7 @@ impl Thread { pending_approval: None, pending_auth: None, pending_messages: VecDeque::new(), + source_channel: source_channel.map(String::from), } } @@ -787,13 +817,13 @@ mod tests { let mut session = Session::new("user-123"); assert!(session.active_thread.is_none()); - session.create_thread(); + session.create_thread(None); assert!(session.active_thread.is_some()); } #[test] fn test_thread_turns() { - let mut thread = Thread::new(Uuid::new_v4()); + let mut thread = Thread::new(Uuid::new_v4(), None); thread.start_turn("Hello"); assert_eq!(thread.state, ThreadState::Processing); @@ -806,7 +836,7 @@ mod tests { #[test] fn test_thread_messages() { - let mut thread = Thread::new(Uuid::new_v4()); + let mut thread = Thread::new(Uuid::new_v4(), None); thread.start_turn("First message"); thread.complete_turn("First response"); @@ -829,7 +859,7 @@ mod tests { #[test] fn test_restore_from_messages() { - let mut thread = Thread::new(Uuid::new_v4()); + let mut thread = Thread::new(Uuid::new_v4(), None); // First add some turns thread.start_turn("Original message"); @@ -855,7 +885,7 @@ mod tests { #[test] fn test_restore_from_messages_incomplete_turn() { - let mut thread = Thread::new(Uuid::new_v4()); + let mut thread = Thread::new(Uuid::new_v4(), None); // Messages with incomplete last turn (no assistant response) let messages = vec![ @@ -874,7 +904,7 @@ mod tests { #[test] fn test_enter_auth_mode() { let before = Utc::now(); - let mut thread = Thread::new(Uuid::new_v4()); + let mut thread = Thread::new(Uuid::new_v4(), None); assert!(thread.pending_auth.is_none()); thread.enter_auth_mode("telegram".to_string()); @@ -887,7 +917,7 @@ mod tests { #[test] fn test_take_pending_auth() { - let mut thread = Thread::new(Uuid::new_v4()); + let mut thread = Thread::new(Uuid::new_v4(), None); thread.enter_auth_mode("notion".to_string()); let pending = thread.take_pending_auth(); @@ -902,7 +932,7 @@ mod tests { #[test] fn test_pending_auth_serialization() { - let mut thread = Thread::new(Uuid::new_v4()); + let mut thread = Thread::new(Uuid::new_v4(), None); thread.enter_auth_mode("openai".to_string()); let json = serde_json::to_string(&thread).expect("should serialize"); @@ -932,7 +962,7 @@ mod tests { #[test] fn test_pending_auth_default_none() { // Deserialization of old data without pending_auth should default to None - let mut thread = Thread::new(Uuid::new_v4()); + let mut thread = Thread::new(Uuid::new_v4(), None); thread.pending_auth = None; let json = serde_json::to_string(&thread).expect("serialize"); @@ -946,7 +976,7 @@ mod tests { fn test_thread_with_id() { let specific_id = Uuid::new_v4(); let session_id = Uuid::new_v4(); - let thread = Thread::with_id(specific_id, session_id); + let thread = Thread::with_id(specific_id, session_id, None); assert_eq!(thread.id, specific_id); assert_eq!(thread.session_id, session_id); @@ -958,7 +988,7 @@ mod tests { fn test_thread_with_id_restore_messages() { let thread_id = Uuid::new_v4(); let session_id = Uuid::new_v4(); - let mut thread = Thread::with_id(thread_id, session_id); + let mut thread = Thread::with_id(thread_id, session_id, None); let messages = vec![ ChatMessage::user("Hello from DB"), @@ -977,7 +1007,7 @@ mod tests { #[test] fn test_restore_from_messages_empty() { - let mut thread = Thread::new(Uuid::new_v4()); + let mut thread = Thread::new(Uuid::new_v4(), None); // Add a turn first, then restore with empty vec thread.start_turn("hello"); @@ -993,7 +1023,7 @@ mod tests { #[test] fn test_restore_from_messages_only_assistant_messages() { - let mut thread = Thread::new(Uuid::new_v4()); + let mut thread = Thread::new(Uuid::new_v4(), None); // Only assistant messages (no user messages to anchor turns) let messages = vec![ @@ -1010,7 +1040,7 @@ mod tests { #[test] fn test_restore_from_messages_multiple_user_messages_in_a_row() { - let mut thread = Thread::new(Uuid::new_v4()); + let mut thread = Thread::new(Uuid::new_v4(), None); // Two user messages with no assistant response between them let messages = vec![ @@ -1037,8 +1067,8 @@ mod tests { fn test_thread_switch() { let mut session = Session::new("user-1"); - let t1_id = session.create_thread().id; - let t2_id = session.create_thread().id; + let t1_id = session.create_thread(None).id; + let t2_id = session.create_thread(None).id; // After creating two threads, active should be the last one assert_eq!(session.active_thread, Some(t2_id)); @@ -1058,8 +1088,8 @@ mod tests { fn test_get_or_create_thread_idempotent() { let mut session = Session::new("user-1"); - let tid1 = session.get_or_create_thread().id; - let tid2 = session.get_or_create_thread().id; + let tid1 = session.get_or_create_thread(None).id; + let tid2 = session.get_or_create_thread(None).id; // Should return the same thread (not create a new one each time) assert_eq!(tid1, tid2); @@ -1068,7 +1098,7 @@ mod tests { #[test] fn test_truncate_turns() { - let mut thread = Thread::new(Uuid::new_v4()); + let mut thread = Thread::new(Uuid::new_v4(), None); for i in 0..5 { thread.start_turn(format!("msg-{}", i)); @@ -1092,7 +1122,7 @@ mod tests { #[test] fn test_truncate_turns_noop_when_fewer() { - let mut thread = Thread::new(Uuid::new_v4()); + let mut thread = Thread::new(Uuid::new_v4(), None); thread.start_turn("only one"); thread.complete_turn("response"); @@ -1104,7 +1134,7 @@ mod tests { #[test] fn test_thread_interrupt_and_resume() { - let mut thread = Thread::new(Uuid::new_v4()); + let mut thread = Thread::new(Uuid::new_v4(), None); thread.start_turn("do something"); assert_eq!(thread.state, ThreadState::Processing); @@ -1122,7 +1152,7 @@ mod tests { #[test] fn test_resume_only_from_interrupted() { - let mut thread = Thread::new(Uuid::new_v4()); + let mut thread = Thread::new(Uuid::new_v4(), None); // Idle thread: resume should be a no-op assert_eq!(thread.state, ThreadState::Idle); @@ -1138,7 +1168,7 @@ mod tests { #[test] fn test_turn_fail() { - let mut thread = Thread::new(Uuid::new_v4()); + let mut thread = Thread::new(Uuid::new_v4(), None); thread.start_turn("risky operation"); thread.fail_turn("connection timed out"); @@ -1154,7 +1184,7 @@ mod tests { #[test] fn test_messages_with_incomplete_last_turn() { - let mut thread = Thread::new(Uuid::new_v4()); + let mut thread = Thread::new(Uuid::new_v4(), None); thread.start_turn("first"); thread.complete_turn("first reply"); @@ -1170,7 +1200,7 @@ mod tests { #[test] fn test_thread_serialization_round_trip() { - let mut thread = Thread::new(Uuid::new_v4()); + let mut thread = Thread::new(Uuid::new_v4(), None); thread.start_turn("hello"); thread.complete_turn("world"); @@ -1188,7 +1218,7 @@ mod tests { #[test] fn test_session_serialization_round_trip() { let mut session = Session::new("user-ser"); - session.create_thread(); + session.create_thread(None); session.auto_approve_tool("echo"); let json = serde_json::to_string(&session).unwrap(); @@ -1226,7 +1256,7 @@ mod tests { #[test] fn test_turn_number_increments() { - let mut thread = Thread::new(Uuid::new_v4()); + let mut thread = Thread::new(Uuid::new_v4(), None); // Before any turns, turn_number() is 1 (1-indexed for display) assert_eq!(thread.turn_number(), 1); @@ -1241,7 +1271,7 @@ mod tests { #[test] fn test_complete_turn_on_empty_thread() { - let mut thread = Thread::new(Uuid::new_v4()); + let mut thread = Thread::new(Uuid::new_v4(), None); // Completing a turn when there are no turns should be a safe no-op thread.complete_turn("phantom response"); @@ -1251,7 +1281,7 @@ mod tests { #[test] fn test_fail_turn_on_empty_thread() { - let mut thread = Thread::new(Uuid::new_v4()); + let mut thread = Thread::new(Uuid::new_v4(), None); // Failing a turn when there are no turns should be a safe no-op thread.fail_turn("phantom error"); @@ -1261,7 +1291,7 @@ mod tests { #[test] fn test_pending_approval_flow() { - let mut thread = Thread::new(Uuid::new_v4()); + let mut thread = Thread::new(Uuid::new_v4(), None); let approval = PendingApproval { request_id: Uuid::new_v4(), @@ -1288,7 +1318,7 @@ mod tests { #[test] fn test_clear_pending_approval() { - let mut thread = Thread::new(Uuid::new_v4()); + let mut thread = Thread::new(Uuid::new_v4(), None); let approval = PendingApproval { request_id: Uuid::new_v4(), @@ -1317,7 +1347,7 @@ mod tests { assert!(session.active_thread().is_none()); assert!(session.active_thread_mut().is_none()); - let tid = session.create_thread().id; + let tid = session.create_thread(None).id; assert!(session.active_thread().is_some()); assert_eq!(session.active_thread().unwrap().id, tid); @@ -1334,7 +1364,7 @@ mod tests { #[test] fn test_messages_includes_tool_calls() { - let mut thread = Thread::new(Uuid::new_v4()); + let mut thread = Thread::new(Uuid::new_v4(), None); thread.start_turn("Search for X"); { @@ -1366,7 +1396,7 @@ mod tests { #[test] fn test_messages_multiple_tool_calls_per_turn() { - let mut thread = Thread::new(Uuid::new_v4()); + let mut thread = Thread::new(Uuid::new_v4(), None); thread.start_turn("Do two things"); { @@ -1393,7 +1423,7 @@ mod tests { #[test] fn test_restore_from_messages_with_tool_calls() { - let mut thread = Thread::new(Uuid::new_v4()); + let mut thread = Thread::new(Uuid::new_v4(), None); // Build a message sequence with tool calls let tc = ToolCall { @@ -1425,7 +1455,7 @@ mod tests { #[test] fn test_restore_from_messages_with_tool_error() { - let mut thread = Thread::new(Uuid::new_v4()); + let mut thread = Thread::new(Uuid::new_v4(), None); let tc = ToolCall { id: "call_0".to_string(), @@ -1456,7 +1486,7 @@ mod tests { fn test_messages_round_trip_with_tools() { // Build a thread with tool calls, get messages(), restore, get messages() again // The two message sequences should be equivalent. - let mut thread = Thread::new(Uuid::new_v4()); + let mut thread = Thread::new(Uuid::new_v4(), None); thread.start_turn("Do search"); { @@ -1469,7 +1499,7 @@ mod tests { let messages_original = thread.messages(); // Restore into a new thread - let mut thread2 = Thread::new(Uuid::new_v4()); + let mut thread2 = Thread::new(Uuid::new_v4(), None); thread2.restore_from_messages(messages_original.clone()); let messages_restored = thread2.messages(); @@ -1491,7 +1521,7 @@ mod tests { #[test] fn test_restore_multi_stage_tool_calls() { - let mut thread = Thread::new(Uuid::new_v4()); + let mut thread = Thread::new(Uuid::new_v4(), None); let tc1 = ToolCall { id: "call_a".to_string(), @@ -1534,7 +1564,7 @@ mod tests { #[test] fn test_messages_truncates_large_tool_results() { - let mut thread = Thread::new(Uuid::new_v4()); + let mut thread = Thread::new(Uuid::new_v4(), None); thread.start_turn("Read big file"); { @@ -1557,7 +1587,7 @@ mod tests { #[test] fn test_thread_message_queue() { - let mut thread = Thread::new(Uuid::new_v4()); + let mut thread = Thread::new(Uuid::new_v4(), None); // Queue is initially empty assert!(thread.pending_messages.is_empty()); @@ -1593,7 +1623,7 @@ mod tests { #[test] fn test_thread_message_queue_serialization() { - let mut thread = Thread::new(Uuid::new_v4()); + let mut thread = Thread::new(Uuid::new_v4(), None); // Empty queue should not appear in serialization (skip_serializing_if) let json = serde_json::to_string(&thread).unwrap(); @@ -1613,7 +1643,7 @@ mod tests { #[test] fn test_thread_message_queue_default_on_old_data() { // Deserialization of old data without pending_messages should default to empty - let thread = Thread::new(Uuid::new_v4()); + let thread = Thread::new(Uuid::new_v4(), None); let json = serde_json::to_string(&thread).unwrap(); // The field is absent (skip_serializing_if), simulating old data @@ -1624,7 +1654,7 @@ mod tests { #[test] fn test_interrupt_clears_pending_messages() { - let mut thread = Thread::new(Uuid::new_v4()); + let mut thread = Thread::new(Uuid::new_v4(), None); // Start a turn so there's something to interrupt thread.start_turn("initial input"); @@ -1643,7 +1673,7 @@ mod tests { #[test] fn test_thread_state_idle_after_full_drain() { - let mut thread = Thread::new(Uuid::new_v4()); + let mut thread = Thread::new(Uuid::new_v4(), None); // Simulate a full drain cycle: start turn, queue messages, complete turn, // then drain all queued messages as a single merged turn (#259). @@ -1671,7 +1701,7 @@ mod tests { #[test] fn test_drain_pending_messages_merges_with_newlines() { - let mut thread = Thread::new(Uuid::new_v4()); + let mut thread = Thread::new(Uuid::new_v4(), None); // Empty queue returns None assert!(thread.drain_pending_messages().is_none()); @@ -1700,7 +1730,7 @@ mod tests { #[test] fn test_requeue_drained_preserves_content_at_front() { - let mut thread = Thread::new(Uuid::new_v4()); + let mut thread = Thread::new(Uuid::new_v4(), None); // Re-queue into empty queue thread.requeue_drained("failed batch".to_string()); @@ -1811,4 +1841,200 @@ mod tests { &serde_json::json!("done") ); } + + #[test] + fn test_thread_new_stores_source_channel() { + let thread = Thread::new(Uuid::new_v4(), Some("telegram")); + assert_eq!(thread.source_channel.as_deref(), Some("telegram")); + } + + #[test] + fn test_thread_new_none_channel() { + let thread = Thread::new(Uuid::new_v4(), None); + assert!(thread.source_channel.is_none()); + } + + #[test] + fn test_source_channel_serde_backcompat() { + // Simulate deserializing a Thread from older DB records that lack source_channel. + let thread = Thread::new(Uuid::new_v4(), Some("cli")); + let json = serde_json::to_string(&thread).unwrap(); + + // Remove the source_channel field to simulate an old record. + let mut value: serde_json::Value = serde_json::from_str(&json).unwrap(); + value.as_object_mut().unwrap().remove("source_channel"); + let old_json = serde_json::to_string(&value).unwrap(); + + let deserialized: Thread = serde_json::from_str(&old_json).unwrap(); + assert!( + deserialized.source_channel.is_none(), + "missing source_channel should deserialize as None" + ); + } + + #[test] + fn test_approval_authorized_same_channel() { + assert!( + is_approval_authorized(Some("telegram"), "telegram"), + "same channel should be authorized" + ); + } + + #[test] + fn test_approval_authorized_different_channel_blocked() { + assert!( + !is_approval_authorized(Some("telegram"), "http"), + "different channel should be blocked" + ); + } + + #[test] + fn test_approval_authorized_web_always_allowed() { + assert!( + is_approval_authorized(Some("telegram"), "web"), + "web channel should always be authorized" + ); + } + + #[test] + fn test_approval_authorized_gateway_always_allowed() { + assert!( + is_approval_authorized(Some("telegram"), "gateway"), + "gateway channel should always be authorized" + ); + } + + #[test] + fn test_approval_authorized_none_denied() { + assert!( + !is_approval_authorized(None, "telegram"), + "None source_channel should be denied (fail-closed)" + ); + assert!( + !is_approval_authorized(None, "web"), + "None source_channel should be denied even for web" + ); + } + + #[test] + fn test_approval_authorized_bootstrap_any_channel() { + assert!( + is_approval_authorized(Some(BOOTSTRAP_SOURCE_CHANNEL), "telegram"), + "__bootstrap__ should be authorized from any channel" + ); + assert!( + is_approval_authorized(Some(BOOTSTRAP_SOURCE_CHANNEL), "http"), + "__bootstrap__ should be authorized from any channel" + ); + assert!( + is_approval_authorized(Some(BOOTSTRAP_SOURCE_CHANNEL), "cli"), + "__bootstrap__ should be authorized from any channel" + ); + } + + #[test] + fn test_approval_authorized_uses_trusted_channels_constant() { + // Every channel in TRUSTED_APPROVAL_CHANNELS should be authorized + // against any source, ensuring the constant drives the logic. + for &trusted in TRUSTED_APPROVAL_CHANNELS { + assert!( + is_approval_authorized(Some("any-source"), trusted), + "TRUSTED_APPROVAL_CHANNELS entry '{}' should always be authorized", + trusted + ); + } + } + + #[test] + fn test_approval_blocks_thread_without_pending_approval() { + // A thread with no pending_approval should not be eligible for + // approval routing. This test verifies the data-level invariant + // that `agent_loop.rs` checks before calling is_approval_authorized. + let thread = Thread::new(Uuid::new_v4(), Some("telegram")); + assert!( + thread.pending_approval.is_none(), + "new thread should have no pending approval" + ); + + // Set up a thread WITH a pending approval to contrast + let mut thread_with_approval = Thread::new(Uuid::new_v4(), Some("telegram")); + thread_with_approval.pending_approval = Some(PendingApproval { + request_id: Uuid::new_v4(), + tool_name: "shell".to_string(), + parameters: serde_json::json!({"cmd": "rm -rf /"}), + display_parameters: serde_json::json!({"cmd": "rm -rf /"}), + description: "run shell command".to_string(), + tool_call_id: "call_1".to_string(), + context_messages: vec![], + deferred_tool_calls: vec![], + user_timezone: None, + allow_always: true, + }); + assert!( + thread_with_approval.pending_approval.is_some(), + "thread with pending approval should be eligible" + ); + + // Authorization check should pass for the thread with pending approval + // (same channel), confirming the two checks compose correctly. + assert!(is_approval_authorized( + thread_with_approval.source_channel.as_deref(), + "telegram" + )); + } + + #[test] + fn test_approval_wasm_channel_cannot_impersonate_trusted() { + // A WASM channel named "web" or "gateway" would bypass authorization. + // This test documents the invariant that WASM setup must reject these + // names (tested separately in wasm/setup.rs). + // Here we verify the authorization logic itself treats them as trusted. + assert!(is_approval_authorized(Some("telegram"), "web")); + assert!(is_approval_authorized(Some("telegram"), "gateway")); + // But a random WASM channel name should NOT be trusted + assert!(!is_approval_authorized(Some("telegram"), "my-wasm-channel")); + } + + #[test] + fn test_approval_bootstrap_sentinel_not_a_normal_channel() { + // If a channel happens to be named __bootstrap__, it should be treated + // as the source (always authorized), NOT as a requesting channel with + // special trust. Only TRUSTED_APPROVAL_CHANNELS get that privilege. + assert!( + !is_approval_authorized(Some("telegram"), BOOTSTRAP_SOURCE_CHANNEL), + "__bootstrap__ as requesting channel should not have special trust" + ); + } + + #[test] + fn test_create_thread_propagates_channel() { + let mut session = Session::new("user-chan"); + let tid = session.create_thread(Some("signal")).id; + let thread = session.threads.get(&tid).unwrap(); + assert_eq!(thread.source_channel.as_deref(), Some("signal")); + } + + #[test] + fn test_get_or_create_thread_propagates_channel() { + let mut session = Session::new("user-chan2"); + // First call creates + let tid = session.get_or_create_thread(Some("http")).id; + assert_eq!( + session.threads.get(&tid).unwrap().source_channel.as_deref(), + Some("http") + ); + // Second call returns existing (channel param ignored) + let tid2 = session.get_or_create_thread(Some("different")).id; + assert_eq!(tid, tid2); + assert_eq!( + session + .threads + .get(&tid2) + .unwrap() + .source_channel + .as_deref(), + Some("http"), + "existing thread should keep its original source_channel" + ); + } } diff --git a/src/agent/session_manager.rs b/src/agent/session_manager.rs index ae98b0b03e9..7736f85f774 100644 --- a/src/agent/session_manager.rs +++ b/src/agent/session_manager.rs @@ -200,7 +200,7 @@ impl SessionManager { // Create new thread (always create a new one for a new key) let thread_id = { let mut sess = session.lock().await; - let thread = sess.create_thread(); + let thread = sess.create_thread(Some(channel)); thread.id }; @@ -476,7 +476,7 @@ mod tests { let session = Arc::new(Mutex::new(Session::new("user-hydrate"))); { let mut sess = session.lock().await; - let thread = Thread::with_id(thread_id, sess.id); + let thread = Thread::with_id(thread_id, sess.id, None); sess.threads.insert(thread_id, thread); sess.active_thread = Some(thread_id); } @@ -600,7 +600,7 @@ mod tests { // Simulate hydration: create thread with a known UUID { let mut sess = session.lock().await; - let thread = Thread::with_id(known_uuid, session_id); + let thread = Thread::with_id(known_uuid, session_id, None); sess.threads.insert(known_uuid, thread); } @@ -627,7 +627,7 @@ mod tests { let session = Arc::new(Mutex::new(Session::new("user-idem"))); { let mut sess = session.lock().await; - let thread = Thread::with_id(tid, sess.id); + let thread = Thread::with_id(tid, sess.id, None); sess.threads.insert(tid, thread); } @@ -656,7 +656,7 @@ mod tests { let session = Arc::new(Mutex::new(Session::new("user-undo"))); { let mut sess = session.lock().await; - let thread = Thread::with_id(tid, sess.id); + let thread = Thread::with_id(tid, sess.id, None); sess.threads.insert(tid, thread); } @@ -680,7 +680,7 @@ mod tests { let session = Arc::new(Mutex::new(Session::new("user-new"))); { let mut sess = session.lock().await; - let thread = Thread::with_id(tid, sess.id); + let thread = Thread::with_id(tid, sess.id, None); sess.threads.insert(tid, thread); } @@ -788,7 +788,7 @@ mod tests { let session = Arc::new(Mutex::new(Session::new("user-cross"))); { let mut sess = session.lock().await; - let thread = Thread::with_id(tid, sess.id); + let thread = Thread::with_id(tid, sess.id, None); sess.threads.insert(tid, thread); } @@ -815,7 +815,7 @@ mod tests { let session = Arc::new(Mutex::new(Session::new("user-cross"))); { let mut sess = session.lock().await; - let thread = Thread::with_id(tid, sess.id); + let thread = Thread::with_id(tid, sess.id, None); sess.threads.insert(tid, thread); } @@ -966,7 +966,7 @@ mod tests { let adopted_id = Uuid::new_v4(); { let mut sess = session1.lock().await; - let thread = Thread::with_id(adopted_id, sess.id); + let thread = Thread::with_id(adopted_id, sess.id, None); sess.threads.insert(adopted_id, thread); } // Resolve with the UUID as external_thread_id -- should adopt it @@ -992,7 +992,7 @@ mod tests { let session = Arc::new(Mutex::new(Session::new("user-direct"))); { let mut sess = session.lock().await; - let thread = Thread::with_id(tid, sess.id); + let thread = Thread::with_id(tid, sess.id, None); sess.threads.insert(tid, thread); } { @@ -1030,7 +1030,7 @@ mod tests { let known_id = Uuid::new_v4(); { let mut sess = session.lock().await; - let thread = Thread::with_id(known_id, sess.id); + let thread = Thread::with_id(known_id, sess.id, None); sess.threads.insert(known_id, thread); } @@ -1057,7 +1057,7 @@ mod tests { let known_id = Uuid::new_v4(); { let mut sess = session.lock().await; - let thread = Thread::with_id(known_id, sess.id); + let thread = Thread::with_id(known_id, sess.id, None); sess.threads.insert(known_id, thread); } @@ -1081,7 +1081,7 @@ mod tests { let known_id = Uuid::new_v4(); { let mut sess = session.lock().await; - let thread = Thread::with_id(known_id, sess.id); + let thread = Thread::with_id(known_id, sess.id, None); sess.threads.insert(known_id, thread); } @@ -1102,4 +1102,19 @@ mod tests { "should NOT adopt UUID when external_thread_id is None" ); } + + #[tokio::test] + async fn test_thread_stores_source_channel() { + let manager = SessionManager::new(); + + let (session, thread_id) = manager.resolve_thread("user-1", "telegram", None).await; + + let sess = session.lock().await; + let thread = sess.threads.get(&thread_id).unwrap(); + assert_eq!( + thread.source_channel.as_deref(), + Some("telegram"), + "resolve_thread should store source_channel from the channel parameter" + ); + } } diff --git a/src/agent/thread_ops.rs b/src/agent/thread_ops.rs index af0bd67f5cc..1797fbcde93 100644 --- a/src/agent/thread_ops.rs +++ b/src/agent/thread_ops.rs @@ -135,13 +135,52 @@ impl Agent { msg_count = 0; } - // Create thread with the historical ID and restore messages + // Create thread with the historical ID and restore messages. + // Read source_channel from DB so the authorization check uses the + // original creator's channel, not the requesting message's channel. + // + // Fail-closed policy: if the DB lookup fails or the conversation has + // no stored source_channel (legacy row), the thread is hydrated with + // source_channel = None. `is_approval_authorized(None, _)` returns + // false, so approvals are denied until the conversation is backfilled + // with a source_channel via an explicit migration or re-creation. + let db_source_channel = if let Some(store) = self.store() { + match store.get_conversation_source_channel(thread_uuid).await { + Ok(sc) => { + if sc.is_none() { + tracing::warn!( + thread_id = %thread_uuid, + "Legacy thread has no stored source_channel; \ + cross-channel approvals will be denied (fail-closed)" + ); + } + sc + } + Err(e) => { + tracing::error!( + thread_id = %thread_uuid, + error = %e, + "Failed to read source_channel from DB; \ + cross-channel approvals will be denied (fail-closed)" + ); + None + } + } + } else { + None + }; + let effective_source_channel = db_source_channel.as_deref(); + let session_id = { let sess = session.lock().await; sess.id }; - let mut thread = crate::agent::session::Thread::with_id(thread_uuid, session_id); + let mut thread = crate::agent::session::Thread::with_id( + thread_uuid, + session_id, + effective_source_channel, + ); if !chat_messages.is_empty() { thread.restore_from_messages(chat_messages); } @@ -636,7 +675,7 @@ impl Agent { user_id: &str, ) -> bool { match store - .ensure_conversation(thread_id, channel, user_id, None) + .ensure_conversation(thread_id, channel, user_id, None, Some(channel)) .await { Ok(true) => true, @@ -1781,7 +1820,7 @@ impl Agent { .get_or_create_session(&message.user_id) .await; let mut sess = session.lock().await; - let thread = sess.create_thread(); + let thread = sess.create_thread(Some(&message.channel)); let thread_id = thread.id; Ok(SubmissionResult::ok_with_message(format!( "New thread: {}", @@ -2117,7 +2156,7 @@ mod tests { let session_id = Uuid::new_v4(); let thread_id = Uuid::new_v4(); - let mut thread = Thread::with_id(thread_id, session_id); + let mut thread = Thread::with_id(thread_id, session_id, None); // Set thread to AwaitingApproval with a pending tool approval let pending = PendingApproval { @@ -2185,7 +2224,7 @@ mod tests { use crate::agent::session::{MAX_PENDING_MESSAGES, Thread, ThreadState}; use uuid::Uuid; - let mut thread = Thread::new(Uuid::new_v4()); + let mut thread = Thread::new(Uuid::new_v4(), None); thread.start_turn("processing something"); assert_eq!(thread.state, ThreadState::Processing); @@ -2211,7 +2250,7 @@ mod tests { use crate::agent::session::{Thread, ThreadState}; use uuid::Uuid; - let mut thread = Thread::new(Uuid::new_v4()); + let mut thread = Thread::new(Uuid::new_v4(), None); thread.start_turn("processing"); thread.queue_message("pending-1".to_string()); @@ -2241,7 +2280,7 @@ mod tests { let thread_id = Uuid::new_v4(); let session_id = Uuid::new_v4(); - let mut thread = Thread::with_id(thread_id, session_id); + let mut thread = Thread::with_id(thread_id, session_id, None); thread.start_turn("working"); assert_eq!(thread.state, ThreadState::Processing); @@ -2268,7 +2307,7 @@ mod tests { let thread_id = Uuid::new_v4(); let session_id = Uuid::new_v4(); - let mut thread = Thread::with_id(thread_id, session_id); + let mut thread = Thread::with_id(thread_id, session_id, None); thread.start_turn("working"); assert_eq!(thread.state, ThreadState::Processing); diff --git a/src/channels/channel.rs b/src/channels/channel.rs index 784b6bcf1f8..99b19a81243 100644 --- a/src/channels/channel.rs +++ b/src/channels/channel.rs @@ -201,28 +201,26 @@ impl IncomingMessage { } /// Extract a channel-specific proactive routing target from message metadata. +/// +/// Checked keys (first match wins): +/// - `signal_target` — Signal phone number or group ID +/// - `chat_id` — Telegram chat ID +/// - `channel_id` — Slack channel/DM ID (used by channel-relay) +/// - `target` — generic fallback pub fn routing_target_from_metadata(metadata: &serde_json::Value) -> Option { - metadata - .get("signal_target") - .and_then(|value| match value { + // Helper to extract a string or numeric value from a JSON key. + let extract = |key: &str| -> Option { + metadata.get(key).and_then(|value| match value { serde_json::Value::String(s) => Some(s.clone()), serde_json::Value::Number(n) => Some(n.to_string()), _ => None, }) - .or_else(|| { - metadata.get("chat_id").and_then(|value| match value { - serde_json::Value::String(s) => Some(s.clone()), - serde_json::Value::Number(n) => Some(n.to_string()), - _ => None, - }) - }) - .or_else(|| { - metadata.get("target").and_then(|value| match value { - serde_json::Value::String(s) => Some(s.clone()), - serde_json::Value::Number(n) => Some(n.to_string()), - _ => None, - }) - }) + }; + + extract("signal_target") + .or_else(|| extract("chat_id")) + .or_else(|| extract("channel_id")) + .or_else(|| extract("target")) } /// Stream of incoming messages. @@ -357,6 +355,20 @@ pub enum StatusUpdate { }, } +/// Shared chat-style approval prompt formatting used by non-web channels. +#[derive(Debug, Clone, PartialEq, Eq)] +pub struct ChatApprovalPrompt { + pub request_id: String, + pub tool_name: String, + pub description: String, + pub parameters: serde_json::Value, + pub allow_always: bool, +} + +const APPROVAL_PARAMETER_PREVIEW_BYTES: usize = 1200; +const APPROVAL_PARAMETER_TRUNCATION_SUFFIX: &str = "\n... [parameters truncated]"; +const APPROVAL_SUMMARY_DESCRIPTION_BYTES: usize = 120; + impl StatusUpdate { /// Build a `ToolCompleted` status with redacted parameters. /// @@ -389,6 +401,131 @@ impl StatusUpdate { } } +impl ChatApprovalPrompt { + /// Build a shared chat approval prompt from a status update. + pub fn from_status(status: &StatusUpdate) -> Option { + let StatusUpdate::ApprovalNeeded { + request_id, + tool_name, + description, + parameters, + allow_always, + } = status + else { + return None; + }; + + Some(Self { + request_id: request_id.clone(), + tool_name: tool_name.clone(), + description: description.clone(), + parameters: parameters.clone(), + allow_always: *allow_always, + }) + } + + fn truncated_text(input: &str, max_bytes: usize, suffix: &str) -> String { + if input.len() <= max_bytes { + return input.to_string(); + } + + let budget = max_bytes.saturating_sub(suffix.len()); + let end = crate::util::floor_char_boundary(input, budget); + format!("{}{}", &input[..end], suffix) + } + + /// Pretty-printed tool parameters for display, bounded for chat channels. + pub fn parameters_preview(&self) -> String { + let rendered = serde_json::to_string_pretty(&self.parameters) + .unwrap_or_else(|_| self.parameters.to_string()); + Self::truncated_text( + &rendered, + APPROVAL_PARAMETER_PREVIEW_BYTES, + APPROVAL_PARAMETER_TRUNCATION_SUFFIX, + ) + } + + /// Shared reply vocabulary summary for compact status surfaces. + pub fn reply_summary(&self) -> &'static str { + if self.allow_always { + "yes (or /approve), no (or /deny), or always (or /always)" + } else { + "yes (or /approve) or no (or /deny)" + } + } + + /// Compact approval summary for fallback/accessibility surfaces. + pub fn summary_text(&self) -> String { + let description = Self::truncated_text( + &self.description.replace('\n', " "), + APPROVAL_SUMMARY_DESCRIPTION_BYTES, + "...", + ); + format!( + "Approval needed for {}: {} (Request ID: {}). Reply with {}.", + self.tool_name, + description, + self.request_id, + self.reply_summary() + ) + } + + fn markdown_parameters_preview(&self) -> String { + self.parameters_preview().replace('`', "\\`") + } + + /// Approval prompt formatted for plain-text chat channels. + pub fn plain_text_message(&self) -> String { + let mut lines = vec![ + format!("Approval needed: {}", self.tool_name), + self.description.clone(), + String::new(), + format!("Request ID: {}", self.request_id), + "Parameters:".to_string(), + self.parameters_preview(), + String::new(), + "Reply with:".to_string(), + "- yes, y, approve, or /approve to approve this request".to_string(), + ]; + + if self.allow_always { + lines.push(format!( + "- always, a, or /always to approve this request and auto-approve future {} requests", + self.tool_name + )); + } + + lines.push("- no, n, deny, or /deny to deny this request".to_string()); + lines.join("\n") + } + + /// Approval prompt formatted for Markdown-capable chat channels. + pub fn markdown_message(&self) -> String { + let mut lines = vec![ + "⚠️ *Approval Required*".to_string(), + String::new(), + format!("*Request ID:* `{}`", self.request_id), + format!("*Tool:* {}", self.tool_name), + format!("*Description:* {}", self.description), + "*Parameters:*".to_string(), + format!("```json\n{}\n```", self.markdown_parameters_preview()), + String::new(), + "Reply with:".to_string(), + "• `yes`, `y`, `approve`, or `/approve` - Approve this request".to_string(), + ]; + + if self.allow_always { + lines.push(format!( + "• `always`, `a`, or `/always` - Approve this request and auto-approve future {} requests", + self.tool_name + )); + } + + lines.push("• `no`, `n`, `deny`, or `/deny` - Deny this request".to_string()); + lines.join("\n") + } +} + /// Trait for message channels. /// /// Channels receive messages from external sources and convert them to @@ -612,4 +749,120 @@ mod tests { let msg = IncomingMessage::new("test", "user1", "hello").with_timezone("America/New_York"); assert_eq!(msg.timezone.as_deref(), Some("America/New_York")); } + + #[test] + fn routing_target_extracts_slack_channel_id() { + // Slack relay messages carry channel_id in metadata — this must be + // picked up for proactive broadcasts to land in the correct channel + // instead of falling back to sender_id (which routes to DMs). + let metadata = serde_json::json!({ + "team_id": "T05CUBCSQPL", + "channel_id": "C088K6C3SQZ", + "sender_id": "UCBGL1WNS", + }); + assert_eq!( + routing_target_from_metadata(&metadata).as_deref(), + Some("C088K6C3SQZ"), + ); + } + + #[test] + fn routing_target_prefers_signal_over_channel_id() { + let metadata = serde_json::json!({ + "signal_target": "+15551234567", + "channel_id": "C088K6C3SQZ", + }); + assert_eq!( + routing_target_from_metadata(&metadata).as_deref(), + Some("+15551234567"), + ); + } + + #[test] + fn routing_target_prefers_chat_id_over_channel_id() { + let metadata = serde_json::json!({ + "chat_id": "123456789", + "channel_id": "C088K6C3SQZ", + }); + assert_eq!( + routing_target_from_metadata(&metadata).as_deref(), + Some("123456789"), + ); + } + + #[test] + fn routing_target_returns_none_for_empty_metadata() { + let metadata = serde_json::json!({}); + assert!(routing_target_from_metadata(&metadata).is_none()); + } + + #[test] + fn chat_approval_prompt_plain_text_includes_all_reply_forms() { + let prompt = ChatApprovalPrompt::from_status(&StatusUpdate::ApprovalNeeded { + request_id: "req-123".into(), + tool_name: "http".into(), + description: "Fetch weather data".into(), + parameters: serde_json::json!({"url": "https://api.weather.test"}), + allow_always: true, + }) + .expect("approval prompt"); + + let text = prompt.plain_text_message(); + assert!(text.contains("Request ID: req-123")); + assert!(text.contains("approve, or /approve")); + assert!(text.contains("always, a, or /always")); + assert!(text.contains("deny, or /deny")); + } + + #[test] + fn chat_approval_prompt_hides_always_when_not_allowed() { + let prompt = ChatApprovalPrompt::from_status(&StatusUpdate::ApprovalNeeded { + request_id: "req-456".into(), + tool_name: "shell".into(), + description: "Run command".into(), + parameters: serde_json::json!({"command": "rm -rf /tmp/demo"}), + allow_always: false, + }) + .expect("approval prompt"); + + let markdown = prompt.markdown_message(); + assert!(markdown.contains("`/approve`")); + assert!(markdown.contains("`/deny`")); + assert!(!markdown.contains("`/always`")); + } + + #[test] + fn chat_approval_prompt_truncates_large_parameters() { + let prompt = ChatApprovalPrompt::from_status(&StatusUpdate::ApprovalNeeded { + request_id: "req-789".into(), + tool_name: "http".into(), + description: "Fetch large payload".into(), + parameters: serde_json::json!({ + "body": "x".repeat(APPROVAL_PARAMETER_PREVIEW_BYTES + 200), + }), + allow_always: true, + }) + .expect("approval prompt"); + + let preview = prompt.parameters_preview(); + assert!(preview.contains("[parameters truncated]")); + assert!(preview.len() <= APPROVAL_PARAMETER_PREVIEW_BYTES); + } + + #[test] + fn chat_approval_prompt_escapes_backticks_in_markdown_parameters() { + let prompt = ChatApprovalPrompt::from_status(&StatusUpdate::ApprovalNeeded { + request_id: "req-999".into(), + tool_name: "shell".into(), + description: "Run command".into(), + parameters: serde_json::json!({ + "command": "printf '```danger```'" + }), + allow_always: true, + }) + .expect("approval prompt"); + + let markdown = prompt.markdown_message(); + assert!(markdown.contains("\\`\\`\\`danger\\`\\`\\`")); + } } diff --git a/src/channels/mod.rs b/src/channels/mod.rs index 46e255145ff..7f0a929240e 100644 --- a/src/channels/mod.rs +++ b/src/channels/mod.rs @@ -38,8 +38,9 @@ pub mod web; mod webhook_server; pub use channel::{ - AttachmentKind, Channel, ChannelSecretUpdater, IncomingAttachment, IncomingMessage, - MessageStream, OutgoingResponse, StatusUpdate, ToolDecision, routing_target_from_metadata, + AttachmentKind, Channel, ChannelSecretUpdater, ChatApprovalPrompt, IncomingAttachment, + IncomingMessage, MessageStream, OutgoingResponse, StatusUpdate, ToolDecision, + routing_target_from_metadata, }; pub use http::{HttpChannel, HttpChannelState}; pub use manager::ChannelManager; diff --git a/src/channels/relay/channel.rs b/src/channels/relay/channel.rs index 3b6c3379686..858c8e446e0 100644 --- a/src/channels/relay/channel.rs +++ b/src/channels/relay/channel.rs @@ -11,7 +11,9 @@ use async_trait::async_trait; use tokio::sync::mpsc; use crate::channels::relay::client::{ChannelEvent, RelayClient}; -use crate::channels::{Channel, IncomingMessage, MessageStream, OutgoingResponse, StatusUpdate}; +use crate::channels::{ + Channel, ChatApprovalPrompt, IncomingMessage, MessageStream, OutgoingResponse, StatusUpdate, +}; use crate::error::ChannelError; /// Default channel name for the Slack relay integration. @@ -126,6 +128,58 @@ impl RelayChannel { .proxy_provider(self.provider.as_str(), team_id, method, body) .await } + + fn build_approval_body( + &self, + channel_id: &str, + thread_id: Option<&str>, + prompt: &ChatApprovalPrompt, + approval_token: &str, + ) -> serde_json::Value { + let value_payload = serde_json::json!({ + "approval_token": approval_token, + }); + let value_str = value_payload.to_string(); + + let blocks = serde_json::json!([ + { + "type": "section", + "text": { + "type": "mrkdwn", + "text": prompt.markdown_message(), + } + }, + { + "type": "actions", + "elements": [ + { + "type": "button", + "text": { "type": "plain_text", "text": "Approve" }, + "style": "primary", + "action_id": "approve_tool", + "value": value_str, + }, + { + "type": "button", + "text": { "type": "plain_text", "text": "Deny" }, + "style": "danger", + "action_id": "deny_tool", + "value": value_str, + } + ] + } + ]); + + let mut body = serde_json::json!({ + "channel": channel_id, + "text": prompt.summary_text(), + "blocks": blocks, + }); + if let Some(tid) = thread_id { + body["thread_ts"] = serde_json::Value::String(tid.to_string()); + } + body + } } #[async_trait] @@ -194,12 +248,18 @@ impl Channel for RelayChannel { "sender_id": event.sender_id, "sender_name": event.display_name(), "event_type": event.event_type, - "thread_id": event.thread_id, + "thread_id": event.thread_id.as_deref().unwrap_or(&event.id), "provider": event.provider, })); + // Use the original thread_id if present (already in a thread), + // otherwise use the message timestamp (event.id) so that + // responses are threaded under the user's message in channels. + // Fall back to channel_id only if event.id is missing. let msg = if let Some(ref thread_id) = event.thread_id { msg.with_thread(thread_id) + } else if !event.id.is_empty() { + msg.with_thread(&event.id) } else { msg.with_thread(&event.channel_id) }; @@ -260,14 +320,7 @@ impl Channel for RelayChannel { metadata: &serde_json::Value, ) -> Result<(), ChannelError> { // Only handle ApprovalNeeded — all other variants are no-ops - let StatusUpdate::ApprovalNeeded { - request_id, - tool_name, - description, - parameters, - allow_always: _, - } = status - else { + let Some(prompt) = ChatApprovalPrompt::from_status(&status) else { return Ok(()); }; @@ -278,7 +331,7 @@ impl Channel for RelayChannel { .unwrap_or(""); if event_type != "direct_message" { tracing::warn!( - tool = %tool_name, + tool = %prompt.tool_name, event_type, "Approval requested in non-DM, skipping buttons" ); @@ -303,60 +356,13 @@ impl Channel for RelayChannel { // The button value contains ONLY the token — no routing fields. let approval_token = self .client - .create_approval(team_id, channel_id, thread_id, &request_id) + .create_approval(team_id, channel_id, thread_id, &prompt.request_id) .await .map_err(|e| ChannelError::SendFailed { name: self.name().to_string(), reason: format!("Failed to register approval: {e}"), })?; - let value_payload = serde_json::json!({ - "approval_token": approval_token, - }); - let value_str = value_payload.to_string(); - - // Parameters are already redacted via redact_params() in dispatcher.rs - let params_display = - serde_json::to_string_pretty(¶meters).unwrap_or_else(|_| parameters.to_string()); - - let blocks = serde_json::json!([ - { - "type": "section", - "text": { - "type": "mrkdwn", - "text": format!( - "*Tool approval required*\n`{tool_name}`: {description}\n```{params_display}```" - ) - } - }, - { - "type": "actions", - "elements": [ - { - "type": "button", - "text": { "type": "plain_text", "text": "Approve" }, - "style": "primary", - "action_id": "approve_tool", - "value": value_str, - }, - { - "type": "button", - "text": { "type": "plain_text", "text": "Deny" }, - "style": "danger", - "action_id": "deny_tool", - "value": value_str, - } - ] - } - ]); - - let mut body = serde_json::json!({ - "channel": channel_id, - "text": format!("Tool approval required: {tool_name} - {description}"), - "blocks": blocks, - }); - if let Some(tid) = thread_id { - body["thread_ts"] = serde_json::Value::String(tid.to_string()); - } + let body = self.build_approval_body(channel_id, thread_id, &prompt, &approval_token); self.proxy_send(team_id, "chat.postMessage", body) .await @@ -497,6 +503,29 @@ mod tests { assert_eq!(body["thread_ts"], "1234567.890"); } + #[test] + fn build_approval_body_includes_chat_reply_instructions() { + let channel = make_channel(); + let prompt = ChatApprovalPrompt { + request_id: "req-1".into(), + tool_name: "http".into(), + description: "HTTP requests to external APIs".into(), + parameters: serde_json::json!({"method": "POST", "url": "https://example.com"}), + allow_always: true, + }; + + let body = channel.build_approval_body("C456", Some("1234567.890"), &prompt, "token-123"); + let text = body["text"].as_str().expect("plain text"); + let block_text = body["blocks"][0]["text"]["text"].as_str().expect("mrkdwn"); + + assert!(text.contains("Request ID: req-1")); + assert!(text.contains("Reply with yes (or /approve)")); + assert!(!text.contains("Parameters:")); + assert!(block_text.contains("`/approve`")); + assert!(block_text.contains("`/always`")); + assert_eq!(body["thread_ts"], "1234567.890"); + } + #[tokio::test] async fn start_processes_events() { let (tx, rx) = mpsc::channel(64); @@ -532,6 +561,41 @@ mod tests { assert_eq!(msg.user_id, "U789"); } + #[tokio::test] + async fn start_threaded_message_preserves_thread_scope() { + let (tx, rx) = mpsc::channel(64); + let channel = + RelayChannel::new(test_client(), "T123".into(), "inst1".into(), tx.clone(), rx); + + let mut stream = channel.start().await.unwrap(); + + tx.send(ChannelEvent { + id: "threaded-1".into(), + event_type: "direct_message".into(), + provider: "slack".into(), + provider_scope: "T123".into(), + channel_id: "D456".into(), + sender_id: "U789".into(), + sender_name: Some("alice".into()), + content: Some("approve".into()), + thread_id: Some("1712345678.123".into()), + raw: serde_json::Value::Null, + timestamp: None, + }) + .await + .unwrap(); + + use futures::StreamExt; + let msg = tokio::time::timeout(std::time::Duration::from_secs(1), stream.next()) + .await + .unwrap() + .unwrap(); + + assert_eq!(msg.content, "approve"); + assert_eq!(msg.thread_id.as_deref(), Some("1712345678.123")); + assert_eq!(msg.conversation_scope(), Some("1712345678.123")); + } + #[tokio::test] async fn start_skips_non_message_events() { let (tx, rx) = mpsc::channel(64); @@ -649,6 +713,53 @@ mod tests { ); } + /// Regression: channel mentions must use the message timestamp (event.id) + /// as thread_id, not the channel_id. Slack requires thread_ts to be a + /// message timestamp for threading to work. + #[tokio::test] + async fn start_uses_message_ts_as_thread_id_for_mentions() { + let (tx, rx) = mpsc::channel(64); + let channel = + RelayChannel::new(test_client(), "T123".into(), "inst1".into(), tx.clone(), rx); + + let mut stream = channel.start().await.unwrap(); + + // Simulate a channel mention (no thread_id, id = message ts) + tx.send(ChannelEvent { + id: "1609459200.000100".into(), + event_type: "mention".into(), + provider: "slack".into(), + provider_scope: "T123".into(), + channel_id: "C456".into(), + sender_id: "U789".into(), + sender_name: Some("alice".into()), + content: Some("hello bot".into()), + thread_id: None, + raw: serde_json::Value::Null, + timestamp: None, + }) + .await + .unwrap(); + + use futures::StreamExt; + let msg = tokio::time::timeout(std::time::Duration::from_secs(1), stream.next()) + .await + .unwrap() + .unwrap(); + + // thread_id should be the message timestamp, NOT the channel_id + assert_eq!( + msg.thread_id.as_deref(), + Some("1609459200.000100"), + "thread_id should be the message ts for threading, not the channel_id" + ); + // metadata should also have the correct thread_id + assert_eq!( + msg.metadata.get("thread_id").and_then(|v| v.as_str()), + Some("1609459200.000100"), + ); + } + #[tokio::test] async fn test_send_status_approval_dm_without_sender_id_is_ok() { let channel = make_channel(); diff --git a/src/channels/relay/client.rs b/src/channels/relay/client.rs index 1bc60a5690a..16f40f66474 100644 --- a/src/channels/relay/client.rs +++ b/src/channels/relay/client.rs @@ -277,9 +277,30 @@ impl RelayClient { }); } - resp.json() + let json: serde_json::Value = resp + .json() .await - .map_err(|e| RelayError::Protocol(e.to_string())) + .map_err(|e| RelayError::Protocol(e.to_string()))?; + + // Slack API always returns HTTP 200 but signals errors via {"ok": false}. + // Surface these as relay errors so callers get actionable feedback. + if json.get("ok") == Some(&serde_json::Value::Bool(false)) { + let slack_error = json + .get("error") + .and_then(|v| v.as_str()) + .unwrap_or("unknown"); + tracing::warn!( + relay_url = %url, + slack_error = %slack_error, + "RelayClient::proxy_provider: Slack API returned ok=false" + ); + return Err(RelayError::Api { + status: 200, + message: format!("Slack API error: {slack_error}"), + }); + } + + Ok(json) } /// Fetch the per-instance callback signing secret from channel-relay. diff --git a/src/channels/signal.rs b/src/channels/signal.rs index 84afccd5fb0..ac5d00d86fd 100644 --- a/src/channels/signal.rs +++ b/src/channels/signal.rs @@ -910,34 +910,10 @@ impl Channel for SignalChannel { } // Send approval prompt to user - if let StatusUpdate::ApprovalNeeded { - request_id, - tool_name, - description: _, - parameters, - allow_always, - } = &status + if let Some(prompt) = crate::channels::ChatApprovalPrompt::from_status(&status) && let Some(target_str) = metadata.get("signal_target").and_then(|v| v.as_str()) { - let params_json = serde_json::to_string_pretty(parameters).unwrap_or_default(); - let always_line = if *allow_always { - format!( - "\n• `always` or `a` - Approve and auto-approve future {} requests", - tool_name - ) - } else { - String::new() - }; - let message = format!( - "⚠️ *Approval Required*\n\n\ - *Request ID:* `{}`\n\ - *Tool:* {}\n\ - *Parameters:*\n```\n{}\n```\n\n\ - Reply with:\n\ - • `yes` or `y` - Approve this request{}\n\ - • `no` or `n` - Deny", - request_id, tool_name, params_json, always_line - ); + let message = prompt.markdown_message(); self.send_status_message(target_str, &message).await; } diff --git a/src/channels/wasm/schema.rs b/src/channels/wasm/schema.rs index f2ca66745da..8ced07defcc 100644 --- a/src/channels/wasm/schema.rs +++ b/src/channels/wasm/schema.rs @@ -702,6 +702,29 @@ mod tests { ); } + #[test] + fn test_telegram_bundled_setup_includes_webhook_secret() { + let json = include_str!("../../../channels-src/telegram/telegram.capabilities.json"); + let file = ChannelCapabilitiesFile::from_json(json).unwrap(); + + let webhook_secret = file + .setup + .required_secrets + .iter() + .find(|secret| secret.name == "telegram_webhook_secret") + .expect("telegram webhook secret should be declared in the bundled manifest"); + + assert!(webhook_secret.optional); + assert_eq!( + webhook_secret + .auto_generate + .as_ref() + .expect("telegram webhook secret should auto-generate") + .length, + 64 + ); + } + // ── Category 5: Discord Capabilities Setup & Configuration ────────── #[test] diff --git a/src/channels/wasm/setup.rs b/src/channels/wasm/setup.rs index 84df615fdf1..638fae56c65 100644 --- a/src/channels/wasm/setup.rs +++ b/src/channels/wasm/setup.rs @@ -34,6 +34,7 @@ pub async fn setup_wasm_channels( secrets_store: &Option>, extension_manager: Option<&Arc>, database: Option<&Arc>, + registered_channel_names: &[String], ) -> Option { let runtime = match WasmChannelRuntime::new(WasmChannelRuntimeConfig::default()) { Ok(r) => Arc::new(r), @@ -71,7 +72,51 @@ pub async fn setup_wasm_channels( let mut channels: Vec<(String, Box)> = Vec::new(); let mut channel_names: Vec = Vec::new(); + // Reserved channel names that WASM modules must not claim. + // A malicious module could otherwise register as a trusted built-in + // channel and bypass cross-channel authorization checks. + // + // This list includes: + // - All built-in channel names (prevent impersonation) + // - Trusted approval channels from session::TRUSTED_APPROVAL_CHANNELS + // - The bootstrap sentinel (universal approval wildcard) + use crate::agent::session::{BOOTSTRAP_SOURCE_CHANNEL, TRUSTED_APPROVAL_CHANNELS}; + + let mut reserved: Vec<&str> = vec![ + "cli", + "repl", + "http", + "signal", + "telegram", + "slack-relay", + "secret_save", + ]; + reserved.extend(TRUSTED_APPROVAL_CHANNELS); + reserved.push(BOOTSTRAP_SOURCE_CHANNEL); + for loaded in results.loaded { + let name_lower = loaded.name().to_ascii_lowercase(); + if reserved.contains(&name_lower.as_str()) { + tracing::warn!( + channel = %loaded.name(), + "Rejected WASM channel with reserved name" + ); + continue; + } + // Also reject any name that collides with an already-registered + // channel to prevent a WASM module from shadowing a channel that + // was registered earlier in the startup sequence. + if registered_channel_names + .iter() + .any(|n| n.to_ascii_lowercase() == name_lower) + { + tracing::warn!( + channel = %loaded.name(), + "Rejected WASM channel that collides with already-registered channel" + ); + continue; + } + let (name, channel) = register_channel( loaded, config, @@ -449,3 +494,75 @@ async fn inject_channel_secrets_into_config( } } } + +#[cfg(test)] +mod tests { + use crate::agent::session::{BOOTSTRAP_SOURCE_CHANNEL, TRUSTED_APPROVAL_CHANNELS}; + + /// Build the same reserved-name list that `setup_wasm_channels` uses. + fn reserved_names() -> Vec<&'static str> { + let mut reserved: Vec<&str> = vec![ + "cli", + "repl", + "http", + "signal", + "telegram", + "slack-relay", + "secret_save", + ]; + reserved.extend(TRUSTED_APPROVAL_CHANNELS); + reserved.push(BOOTSTRAP_SOURCE_CHANNEL); + reserved + } + + #[test] + fn reserved_names_include_trusted_approval_channels() { + let reserved = reserved_names(); + for &trusted in TRUSTED_APPROVAL_CHANNELS { + assert!( + reserved.contains(&trusted), + "trusted approval channel '{}' must be in WASM reserved names", + trusted + ); + } + } + + #[test] + fn reserved_names_include_bootstrap_sentinel() { + let reserved = reserved_names(); + assert!( + reserved.contains(&BOOTSTRAP_SOURCE_CHANNEL), + "__bootstrap__ sentinel must be in WASM reserved names" + ); + } + + #[test] + fn reserved_names_reject_case_insensitive() { + // The setup logic lowercases the WASM channel name before checking. + // Verify that "Web" or "GATEWAY" would be caught. + let reserved = reserved_names(); + let test_cases = ["Web", "GATEWAY", "CLI", "Repl", "__BOOTSTRAP__"]; + for name in test_cases { + let lowered = name.to_ascii_lowercase(); + assert!( + reserved.contains(&lowered.as_str()), + "'{}' (lowercased to '{}') should match a reserved name", + name, + lowered + ); + } + } + + #[test] + fn non_reserved_names_allowed() { + let reserved = reserved_names(); + let allowed = ["discord", "my-custom-channel", "slack-bot"]; + for name in allowed { + assert!( + !reserved.contains(&name), + "'{}' should NOT be reserved", + name + ); + } + } +} diff --git a/src/channels/wasm/wrapper.rs b/src/channels/wasm/wrapper.rs index c502112700d..10d3a0b6184 100644 --- a/src/channels/wasm/wrapper.rs +++ b/src/channels/wasm/wrapper.rs @@ -2298,63 +2298,16 @@ impl WasmChannel { StatusUpdate::StreamChunk(_) => { // No-op, too noisy } - StatusUpdate::ApprovalNeeded { - tool_name, - description, - parameters, - allow_always, - .. - } => { + StatusUpdate::ApprovalNeeded { .. } => { // WASM channels (Telegram, Slack, etc.) cannot render // interactive approval overlays. Send the approval prompt // as an actual message so the user can reply yes/no. self.cancel_typing_task().await; - let params_preview = parameters - .as_object() - .map(|obj| { - obj.iter() - .map(|(k, v)| { - let val = match v { - serde_json::Value::String(s) => { - if s.chars().count() > 80 { - let truncated: String = s.chars().take(77).collect(); - format!("\"{}...\"", truncated) - } else { - format!("\"{}\"", s) - } - } - other => { - let s = other.to_string(); - if s.chars().count() > 80 { - let truncated: String = s.chars().take(77).collect(); - format!("{}...", truncated) - } else { - s - } - } - }; - format!(" {}: {}", k, val) - }) - .collect::>() - .join("\n") - }) - .unwrap_or_default(); - - let reply_hint = if *allow_always { - "Reply \"yes\" to approve, \"no\" to deny, or \"always\" to auto-approve." - } else { - "Reply \"yes\" to approve or \"no\" to deny." + let Some(prompt) = crate::channels::ChatApprovalPrompt::from_status(&status) else { + return Ok(()); }; - let prompt = format!( - "Approval needed: {tool_name}\n\ - {description}\n\ - \n\ - Parameters:\n\ - {params_preview}\n\ - \n\ - {reply_hint}" - ); + let prompt = prompt.plain_text_message(); let metadata_json = serde_json::to_string(metadata).unwrap_or_default(); if let Err(e) = self @@ -3813,24 +3766,11 @@ fn status_to_wit( metadata_json, } } - StatusUpdate::ApprovalNeeded { - request_id, - tool_name, - description, - allow_always, - .. - } => { - let reply_hint = if *allow_always { - "yes (or /approve), no (or /deny), or always (or /always)" - } else { - "yes (or /approve) or no (or /deny)" - }; + StatusUpdate::ApprovalNeeded { .. } => { + let prompt = crate::channels::ChatApprovalPrompt::from_status(status)?; wit_channel::StatusUpdate { status: wit_channel::StatusType::ApprovalNeeded, - message: format!( - "Approval needed for tool '{}'. {}\nRequest ID: {}\nReply with: {}.", - tool_name, description, request_id, reply_hint - ), + message: prompt.plain_text_message(), metadata_json, } } diff --git a/src/channels/web/handlers/chat.rs b/src/channels/web/handlers/chat.rs index d1580f5cf25..e0c82957c49 100644 --- a/src/channels/web/handlers/chat.rs +++ b/src/channels/web/handlers/chat.rs @@ -574,7 +574,7 @@ pub async fn chat_new_thread_handler( .await; let (thread_id, info) = { let mut sess = session.lock().await; - let thread = sess.create_thread(); + let thread = sess.create_thread(Some("web")); let id = thread.id; let info = ThreadInfo { id: thread.id, @@ -593,7 +593,13 @@ pub async fn chat_new_thread_handler( // so that the subsequent loadThreads() call from the frontend sees it. if let Some(ref store) = state.store { match store - .ensure_conversation(thread_id, "gateway", &identity.user_id, None) + .ensure_conversation( + thread_id, + "gateway", + &identity.user_id, + None, + Some("gateway"), + ) .await { Ok(true) => {} diff --git a/src/channels/web/handlers/jobs.rs b/src/channels/web/handlers/jobs.rs index aca3e97cdb8..eb54c9d5ac1 100644 --- a/src/channels/web/handlers/jobs.rs +++ b/src/channels/web/handlers/jobs.rs @@ -456,7 +456,10 @@ pub async fn jobs_restart_handler( &task, Some(project_dir), mode, - credential_grants, + crate::orchestrator::job_manager::JobCreationParams { + credential_grants, + ..Default::default() + }, ) .await .map_err(|e| { diff --git a/src/channels/web/server.rs b/src/channels/web/server.rs index 710c5a587fc..a8c6e4d3b1a 100644 --- a/src/channels/web/server.rs +++ b/src/channels/web/server.rs @@ -2063,7 +2063,7 @@ async fn chat_new_thread_handler( let session = session_manager.get_or_create_session(&user.user_id).await; let (thread_id, info) = { let mut sess = session.lock().await; - let thread = sess.create_thread(); + let thread = sess.create_thread(Some("gateway")); let id = thread.id; let info = ThreadInfo { id: thread.id, @@ -2082,7 +2082,7 @@ async fn chat_new_thread_handler( // so that the subsequent loadThreads() call from the frontend sees it. if let Some(ref store) = state.store { match store - .ensure_conversation(thread_id, "gateway", &user.user_id, None) + .ensure_conversation(thread_id, "gateway", &user.user_id, None, Some("gateway")) .await { Ok(true) => {} diff --git a/src/cli/doctor.rs b/src/cli/doctor.rs index 023ac4e1dc1..dca408d73fd 100644 --- a/src/cli/doctor.rs +++ b/src/cli/doctor.rs @@ -80,7 +80,7 @@ pub async fn run_doctor_command() -> anyhow::Result<()> { check( "Routines config", - check_routines_config(), + check_routines_config(&settings), &mut passed, &mut failed, &mut skipped, @@ -434,8 +434,8 @@ fn check_embeddings(settings: &Settings) -> CheckResult { // ── Routines config ───────────────────────────────────────── -fn check_routines_config() -> CheckResult { - match crate::config::RoutineConfig::resolve() { +fn check_routines_config(settings: &Settings) -> CheckResult { + match crate::config::RoutineConfig::resolve(settings) { Ok(config) => { if config.enabled { CheckResult::Pass(format!( @@ -737,7 +737,8 @@ mod tests { #[test] fn check_routines_config_does_not_panic() { - let result = check_routines_config(); + let settings = Settings::default(); + let result = check_routines_config(&settings); match result { CheckResult::Pass(_) | CheckResult::Fail(_) | CheckResult::Skip(_) => {} } @@ -866,7 +867,8 @@ mod tests { unsafe { std::env::remove_var("ROUTINES_ENABLED"); } - match check_routines_config() { + let settings = Settings::default(); + match check_routines_config(&settings) { CheckResult::Pass(msg) => { assert!( msg.contains("enabled"), diff --git a/src/cli/mod.rs b/src/cli/mod.rs index 5da12e0bfa0..ebc97ac5b8e 100644 --- a/src/cli/mod.rs +++ b/src/cli/mod.rs @@ -288,7 +288,7 @@ pub enum Command { orchestrator_url: String, /// Maximum iterations before stopping. - #[arg(long, default_value = "50")] + #[arg(long, env = "IRONCLAW_MAX_ITERATIONS", default_value = "50")] max_iterations: u32, }, diff --git a/src/config/agent.rs b/src/config/agent.rs index a74eb4d6c75..261cf2d12ea 100644 --- a/src/config/agent.rs +++ b/src/config/agent.rs @@ -1,6 +1,8 @@ use std::time::Duration; -use crate::config::helpers::{parse_bool_env, parse_option_env, parse_optional_env}; +use crate::config::helpers::{ + db_first_bool, db_first_or_default, parse_bool_env, parse_option_env, +}; use crate::error::ConfigError; use crate::settings::Settings; @@ -73,49 +75,65 @@ impl AgentConfig { } pub(crate) fn resolve(settings: &Settings) -> Result { + let defaults = crate::settings::AgentSettings::default(); + Ok(Self { - name: parse_optional_env("AGENT_NAME", settings.agent.name.clone())?, - max_parallel_jobs: parse_optional_env( + name: db_first_or_default(&settings.agent.name, &defaults.name, "AGENT_NAME")?, + // Settings stores u32, config uses usize — cast for comparison. + max_parallel_jobs: db_first_or_default( + &(settings.agent.max_parallel_jobs as usize), + &(defaults.max_parallel_jobs as usize), "AGENT_MAX_PARALLEL_JOBS", - settings.agent.max_parallel_jobs as usize, )?, - job_timeout: Duration::from_secs(parse_optional_env( + job_timeout: Duration::from_secs(db_first_or_default( + &settings.agent.job_timeout_secs, + &defaults.job_timeout_secs, "AGENT_JOB_TIMEOUT_SECS", - settings.agent.job_timeout_secs, )?), - stuck_threshold: Duration::from_secs(parse_optional_env( + stuck_threshold: Duration::from_secs(db_first_or_default( + &settings.agent.stuck_threshold_secs, + &defaults.stuck_threshold_secs, "AGENT_STUCK_THRESHOLD_SECS", - settings.agent.stuck_threshold_secs, )?), - repair_check_interval: Duration::from_secs(parse_optional_env( + repair_check_interval: Duration::from_secs(db_first_or_default( + &settings.agent.repair_check_interval_secs, + &defaults.repair_check_interval_secs, "SELF_REPAIR_CHECK_INTERVAL_SECS", - settings.agent.repair_check_interval_secs, )?), - max_repair_attempts: parse_optional_env( + max_repair_attempts: db_first_or_default( + &settings.agent.max_repair_attempts, + &defaults.max_repair_attempts, "SELF_REPAIR_MAX_ATTEMPTS", - settings.agent.max_repair_attempts, )?, - use_planning: parse_bool_env("AGENT_USE_PLANNING", settings.agent.use_planning)?, - session_idle_timeout: Duration::from_secs(parse_optional_env( + use_planning: db_first_bool( + settings.agent.use_planning, + defaults.use_planning, + "AGENT_USE_PLANNING", + )?, + session_idle_timeout: Duration::from_secs(db_first_or_default( + &settings.agent.session_idle_timeout_secs, + &defaults.session_idle_timeout_secs, "SESSION_IDLE_TIMEOUT_SECS", - settings.agent.session_idle_timeout_secs, )?), allow_local_tools: parse_bool_env("ALLOW_LOCAL_TOOLS", false)?, max_cost_per_day_cents: parse_option_env("MAX_COST_PER_DAY_CENTS")?, max_actions_per_hour: parse_option_env("MAX_ACTIONS_PER_HOUR")?, max_cost_per_user_per_day_cents: parse_option_env("MAX_COST_PER_USER_PER_DAY_CENTS")?, - max_tool_iterations: parse_optional_env( + max_tool_iterations: db_first_or_default( + &settings.agent.max_tool_iterations, + &defaults.max_tool_iterations, "AGENT_MAX_TOOL_ITERATIONS", - settings.agent.max_tool_iterations, )?, - auto_approve_tools: parse_bool_env( - "AGENT_AUTO_APPROVE_TOOLS", + auto_approve_tools: db_first_bool( settings.agent.auto_approve_tools, + defaults.auto_approve_tools, + "AGENT_AUTO_APPROVE_TOOLS", )?, default_timezone: { - let tz: String = parse_optional_env( + let tz: String = db_first_or_default( + &settings.agent.default_timezone, + &defaults.default_timezone, "DEFAULT_TIMEZONE", - settings.agent.default_timezone.clone(), )?; if crate::timezone::parse_timezone(&tz).is_none() { return Err(ConfigError::InvalidValue { @@ -126,9 +144,10 @@ impl AgentConfig { tz }, max_jobs_per_user: parse_option_env("MAX_JOBS_PER_USER")?, - max_tokens_per_job: parse_optional_env( + max_tokens_per_job: db_first_or_default( + &settings.agent.max_tokens_per_job, + &defaults.max_tokens_per_job, "AGENT_MAX_TOKENS_PER_JOB", - settings.agent.max_tokens_per_job, )?, multi_tenant: parse_bool_env("AGENT_MULTI_TENANT", false)?, max_llm_concurrent_per_user: parse_option_env("TENANT_MAX_LLM_CONCURRENT")?, diff --git a/src/config/builder.rs b/src/config/builder.rs index f7bad12c9c6..73c8930ad12 100644 --- a/src/config/builder.rs +++ b/src/config/builder.rs @@ -1,7 +1,7 @@ use std::path::PathBuf; use std::time::Duration; -use crate::config::helpers::{optional_env, parse_bool_env, parse_optional_env}; +use crate::config::helpers::{db_first_bool, db_first_or_default, optional_env}; use crate::error::ConfigError; /// Builder mode configuration. @@ -34,14 +34,29 @@ impl Default for BuilderModeConfig { impl BuilderModeConfig { pub(crate) fn resolve(settings: &crate::settings::Settings) -> Result { let bs = &settings.builder; + let defaults = crate::settings::BuilderSettings::default(); Ok(Self { - enabled: parse_bool_env("BUILDER_ENABLED", bs.enabled)?, - build_dir: optional_env("BUILDER_DIR")? - .map(PathBuf::from) - .or_else(|| bs.build_dir.clone()), - max_iterations: parse_optional_env("BUILDER_MAX_ITERATIONS", bs.max_iterations)?, - timeout_secs: parse_optional_env("BUILDER_TIMEOUT_SECS", bs.timeout_secs)?, - auto_register: parse_bool_env("BUILDER_AUTO_REGISTER", bs.auto_register)?, + enabled: db_first_bool(bs.enabled, defaults.enabled, "BUILDER_ENABLED")?, + build_dir: if let Some(ref dir) = bs.build_dir { + Some(dir.clone()) + } else { + optional_env("BUILDER_DIR")?.map(PathBuf::from) + }, + max_iterations: db_first_or_default( + &bs.max_iterations, + &defaults.max_iterations, + "BUILDER_MAX_ITERATIONS", + )?, + timeout_secs: db_first_or_default( + &bs.timeout_secs, + &defaults.timeout_secs, + "BUILDER_TIMEOUT_SECS", + )?, + auto_register: db_first_bool( + bs.auto_register, + defaults.auto_register, + "BUILDER_AUTO_REGISTER", + )?, }) } @@ -79,7 +94,7 @@ mod tests { } #[test] - fn env_overrides_settings() { + fn db_settings_override_env() { let _guard = lock_env(); let mut settings = Settings::default(); settings.builder.timeout_secs = 123; @@ -89,6 +104,22 @@ mod tests { let cfg = BuilderModeConfig::resolve(&settings).expect("resolve"); unsafe { std::env::remove_var("BUILDER_TIMEOUT_SECS") }; - assert_eq!(cfg.timeout_secs, 3); + assert_eq!(cfg.timeout_secs, 123, "DB setting should win over env"); + } + + #[test] + fn env_used_when_no_db_setting() { + let _guard = lock_env(); + let settings = Settings::default(); + + // SAFETY: Under ENV_MUTEX, no concurrent env access. + unsafe { std::env::set_var("BUILDER_TIMEOUT_SECS", "42") }; + let cfg = BuilderModeConfig::resolve(&settings).expect("resolve"); + unsafe { std::env::remove_var("BUILDER_TIMEOUT_SECS") }; + + assert_eq!( + cfg.timeout_secs, 42, + "env should be used when DB has the default value" + ); } } diff --git a/src/config/channels.rs b/src/config/channels.rs index 74f98dbfca4..953feb0f60c 100644 --- a/src/config/channels.rs +++ b/src/config/channels.rs @@ -2,9 +2,12 @@ use std::collections::HashMap; use std::path::PathBuf; use crate::bootstrap::ironclaw_base_dir; -use crate::config::helpers::{optional_env, parse_bool_env, parse_optional_env}; +use crate::config::helpers::{ + db_first_bool, db_first_optional_string, db_first_or_default, optional_env, parse_bool_env, + parse_optional_env, +}; use crate::error::ConfigError; -use crate::settings::Settings; +use crate::settings::{ChannelSettings, Settings}; use secrecy::SecretString; /// Channel configurations. @@ -120,15 +123,24 @@ pub struct SignalConfig { impl ChannelsConfig { pub(crate) fn resolve(settings: &Settings, owner_id: &str) -> Result { let cs = &settings.channels; + let defaults = ChannelSettings::default(); let http_enabled_by_env = optional_env("HTTP_PORT")?.is_some() || optional_env("HTTP_HOST")?.is_some(); - let http = if http_enabled_by_env || cs.http_enabled { + let http_enabled_by_db = + db_first_bool(cs.http_enabled, defaults.http_enabled, "HTTP_ENABLED")?; + let http = if http_enabled_by_env || http_enabled_by_db { Some(HttpConfig { - host: optional_env("HTTP_HOST")? - .or_else(|| cs.http_host.clone()) + host: db_first_optional_string(&cs.http_host, "HTTP_HOST")? .unwrap_or_else(|| "0.0.0.0".to_string()), - port: parse_optional_env("HTTP_PORT", cs.http_port.unwrap_or(8080))?, + port: { + // defaults.http_port is None, so any Some(..) is an explicit DB override. + if let Some(ref db_port) = cs.http_port { + db_first_or_default(db_port, &8080, "HTTP_PORT")? + } else { + parse_optional_env("HTTP_PORT", 8080)? + } + }, webhook_secret: optional_env("HTTP_WEBHOOK_SECRET")?.map(SecretString::from), user_id: owner_id.to_string(), }) @@ -136,7 +148,11 @@ impl ChannelsConfig { None }; - let gateway_enabled = parse_bool_env("GATEWAY_ENABLED", cs.gateway_enabled)?; + let gateway_enabled = db_first_bool( + cs.gateway_enabled, + defaults.gateway_enabled, + "GATEWAY_ENABLED", + )?; let gateway = if gateway_enabled { let memory_layers: Vec = match optional_env("MEMORY_LAYERS")? { @@ -234,15 +250,26 @@ impl ChannelsConfig { }; Some(GatewayConfig { - host: optional_env("GATEWAY_HOST")? - .or_else(|| cs.gateway_host.clone()) + host: db_first_optional_string(&cs.gateway_host, "GATEWAY_HOST")? .unwrap_or_else(|| "127.0.0.1".to_string()), - port: parse_optional_env( - "GATEWAY_PORT", - cs.gateway_port.unwrap_or(DEFAULT_GATEWAY_PORT), - )?, - auth_token: optional_env("GATEWAY_AUTH_TOKEN")? - .or_else(|| cs.gateway_auth_token.clone()), + port: { + // defaults.gateway_port is None, so any Some(..) is an explicit DB override. + if let Some(ref db_port) = cs.gateway_port { + db_first_or_default(db_port, &DEFAULT_GATEWAY_PORT, "GATEWAY_PORT")? + } else { + parse_optional_env("GATEWAY_PORT", DEFAULT_GATEWAY_PORT)? + } + }, + // Security: auth token is env-only — never read from DB settings. + auth_token: { + if cs.gateway_auth_token.is_some() { + tracing::warn!( + "gateway_auth_token is set in DB/TOML but is now env-only \ + (GATEWAY_AUTH_TOKEN). Remove it from DB/TOML settings." + ); + } + optional_env("GATEWAY_AUTH_TOKEN")? + }, workspace_read_scopes, memory_layers, oidc, @@ -251,16 +278,24 @@ impl ChannelsConfig { None }; - let signal_url = optional_env("SIGNAL_HTTP_URL")?.or_else(|| cs.signal_http_url.clone()); - let signal = if let Some(http_url) = signal_url { - let account = optional_env("SIGNAL_ACCOUNT")? - .or_else(|| cs.signal_account.clone()) - .ok_or(ConfigError::InvalidValue { + let signal_enabled = + db_first_bool(cs.signal_enabled, defaults.signal_enabled, "SIGNAL_ENABLED")?; + let signal_url = db_first_optional_string(&cs.signal_http_url, "SIGNAL_HTTP_URL")?; + let signal = if signal_enabled || signal_url.is_some() { + let http_url = signal_url.ok_or(ConfigError::InvalidValue { + key: "SIGNAL_HTTP_URL".to_string(), + message: "SIGNAL_HTTP_URL is required when signal_enabled is set in DB/TOML \ + or SIGNAL_ENABLED env var is true" + .to_string(), + })?; + let account = db_first_optional_string(&cs.signal_account, "SIGNAL_ACCOUNT")?.ok_or( + ConfigError::InvalidValue { key: "SIGNAL_ACCOUNT".to_string(), message: "SIGNAL_ACCOUNT is required when SIGNAL_HTTP_URL is set".to_string(), - })?; + }, + )?; let allow_from = - match optional_env("SIGNAL_ALLOW_FROM")?.or_else(|| cs.signal_allow_from.clone()) { + match db_first_optional_string(&cs.signal_allow_from, "SIGNAL_ALLOW_FROM")? { None => vec![account.clone()], Some(s) => s .split(',') @@ -268,36 +303,39 @@ impl ChannelsConfig { .filter(|s| !s.is_empty()) .collect(), }; - let dm_policy = optional_env("SIGNAL_DM_POLICY")? - .or_else(|| cs.signal_dm_policy.clone()) + let dm_policy = db_first_optional_string(&cs.signal_dm_policy, "SIGNAL_DM_POLICY")? .unwrap_or_else(|| "pairing".to_string()); - let group_policy = optional_env("SIGNAL_GROUP_POLICY")? - .or_else(|| cs.signal_group_policy.clone()) - .unwrap_or_else(|| "allowlist".to_string()); + let group_policy = + db_first_optional_string(&cs.signal_group_policy, "SIGNAL_GROUP_POLICY")? + .unwrap_or_else(|| "allowlist".to_string()); Some(SignalConfig { http_url, account, allow_from, - allow_from_groups: optional_env("SIGNAL_ALLOW_FROM_GROUPS")? - .or_else(|| cs.signal_allow_from_groups.clone()) - .map(|s| { - s.split(',') - .map(|e| e.trim().to_string()) - .filter(|s| !s.is_empty()) - .collect() - }) - .unwrap_or_default(), + allow_from_groups: db_first_optional_string( + &cs.signal_allow_from_groups, + "SIGNAL_ALLOW_FROM_GROUPS", + )? + .map(|s| { + s.split(',') + .map(|e| e.trim().to_string()) + .filter(|s| !s.is_empty()) + .collect() + }) + .unwrap_or_default(), dm_policy, group_policy, - group_allow_from: optional_env("SIGNAL_GROUP_ALLOW_FROM")? - .or_else(|| cs.signal_group_allow_from.clone()) - .map(|s| { - s.split(',') - .map(|e| e.trim().to_string()) - .filter(|s| !s.is_empty()) - .collect() - }) - .unwrap_or_default(), + group_allow_from: db_first_optional_string( + &cs.signal_group_allow_from, + "SIGNAL_GROUP_ALLOW_FROM", + )? + .map(|s| { + s.split(',') + .map(|e| e.trim().to_string()) + .filter(|s| !s.is_empty()) + .collect() + }) + .unwrap_or_default(), ignore_attachments: optional_env("SIGNAL_IGNORE_ATTACHMENTS")? .map(|s| s.to_lowercase() == "true" || s == "1") .unwrap_or(false), @@ -309,7 +347,7 @@ impl ChannelsConfig { None }; - let cli_enabled = parse_bool_env("CLI_ENABLED", cs.cli_enabled)?; + let cli_enabled = db_first_bool(cs.cli_enabled, defaults.cli_enabled, "CLI_ENABLED")?; Ok(Self { cli: CliConfig { @@ -318,13 +356,21 @@ impl ChannelsConfig { http, gateway, signal, - wasm_channels_dir: optional_env("WASM_CHANNELS_DIR")? - .map(PathBuf::from) - .or_else(|| cs.wasm_channels_dir.clone()) - .unwrap_or_else(default_channels_dir), - wasm_channels_enabled: parse_bool_env( - "WASM_CHANNELS_ENABLED", + wasm_channels_dir: { + // DB-first: use settings if explicitly set, else env, else default. + // defaults.wasm_channels_dir is None, so any Some(..) is an explicit DB override. + if let Some(ref db_dir) = cs.wasm_channels_dir { + db_dir.clone() + } else { + optional_env("WASM_CHANNELS_DIR")? + .map(PathBuf::from) + .unwrap_or_else(default_channels_dir) + } + }, + wasm_channels_enabled: db_first_bool( cs.wasm_channels_enabled, + defaults.wasm_channels_enabled, + "WASM_CHANNELS_ENABLED", )?, wasm_channel_owner_ids: { let mut ids = cs.wasm_channel_owner_ids.clone(); @@ -526,7 +572,9 @@ mod tests { settings.channels.gateway_enabled = true; settings.channels.gateway_host = Some("127.0.0.3".to_string()); settings.channels.gateway_port = Some(9191); - settings.channels.gateway_auth_token = Some("tok".to_string()); + // auth_token is env-only (security), set via env var + // SAFETY: under ENV_MUTEX + unsafe { std::env::set_var("GATEWAY_AUTH_TOKEN", "tok") }; settings.channels.signal_http_url = Some("http://127.0.0.1:8080".to_string()); settings.channels.signal_account = Some("+15551234567".to_string()); settings.channels.signal_allow_from = Some("+15551234567,+15557654321".to_string()); @@ -554,5 +602,8 @@ mod tests { PathBuf::from("/tmp/settings-channels") ); assert!(!cfg.wasm_channels_enabled); + + // SAFETY: under ENV_MUTEX + unsafe { std::env::remove_var("GATEWAY_AUTH_TOKEN") }; } } diff --git a/src/config/embeddings.rs b/src/config/embeddings.rs index 981839762f6..b44e00777da 100644 --- a/src/config/embeddings.rs +++ b/src/config/embeddings.rs @@ -2,7 +2,9 @@ use std::sync::Arc; use secrecy::{ExposeSecret, SecretString}; -use crate::config::helpers::{optional_env, parse_bool_env, parse_optional_env, validate_base_url}; +use crate::config::helpers::{ + db_first_bool, db_first_or_default, optional_env, parse_optional_env, validate_base_url, +}; use crate::error::ConfigError; use crate::llm::SessionManager; use crate::settings::Settings; @@ -71,22 +73,44 @@ pub(crate) fn default_dimension_for_model(model: &str) -> usize { impl EmbeddingsConfig { pub(crate) fn resolve(settings: &Settings) -> Result { - let openai_api_key = optional_env("OPENAI_API_KEY")?.map(SecretString::from); - - let provider = optional_env("EMBEDDING_PROVIDER")? - .unwrap_or_else(|| settings.embeddings.provider.clone()); + let defaults = crate::settings::EmbeddingsSettings::default(); - let model = - optional_env("EMBEDDING_MODEL")?.unwrap_or_else(|| settings.embeddings.model.clone()); + let openai_api_key = optional_env("OPENAI_API_KEY")?.map(SecretString::from); - let ollama_base_url = optional_env("OLLAMA_BASE_URL")? - .or_else(|| settings.ollama_base_url.clone()) - .unwrap_or_else(|| "http://localhost:11434".to_string()); + let provider = db_first_or_default( + &settings.embeddings.provider, + &defaults.provider, + "EMBEDDING_PROVIDER", + )?; + + let model = db_first_or_default( + &settings.embeddings.model, + &defaults.model, + "EMBEDDING_MODEL", + )?; + + // ollama_base_url lives on the top-level Settings, not the embeddings + // sub-struct. Use a manual DB > env > default chain. + let default_ollama_url = "http://localhost:11434".to_string(); + let ollama_base_url = match settings + .ollama_base_url + .as_ref() + .filter(|s| !s.is_empty()) + .cloned() + { + Some(url) => url, + None => optional_env("OLLAMA_BASE_URL")?.unwrap_or(default_ollama_url), + }; + // Dimension depends on the resolved model, not on a DB setting — env-only. let dimension = parse_optional_env("EMBEDDING_DIMENSION", default_dimension_for_model(&model))?; - let enabled = parse_bool_env("EMBEDDING_ENABLED", settings.embeddings.enabled)?; + let enabled = db_first_bool( + settings.embeddings.enabled, + defaults.enabled, + "EMBEDDING_ENABLED", + )?; let openai_base_url = optional_env("EMBEDDING_BASE_URL")?; @@ -207,9 +231,11 @@ mod tests { std::env::remove_var("EMBEDDING_ENABLED"); std::env::remove_var("EMBEDDING_PROVIDER"); std::env::remove_var("EMBEDDING_MODEL"); + std::env::remove_var("EMBEDDING_DIMENSION"); std::env::remove_var("OPENAI_API_KEY"); std::env::remove_var("EMBEDDING_BASE_URL"); std::env::remove_var("EMBEDDING_CACHE_SIZE"); + std::env::remove_var("OLLAMA_BASE_URL"); } } @@ -264,18 +290,21 @@ mod tests { } #[test] - fn embeddings_env_override_takes_precedence() { + fn db_settings_override_env() { let _guard = lock_env(); clear_embedding_env(); // SAFETY: Under ENV_MUTEX. unsafe { - std::env::set_var("EMBEDDING_ENABLED", "true"); + std::env::set_var("EMBEDDING_ENABLED", "false"); + std::env::set_var("EMBEDDING_PROVIDER", "ollama"); + std::env::set_var("EMBEDDING_MODEL", "all-minilm"); } let settings = Settings { embeddings: EmbeddingsSettings { - enabled: false, - ..Default::default() + enabled: true, + provider: "openai".to_string(), + model: "text-embedding-3-large".to_string(), }, ..Default::default() }; @@ -283,12 +312,55 @@ mod tests { let config = EmbeddingsConfig::resolve(&settings).expect("resolve should succeed"); assert!( config.enabled, - "EMBEDDING_ENABLED=true env var should override settings" + "DB enabled=true should win over env EMBEDDING_ENABLED=false" + ); + assert_eq!(config.provider, "openai", "DB provider should win over env"); + assert_eq!( + config.model, "text-embedding-3-large", + "DB model should win over env" + ); + + // SAFETY: Under ENV_MUTEX. + unsafe { + std::env::remove_var("EMBEDDING_ENABLED"); + std::env::remove_var("EMBEDDING_PROVIDER"); + std::env::remove_var("EMBEDDING_MODEL"); + } + } + + #[test] + fn env_used_when_no_db_setting() { + let _guard = lock_env(); + clear_embedding_env(); + // SAFETY: Under ENV_MUTEX. + unsafe { + std::env::set_var("EMBEDDING_ENABLED", "true"); + std::env::set_var("EMBEDDING_PROVIDER", "ollama"); + std::env::set_var("EMBEDDING_MODEL", "nomic-embed-text"); + } + + // Settings left at defaults — no explicit DB/TOML override + let settings = Settings::default(); + + let config = EmbeddingsConfig::resolve(&settings).expect("resolve should succeed"); + assert!( + config.enabled, + "env EMBEDDING_ENABLED should be used when settings at default" + ); + assert_eq!( + config.provider, "ollama", + "env EMBEDDING_PROVIDER should be used when settings at default" + ); + assert_eq!( + config.model, "nomic-embed-text", + "env EMBEDDING_MODEL should be used when settings at default" ); // SAFETY: Under ENV_MUTEX. unsafe { std::env::remove_var("EMBEDDING_ENABLED"); + std::env::remove_var("EMBEDDING_PROVIDER"); + std::env::remove_var("EMBEDDING_MODEL"); } } diff --git a/src/config/heartbeat.rs b/src/config/heartbeat.rs index ecf25333349..639c3d4c416 100644 --- a/src/config/heartbeat.rs +++ b/src/config/heartbeat.rs @@ -1,4 +1,7 @@ -use crate::config::helpers::{optional_env, parse_bool_env, parse_option_env, parse_optional_env}; +use crate::config::helpers::{ + db_first_bool, db_first_option, db_first_optional_string, db_first_or_default, optional_env, + parse_bool_env, +}; use crate::error::ConfigError; use crate::settings::Settings; @@ -44,8 +47,11 @@ impl Default for HeartbeatConfig { impl HeartbeatConfig { pub(crate) fn resolve(settings: &Settings) -> Result { + let defaults = crate::settings::HeartbeatSettings::default(); + + // fire_at: DB > env, then parse into NaiveTime let fire_at_str = - optional_env("HEARTBEAT_FIRE_AT")?.or_else(|| settings.heartbeat.fire_at.clone()); + db_first_optional_string(&settings.heartbeat.fire_at, "HEARTBEAT_FIRE_AT")?; let fire_at = fire_at_str .map(|s| { chrono::NaiveTime::parse_from_str(&s, "%H:%M").map_err(|e| { @@ -57,31 +63,24 @@ impl HeartbeatConfig { }) .transpose()?; - Ok(Self { - enabled: parse_bool_env("HEARTBEAT_ENABLED", settings.heartbeat.enabled)?, - interval_secs: parse_optional_env( - "HEARTBEAT_INTERVAL_SECS", - settings.heartbeat.interval_secs, - )?, - notify_channel: optional_env("HEARTBEAT_NOTIFY_CHANNEL")? - .or_else(|| settings.heartbeat.notify_channel.clone()), - notify_user: optional_env("HEARTBEAT_NOTIFY_USER")? - .or_else(|| settings.heartbeat.notify_user.clone()), - fire_at, - quiet_hours_start: parse_option_env::("HEARTBEAT_QUIET_START")? - .or(settings.heartbeat.quiet_hours_start) - .map(|h| { - if h > 23 { - return Err(ConfigError::InvalidValue { - key: "HEARTBEAT_QUIET_START".into(), - message: "must be 0-23".into(), - }); - } - Ok(h) - }) - .transpose()?, - quiet_hours_end: parse_option_env::("HEARTBEAT_QUIET_END")? - .or(settings.heartbeat.quiet_hours_end) + // quiet_hours: DB > env (using db_first_option for shadow warnings) + let quiet_hours_start = db_first_option( + &settings.heartbeat.quiet_hours_start, + "HEARTBEAT_QUIET_START", + )? + .map(|h| { + if h > 23 { + return Err(ConfigError::InvalidValue { + key: "HEARTBEAT_QUIET_START".into(), + message: "must be 0-23".into(), + }); + } + Ok(h) + }) + .transpose()?; + + let quiet_hours_end = + db_first_option(&settings.heartbeat.quiet_hours_end, "HEARTBEAT_QUIET_END")? .map(|h| { if h > 23 { return Err(ConfigError::InvalidValue { @@ -91,10 +90,33 @@ impl HeartbeatConfig { } Ok(h) }) - .transpose()?, + .transpose()?; + + Ok(Self { + enabled: db_first_bool( + settings.heartbeat.enabled, + defaults.enabled, + "HEARTBEAT_ENABLED", + )?, + interval_secs: db_first_or_default( + &settings.heartbeat.interval_secs, + &defaults.interval_secs, + "HEARTBEAT_INTERVAL_SECS", + )?, + notify_channel: db_first_optional_string( + &settings.heartbeat.notify_channel, + "HEARTBEAT_NOTIFY_CHANNEL", + )?, + notify_user: db_first_optional_string( + &settings.heartbeat.notify_user, + "HEARTBEAT_NOTIFY_USER", + )?, + fire_at, + quiet_hours_start, + quiet_hours_end, timezone: { - let tz = optional_env("HEARTBEAT_TIMEZONE")? - .or_else(|| settings.heartbeat.timezone.clone()); + let tz = + db_first_optional_string(&settings.heartbeat.timezone, "HEARTBEAT_TIMEZONE")?; if let Some(ref tz_str) = tz && crate::timezone::parse_timezone(tz_str).is_none() { @@ -105,7 +127,12 @@ impl HeartbeatConfig { } tz }, - multi_tenant: parse_bool_env("HEARTBEAT_MULTI_TENANT", false)?, + // Auto-detect multi-tenant mode from GATEWAY_USER_TOKENS presence, + // or allow explicit override via HEARTBEAT_MULTI_TENANT. Stays env-only. + multi_tenant: parse_bool_env( + "HEARTBEAT_MULTI_TENANT", + optional_env("GATEWAY_USER_TOKENS")?.is_some(), + )?, }) } } @@ -113,10 +140,11 @@ impl HeartbeatConfig { #[cfg(test)] mod tests { use super::*; + use crate::config::helpers::lock_env; #[test] - fn test_quiet_hours_settings_fallback() { - // When env vars are not set, settings values should be used + fn test_quiet_hours_settings_have_priority() { + // DB/settings values should take priority over env let mut settings = Settings::default(); settings.heartbeat.quiet_hours_start = Some(22); settings.heartbeat.quiet_hours_end = Some(6); @@ -163,4 +191,116 @@ mod tests { let config = HeartbeatConfig::resolve(&settings).expect("resolve"); assert_eq!(config.timezone.as_deref(), Some("America/New_York")); } + + #[test] + fn test_db_first_enabled_beats_env() { + let _guard = lock_env(); + // SAFETY: under ENV_MUTEX + unsafe { std::env::set_var("HEARTBEAT_ENABLED", "false") }; + + let mut settings = Settings::default(); + settings.heartbeat.enabled = true; // DB says enabled + + let config = HeartbeatConfig::resolve(&settings).expect("resolve"); + assert!(config.enabled, "DB value (true) should beat env (false)"); + + unsafe { std::env::remove_var("HEARTBEAT_ENABLED") }; + } + + #[test] + fn test_db_first_interval_beats_env() { + let _guard = lock_env(); + unsafe { std::env::set_var("HEARTBEAT_INTERVAL_SECS", "999") }; + + let mut settings = Settings::default(); + settings.heartbeat.interval_secs = 600; // DB says 600 + + let config = HeartbeatConfig::resolve(&settings).expect("resolve"); + assert_eq!(config.interval_secs, 600, "DB value should beat env"); + + unsafe { std::env::remove_var("HEARTBEAT_INTERVAL_SECS") }; + } + + #[test] + fn test_db_first_notify_channel_beats_env() { + let _guard = lock_env(); + unsafe { std::env::set_var("HEARTBEAT_NOTIFY_CHANNEL", "env-channel") }; + + let mut settings = Settings::default(); + settings.heartbeat.notify_channel = Some("db-channel".to_string()); + + let config = HeartbeatConfig::resolve(&settings).expect("resolve"); + assert_eq!( + config.notify_channel.as_deref(), + Some("db-channel"), + "DB value should beat env" + ); + + unsafe { std::env::remove_var("HEARTBEAT_NOTIFY_CHANNEL") }; + } + + #[test] + fn test_env_fallback_when_db_at_default() { + let _guard = lock_env(); + unsafe { std::env::set_var("HEARTBEAT_INTERVAL_SECS", "999") }; + + // Settings at default => env should win + let settings = Settings::default(); + + let config = HeartbeatConfig::resolve(&settings).expect("resolve"); + assert_eq!( + config.interval_secs, 999, + "env should win when DB at default" + ); + + unsafe { std::env::remove_var("HEARTBEAT_INTERVAL_SECS") }; + } + + #[test] + fn test_fire_at_db_first() { + let _guard = lock_env(); + unsafe { std::env::set_var("HEARTBEAT_FIRE_AT", "08:00") }; + + let mut settings = Settings::default(); + settings.heartbeat.fire_at = Some("14:30".to_string()); + + let config = HeartbeatConfig::resolve(&settings).expect("resolve"); + assert_eq!( + config.fire_at, + Some(chrono::NaiveTime::from_hms_opt(14, 30, 0).unwrap()), + "DB fire_at should beat env" + ); + + unsafe { std::env::remove_var("HEARTBEAT_FIRE_AT") }; + } + + #[test] + fn test_timezone_db_first() { + let _guard = lock_env(); + unsafe { std::env::set_var("HEARTBEAT_TIMEZONE", "UTC") }; + + let mut settings = Settings::default(); + settings.heartbeat.timezone = Some("America/New_York".to_string()); + + let config = HeartbeatConfig::resolve(&settings).expect("resolve"); + assert_eq!( + config.timezone.as_deref(), + Some("America/New_York"), + "DB timezone should beat env" + ); + + unsafe { std::env::remove_var("HEARTBEAT_TIMEZONE") }; + } + + #[test] + fn test_multi_tenant_stays_env_only() { + let _guard = lock_env(); + unsafe { std::env::set_var("HEARTBEAT_MULTI_TENANT", "true") }; + + let settings = Settings::default(); + let config = HeartbeatConfig::resolve(&settings).expect("resolve"); + assert!(config.multi_tenant, "multi_tenant should read from env"); + + unsafe { std::env::remove_var("HEARTBEAT_MULTI_TENANT") }; + } } diff --git a/src/config/helpers.rs b/src/config/helpers.rs index ff5ee70629b..278e9141d24 100644 --- a/src/config/helpers.rs +++ b/src/config/helpers.rs @@ -331,6 +331,99 @@ pub(crate) fn validate_base_url(url: &str, field_name: &str) -> Result<(), Confi Ok(()) } +// --------------------------------------------------------------------------- +// DB-first resolution helpers (DB > env > default) +// --------------------------------------------------------------------------- + +/// Log a warning when a DB/TOML setting shadows a set env var. +/// +/// Only checks real env vars (`std::env::var`), not runtime overrides or +/// injected vars — those are internal and don't warrant operator warnings. +/// Values are intentionally NOT logged to avoid leaking sensitive data. +fn warn_if_db_shadows_env(env_key: &str) { + if std::env::var(env_key).is_ok_and(|v| !v.is_empty()) { + tracing::warn!( + env_key = %env_key, + "{env_key} env var is set but a DB or TOML setting takes priority. \ + Remove the setting from DB/TOML to use the env var." + ); + } +} + +/// Resolve with DB > env > default priority for concrete settings fields. +/// +/// If `settings_val != default_val`, the settings value wins (it was explicitly +/// set in DB or TOML). Otherwise falls back to `optional_env(env_key)`, then +/// `default_val`. +/// +/// **Limitation:** Uses `settings_val != default_val` as a heuristic for +/// "was this field explicitly set." If a user deliberately sets a DB value +/// equal to the default, it's indistinguishable from "unset" and the env +/// var will win. This matches `merge_from()` semantics and is acceptable +/// since setting a value to its default is effectively a no-op. +pub(crate) fn db_first_or_default( + settings_val: &T, + default_val: &T, + env_key: &str, +) -> Result +where + T: std::str::FromStr + Clone + PartialEq + std::fmt::Display, + T::Err: std::fmt::Display, +{ + if settings_val != default_val { + warn_if_db_shadows_env(env_key); + return Ok(settings_val.clone()); + } + parse_optional_env(env_key, default_val.clone()) +} + +/// Resolve a bool with DB > env > default priority. +pub(crate) fn db_first_bool( + settings_val: bool, + default_val: bool, + env_key: &str, +) -> Result { + if settings_val != default_val { + warn_if_db_shadows_env(env_key); + return Ok(settings_val); + } + parse_bool_env(env_key, default_val) +} + +/// Resolve an `Option` with DB > env priority (no hardcoded default). +/// +/// Non-empty `Some` means DB set it; `None` or empty falls back to env. +pub(crate) fn db_first_optional_string( + settings_val: &Option, + env_key: &str, +) -> Result, ConfigError> { + if let Some(val) = settings_val + && !val.is_empty() + { + warn_if_db_shadows_env(env_key); + return Ok(Some(val.clone())); + } + optional_env(env_key) +} + +/// Resolve an `Option` with DB > env priority (no hardcoded default). +/// +/// `Some(v)` means DB set it; `None` falls back to env. +pub(crate) fn db_first_option( + settings_val: &Option, + env_key: &str, +) -> Result, ConfigError> +where + T: std::str::FromStr + Clone + std::fmt::Display, + T::Err: std::fmt::Display, +{ + if let Some(val) = settings_val { + warn_if_db_shadows_env(env_key); + return Ok(Some(val.clone())); + } + parse_option_env(env_key) +} + #[cfg(test)] mod tests { use super::*; @@ -519,4 +612,144 @@ mod tests { "Expected DNS resolution failure, got: {err}" ); } + + // --- db_first_* helper tests --- + + #[test] + fn db_first_or_default_prefers_settings_over_env() { + let _guard = lock_env(); + let key = "IRONCLAW_TEST_DB_FIRST_1"; + // SAFETY: under ENV_MUTEX + unsafe { std::env::set_var(key, "from-env") }; + + let result: String = + db_first_or_default(&"from-db".to_string(), &"default".to_string(), key) + .expect("should resolve"); + assert_eq!(result, "from-db", "DB value should win over env"); + + unsafe { std::env::remove_var(key) }; + } + + #[test] + fn db_first_or_default_falls_back_to_env() { + let _guard = lock_env(); + let key = "IRONCLAW_TEST_DB_FIRST_2"; + unsafe { std::env::set_var(key, "from-env") }; + + // settings_val == default_val → treated as "unset" + let result: String = + db_first_or_default(&"default".to_string(), &"default".to_string(), key) + .expect("should resolve"); + assert_eq!( + result, "from-env", + "env should win when settings at default" + ); + + unsafe { std::env::remove_var(key) }; + } + + #[test] + fn db_first_or_default_uses_default_when_neither_set() { + let _guard = lock_env(); + let key = "IRONCLAW_TEST_DB_FIRST_3"; + unsafe { std::env::remove_var(key) }; + + let result: String = + db_first_or_default(&"default".to_string(), &"default".to_string(), key) + .expect("should resolve"); + assert_eq!(result, "default"); + } + + #[test] + fn db_first_bool_prefers_settings() { + let _guard = lock_env(); + let key = "IRONCLAW_TEST_DB_FIRST_BOOL_1"; + unsafe { std::env::set_var(key, "false") }; + + let result = db_first_bool(true, false, key).expect("should resolve"); + assert!(result, "DB true should win over env false"); + + unsafe { std::env::remove_var(key) }; + } + + #[test] + fn db_first_bool_falls_back_to_env() { + let _guard = lock_env(); + let key = "IRONCLAW_TEST_DB_FIRST_BOOL_2"; + unsafe { std::env::set_var(key, "true") }; + + // settings == default → falls back to env + let result = db_first_bool(false, false, key).expect("should resolve"); + assert!(result, "env should win when settings at default"); + + unsafe { std::env::remove_var(key) }; + } + + #[test] + fn db_first_optional_string_prefers_settings() { + let _guard = lock_env(); + let key = "IRONCLAW_TEST_DB_FIRST_OPT_1"; + unsafe { std::env::set_var(key, "from-env") }; + + let val = Some("from-db".to_string()); + let result = db_first_optional_string(&val, key).expect("should resolve"); + assert_eq!(result, Some("from-db".to_string())); + + unsafe { std::env::remove_var(key) }; + } + + #[test] + fn db_first_optional_string_falls_back_to_env() { + let _guard = lock_env(); + let key = "IRONCLAW_TEST_DB_FIRST_OPT_2"; + unsafe { std::env::set_var(key, "from-env") }; + + let result = db_first_optional_string(&None, key).expect("should resolve"); + assert_eq!(result, Some("from-env".to_string())); + + unsafe { std::env::remove_var(key) }; + } + + #[test] + fn db_first_optional_string_empty_treated_as_unset() { + let _guard = lock_env(); + let key = "IRONCLAW_TEST_DB_FIRST_OPT_3"; + unsafe { std::env::set_var(key, "from-env") }; + + let val = Some(String::new()); + let result = db_first_optional_string(&val, key).expect("should resolve"); + assert_eq!( + result, + Some("from-env".to_string()), + "empty string should be treated as unset" + ); + + unsafe { std::env::remove_var(key) }; + } + + #[test] + fn db_first_option_prefers_settings() { + let _guard = lock_env(); + let key = "IRONCLAW_TEST_DB_FIRST_OPT_T_1"; + unsafe { std::env::set_var(key, "99") }; + + let val: Option = Some(42); + let result = db_first_option(&val, key).expect("should resolve"); + assert_eq!(result, Some(42)); + + unsafe { std::env::remove_var(key) }; + } + + #[test] + fn db_first_option_falls_back_to_env() { + let _guard = lock_env(); + let key = "IRONCLAW_TEST_DB_FIRST_OPT_T_2"; + unsafe { std::env::set_var(key, "99") }; + + let val: Option = None; + let result = db_first_option(&val, key).expect("should resolve"); + assert_eq!(result, Some(99)); + + unsafe { std::env::remove_var(key) }; + } } diff --git a/src/config/hygiene.rs b/src/config/hygiene.rs index b510933af50..f408de4563b 100644 --- a/src/config/hygiene.rs +++ b/src/config/hygiene.rs @@ -1,6 +1,7 @@ use crate::bootstrap::ironclaw_base_dir; -use crate::config::helpers::{parse_bool_env, parse_optional_env}; +use crate::config::helpers::{db_first_bool, db_first_or_default}; use crate::error::ConfigError; +use crate::settings::Settings; /// Memory hygiene configuration. /// @@ -30,15 +31,27 @@ impl Default for HygieneConfig { } impl HygieneConfig { - pub(crate) fn resolve() -> Result { + pub(crate) fn resolve(settings: &Settings) -> Result { + let defaults = crate::settings::HygieneSettings::default(); + let hs = &settings.hygiene; + Ok(Self { - enabled: parse_bool_env("MEMORY_HYGIENE_ENABLED", true)?, - daily_retention_days: parse_optional_env("MEMORY_HYGIENE_DAILY_RETENTION_DAYS", 30)?, - conversation_retention_days: parse_optional_env( + enabled: db_first_bool(hs.enabled, defaults.enabled, "MEMORY_HYGIENE_ENABLED")?, + daily_retention_days: db_first_or_default( + &hs.daily_retention_days, + &defaults.daily_retention_days, + "MEMORY_HYGIENE_DAILY_RETENTION_DAYS", + )?, + conversation_retention_days: db_first_or_default( + &hs.conversation_retention_days, + &defaults.conversation_retention_days, "MEMORY_HYGIENE_CONVERSATION_RETENTION_DAYS", - 7, )?, - cadence_hours: parse_optional_env("MEMORY_HYGIENE_CADENCE_HOURS", 12)?, + cadence_hours: db_first_or_default( + &hs.cadence_hours, + &defaults.cadence_hours, + "MEMORY_HYGIENE_CADENCE_HOURS", + )?, }) } diff --git a/src/config/mod.rs b/src/config/mod.rs index ed9b6a5fff9..54d639d5524 100644 --- a/src/config/mod.rs +++ b/src/config/mod.rs @@ -1,10 +1,19 @@ //! Configuration for IronClaw. //! -//! Settings are loaded from env vars, the DB settings table, TOML config, -//! and built-in defaults. Priority varies by subsystem: +//! Settings are loaded with priority: **DB/TOML > env > default**. //! -//! - **LLM settings** (backend, model, api_key, base_url): DB > env > default -//! - **Most other settings** (agent, channels, tunnel, …): env > DB > default +//! DB and TOML are merged into a single `Settings` struct before +//! resolution (DB wins over TOML when both set the same field). +//! Resolvers then check settings before env vars. +//! +//! For concrete (non-`Option`) fields, a settings value equal to the +//! built-in default is treated as "unset" and falls through to env. +//! +//! Exceptions: +//! - Bootstrap configs (database, secrets): env-only (DB not yet available) +//! - Security-sensitive fields (allow_local_tools, allow_full_access, +//! cost limits, auth tokens): env-only +//! - API keys: env/secrets store only //! //! `DATABASE_URL` lives in `~/.ironclaw/.env` (loaded via dotenvy early //! in startup). @@ -191,9 +200,9 @@ impl Config { /// Load configuration from environment variables and the database. /// - /// 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. + /// Priority: DB/TOML > env > default. TOML is loaded first as a + /// base, then DB values are merged on top. Subsystem resolvers check + /// the merged settings before env vars (except bootstrap/security fields). pub async fn from_db( store: &(dyn crate::db::SettingsStore + Sync), user_id: &str, @@ -203,9 +212,8 @@ 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. + /// Priority: DB/TOML > env > default. TOML is loaded as the base, + /// then DB values are merged on top. See module docs for exceptions. pub async fn from_db_with_toml( store: &(dyn crate::db::SettingsStore + Sync), user_id: &str, @@ -366,13 +374,13 @@ impl Config { secrets: SecretsConfig::resolve().await?, builder: BuilderModeConfig::resolve(settings)?, heartbeat: HeartbeatConfig::resolve(settings)?, - hygiene: HygieneConfig::resolve()?, - routines: RoutineConfig::resolve()?, + hygiene: HygieneConfig::resolve(settings)?, + routines: RoutineConfig::resolve(settings)?, sandbox: SandboxModeConfig::resolve(settings)?, claude_code: ClaudeCodeConfig::resolve(settings)?, - skills: SkillsConfig::resolve()?, + skills: SkillsConfig::resolve(settings)?, transcription: TranscriptionConfig::resolve(settings)?, - search: WorkspaceSearchConfig::resolve()?, + search: WorkspaceSearchConfig::resolve(settings)?, workspace, observability: crate::observability::ObservabilityConfig { backend: std::env::var("OBSERVABILITY_BACKEND").unwrap_or_else(|_| "none".into()), diff --git a/src/config/routines.rs b/src/config/routines.rs index c82aa8b54d2..d28c2f50f03 100644 --- a/src/config/routines.rs +++ b/src/config/routines.rs @@ -1,5 +1,6 @@ -use crate::config::helpers::{parse_bool_env, parse_optional_env}; +use crate::config::helpers::{db_first_bool, db_first_or_default}; use crate::error::ConfigError; +use crate::settings::Settings; /// Routines configuration. #[derive(Debug, Clone)] @@ -35,15 +36,42 @@ impl Default for RoutineConfig { } impl RoutineConfig { - pub(crate) fn resolve() -> Result { - let max_iterations: u32 = parse_optional_env("ROUTINES_LIGHTWEIGHT_MAX_ITERATIONS", 3)?; + pub(crate) fn resolve(settings: &Settings) -> Result { + let defaults = crate::settings::RoutineSettings::default(); + let rs = &settings.routines; + + let max_iterations: u32 = db_first_or_default( + &rs.lightweight_max_iterations, + &defaults.lightweight_max_iterations, + "ROUTINES_LIGHTWEIGHT_MAX_ITERATIONS", + )?; Ok(Self { - enabled: parse_bool_env("ROUTINES_ENABLED", true)?, - cron_check_interval_secs: parse_optional_env("ROUTINES_CRON_INTERVAL", 15)?, - max_concurrent_routines: parse_optional_env("ROUTINES_MAX_CONCURRENT", 10)?, - default_cooldown_secs: parse_optional_env("ROUTINES_DEFAULT_COOLDOWN", 300)?, - max_lightweight_tokens: parse_optional_env("ROUTINES_MAX_TOKENS", 4096)?, - lightweight_tools_enabled: parse_bool_env("ROUTINES_LIGHTWEIGHT_TOOLS", true)?, + enabled: db_first_bool(rs.enabled, defaults.enabled, "ROUTINES_ENABLED")?, + cron_check_interval_secs: db_first_or_default( + &rs.cron_check_interval_secs, + &defaults.cron_check_interval_secs, + "ROUTINES_CRON_INTERVAL", + )?, + max_concurrent_routines: db_first_or_default( + &rs.max_concurrent_routines, + &defaults.max_concurrent_routines, + "ROUTINES_MAX_CONCURRENT", + )?, + default_cooldown_secs: db_first_or_default( + &rs.default_cooldown_secs, + &defaults.default_cooldown_secs, + "ROUTINES_DEFAULT_COOLDOWN", + )?, + max_lightweight_tokens: db_first_or_default( + &rs.max_lightweight_tokens, + &defaults.max_lightweight_tokens, + "ROUTINES_MAX_TOKENS", + )?, + lightweight_tools_enabled: db_first_bool( + rs.lightweight_tools_enabled, + defaults.lightweight_tools_enabled, + "ROUTINES_LIGHTWEIGHT_TOOLS", + )?, lightweight_max_iterations: max_iterations.min(5), // cap at 5 }) } diff --git a/src/config/safety.rs b/src/config/safety.rs index edeceee01d4..953256d8a61 100644 --- a/src/config/safety.rs +++ b/src/config/safety.rs @@ -1,4 +1,4 @@ -use crate::config::helpers::{parse_bool_env, parse_optional_env}; +use crate::config::helpers::{db_first_bool, db_first_or_default}; use crate::error::ConfigError; pub use ironclaw_safety::SafetyConfig; @@ -7,11 +7,17 @@ pub(crate) fn resolve_safety_config( settings: &crate::settings::Settings, ) -> Result { let ss = &settings.safety; + let defaults = crate::settings::SafetySettings::default(); Ok(SafetyConfig { - max_output_length: parse_optional_env("SAFETY_MAX_OUTPUT_LENGTH", ss.max_output_length)?, - injection_check_enabled: parse_bool_env( - "SAFETY_INJECTION_CHECK_ENABLED", + max_output_length: db_first_or_default( + &ss.max_output_length, + &defaults.max_output_length, + "SAFETY_MAX_OUTPUT_LENGTH", + )?, + injection_check_enabled: db_first_bool( ss.injection_check_enabled, + defaults.injection_check_enabled, + "SAFETY_INJECTION_CHECK_ENABLED", )?, }) } @@ -35,9 +41,10 @@ mod tests { } #[test] - fn env_overrides_settings() { + fn db_settings_override_env() { let _guard = lock_env(); let mut settings = Settings::default(); + // Non-default value simulates an explicit DB/TOML setting settings.safety.max_output_length = 42; // SAFETY: Under ENV_MUTEX, no concurrent env access. @@ -45,6 +52,25 @@ mod tests { let cfg = resolve_safety_config(&settings).expect("resolve"); unsafe { std::env::remove_var("SAFETY_MAX_OUTPUT_LENGTH") }; + // DB value (42) wins over env value (7) + assert_eq!(cfg.max_output_length, 42); + } + + #[test] + fn env_used_when_no_db_setting() { + let _guard = lock_env(); + // Settings left at defaults — no explicit DB/TOML override + let settings = Settings::default(); + + // SAFETY: Under ENV_MUTEX, no concurrent env access. + unsafe { std::env::set_var("SAFETY_MAX_OUTPUT_LENGTH", "7") }; + unsafe { std::env::set_var("SAFETY_INJECTION_CHECK_ENABLED", "false") }; + let cfg = resolve_safety_config(&settings).expect("resolve"); + unsafe { std::env::remove_var("SAFETY_MAX_OUTPUT_LENGTH") }; + unsafe { std::env::remove_var("SAFETY_INJECTION_CHECK_ENABLED") }; + + // Env values win when settings are at their defaults assert_eq!(cfg.max_output_length, 7); + assert!(!cfg.injection_check_enabled); } } diff --git a/src/config/sandbox.rs b/src/config/sandbox.rs index 01a8c327efc..7f5adb7a01b 100644 --- a/src/config/sandbox.rs +++ b/src/config/sandbox.rs @@ -1,4 +1,7 @@ -use crate::config::helpers::{optional_env, parse_bool_env, parse_optional_env, parse_string_env}; +use crate::config::helpers::{ + db_first_bool, db_first_or_default, optional_env, parse_bool_env, parse_optional_env, + parse_string_env, +}; use crate::error::ConfigError; /// Docker sandbox configuration. @@ -54,16 +57,16 @@ impl Default for SandboxModeConfig { impl SandboxModeConfig { pub(crate) fn resolve(settings: &crate::settings::Settings) -> Result { let ss = &settings.sandbox; - - let extra_domains = optional_env("SANDBOX_EXTRA_DOMAINS")? - .map(|s| s.split(',').map(|d| d.trim().to_string()).collect()) - .unwrap_or_else(|| { - if ss.extra_allowed_domains.is_empty() { - Vec::new() - } else { - ss.extra_allowed_domains.clone() - } - }); + let defaults = crate::settings::SandboxSettings::default(); + + // extra_allowed_domains: DB wins if non-empty, otherwise env, otherwise empty. + let extra_domains = if !ss.extra_allowed_domains.is_empty() { + ss.extra_allowed_domains.clone() + } else { + optional_env("SANDBOX_EXTRA_DOMAINS")? + .map(|s| s.split(',').map(|d| d.trim().to_string()).collect()) + .unwrap_or_default() + }; // reaper/orphan fields have no Settings counterpart — env > default only. let reaper_interval_secs: u64 = parse_optional_env("SANDBOX_REAPER_INTERVAL_SECS", 300)?; @@ -85,15 +88,31 @@ impl SandboxModeConfig { } Ok(Self { - enabled: parse_bool_env("SANDBOX_ENABLED", ss.enabled)?, - policy: parse_string_env("SANDBOX_POLICY", ss.policy.clone())?, - // allow_full_access has no Settings counterpart — env > default only. + enabled: db_first_bool(ss.enabled, defaults.enabled, "SANDBOX_ENABLED")?, + policy: db_first_or_default(&ss.policy, &defaults.policy, "SANDBOX_POLICY")?, + // allow_full_access has no Settings counterpart — env > default only (security). allow_full_access: parse_bool_env("SANDBOX_ALLOW_FULL_ACCESS", false)?, - timeout_secs: parse_optional_env("SANDBOX_TIMEOUT_SECS", ss.timeout_secs)?, - memory_limit_mb: parse_optional_env("SANDBOX_MEMORY_LIMIT_MB", ss.memory_limit_mb)?, - cpu_shares: parse_optional_env("SANDBOX_CPU_SHARES", ss.cpu_shares)?, - image: parse_string_env("SANDBOX_IMAGE", ss.image.clone())?, - auto_pull_image: parse_bool_env("SANDBOX_AUTO_PULL", ss.auto_pull_image)?, + timeout_secs: db_first_or_default( + &ss.timeout_secs, + &defaults.timeout_secs, + "SANDBOX_TIMEOUT_SECS", + )?, + memory_limit_mb: db_first_or_default( + &ss.memory_limit_mb, + &defaults.memory_limit_mb, + "SANDBOX_MEMORY_LIMIT_MB", + )?, + cpu_shares: db_first_or_default( + &ss.cpu_shares, + &defaults.cpu_shares, + "SANDBOX_CPU_SHARES", + )?, + image: db_first_or_default(&ss.image, &defaults.image, "SANDBOX_IMAGE")?, + auto_pull_image: db_first_bool( + ss.auto_pull_image, + defaults.auto_pull_image, + "SANDBOX_AUTO_PULL", + )?, extra_allowed_domains: extra_domains, reaper_interval_secs, orphan_threshold_secs, @@ -264,19 +283,28 @@ impl ClaudeCodeConfig { } pub(crate) fn resolve(settings: &crate::settings::Settings) -> Result { + let ss = &settings.sandbox; let defaults = Self::default(); Ok(Self { - // Use settings.sandbox.claude_code_enabled as fallback (written by setup wizard). - enabled: parse_bool_env("CLAUDE_CODE_ENABLED", settings.sandbox.claude_code_enabled)?, + enabled: db_first_bool( + ss.claude_code_enabled, + defaults.enabled, + "CLAUDE_CODE_ENABLED", + )?, + // config_dir has no Settings counterpart — env > default only. config_dir: optional_env("CLAUDE_CONFIG_DIR")? .map(std::path::PathBuf::from) .unwrap_or(defaults.config_dir), + // model has no Settings counterpart — env > default only. model: parse_string_env("CLAUDE_CODE_MODEL", defaults.model)?, + // max_turns has no Settings counterpart — env > default only. max_turns: parse_optional_env("CLAUDE_CODE_MAX_TURNS", defaults.max_turns)?, + // memory_limit_mb has no Settings counterpart — env > default only. memory_limit_mb: parse_optional_env( "CLAUDE_CODE_MEMORY_LIMIT_MB", defaults.memory_limit_mb, )?, + // allowed_tools has no Settings counterpart — env > default only. allowed_tools: optional_env("CLAUDE_CODE_ALLOWED_TOOLS")? .map(|s| { s.split(',') @@ -607,7 +635,7 @@ mod tests { } #[test] - fn sandbox_env_overrides_settings() { + fn sandbox_db_settings_override_env() { let _guard = crate::config::helpers::lock_env(); let mut settings = crate::settings::Settings::default(); settings.sandbox.timeout_secs = 999; @@ -617,7 +645,26 @@ mod tests { let cfg = SandboxModeConfig::resolve(&settings).expect("resolve"); unsafe { std::env::remove_var("SANDBOX_TIMEOUT_SECS") }; - assert_eq!(cfg.timeout_secs, 5); + // DB value (999) wins over env (5) under DB-first priority. + assert_eq!(cfg.timeout_secs, 999); + } + + #[test] + fn sandbox_env_used_when_no_db_setting() { + let _guard = crate::config::helpers::lock_env(); + // Default settings — all fields at their defaults, so DB is "unset". + let settings = crate::settings::Settings::default(); + + // SAFETY: Under ENV_MUTEX, no concurrent env access. + unsafe { std::env::set_var("SANDBOX_TIMEOUT_SECS", "42") }; + unsafe { std::env::set_var("SANDBOX_MEMORY_LIMIT_MB", "512") }; + let cfg = SandboxModeConfig::resolve(&settings).expect("resolve"); + unsafe { std::env::remove_var("SANDBOX_TIMEOUT_SECS") }; + unsafe { std::env::remove_var("SANDBOX_MEMORY_LIMIT_MB") }; + + // Env values win when settings are at their defaults. + assert_eq!(cfg.timeout_secs, 42); + assert_eq!(cfg.memory_limit_mb, 512); } // ── ClaudeCodeConfig settings fallback tests ──────────────────── @@ -641,7 +688,7 @@ mod tests { } #[test] - fn claude_code_env_overrides_settings() { + fn claude_code_db_settings_override_env() { let _guard = crate::config::helpers::lock_env(); let mut settings = crate::settings::Settings::default(); settings.sandbox.claude_code_enabled = true; @@ -651,7 +698,8 @@ mod tests { let cfg = ClaudeCodeConfig::resolve(&settings).expect("resolve"); unsafe { std::env::remove_var("CLAUDE_CODE_ENABLED") }; - assert!(!cfg.enabled); + // DB value (true) wins over env (false) under DB-first priority. + assert!(cfg.enabled); } #[test] diff --git a/src/config/search.rs b/src/config/search.rs index e6b663cf8b7..363c5d678ff 100644 --- a/src/config/search.rs +++ b/src/config/search.rs @@ -1,5 +1,6 @@ -use crate::config::helpers::{optional_env, parse_optional_env}; +use crate::config::helpers::{db_first_option, db_first_or_default, parse_optional_env}; use crate::error::ConfigError; +use crate::settings::Settings; use crate::workspace::FusionStrategy; /// Workspace search configuration resolved from environment variables. @@ -33,30 +34,41 @@ impl Default for WorkspaceSearchConfig { } impl WorkspaceSearchConfig { - pub(crate) fn resolve() -> Result { - let fusion_strategy = match optional_env("SEARCH_FUSION_STRATEGY")? { - Some(s) => match s.to_lowercase().as_str() { - "rrf" => FusionStrategy::Rrf, - "weighted" => FusionStrategy::WeightedScore, - other => { - return Err(ConfigError::InvalidValue { - key: "SEARCH_FUSION_STRATEGY".to_string(), - message: format!("must be 'rrf' or 'weighted', got '{other}'"), - }); - } - }, - None => FusionStrategy::default(), + pub(crate) fn resolve(settings: &Settings) -> Result { + let defaults = crate::settings::SearchSettings::default(); + let ss = &settings.search; + + // Resolve fusion_strategy string via DB-first, then parse into enum. + let strategy_str = db_first_or_default( + &ss.fusion_strategy, + &defaults.fusion_strategy, + "SEARCH_FUSION_STRATEGY", + )?; + let fusion_strategy = match strategy_str.to_lowercase().as_str() { + "rrf" => FusionStrategy::Rrf, + "weighted" => FusionStrategy::WeightedScore, + other => { + return Err(ConfigError::InvalidValue { + key: "SEARCH_FUSION_STRATEGY".to_string(), + message: format!("must be 'rrf' or 'weighted', got '{other}'"), + }); + } }; - let rrf_k = parse_optional_env("SEARCH_RRF_K", 60u32)?; + let rrf_k = db_first_or_default(&ss.rrf_k, &defaults.rrf_k, "SEARCH_RRF_K")?; // Per-strategy weight defaults: RRF uses 0.5/0.5, weighted uses 0.3/0.7 (vector-biased). let (default_fts, default_vec) = match fusion_strategy { FusionStrategy::Rrf => (0.5f32, 0.5f32), FusionStrategy::WeightedScore => (0.3f32, 0.7f32), }; - let fts_weight = parse_optional_env("SEARCH_FTS_WEIGHT", default_fts)?; - let vector_weight = parse_optional_env("SEARCH_VECTOR_WEIGHT", default_vec)?; + + // Weights: DB (Some) > env > per-strategy default. + // Uses db_first_option for shadow warnings when DB overrides env. + let fts_weight = db_first_option(&ss.fts_weight, "SEARCH_FTS_WEIGHT")? + .unwrap_or(parse_optional_env("SEARCH_FTS_WEIGHT", default_fts)?); + let vector_weight = db_first_option(&ss.vector_weight, "SEARCH_VECTOR_WEIGHT")? + .unwrap_or(parse_optional_env("SEARCH_VECTOR_WEIGHT", default_vec)?); if !fts_weight.is_finite() || fts_weight < 0.0 { return Err(ConfigError::InvalidValue { @@ -109,7 +121,8 @@ mod tests { let _guard = lock_env(); clear_search_env(); - let config = WorkspaceSearchConfig::resolve().expect("should resolve"); + let settings = Settings::default(); + let config = WorkspaceSearchConfig::resolve(&settings).expect("should resolve"); assert_eq!(config.fusion_strategy, FusionStrategy::Rrf); assert_eq!(config.rrf_k, 60); assert!((config.fts_weight - 0.5).abs() < 0.001); @@ -117,7 +130,35 @@ mod tests { } #[test] - fn env_overrides() { + fn db_settings_override_env() { + let _guard = lock_env(); + clear_search_env(); + + // SAFETY: Under ENV_MUTEX. + unsafe { + std::env::set_var("SEARCH_FUSION_STRATEGY", "rrf"); + std::env::set_var("SEARCH_RRF_K", "30"); + std::env::set_var("SEARCH_FTS_WEIGHT", "0.9"); + std::env::set_var("SEARCH_VECTOR_WEIGHT", "0.1"); + } + + let mut settings = Settings::default(); + settings.search.fusion_strategy = "weighted".to_string(); + settings.search.rrf_k = 42; + settings.search.fts_weight = Some(0.4); + settings.search.vector_weight = Some(0.6); + + let config = WorkspaceSearchConfig::resolve(&settings).expect("should resolve"); + assert_eq!(config.fusion_strategy, FusionStrategy::WeightedScore); + assert_eq!(config.rrf_k, 42); + assert!((config.fts_weight - 0.4).abs() < 0.001); + assert!((config.vector_weight - 0.6).abs() < 0.001); + + clear_search_env(); + } + + #[test] + fn env_fallback_when_settings_at_default() { let _guard = lock_env(); clear_search_env(); @@ -129,7 +170,8 @@ mod tests { std::env::set_var("SEARCH_VECTOR_WEIGHT", "0.1"); } - let config = WorkspaceSearchConfig::resolve().expect("should resolve"); + let settings = Settings::default(); + let config = WorkspaceSearchConfig::resolve(&settings).expect("should resolve"); assert_eq!(config.fusion_strategy, FusionStrategy::WeightedScore); assert_eq!(config.rrf_k, 30); assert!((config.fts_weight - 0.9).abs() < 0.001); @@ -148,7 +190,8 @@ mod tests { std::env::set_var("SEARCH_FUSION_STRATEGY", "bm25"); } - let result = WorkspaceSearchConfig::resolve(); + let settings = Settings::default(); + let result = WorkspaceSearchConfig::resolve(&settings); assert!(result.is_err()); clear_search_env(); @@ -164,7 +207,8 @@ mod tests { std::env::set_var("SEARCH_FUSION_STRATEGY", "weighted"); } - let config = WorkspaceSearchConfig::resolve().expect("should resolve"); + let settings = Settings::default(); + let config = WorkspaceSearchConfig::resolve(&settings).expect("should resolve"); assert_eq!(config.fusion_strategy, FusionStrategy::WeightedScore); // Weighted mode should default to 0.3 FTS / 0.7 vector assert!((config.fts_weight - 0.3).abs() < 0.001); @@ -185,7 +229,8 @@ mod tests { std::env::set_var("SEARCH_VECTOR_WEIGHT", "0.0"); } - let result = WorkspaceSearchConfig::resolve(); + let settings = Settings::default(); + let result = WorkspaceSearchConfig::resolve(&settings); assert!(result.is_err()); clear_search_env(); @@ -203,7 +248,8 @@ mod tests { } // RRF ignores weights, so both=0 is fine - let config = WorkspaceSearchConfig::resolve().expect("should resolve"); + let settings = Settings::default(); + let config = WorkspaceSearchConfig::resolve(&settings).expect("should resolve"); assert_eq!(config.fusion_strategy, FusionStrategy::Rrf); clear_search_env(); diff --git a/src/config/skills.rs b/src/config/skills.rs index e893d2316e7..596655f0f5d 100644 --- a/src/config/skills.rs +++ b/src/config/skills.rs @@ -1,8 +1,11 @@ use std::path::PathBuf; use crate::bootstrap::ironclaw_base_dir; -use crate::config::helpers::{optional_env, parse_bool_env, parse_optional_env}; +use crate::config::helpers::{ + db_first_bool, db_first_or_default, optional_env, parse_optional_env, +}; use crate::error::ConfigError; +use crate::settings::Settings; /// Skills system configuration. #[derive(Debug, Clone)] @@ -48,17 +51,29 @@ fn default_installed_skills_dir() -> PathBuf { } impl SkillsConfig { - pub(crate) fn resolve() -> Result { + pub(crate) fn resolve(settings: &Settings) -> Result { + let defaults = crate::settings::SkillsSettings::default(); + let ss = &settings.skills; + Ok(Self { - enabled: parse_bool_env("SKILLS_ENABLED", true)?, + enabled: db_first_bool(ss.enabled, defaults.enabled, "SKILLS_ENABLED")?, + // local_dir and installed_dir are env-only (filesystem paths, no settings counterpart) local_dir: optional_env("SKILLS_DIR")? .map(PathBuf::from) .unwrap_or_else(default_skills_dir), installed_dir: optional_env("SKILLS_INSTALLED_DIR")? .map(PathBuf::from) .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_active_skills: db_first_or_default( + &ss.max_active_skills, + &defaults.max_active_skills, + "SKILLS_MAX_ACTIVE", + )?, + max_context_tokens: db_first_or_default( + &ss.max_context_tokens, + &defaults.max_context_tokens, + "SKILLS_MAX_CONTEXT_TOKENS", + )?, max_scan_depth: parse_optional_env("SKILLS_MAX_SCAN_DEPTH", 3)?, }) } diff --git a/src/config/transcription.rs b/src/config/transcription.rs index 191d2a02fdd..ed72832ab76 100644 --- a/src/config/transcription.rs +++ b/src/config/transcription.rs @@ -39,10 +39,11 @@ impl Default for TranscriptionConfig { impl TranscriptionConfig { pub(crate) fn resolve(settings: &Settings) -> Result { - let enabled = parse_bool_env( - "TRANSCRIPTION_ENABLED", - settings.transcription.as_ref().is_some_and(|t| t.enabled), - )?; + // Tri-state: Some(true/false) = explicit DB value, None = unset (fall back to env). + let enabled = match settings.transcription.as_ref().map(|t| t.enabled) { + Some(db_enabled) => db_enabled, + None => parse_bool_env("TRANSCRIPTION_ENABLED", false)?, + }; let provider = optional_env("TRANSCRIPTION_PROVIDER")?.unwrap_or_else(|| "openai".to_string()); diff --git a/src/config/tunnel.rs b/src/config/tunnel.rs index 1481a77356e..9bc9555147c 100644 --- a/src/config/tunnel.rs +++ b/src/config/tunnel.rs @@ -1,12 +1,14 @@ -use crate::config::helpers::optional_env; +use crate::config::helpers::{db_first_bool, db_first_optional_string, optional_env}; use crate::error::ConfigError; -use crate::settings::Settings; +use crate::settings::{Settings, TunnelSettings}; /// Tunnel configuration for exposing the agent to the internet. /// /// Used by channels and tools that need public webhook endpoints. /// The tunnel URL is shared across all channels (Telegram, Slack, etc.). /// +/// Resolution priority: DB/settings > env var > default. +/// /// Two modes: /// - **Static URL** (`TUNNEL_URL`): set the public URL directly (manual tunnel) /// - **Managed provider** (`TUNNEL_PROVIDER`): lifecycle-managed tunnel process @@ -25,8 +27,10 @@ pub struct TunnelConfig { impl TunnelConfig { pub(crate) fn resolve(settings: &Settings) -> Result { - let public_url = optional_env("TUNNEL_URL")? - .or_else(|| settings.tunnel.public_url.clone().filter(|s| !s.is_empty())); + let defaults = TunnelSettings::default(); + + // Priority: DB/settings > env > default. + let public_url = db_first_optional_string(&settings.tunnel.public_url, "TUNNEL_URL")?; if let Some(ref url) = public_url && !url.starts_with("https://") @@ -38,9 +42,8 @@ impl TunnelConfig { } // Resolve managed tunnel provider config. - // Priority: env var > settings > default (none). - let provider_name = optional_env("TUNNEL_PROVIDER")? - .or_else(|| settings.tunnel.provider.clone()) + // Priority: DB/settings > env > default (none). + let provider_name = db_first_optional_string(&settings.tunnel.provider, "TUNNEL_PROVIDER")? .unwrap_or_default(); let provider = if provider_name.is_empty() || provider_name == "none" { @@ -48,38 +51,64 @@ impl TunnelConfig { } else { Some(crate::tunnel::TunnelProviderConfig { provider: provider_name.clone(), - cloudflare: optional_env("TUNNEL_CF_TOKEN")? - .or_else(|| settings.tunnel.cf_token.clone()) - .map(|token| crate::tunnel::CloudflareTunnelConfig { token }), + // Security: tunnel auth tokens are env-only (sensitive credentials). + cloudflare: { + if settings.tunnel.cf_token.is_some() { + tracing::warn!( + "tunnel.cf_token is set in DB/TOML but is now env-only \ + (TUNNEL_CF_TOKEN). Remove it from DB/TOML settings." + ); + } + optional_env("TUNNEL_CF_TOKEN")? + .map(|token| crate::tunnel::CloudflareTunnelConfig { token }) + }, tailscale: Some(crate::tunnel::TailscaleTunnelConfig { - funnel: optional_env("TUNNEL_TS_FUNNEL")? - .map(|s| s == "true" || s == "1") - .unwrap_or(settings.tunnel.ts_funnel), - hostname: optional_env("TUNNEL_TS_HOSTNAME")? - .or_else(|| settings.tunnel.ts_hostname.clone()), + funnel: db_first_bool( + settings.tunnel.ts_funnel, + defaults.ts_funnel, + "TUNNEL_TS_FUNNEL", + )?, + hostname: db_first_optional_string( + &settings.tunnel.ts_hostname, + "TUNNEL_TS_HOSTNAME", + )?, }), ngrok: { - let ngrok_domain = optional_env("TUNNEL_NGROK_DOMAIN")? - .or_else(|| settings.tunnel.ngrok_domain.clone()); - optional_env("TUNNEL_NGROK_TOKEN")? - .or_else(|| settings.tunnel.ngrok_token.clone()) - .map(|auth_token| crate::tunnel::NgrokTunnelConfig { + let ngrok_domain = db_first_optional_string( + &settings.tunnel.ngrok_domain, + "TUNNEL_NGROK_DOMAIN", + )?; + if settings.tunnel.ngrok_token.is_some() { + tracing::warn!( + "tunnel.ngrok_token is set in DB/TOML but is now env-only \ + (TUNNEL_NGROK_TOKEN). Remove it from DB/TOML settings." + ); + } + optional_env("TUNNEL_NGROK_TOKEN")?.map(|auth_token| { + crate::tunnel::NgrokTunnelConfig { auth_token, domain: ngrok_domain, - }) + } + }) }, custom: { - let health_url = optional_env("TUNNEL_CUSTOM_HEALTH_URL")? - .or_else(|| settings.tunnel.custom_health_url.clone()); - let url_pattern = optional_env("TUNNEL_CUSTOM_URL_PATTERN")? - .or_else(|| settings.tunnel.custom_url_pattern.clone()); - optional_env("TUNNEL_CUSTOM_COMMAND")? - .or_else(|| settings.tunnel.custom_command.clone()) - .map(|start_command| crate::tunnel::CustomTunnelConfig { - start_command, - health_url, - url_pattern, - }) + let health_url = db_first_optional_string( + &settings.tunnel.custom_health_url, + "TUNNEL_CUSTOM_HEALTH_URL", + )?; + let url_pattern = db_first_optional_string( + &settings.tunnel.custom_url_pattern, + "TUNNEL_CUSTOM_URL_PATTERN", + )?; + db_first_optional_string( + &settings.tunnel.custom_command, + "TUNNEL_CUSTOM_COMMAND", + )? + .map(|start_command| crate::tunnel::CustomTunnelConfig { + start_command, + health_url, + url_pattern, + }) }, }) }; diff --git a/src/config/wasm.rs b/src/config/wasm.rs index 4c494a38e00..7db3160af46 100644 --- a/src/config/wasm.rs +++ b/src/config/wasm.rs @@ -2,7 +2,7 @@ use std::path::PathBuf; use std::time::Duration; use crate::bootstrap::ironclaw_base_dir; -use crate::config::helpers::{optional_env, parse_bool_env, parse_optional_env}; +use crate::config::helpers::{db_first_bool, db_first_or_default, optional_env}; use crate::error::ConfigError; /// WASM sandbox configuration. @@ -46,28 +46,41 @@ fn default_tools_dir() -> PathBuf { impl WasmConfig { pub(crate) fn resolve(settings: &crate::settings::Settings) -> Result { let ws = &settings.wasm; + let defaults = crate::settings::WasmSettings::default(); Ok(Self { - enabled: parse_bool_env("WASM_ENABLED", ws.enabled)?, - tools_dir: optional_env("WASM_TOOLS_DIR")? - .map(PathBuf::from) - .or_else(|| ws.tools_dir.clone()) - .unwrap_or_else(default_tools_dir), - default_memory_limit: parse_optional_env( + enabled: db_first_bool(ws.enabled, defaults.enabled, "WASM_ENABLED")?, + tools_dir: if let Some(ref dir) = ws.tools_dir { + dir.clone() + } else { + optional_env("WASM_TOOLS_DIR")? + .map(PathBuf::from) + .unwrap_or_else(default_tools_dir) + }, + default_memory_limit: db_first_or_default( + &ws.default_memory_limit, + &defaults.default_memory_limit, "WASM_DEFAULT_MEMORY_LIMIT", - ws.default_memory_limit, )?, - default_timeout_secs: parse_optional_env( + default_timeout_secs: db_first_or_default( + &ws.default_timeout_secs, + &defaults.default_timeout_secs, "WASM_DEFAULT_TIMEOUT_SECS", - ws.default_timeout_secs, )?, - default_fuel_limit: parse_optional_env( + default_fuel_limit: db_first_or_default( + &ws.default_fuel_limit, + &defaults.default_fuel_limit, "WASM_DEFAULT_FUEL_LIMIT", - ws.default_fuel_limit, )?, - cache_compiled: parse_bool_env("WASM_CACHE_COMPILED", ws.cache_compiled)?, - cache_dir: optional_env("WASM_CACHE_DIR")? - .map(PathBuf::from) - .or_else(|| ws.cache_dir.clone()), + cache_compiled: db_first_bool( + ws.cache_compiled, + defaults.cache_compiled, + "WASM_CACHE_COMPILED", + )?, + cache_dir: if let Some(ref dir) = ws.cache_dir { + Some(dir.clone()) + } else { + optional_env("WASM_CACHE_DIR")?.map(PathBuf::from) + }, }) } @@ -111,7 +124,7 @@ mod tests { } #[test] - fn env_overrides_settings() { + fn db_settings_override_env() { let _guard = lock_env(); let mut settings = Settings::default(); settings.wasm.default_fuel_limit = 42; @@ -121,6 +134,19 @@ mod tests { let cfg = WasmConfig::resolve(&settings).expect("resolve"); unsafe { std::env::remove_var("WASM_DEFAULT_FUEL_LIMIT") }; + assert_eq!(cfg.default_fuel_limit, 42); + } + + #[test] + fn env_used_when_no_db_setting() { + let _guard = lock_env(); + let settings = Settings::default(); + + // SAFETY: Under ENV_MUTEX, no concurrent env access. + unsafe { std::env::set_var("WASM_DEFAULT_FUEL_LIMIT", "7") }; + let cfg = WasmConfig::resolve(&settings).expect("resolve"); + unsafe { std::env::remove_var("WASM_DEFAULT_FUEL_LIMIT") }; + assert_eq!(cfg.default_fuel_limit, 7); } } diff --git a/src/db/libsql/conversations.rs b/src/db/libsql/conversations.rs index 4f9f1079309..00bc2ef23e1 100644 --- a/src/db/libsql/conversations.rs +++ b/src/db/libsql/conversations.rs @@ -67,19 +67,20 @@ impl ConversationStore for LibSqlBackend { channel: &str, user_id: &str, thread_id: Option<&str>, + source_channel: Option<&str>, ) -> Result { let conn = self.connect().await?; let now = fmt_ts(&Utc::now()); let affected = conn .execute( r#" - INSERT INTO conversations (id, channel, user_id, thread_id, started_at, last_activity) - VALUES (?1, ?2, ?3, ?4, ?5, ?5) + INSERT INTO conversations (id, channel, user_id, thread_id, source_channel, started_at, last_activity) + VALUES (?1, ?2, ?3, ?4, ?5, ?6, ?6) ON CONFLICT (id) DO UPDATE SET last_activity = excluded.last_activity WHERE conversations.user_id = excluded.user_id AND conversations.channel = excluded.channel "#, - params![id.to_string(), channel, user_id, opt_text(thread_id), now], + params![id.to_string(), channel, user_id, opt_text(thread_id), opt_text(source_channel), now], ) .await .map_err(|e| DatabaseError::Query(e.to_string()))?; @@ -600,6 +601,28 @@ impl ConversationStore for LibSqlBackend { .map_err(|e| DatabaseError::Query(e.to_string()))?; Ok(found.is_some()) } + + async fn get_conversation_source_channel( + &self, + conversation_id: Uuid, + ) -> Result, DatabaseError> { + let conn = self.connect().await?; + let mut rows = conn + .query( + "SELECT source_channel FROM conversations WHERE id = ?1", + params![conversation_id.to_string()], + ) + .await + .map_err(|e| DatabaseError::Query(e.to_string()))?; + match rows + .next() + .await + .map_err(|e| DatabaseError::Query(e.to_string()))? + { + Some(row) => Ok(get_opt_text(&row, 0)), + None => Ok(None), + } + } } #[cfg(test)] @@ -725,4 +748,111 @@ mod tests { "Expected same heartbeat conversation on repeated calls" ); } + + #[tokio::test] + async fn test_source_channel_round_trip() { + let dir = tempfile::tempdir().unwrap(); + let db_path = dir.path().join("test_source_channel.db"); + let backend = LibSqlBackend::new_local(&db_path).await.unwrap(); + backend.run_migrations().await.unwrap(); + + let conv_id = Uuid::new_v4(); + let user_id = "user-src-chan"; + + // Create conversation with a source_channel + let created = backend + .ensure_conversation(conv_id, "telegram", user_id, None, Some("telegram")) + .await + .unwrap(); + assert!(created, "first ensure should create"); + + // Read it back + let source = backend + .get_conversation_source_channel(conv_id) + .await + .unwrap(); + assert_eq!( + source.as_deref(), + Some("telegram"), + "source_channel should round-trip through DB" + ); + } + + #[tokio::test] + async fn test_source_channel_none_round_trip() { + let dir = tempfile::tempdir().unwrap(); + let db_path = dir.path().join("test_source_channel_none.db"); + let backend = LibSqlBackend::new_local(&db_path).await.unwrap(); + backend.run_migrations().await.unwrap(); + + let conv_id = Uuid::new_v4(); + let user_id = "user-no-src"; + + // Create conversation without source_channel + backend + .ensure_conversation(conv_id, "http", user_id, None, None) + .await + .unwrap(); + + let source = backend + .get_conversation_source_channel(conv_id) + .await + .unwrap(); + assert!( + source.is_none(), + "None source_channel should persist as NULL" + ); + } + + #[tokio::test] + async fn test_source_channel_not_found() { + let dir = tempfile::tempdir().unwrap(); + let db_path = dir.path().join("test_source_channel_404.db"); + let backend = LibSqlBackend::new_local(&db_path).await.unwrap(); + backend.run_migrations().await.unwrap(); + + let source = backend + .get_conversation_source_channel(Uuid::new_v4()) + .await + .unwrap(); + assert!( + source.is_none(), + "non-existent conversation should return None" + ); + } + + #[tokio::test] + async fn test_source_channel_not_overwritten_on_upsert() { + let dir = tempfile::tempdir().unwrap(); + let db_path = dir.path().join("test_source_channel_upsert.db"); + let backend = LibSqlBackend::new_local(&db_path).await.unwrap(); + backend.run_migrations().await.unwrap(); + + let conv_id = Uuid::new_v4(); + let user_id = "user-upsert"; + + // First insert with source_channel = "telegram" + backend + .ensure_conversation(conv_id, "telegram", user_id, None, Some("telegram")) + .await + .unwrap(); + + // Upsert same conversation (same user/channel) — source_channel should + // NOT be overwritten because the ON CONFLICT clause only updates + // last_activity. + backend + .ensure_conversation(conv_id, "telegram", user_id, None, Some("different")) + .await + .unwrap(); + + let source = backend + .get_conversation_source_channel(conv_id) + .await + .unwrap(); + assert_eq!( + source.as_deref(), + Some("telegram"), + "upsert should not overwrite original source_channel" + ); + } } diff --git a/src/db/libsql_migrations.rs b/src/db/libsql_migrations.rs index 2a4fa5c5181..6bdaac1703e 100644 --- a/src/db/libsql_migrations.rs +++ b/src/db/libsql_migrations.rs @@ -785,10 +785,48 @@ CREATE TABLE IF NOT EXISTS api_tokens ( ); CREATE INDEX IF NOT EXISTS idx_api_tokens_user ON api_tokens(user_id); CREATE INDEX IF NOT EXISTS idx_api_tokens_hash ON api_tokens(token_hash); +"#, + ), + ( + 16, + "conversation_source_channel", + // Add source_channel to conversations for cross-channel approval authorization. + // Marked as idempotent (see IDEMPOTENT_ADD_COLUMN_MIGRATIONS below) + // because SQLite does not support IF NOT EXISTS for ADD COLUMN. + // The runner checks pragma_table_info before executing the ALTER. + r#" +ALTER TABLE conversations ADD COLUMN source_channel TEXT; "#, ), ]; +/// Migrations whose ADD COLUMN should be skipped when the column already +/// exists (e.g. because the base SCHEMA was updated to include it). +/// Each entry is `(version, table_name, column_name)`. +const IDEMPOTENT_ADD_COLUMN_MIGRATIONS: &[(i64, &str, &str)] = + &[(16, "conversations", "source_channel")]; + +/// Check whether `table` already contains `column` via `pragma_table_info`. +async fn column_exists( + conn: &libsql::Connection, + table: &str, + column: &str, +) -> Result { + use crate::error::DatabaseError; + + let sql = format!( + "SELECT 1 FROM pragma_table_info('{}') WHERE name = ?1", + table + ); + let mut rows = conn + .query(&sql, libsql::params![column]) + .await + .map_err(|e| { + DatabaseError::Migration(format!("Failed to check column {table}.{column}: {e}")) + })?; + Ok(rows.next().await.ok().flatten().is_some()) +} + /// Run incremental migrations that haven't been applied yet. /// /// Each migration is wrapped in a transaction. On success the version is @@ -813,6 +851,18 @@ pub async fn run_incremental(conn: &libsql::Connection) -> Result<(), crate::err continue; // Already applied } + // For ADD COLUMN migrations, skip the ALTER if the column already + // exists (e.g. because the base SCHEMA was updated to include it) + // and just record the migration as applied. + let skip_sql = if let Some(&(_, table, column)) = IDEMPOTENT_ADD_COLUMN_MIGRATIONS + .iter() + .find(|(v, _, _)| *v == version) + { + column_exists(conn, table, column).await? + } else { + false + }; + // Wrap migration + recording in a transaction for atomicity. // If the process crashes mid-migration, the transaction rolls back // and the migration will be retried on next startup. @@ -822,9 +872,19 @@ pub async fn run_incremental(conn: &libsql::Connection) -> Result<(), crate::err )) })?; - tx.execute_batch(sql).await.map_err(|e| { - DatabaseError::Migration(format!("libSQL migration V{version} ({name}) failed: {e}")) - })?; + if skip_sql { + tracing::debug!( + version, + name, + "libSQL: column already exists, recording migration as applied" + ); + } else { + tx.execute_batch(sql).await.map_err(|e| { + DatabaseError::Migration(format!( + "libSQL migration V{version} ({name}) failed: {e}" + )) + })?; + } // Record as applied (inside the same transaction) tx.execute( diff --git a/src/db/mod.rs b/src/db/mod.rs index e2a81412c18..6edae6082e2 100644 --- a/src/db/mod.rs +++ b/src/db/mod.rs @@ -373,6 +373,7 @@ pub trait ConversationStore: Send + Sync { channel: &str, user_id: &str, thread_id: Option<&str>, + source_channel: Option<&str>, ) -> Result; async fn list_conversations_with_preview( &self, @@ -438,6 +439,11 @@ pub trait ConversationStore: Send + Sync { conversation_id: Uuid, user_id: &str, ) -> Result; + /// Get the source_channel for a conversation (the channel that created it). + async fn get_conversation_source_channel( + &self, + conversation_id: Uuid, + ) -> Result, DatabaseError>; } #[async_trait] diff --git a/src/db/postgres.rs b/src/db/postgres.rs index 462c3c46020..a00675eff9e 100644 --- a/src/db/postgres.rs +++ b/src/db/postgres.rs @@ -99,9 +99,10 @@ impl ConversationStore for PgBackend { channel: &str, user_id: &str, thread_id: Option<&str>, + source_channel: Option<&str>, ) -> Result { self.store - .ensure_conversation(id, channel, user_id, thread_id) + .ensure_conversation(id, channel, user_id, thread_id, source_channel) .await } @@ -222,6 +223,15 @@ impl ConversationStore for PgBackend { .conversation_belongs_to_user(conversation_id, user_id) .await } + + async fn get_conversation_source_channel( + &self, + conversation_id: Uuid, + ) -> Result, DatabaseError> { + self.store + .get_conversation_source_channel(conversation_id) + .await + } } // ==================== JobStore ==================== diff --git a/src/history/store.rs b/src/history/store.rs index 1a1fecbc58f..2ed35f9b7e6 100644 --- a/src/history/store.rs +++ b/src/history/store.rs @@ -1585,19 +1585,20 @@ impl Store { channel: &str, user_id: &str, thread_id: Option<&str>, + source_channel: Option<&str>, ) -> Result { let conn = self.conn().await?; let affected = conn .execute( r#" - INSERT INTO conversations (id, channel, user_id, thread_id) - VALUES ($1, $2, $3, $4) + INSERT INTO conversations (id, channel, user_id, thread_id, source_channel) + VALUES ($1, $2, $3, $4, $5) ON CONFLICT (id) DO UPDATE SET last_activity = NOW() WHERE conversations.user_id = EXCLUDED.user_id AND conversations.channel = EXCLUDED.channel "#, - &[&id, &channel, &user_id, &thread_id], + &[&id, &channel, &user_id, &thread_id, &source_channel], ) .await?; Ok(affected > 0) @@ -1917,6 +1918,21 @@ impl Store { Ok(row.is_some()) } + /// Get the source_channel for a conversation. + pub async fn get_conversation_source_channel( + &self, + conversation_id: Uuid, + ) -> Result, DatabaseError> { + let conn = self.conn().await?; + let row = conn + .query_opt( + "SELECT source_channel FROM conversations WHERE id = $1", + &[&conversation_id], + ) + .await?; + Ok(row.and_then(|r| r.get::<_, Option>(0))) + } + /// Load messages for a conversation with cursor-based pagination. /// /// Returns `(messages_oldest_first, has_more)`. diff --git a/src/llm/gemini_oauth.rs b/src/llm/gemini_oauth.rs index a12e5f82fba..9e035be99c4 100644 --- a/src/llm/gemini_oauth.rs +++ b/src/llm/gemini_oauth.rs @@ -1,3 +1,4 @@ +use std::collections::HashMap; use std::net::TcpListener; use std::path::{Path, PathBuf}; use std::time::Duration; @@ -921,8 +922,15 @@ pub struct GeminiOauthProvider { http_client: Client, /// Latest response metadata (updated after each request). last_response_meta: std::sync::Mutex, + /// Captured thought signatures keyed by tool-call ID. Gemini 3.x models + /// require these echoed back on `functionCall` parts when replaying history. + /// Populated from responses, consumed when building the next request. + thought_signatures: std::sync::Mutex>, } +/// Parsed Gemini response: (completion, tool_calls, thought_signatures_by_call_id). +type GeminiParsedResponse = (CompletionResponse, Vec, HashMap); + impl GeminiOauthProvider { pub fn new(config: GeminiOauthConfig) -> Result { let cred_manager = CredentialManager::new(&config.credentials_path)?; @@ -939,6 +947,7 @@ impl GeminiOauthProvider { cred_manager, http_client, last_response_meta: std::sync::Mutex::new(GeminiResponseMeta::default()), + thought_signatures: std::sync::Mutex::new(HashMap::new()), }) } @@ -1062,8 +1071,21 @@ impl GeminiOauthProvider { /// Count tokens for the given messages using the Gemini countTokens API. pub async fn count_tokens(&self, messages: &[ChatMessage]) -> Result { - let req = - Self::to_gemini_request(messages, None, None, None, None, None, &self.config.model); + let sigs = self + .thought_signatures + .lock() + .unwrap_or_else(|e| e.into_inner()) + .clone(); + let req = Self::to_gemini_request( + messages, + None, + None, + None, + None, + None, + &self.config.model, + &sigs, + ); let contents = req .get("contents") .cloned() @@ -1589,6 +1611,7 @@ impl GeminiOauthProvider { } } + #[allow(clippy::too_many_arguments)] fn to_gemini_request( messages: &[ChatMessage], tools: Option<&[ToolDefinition]>, @@ -1597,6 +1620,7 @@ impl GeminiOauthProvider { stop_sequences: Option<&[String]>, tool_choice: Option<&str>, model: &str, + thought_sigs: &HashMap, ) -> serde_json::Value { let mut contents = Vec::new(); @@ -1621,12 +1645,24 @@ impl GeminiOauthProvider { } if let Some(ref calls) = msg.tool_calls { for call in calls { - parts.push(serde_json::json!({ + let mut part = serde_json::json!({ "functionCall": { "name": call.name, "args": call.arguments } - })); + }); + // Echo back the real thoughtSignature if captured from a + // prior Gemini response. ensure_thought_signatures() will + // fill in synthetic placeholders for any gaps. + if let Some(sig) = thought_sigs.get(&call.id) + && let Some(obj) = part.as_object_mut() + { + obj.insert( + "thoughtSignature".to_string(), + serde_json::Value::String(sig.clone()), + ); + } + parts.push(part); } } // Fallback: if no parts at all, add empty text to avoid @@ -1855,9 +1891,8 @@ impl GeminiOauthProvider { req } - fn from_gemini_response( - body: serde_json::Value, - ) -> Result<(CompletionResponse, Vec), LlmError> { + /// Parsed Gemini response: (completion, tool_calls, thought_signatures_by_call_id). + fn from_gemini_response(body: serde_json::Value) -> Result { let candidate = body .get("candidates") .and_then(|c| c.as_array()) @@ -1874,6 +1909,7 @@ impl GeminiOauthProvider { let mut text_content = String::new(); let mut tool_calls = Vec::new(); + let mut thought_sigs = HashMap::new(); if let Some(parts) = parts { for part in parts { @@ -1892,6 +1928,11 @@ impl GeminiOauthProvider { .and_then(|i| i.as_str()) .map(|s| s.to_string()) .unwrap_or_else(|| uuid::Uuid::new_v4().to_string()); + // Capture thoughtSignature (sibling of functionCall in the part) + // so it can be echoed back when replaying history. + if let Some(sig) = part.get("thoughtSignature").and_then(|s| s.as_str()) { + thought_sigs.insert(id.clone(), sig.to_string()); + } tool_calls.push(ToolCall { id, @@ -2003,6 +2044,7 @@ impl GeminiOauthProvider { cache_creation_input_tokens: 0, }, tool_calls, + thought_sigs, )) } } @@ -2041,6 +2083,11 @@ impl LlmProvider for GeminiOauthProvider { } async fn complete(&self, request: CompletionRequest) -> Result { + let sigs = self + .thought_signatures + .lock() + .unwrap_or_else(|e| e.into_inner()) + .clone(); let req_json = Self::to_gemini_request( &request.messages, None, @@ -2049,9 +2096,10 @@ impl LlmProvider for GeminiOauthProvider { request.stop_sequences.as_deref(), None, &self.config.model, + &sigs, ); let resp_json = self.send_request(&req_json).await?; - let (response, _tool_calls) = Self::from_gemini_response(resp_json)?; + let (response, _tool_calls, _new_sigs) = Self::from_gemini_response(resp_json)?; Ok(response) } @@ -2065,6 +2113,11 @@ impl LlmProvider for GeminiOauthProvider { Some(request.tools.as_slice()) }; + let sigs = self + .thought_signatures + .lock() + .unwrap_or_else(|e| e.into_inner()) + .clone(); let req_json = Self::to_gemini_request( &request.messages, tool_defs, @@ -2073,9 +2126,29 @@ impl LlmProvider for GeminiOauthProvider { request.stop_sequences.as_deref(), request.tool_choice.as_deref(), &self.config.model, + &sigs, ); let resp_json = self.send_request(&req_json).await?; - let (response, tool_calls) = Self::from_gemini_response(resp_json)?; + let (response, tool_calls, new_sigs) = Self::from_gemini_response(resp_json)?; + // Store captured thought signatures, pruning stale entries to prevent + // unbounded growth over long-running processes. Only keep IDs that + // appear in the conversation history or the just-received response. + { + let mut sigs = self + .thought_signatures + .lock() + .unwrap_or_else(|e| e.into_inner()); + sigs.extend(new_sigs); + let live_ids: std::collections::HashSet<&str> = request + .messages + .iter() + .filter_map(|m| m.tool_calls.as_ref()) + .flatten() + .map(|tc| tc.id.as_str()) + .chain(tool_calls.iter().map(|tc| tc.id.as_str())) + .collect(); + sigs.retain(|id, _| live_ids.contains(id.as_str())); + } Ok(crate::llm::provider::ToolCompletionResponse { content: if response.content.is_empty() { @@ -2218,6 +2291,7 @@ mod tests { None, None, "gemini-2.0-flash", + &HashMap::new(), ); let decls = &req["tools"][0]["functionDeclarations"]; @@ -2240,6 +2314,7 @@ mod tests { None, None, "gemini-2.0-flash", + &HashMap::new(), ); let contents = req["contents"].as_array().unwrap(); @@ -2265,7 +2340,7 @@ mod tests { } }); - let (resp, tool_calls) = GeminiOauthProvider::from_gemini_response(body).unwrap(); + let (resp, tool_calls, _sigs) = GeminiOauthProvider::from_gemini_response(body).unwrap(); assert_eq!(resp.content, "Hello world"); assert_eq!(resp.input_tokens, 10); @@ -2293,7 +2368,7 @@ mod tests { } }); - let (resp, tool_calls) = GeminiOauthProvider::from_gemini_response(body).unwrap(); + let (resp, tool_calls, _sigs) = GeminiOauthProvider::from_gemini_response(body).unwrap(); assert!(resp.content.is_empty()); assert_eq!(tool_calls.len(), 1); @@ -2313,6 +2388,7 @@ mod tests { None, None, "gemini-2.0-flash", + &HashMap::new(), ); let gen_cfg = &req["generationConfig"]; @@ -2333,6 +2409,7 @@ mod tests { None, None, "gemini-3-flash-preview", + &HashMap::new(), ); let thinking = &req["generationConfig"]["thinkingConfig"]; @@ -2353,6 +2430,7 @@ mod tests { None, None, "gemini-2.5-flash-thinking", + &HashMap::new(), ); let thinking = &req["generationConfig"]["thinkingConfig"]; @@ -2376,6 +2454,7 @@ mod tests { Some(&stops), None, "gemini-2.5-flash", + &HashMap::new(), ); let gen_cfg = &req["generationConfig"]; @@ -2403,6 +2482,7 @@ mod tests { None, Some("auto"), "gemini-2.0-flash", + &HashMap::new(), ); assert_eq!( req_auto["toolConfig"]["functionCallingConfig"]["mode"], @@ -2417,6 +2497,7 @@ mod tests { None, Some("required"), "gemini-2.0-flash", + &HashMap::new(), ); assert_eq!( req_req["toolConfig"]["functionCallingConfig"]["mode"], @@ -2431,6 +2512,7 @@ mod tests { None, Some("none"), "gemini-2.0-flash", + &HashMap::new(), ); assert_eq!( req_none["toolConfig"]["functionCallingConfig"]["mode"], @@ -2500,6 +2582,7 @@ mod tests { None, None, "gemini-1.5-flash", + &HashMap::new(), ); let system_instruction = req @@ -2614,4 +2697,133 @@ mod tests { assert_eq!(signed_calls, 2); // safety: test-only assertion } + + #[test] + fn test_from_gemini_response_captures_thought_signature() { + let body = serde_json::json!({ + "candidates": [{ + "content": { + "parts": [{ + "thoughtSignature": "abc123sig", + "functionCall": { + "name": "read_file", + "args": { "path": "/tmp/test.txt" } + } + }] + }, + "finishReason": "STOP" + }], + "usageMetadata": { + "promptTokenCount": 10, + "candidatesTokenCount": 5 + } + }); + + let (_resp, tool_calls, sigs) = GeminiOauthProvider::from_gemini_response(body).unwrap(); + assert_eq!(tool_calls.len(), 1); + assert_eq!( + sigs.get(&tool_calls[0].id).map(|s| s.as_str()), + Some("abc123sig") + ); + } + + #[test] + fn test_from_gemini_response_no_thought_signature_yields_none() { + let body = serde_json::json!({ + "candidates": [{ + "content": { + "parts": [{ + "functionCall": { + "name": "echo", + "args": {} + } + }] + }, + "finishReason": "STOP" + }], + "usageMetadata": { + "promptTokenCount": 5, + "candidatesTokenCount": 3 + } + }); + + let (_resp, tool_calls, sigs) = GeminiOauthProvider::from_gemini_response(body).unwrap(); + assert_eq!(tool_calls.len(), 1); + assert!(!sigs.contains_key(&tool_calls[0].id)); + } + + #[test] + fn test_to_gemini_request_echoes_thought_signature_on_function_call() { + let messages = vec![ + ChatMessage::user("call a tool"), + ChatMessage::assistant_with_tool_calls( + None, + vec![ToolCall { + id: "call_1".to_string(), + name: "read_file".to_string(), + arguments: serde_json::json!({"path": "/tmp/x"}), + reasoning: None, + }], + ), + ChatMessage::tool_result("call_1", "read_file", r#"{"output":"hello"}"#), + ]; + + let mut sigs = HashMap::new(); + sigs.insert("call_1".to_string(), "sig_from_gemini".to_string()); + + let req = GeminiOauthProvider::to_gemini_request( + &messages, + None, + None, + None, + None, + None, + "gemini-3-flash-preview", + &sigs, + ); + + let contents = req["contents"].as_array().unwrap(); + // The model turn (index 1) should have the thoughtSignature on its functionCall part. + let model_turn = &contents[1]; + assert_eq!(model_turn["role"], "model"); + let fc_part = &model_turn["parts"][0]; + assert!(fc_part.get("functionCall").is_some()); + assert_eq!(fc_part["thoughtSignature"], "sig_from_gemini"); + } + + #[test] + fn test_to_gemini_request_omits_thought_signature_when_none() { + let messages = vec![ + ChatMessage::user("call a tool"), + ChatMessage::assistant_with_tool_calls( + None, + vec![ToolCall { + id: "call_1".to_string(), + name: "echo".to_string(), + arguments: serde_json::json!({}), + reasoning: None, + }], + ), + ChatMessage::tool_result("call_1", "echo", r#"{"output":"ok"}"#), + ]; + + let empty_sigs = HashMap::new(); + let req = GeminiOauthProvider::to_gemini_request( + &messages, + None, + None, + None, + None, + None, + "gemini-2.0-flash", + &empty_sigs, + ); + + let contents = req["contents"].as_array().unwrap(); + let model_turn = &contents[1]; + let fc_part = &model_turn["parts"][0]; + assert!(fc_part.get("functionCall").is_some()); + // No thoughtSignature should be present when there is no captured signature for this call ID. + assert!(fc_part.get("thoughtSignature").is_none()); + } } diff --git a/src/main.rs b/src/main.rs index 22dfdcb0bf1..6f4a37c570e 100644 --- a/src/main.rs +++ b/src/main.rs @@ -449,6 +449,7 @@ async fn async_main() -> anyhow::Result<()> { &components.secrets_store, components.extension_manager.as_ref(), components.db.as_ref(), + &channel_names, ) .await; diff --git a/src/orchestrator/job_manager.rs b/src/orchestrator/job_manager.rs index 34b9f37337a..222d24650aa 100644 --- a/src/orchestrator/job_manager.rs +++ b/src/orchestrator/job_manager.rs @@ -16,6 +16,11 @@ use crate::error::OrchestratorError; use crate::orchestrator::auth::{CredentialGrant, TokenStore}; use crate::sandbox::connect_docker; +/// Path to the master worker MCP config on the host. +const WORKER_MCP_CONFIG_PATH: &str = "/opt/ironclaw/config/worker/mcp-servers.json"; + +use ironclaw_common::MAX_WORKER_ITERATIONS; + /// Which mode a sandbox container runs in. #[derive(Debug, Clone, Copy, PartialEq, Eq)] pub enum JobMode { @@ -40,6 +45,19 @@ impl std::fmt::Display for JobMode { } } +/// Parameters for creating a container job, bundled to avoid positional +/// argument proliferation on `create_job` / `execute_sandbox`. +#[derive(Debug, Clone, Default)] +pub struct JobCreationParams { + /// Credential grants for the worker (served via `/credentials`). + pub credential_grants: Vec, + /// Optional filter: which MCP servers to mount into the container. + /// `None` = full master config, `Some([])` = no MCP, `Some(["name"])` = filtered. + pub mcp_servers: Option>, + /// Optional cap on worker agent loop iterations (clamped to 1..=500 server-side). + pub max_iterations: Option, +} + /// Configuration for the container job manager. #[derive(Debug, Clone)] pub struct ContainerJobConfig { @@ -66,6 +84,9 @@ pub struct ContainerJobConfig { pub claude_code_memory_limit_mb: u64, /// Allowed tool patterns for Claude Code (passed as CLAUDE_CODE_ALLOWED_TOOLS env var). pub claude_code_allowed_tools: Vec, + /// Whether per-job MCP server filtering is enabled. + /// When false, `mcp_servers` param on `create_job` is ignored. + pub mcp_per_job_enabled: bool, } impl Default for ContainerJobConfig { @@ -81,6 +102,7 @@ impl Default for ContainerJobConfig { claude_code_max_turns: 50, claude_code_memory_limit_mb: 4096, claude_code_allowed_tools: crate::config::ClaudeCodeConfig::default().allowed_tools, + mcp_per_job_enabled: false, } } } @@ -254,14 +276,14 @@ impl ContainerJobManager { task: &str, project_dir: Option, mode: JobMode, - credential_grants: Vec, + params: JobCreationParams, ) -> Result { // Generate auth token (stored in TokenStore, never logged) let token = self.token_store.create_token(job_id).await; // Store credential grants (revoked automatically when the token is revoked) self.token_store - .store_grants(job_id, credential_grants) + .store_grants(job_id, params.credential_grants) .await; // Record the handle @@ -282,7 +304,14 @@ impl ContainerJobManager { // Run the actual container creation. On any failure, revoke the token // and remove the handle so we don't leak resources. match self - .create_job_inner(job_id, &token, project_dir, mode) + .create_job_inner( + job_id, + &token, + project_dir, + mode, + params.mcp_servers, + params.max_iterations, + ) .await { Ok(()) => Ok(token), @@ -301,6 +330,8 @@ impl ContainerJobManager { token: &str, project_dir: Option, mode: JobMode, + mcp_servers: Option>, + max_iterations: Option, ) -> Result<(), OrchestratorError> { // Connect to Docker (reuses cached connection) let docker = self.docker().await?; @@ -331,6 +362,42 @@ impl ContainerJobManager { env_vec.push("IRONCLAW_WORKSPACE=/workspace".to_string()); } + // Inject max_iterations if specified (only for Worker mode — ClaudeCode uses max_turns). + // Server-side clamp ensures the cap is enforced even if the tool parsing + // layer is bypassed (e.g., direct API call via the web restart handler). + if let Some(iters) = max_iterations + && mode == JobMode::Worker + { + let capped = iters.clamp(1, MAX_WORKER_ITERATIONS); + env_vec.push(format!("IRONCLAW_MAX_ITERATIONS={}", capped)); + } + + // Mount per-job MCP config when the feature is enabled. + if self.config.mcp_per_job_enabled { + let mcp_config_host = std::path::Path::new(WORKER_MCP_CONFIG_PATH); + match generate_worker_mcp_config(mcp_config_host, mcp_servers.as_deref(), job_id) + .await? + { + Some(config_path) => { + binds.push(format!( + "{}:/home/sandbox/.ironclaw/mcp-servers.json:ro", + config_path.display() + )); + tracing::debug!( + job_id = %job_id, + filtered = mcp_servers.is_some(), + "Mounted MCP config into container" + ); + } + None => { + tracing::debug!( + job_id = %job_id, + "No MCP config to mount (master missing or empty filter list)" + ); + } + } + } + // Claude Code mode: auth + tool allowlist. // // Auth strategies (first match wins): @@ -579,6 +646,23 @@ impl ContainerJobManager { /// Remove a completed job handle from memory (called after result is read). pub async fn cleanup_job(&self, job_id: Uuid) { + // Clean up per-job MCP config temp file if one was written. + // Use remove_file directly — avoids TOCTOU race with exists() check. + let tmp_path = std::env::temp_dir() + .join("ironclaw-mcp-configs") + .join(format!("{}.json", job_id)); + match std::fs::remove_file(&tmp_path) { + Ok(()) => {} + Err(e) if e.kind() == std::io::ErrorKind::NotFound => {} // No temp file — normal + Err(e) => { + tracing::warn!( + job_id = %job_id, + error = %e, + "Failed to remove per-job MCP config temp file" + ); + } + } + self.containers.write().await.remove(&job_id); } @@ -611,6 +695,132 @@ impl ContainerJobManager { } } +/// Generate a per-job MCP config file, optionally filtering to specific servers. +/// +/// - `None` → mount the full master config as-is +/// - `Some([])` → no MCP config (no mount) +/// - `Some(["serpstat"])` → filtered config with only matching servers +/// +/// Temp files are written to `/ironclaw-mcp-configs/` and cleaned up +/// in `cleanup_job`. +async fn generate_worker_mcp_config( + master_path: &std::path::Path, + server_names: Option<&[String]>, + job_id: Uuid, +) -> Result, OrchestratorError> { + if !tokio::fs::try_exists(master_path).await.unwrap_or(false) { + return Ok(None); + } + + match server_names { + // No filter → use master config as-is + None => Ok(Some(master_path.to_path_buf())), + + // Empty list → no MCP + Some([]) => Ok(None), + + // Filter to specific servers + Some(names) => { + // Validate server names: reject path separators, null bytes, and + // excessively long names to prevent misuse if names are ever used + // in file paths or shell commands. + for name in names { + if name.len() > 128 + || name.contains('/') + || name.contains('\\') + || name.contains('\0') + { + return Err(OrchestratorError::ContainerCreationFailed { + job_id, + reason: format!("invalid MCP server name: {:?}", name), + }); + } + } + + let content = tokio::fs::read_to_string(master_path).await.map_err(|e| { + OrchestratorError::ContainerCreationFailed { + job_id, + reason: format!("failed to read master MCP config: {e}"), + } + })?; + + let master: serde_json::Value = serde_json::from_str(&content).map_err(|e| { + OrchestratorError::ContainerCreationFailed { + job_id, + reason: format!("failed to parse master MCP config: {e}"), + } + })?; + + let servers = master["servers"] + .as_array() + .cloned() + .unwrap_or_default() + .into_iter() + .filter(|s| { + let name_matches = s["name"] + .as_str() + .map(|n| names.iter().any(|req| req.eq_ignore_ascii_case(n))) + .unwrap_or(false); + let is_enabled = s["enabled"].as_bool().unwrap_or(true); + name_matches && is_enabled + }) + .collect::>(); + + if servers.is_empty() { + tracing::warn!( + job_id = %job_id, + requested = ?names, + "No matching MCP servers found in master config; skipping MCP mount" + ); + return Ok(None); + } + + let schema_version = master + .get("schema_version") + .cloned() + .unwrap_or(serde_json::json!(1)); + let filtered = serde_json::json!({ + "servers": servers, + "schema_version": schema_version + }); + + let tmp_dir = std::env::temp_dir().join("ironclaw-mcp-configs"); + tokio::fs::create_dir_all(&tmp_dir).await.map_err(|e| { + OrchestratorError::ContainerCreationFailed { + job_id, + reason: format!("failed to create MCP config temp dir: {e}"), + } + })?; + + // Restrict directory permissions to owner-only (0o700) to prevent + // other users on the host from reading filtered MCP configs. + #[cfg(unix)] + { + use std::os::unix::fs::PermissionsExt; + let _ = + tokio::fs::set_permissions(&tmp_dir, std::fs::Permissions::from_mode(0o700)) + .await; + } + + let tmp_path = tmp_dir.join(format!("{}.json", job_id)); + let config_json = serde_json::to_string_pretty(&filtered).map_err(|e| { + OrchestratorError::ContainerCreationFailed { + job_id, + reason: format!("failed to serialize filtered MCP config: {e}"), + } + })?; + tokio::fs::write(&tmp_path, config_json) + .await + .map_err(|e| OrchestratorError::ContainerCreationFailed { + job_id, + reason: format!("failed to write per-job MCP config: {e}"), + })?; + + Ok(Some(tmp_path)) + } + } +} + #[cfg(test)] mod tests { use super::*; @@ -705,4 +915,308 @@ mod tests { assert_eq!(handle.worker_iteration, 3); assert_eq!(handle.last_worker_status.as_deref(), Some("Iteration 3")); } + + // ── generate_worker_mcp_config tests ──────────────────────────── + + #[tokio::test] + async fn test_mcp_config_none_filter_returns_master_path() { + let tmp = tempfile::NamedTempFile::new().unwrap(); + std::fs::write(tmp.path(), r#"{"servers":[]}"#).unwrap(); + let result = generate_worker_mcp_config(tmp.path(), None, Uuid::new_v4()).await; + assert_eq!(result.unwrap(), Some(tmp.path().to_path_buf())); + } + + #[tokio::test] + async fn test_mcp_config_empty_filter_returns_none() { + let tmp = tempfile::NamedTempFile::new().unwrap(); + std::fs::write(tmp.path(), r#"{"servers":[]}"#).unwrap(); + let result = generate_worker_mcp_config(tmp.path(), Some(&[]), Uuid::new_v4()).await; + assert_eq!(result.unwrap(), None); + } + + #[tokio::test] + async fn test_mcp_config_missing_master_returns_none() { + let result = generate_worker_mcp_config( + std::path::Path::new("/nonexistent/mcp.json"), + None, + Uuid::new_v4(), + ) + .await; + assert_eq!(result.unwrap(), None); + } + + #[tokio::test] + async fn test_mcp_config_filters_to_named_servers() { + let job_id = Uuid::new_v4(); + let tmp = tempfile::NamedTempFile::new().unwrap(); + std::fs::write( + tmp.path(), + r#"{"schema_version":1,"servers":[ + {"name":"serpstat","enabled":true,"url":"http://localhost:8062"}, + {"name":"notion","enabled":true,"url":"http://localhost:8063"}, + {"name":"disabled","enabled":false,"url":"http://localhost:9999"} + ]}"#, + ) + .unwrap(); + + let names = vec!["serpstat".to_string(), "disabled".to_string()]; + let result = generate_worker_mcp_config(tmp.path(), Some(&names), job_id).await; + let out_path = result.unwrap().expect("should produce a filtered config"); + + let content: serde_json::Value = + serde_json::from_str(&std::fs::read_to_string(&out_path).unwrap()).unwrap(); + let servers = content["servers"].as_array().unwrap(); + + // "disabled" should be excluded because enabled=false + assert_eq!(servers.len(), 1); + assert_eq!(servers[0]["name"], "serpstat"); + assert_eq!(content["schema_version"], 1); + + // cleanup + let _ = std::fs::remove_file(&out_path); + } + + #[tokio::test] + async fn test_mcp_config_no_match_returns_none() { + let tmp = tempfile::NamedTempFile::new().unwrap(); + std::fs::write( + tmp.path(), + r#"{"servers":[{"name":"serpstat","enabled":true}]}"#, + ) + .unwrap(); + + let names = vec!["nonexistent".to_string()]; + let result = generate_worker_mcp_config(tmp.path(), Some(&names), Uuid::new_v4()).await; + assert_eq!(result.unwrap(), None); + } + + #[tokio::test] + async fn test_mcp_config_case_insensitive_match() { + let job_id = Uuid::new_v4(); + let tmp = tempfile::NamedTempFile::new().unwrap(); + std::fs::write( + tmp.path(), + r#"{"servers":[{"name":"Serpstat","enabled":true}]}"#, + ) + .unwrap(); + + let names = vec!["serpstat".to_string()]; + let result = generate_worker_mcp_config(tmp.path(), Some(&names), job_id).await; + let out_path = result.unwrap().expect("case-insensitive match should work"); + let _ = std::fs::remove_file(&out_path); + } + + #[test] + fn test_max_iterations_env_var_injected() { + // Verify the IRONCLAW_MAX_ITERATIONS env var name matches what the + // worker CLI reads via clap's `env` attribute. + let config = ContainerJobConfig::default(); + let mgr = ContainerJobManager::new(config, TokenStore::new()); + // We can't test actual container creation without Docker, but we can + // verify the env var name matches the clap definition. + // The clap definition uses: #[arg(long, env = "IRONCLAW_MAX_ITERATIONS")] + // The create_job_inner injects: format!("IRONCLAW_MAX_ITERATIONS={}", iters) + // This test ensures the constant isn't accidentally changed in either place. + let env_var_in_job = "IRONCLAW_MAX_ITERATIONS"; + let source = include_str!("../cli/mod.rs"); + assert!( + source.contains(&format!("env = \"{}\"", env_var_in_job)), + "cli/mod.rs must have env = \"IRONCLAW_MAX_ITERATIONS\" on the max_iterations arg" + ); + drop(mgr); + } + + #[test] + fn test_max_iterations_not_injected_for_claude_code() { + // ClaudeCode mode uses its own `max_turns`, not IRONCLAW_MAX_ITERATIONS. + // Verify the gate in create_job_inner only injects for Worker mode. + let source = include_str!("job_manager.rs"); + assert!( + source.contains("mode == JobMode::Worker"), + "create_job_inner must gate IRONCLAW_MAX_ITERATIONS on JobMode::Worker \ + (ClaudeCode has its own max_turns)" + ); + } + + #[test] + fn test_server_side_max_iterations_clamp() { + // Verify the server-side clamp uses the same constant as worker/job.rs + let source = include_str!("job_manager.rs"); + assert!( + source.contains("iters.clamp(1, MAX_WORKER_ITERATIONS)"), + "create_job_inner must clamp max_iterations server-side using MAX_WORKER_ITERATIONS" + ); + } + + #[tokio::test] + async fn test_mcp_server_name_validation_rejects_path_separators() { + let job_id = Uuid::new_v4(); + let tmp = tempfile::NamedTempFile::new().unwrap(); + std::fs::write( + tmp.path(), + r#"{"servers":[{"name":"test","enabled":true}]}"#, + ) + .unwrap(); + + // Path separator should be rejected + let names = vec!["../../etc/passwd".to_string()]; + assert!( + generate_worker_mcp_config(tmp.path(), Some(&names), job_id) + .await + .is_err() + ); + + // Null byte should be rejected + let names = vec!["test\0evil".to_string()]; + assert!( + generate_worker_mcp_config(tmp.path(), Some(&names), job_id) + .await + .is_err() + ); + + // Excessively long name should be rejected + let names = vec!["a".repeat(129)]; + assert!( + generate_worker_mcp_config(tmp.path(), Some(&names), job_id) + .await + .is_err() + ); + + // Valid name should pass + let names = vec!["test".to_string()]; + let result = generate_worker_mcp_config(tmp.path(), Some(&names), job_id).await; + assert!(result.is_ok()); + } + + // ── Regression tests (CI-required) ──────────────────────────────── + + #[tokio::test] + async fn test_filtered_config_contains_only_requested_server() { + let job_id = Uuid::new_v4(); + let tmp = tempfile::NamedTempFile::new().unwrap(); + std::fs::write( + tmp.path(), + r#"{"schema_version":2,"servers":[ + {"name":"serpstat","enabled":true,"url":"http://localhost:8062"}, + {"name":"notion","enabled":true,"url":"http://localhost:8063"}, + {"name":"archon","enabled":true,"url":"http://localhost:8064"} + ]}"#, + ) + .unwrap(); + + let names = vec!["serpstat".to_string()]; + let result = generate_worker_mcp_config(tmp.path(), Some(&names), job_id).await; + let out_path = result.unwrap().expect("should produce a filtered config"); + + let content: serde_json::Value = + serde_json::from_str(&std::fs::read_to_string(&out_path).unwrap()).unwrap(); + let servers = content["servers"].as_array().unwrap(); + + assert_eq!(servers.len(), 1, "only serpstat should be present"); + assert_eq!(servers[0]["name"], "serpstat"); + assert!( + !servers.iter().any(|s| s["name"] == "notion"), + "notion must not leak into filtered config" + ); + assert!( + !servers.iter().any(|s| s["name"] == "archon"), + "archon must not leak into filtered config" + ); + assert_eq!( + content["schema_version"], 2, + "schema_version must be preserved" + ); + + let _ = std::fs::remove_file(&out_path); + } + + #[tokio::test] + async fn test_feature_flag_disabled_skips_mcp_filtering() { + // When MCP_PER_JOB_ENABLED is false (the default), the mcp_servers + // parameter should be ignored and no filtered config should be created. + let config = ContainerJobConfig::default(); + assert!( + !config.mcp_per_job_enabled, + "mcp_per_job_enabled must default to false" + ); + + // Verify the gate in create_job_inner: the mcp_per_job_enabled field + // controls whether generate_worker_mcp_config is called at all. + let source = include_str!("job_manager.rs"); + assert!( + source.contains("if self.config.mcp_per_job_enabled"), + "create_job_inner must gate MCP filtering on config.mcp_per_job_enabled" + ); + } + + #[tokio::test] + async fn test_temp_file_cleanup_removes_per_job_config() { + let job_id = Uuid::new_v4(); + let tmp = tempfile::NamedTempFile::new().unwrap(); + std::fs::write( + tmp.path(), + r#"{"servers":[{"name":"serpstat","enabled":true}]}"#, + ) + .unwrap(); + + let names = vec!["serpstat".to_string()]; + let result = generate_worker_mcp_config(tmp.path(), Some(&names), job_id).await; + let out_path = result.unwrap().expect("should produce a filtered config"); + assert!( + out_path.exists(), + "temp config file should exist after creation" + ); + + // Simulate what cleanup_job does + let expected_path = std::env::temp_dir() + .join("ironclaw-mcp-configs") + .join(format!("{}.json", job_id)); + assert_eq!( + out_path, expected_path, + "temp path must match cleanup expectation" + ); + std::fs::remove_file(&out_path).unwrap(); + assert!(!out_path.exists(), "temp file should be gone after cleanup"); + } + + #[tokio::test] + async fn test_cleanup_job_is_idempotent() { + let config = ContainerJobConfig::default(); + let mgr = ContainerJobManager::new(config, TokenStore::new()); + let job_id = Uuid::new_v4(); + + // cleanup_job should not panic or error when called for a job + // that has no temp file and no container handle. + mgr.cleanup_job(job_id).await; + // Second call should also be fine (idempotent). + mgr.cleanup_job(job_id).await; + } + + #[cfg(unix)] + #[tokio::test] + async fn test_temp_dir_has_restrictive_permissions() { + use std::os::unix::fs::PermissionsExt; + + let job_id = Uuid::new_v4(); + let tmp = tempfile::NamedTempFile::new().unwrap(); + std::fs::write( + tmp.path(), + r#"{"servers":[{"name":"test","enabled":true}]}"#, + ) + .unwrap(); + + let names = vec!["test".to_string()]; + let result = generate_worker_mcp_config(tmp.path(), Some(&names), job_id).await; + let out_path = result.unwrap().expect("should produce a filtered config"); + + let dir_path = out_path.parent().unwrap(); + let mode = std::fs::metadata(dir_path).unwrap().permissions().mode() & 0o777; + assert_eq!( + mode, 0o700, + "ironclaw-mcp-configs dir must be 0700, got {:o}", + mode + ); + + let _ = std::fs::remove_file(&out_path); + } } diff --git a/src/orchestrator/mod.rs b/src/orchestrator/mod.rs index 8d09dc53bf2..d6ecb61c2f1 100644 --- a/src/orchestrator/mod.rs +++ b/src/orchestrator/mod.rs @@ -122,6 +122,9 @@ pub async fn setup_orchestrator( claude_code_max_turns: config.claude_code.max_turns, claude_code_memory_limit_mb: config.claude_code.memory_limit_mb, claude_code_allowed_tools: config.claude_code.allowed_tools.clone(), + mcp_per_job_enabled: std::env::var("MCP_PER_JOB_ENABLED") + .map(|v| v.eq_ignore_ascii_case("true") || v == "1") + .unwrap_or(false), }; let jm = Arc::new(ContainerJobManager::new(job_config, token_store.clone())); diff --git a/src/sandbox/container.rs b/src/sandbox/container.rs index a5ef12ab5c6..f5ef014702c 100644 --- a/src/sandbox/container.rs +++ b/src/sandbox/container.rs @@ -98,6 +98,18 @@ impl ContainerRunner { } } + /// Create a runner for image-only operations (exists / pull / build). + /// + /// `proxy_port` is unused for image operations, so this avoids requiring + /// a real port number when the caller only needs to check or build images. + pub fn for_image_ops(docker: Docker, image: String) -> Self { + Self { + docker, + image, + proxy_port: 0, + } + } + /// Check if the Docker daemon is available. pub async fn is_available(&self) -> bool { self.docker.ping().await.is_ok() @@ -140,6 +152,141 @@ impl ContainerRunner { Ok(()) } + /// Build the sandbox image from a Dockerfile. + /// + /// This is used when the image is not available from a registry and needs + /// to be built locally from source. + /// + /// # Security + /// + /// The `dockerfile_path` MUST point to a trusted Dockerfile. Docker builds + /// execute arbitrary `RUN` commands from the Dockerfile, which is a code + /// execution vector. Callers must ensure the path is not user-controlled + /// and points to a known, safe Dockerfile (e.g., bundled with the application). + pub async fn build_image(&self, dockerfile_path: &Path) -> Result<()> { + use tokio::io::AsyncBufReadExt; + use tokio::process::Command; + + const MAX_STDERR_CAPTURE: usize = 4096; + + // Canonicalize so the -f path is absolute and context_dir is its parent. + // This avoids the bug where a relative path like "docker/sandbox.Dockerfile" + // would be resolved twice (once for context_dir, once by docker -f). + let canonical = + dockerfile_path + .canonicalize() + .map_err(|e| SandboxError::ContainerCreationFailed { + reason: format!( + "cannot resolve Dockerfile path '{}': {}", + dockerfile_path.display(), + e + ), + })?; + + let context_dir = + canonical + .parent() + .ok_or_else(|| SandboxError::ContainerCreationFailed { + reason: format!( + "Dockerfile path '{}' has no parent directory", + canonical.display() + ), + })?; + + tracing::info!( + "Building sandbox image from {}: {}", + canonical.display(), + self.image + ); + + let mut child = Command::new("docker") + .arg("build") + .arg("-f") + .arg(&canonical) + .arg("-t") + .arg(&self.image) + .arg(".") + .current_dir(context_dir) + .stdout(std::process::Stdio::piped()) + .stderr(std::process::Stdio::piped()) + .spawn() + .map_err(|e| SandboxError::ContainerCreationFailed { + reason: format!("failed to run docker build: {}", e), + })?; + + // Both streams are piped above, so take() returns Some. + let mut stdout_lines = tokio::io::BufReader::new(child.stdout.take().ok_or_else(|| { + SandboxError::ContainerCreationFailed { + reason: "stdout pipe missing".to_string(), + } + })?) + .lines(); + let mut stderr_lines = tokio::io::BufReader::new(child.stderr.take().ok_or_else(|| { + SandboxError::ContainerCreationFailed { + reason: "stderr pipe missing".to_string(), + } + })?) + .lines(); + + let mut stderr_capture = String::new(); + let mut stdout_done = false; + let mut stderr_done = false; + + while !stdout_done || !stderr_done { + tokio::select! { + line = stdout_lines.next_line(), if !stdout_done => { + match line { + Ok(Some(line)) => tracing::info!("[docker build] {}", line), + Ok(None) => stdout_done = true, + Err(e) => { + tracing::warn!("Error reading docker build stdout: {}", e); + stdout_done = true; + } + } + }, + line = stderr_lines.next_line(), if !stderr_done => { + match line { + Ok(Some(line)) => { + tracing::info!("[docker build] {}", line); + if stderr_capture.len() < MAX_STDERR_CAPTURE { + stderr_capture.push_str(&line); + stderr_capture.push('\n'); + } + } + Ok(None) => stderr_done = true, + Err(e) => { + tracing::warn!("Error reading docker build stderr: {}", e); + stderr_done = true; + } + } + }, + } + } + + let status = child + .wait() + .await + .map_err(|e| SandboxError::ContainerCreationFailed { + reason: format!("docker build wait failed: {}", e), + })?; + + if !status.success() { + let code = status + .code() + .map_or("unknown".to_string(), |c| c.to_string()); + return Err(SandboxError::ContainerCreationFailed { + reason: format!( + "docker build failed (exit {}): {}", + code, + stderr_capture.trim_end() + ), + }); + } + + tracing::info!("Successfully built image: {}", self.image); + Ok(()) + } + /// Execute a command in a new container. pub async fn execute( &self, @@ -617,6 +764,22 @@ mod tests { assert!(candidates.contains(&PathBuf::from("/run/user/1000/docker.sock"))); } + #[tokio::test] + async fn build_image_rejects_nonexistent_dockerfile() { + // The behavior under test (Dockerfile path canonicalization) happens + // before any Docker daemon call, so we don't need a reachable daemon. + let docker = Docker::connect_with_http_defaults().unwrap(); + let runner = ContainerRunner::for_image_ops(docker, "test-nonexistent:latest".to_string()); + let dir = tempfile::tempdir().unwrap(); + let bad_path = dir.path().join("definitely-does-not-exist.Dockerfile"); + let err = runner.build_image(&bad_path).await.unwrap_err(); + let msg = err.to_string(); + assert!( + msg.contains("cannot resolve Dockerfile path"), + "expected path resolution error, got: {msg}" + ); + } + #[tokio::test] async fn test_docker_connection() { // This test requires Docker to be running diff --git a/src/settings.rs b/src/settings.rs index f549557b485..0ac5d534223 100644 --- a/src/settings.rs +++ b/src/settings.rs @@ -201,6 +201,22 @@ pub struct Settings { #[serde(default)] pub builder: BuilderSettings, + /// Routine scheduling and execution configuration. + #[serde(default)] + pub routines: RoutineSettings, + + /// Skills system configuration. + #[serde(default)] + pub skills: SkillsSettings, + + /// Memory hygiene configuration. + #[serde(default)] + pub hygiene: HygieneSettings, + + /// Workspace search fusion configuration. + #[serde(default)] + pub search: SearchSettings, + /// Transcription configuration. #[serde(default)] pub transcription: Option, @@ -791,6 +807,188 @@ impl Default for BuilderSettings { } } +/// Routine scheduling and execution configuration. +#[derive(Debug, Clone, Serialize, Deserialize)] +pub struct RoutineSettings { + /// Whether the routines system is enabled. + #[serde(default = "default_true")] + pub enabled: bool, + + /// How often (seconds) to poll for cron routines that need firing. + #[serde(default = "default_routine_cron_interval")] + pub cron_check_interval_secs: u64, + + /// Max routines executing concurrently. + #[serde(default = "default_routine_max_concurrent")] + pub max_concurrent_routines: usize, + + /// Default cooldown between fires (seconds). + #[serde(default = "default_routine_cooldown")] + pub default_cooldown_secs: u64, + + /// Max output tokens for lightweight routine LLM calls. + #[serde(default = "default_routine_max_tokens")] + pub max_lightweight_tokens: u32, + + /// Enable tool execution in lightweight routines. + #[serde(default = "default_true")] + pub lightweight_tools_enabled: bool, + + /// Max tool iterations for lightweight routines. + #[serde(default = "default_routine_max_iterations")] + pub lightweight_max_iterations: u32, +} + +fn default_routine_cron_interval() -> u64 { + 15 +} + +fn default_routine_max_concurrent() -> usize { + 10 +} + +fn default_routine_cooldown() -> u64 { + 300 +} + +fn default_routine_max_tokens() -> u32 { + 4096 +} + +fn default_routine_max_iterations() -> u32 { + 3 +} + +impl Default for RoutineSettings { + fn default() -> Self { + Self { + enabled: true, + cron_check_interval_secs: default_routine_cron_interval(), + max_concurrent_routines: default_routine_max_concurrent(), + default_cooldown_secs: default_routine_cooldown(), + max_lightweight_tokens: default_routine_max_tokens(), + lightweight_tools_enabled: true, + lightweight_max_iterations: default_routine_max_iterations(), + } + } +} + +/// Skills system configuration. +#[derive(Debug, Clone, Serialize, Deserialize)] +pub struct SkillsSettings { + /// Whether the skills system is enabled. + #[serde(default = "default_true")] + pub enabled: bool, + + /// Maximum number of skills that can be active simultaneously. + #[serde(default = "default_skills_max_active")] + pub max_active_skills: usize, + + /// Maximum total context tokens allocated to skill prompts. + #[serde(default = "default_skills_max_context_tokens")] + pub max_context_tokens: usize, +} + +fn default_skills_max_active() -> usize { + 3 +} + +fn default_skills_max_context_tokens() -> usize { + 4000 +} + +impl Default for SkillsSettings { + fn default() -> Self { + Self { + enabled: true, + max_active_skills: default_skills_max_active(), + max_context_tokens: default_skills_max_context_tokens(), + } + } +} + +/// Memory hygiene configuration. +#[derive(Debug, Clone, Serialize, Deserialize)] +pub struct HygieneSettings { + /// Whether hygiene is enabled. + #[serde(default = "default_true")] + pub enabled: bool, + + /// Days before `daily/` documents are deleted. + #[serde(default = "default_hygiene_daily_retention")] + pub daily_retention_days: u32, + + /// Days before `conversations/` documents are deleted. + #[serde(default = "default_hygiene_conversation_retention")] + pub conversation_retention_days: u32, + + /// Minimum hours between hygiene passes. + #[serde(default = "default_hygiene_cadence_hours")] + pub cadence_hours: u32, +} + +fn default_hygiene_daily_retention() -> u32 { + 30 +} + +fn default_hygiene_conversation_retention() -> u32 { + 7 +} + +fn default_hygiene_cadence_hours() -> u32 { + 12 +} + +impl Default for HygieneSettings { + fn default() -> Self { + Self { + enabled: true, + daily_retention_days: default_hygiene_daily_retention(), + conversation_retention_days: default_hygiene_conversation_retention(), + cadence_hours: default_hygiene_cadence_hours(), + } + } +} + +/// Workspace search fusion configuration. +#[derive(Debug, Clone, Serialize, Deserialize)] +pub struct SearchSettings { + /// Fusion strategy: "rrf" or "weighted". + #[serde(default = "default_search_fusion_strategy")] + pub fusion_strategy: String, + + /// RRF constant k. + #[serde(default = "default_search_rrf_k")] + pub rrf_k: u32, + + /// FTS weight for fusion. `None` = use per-strategy default. + #[serde(default)] + pub fts_weight: Option, + + /// Vector weight for fusion. `None` = use per-strategy default. + #[serde(default)] + pub vector_weight: Option, +} + +fn default_search_fusion_strategy() -> String { + "rrf".to_string() +} + +fn default_search_rrf_k() -> u32 { + 60 +} + +impl Default for SearchSettings { + fn default() -> Self { + Self { + fusion_strategy: default_search_fusion_strategy(), + rrf_k: default_search_rrf_k(), + fts_weight: None, + vector_weight: None, + } + } +} + /// Transcription pipeline settings. #[derive(Debug, Clone, Serialize, Deserialize)] pub struct TranscriptionSettings { @@ -907,8 +1105,9 @@ impl Settings { let content = format!( "# IronClaw configuration file.\n\ #\n\ - # Priority varies by subsystem. LLM: DB > env > this file > defaults.\n\ - # Most others: env > DB > this file > defaults.\n\ + # Priority: DB settings > env vars > this file > defaults.\n\ + # A DB value equal to the built-in default is treated as unset.\n\ + # Exceptions: bootstrap and security-sensitive fields are env-only.\n\ # Uncomment and edit values to override defaults.\n\ # Run `ironclaw config init` to regenerate this file.\n\ #\n\ diff --git a/src/setup/README.md b/src/setup/README.md index c1060cbcbea..58328d56b39 100644 --- a/src/setup/README.md +++ b/src/setup/README.md @@ -349,7 +349,7 @@ key first, then falls back to the standard env var. **Telegram special case** (`setup_telegram`): - Validates bot token via Telegram `getMe` API - Owner binding: polls `getUpdates` for 120s to capture sender's user ID -- Optional webhook secret generation +- Optional webhook secret auto-generation for webhook mode **SecretsContext creation** (`init_secrets_context`): 1. Check `self.secrets_crypto` (set in Step 2) → use if available diff --git a/src/setup/wizard.rs b/src/setup/wizard.rs index 7ad86610905..e8b25935721 100644 --- a/src/setup/wizard.rs +++ b/src/setup/wizard.rs @@ -65,6 +65,9 @@ pub enum SetupError { #[error("User cancelled")] Cancelled, + + #[error("Sandbox error: {0}")] + Sandbox(#[from] crate::sandbox::error::SandboxError), } impl From for SetupError { @@ -2859,6 +2862,9 @@ impl SetupWizard { crate::sandbox::detect::DockerStatus::Available => { self.settings.sandbox.enabled = true; print_success("Docker is installed and running. Sandbox enabled."); + + // Check if the worker image exists + self.ensure_worker_image().await?; } crate::sandbox::detect::DockerStatus::NotInstalled | crate::sandbox::detect::DockerStatus::NotRunning => { @@ -2888,6 +2894,8 @@ impl SetupWizard { } else { "Docker is now running. Sandbox enabled." }); + // Check if the worker image exists + self.ensure_worker_image().await?; } else { self.settings.sandbox.enabled = false; print_info(if not_installed { @@ -2974,6 +2982,113 @@ impl SetupWizard { Ok(()) } + /// Ensure the sandbox worker Docker image exists, building it if necessary. + async fn ensure_worker_image(&mut self) -> Result<(), SetupError> { + use crate::sandbox::container::{ContainerRunner, connect_docker}; + + let image_name = self.settings.sandbox.image.clone(); + let docker = match connect_docker().await { + Ok(d) => d, + Err(e) => { + // check_docker() may report Available (via CLI fallback) even when + // connect_docker() fails (e.g. on Windows). Don't hard-fail setup. + print_info(&format!( + "Could not connect to Docker API to verify image: {}", + e + )); + print_info("Image check skipped. The image will be pulled at first job run."); + return Ok(()); + } + }; + let runner = ContainerRunner::for_image_ops(docker, image_name.clone()); + + if runner.image_exists().await { + print_success(&format!("Worker image '{}' found.", image_name)); + return Ok(()); + } + + println!(); + print_info(&format!("Worker image '{}' not found.", image_name)); + print_info("This image is required for sandboxed job execution."); + println!(); + + // Images that contain '/' look like registry references (e.g. + // "ghcr.io/nearai/ironclaw-worker:v1"). For those, or when + // auto_pull_image is enabled, attempt a pull before offering a + // local build — the runtime would do the same thing via + // SandboxManager::ensure_ready(). + let is_registry_image = image_name.contains('/'); + if is_registry_image || self.settings.sandbox.auto_pull_image { + print_info(&format!("Attempting to pull '{}'...", image_name)); + match runner.pull_image().await { + Ok(()) => { + print_success(&format!("Successfully pulled image '{}'.", image_name)); + return Ok(()); + } + Err(e) => { + if is_registry_image { + // Registry image that can't be pulled — don't offer local build. + print_error(&format!("Failed to pull image: {}", e)); + print_info("Ensure the image is published and accessible, or set"); + print_info("SANDBOX_IMAGE to a local image name and try again."); + return Ok(()); + } + print_info(&format!( + "Pull failed ({}). Checking for local Dockerfile...", + e + )); + } + } + } + + // Only offer local build for default-style local images. + let dockerfile_path = std::path::PathBuf::from("Dockerfile.worker"); + + if dockerfile_path.exists() { + print_info(&format!( + "Found Dockerfile at: {}", + dockerfile_path.display() + )); + if confirm( + "Build the worker image now? (this may take a few minutes)", + true, + ) + .map_err(SetupError::Io)? + { + print_info("Building worker image... This may take a few minutes."); + match runner.build_image(&dockerfile_path).await { + Ok(()) => { + print_success(&format!("Successfully built image '{}'.", image_name)); + } + Err(e) => { + print_error(&format!("Failed to build image: {}", e)); + print_info("You can build it manually later with:"); + print_info(&format!( + " docker build -f Dockerfile.worker -t {} .", + image_name + )); + } + } + } else { + print_info("Skipped image build. Build it manually with:"); + print_info(&format!( + " docker build -f Dockerfile.worker -t {} .", + image_name + )); + } + } else { + print_info("No Dockerfile.worker found in current directory."); + print_info("To use Docker sandbox, build the worker image manually:"); + print_info(&format!( + " docker build -f Dockerfile.worker -t {} .", + image_name + )); + print_info("or clone the IronClaw repository and build from source."); + } + + Ok(()) + } + /// Step 9: Heartbeat configuration. fn step_heartbeat(&mut self) -> Result<(), SetupError> { print_info("Heartbeat runs periodic background tasks (e.g., checking your calendar,"); diff --git a/src/tenant.rs b/src/tenant.rs index b8af5d84dd5..a73d954ce7a 100644 --- a/src/tenant.rs +++ b/src/tenant.rs @@ -278,7 +278,7 @@ impl TenantScope { thread_id: Option<&str>, ) -> Result { self.inner - .ensure_conversation(id, channel, &self.user_id, thread_id) + .ensure_conversation(id, channel, &self.user_id, thread_id, Some(channel)) .await } diff --git a/src/testing/mod.rs b/src/testing/mod.rs index dfff4b10f74..bd1cc52f60b 100644 --- a/src/testing/mod.rs +++ b/src/testing/mod.rs @@ -753,7 +753,7 @@ mod tests { // ensure_conversation should create the row. assert!( - db.ensure_conversation(conv_id, "web", "carol", None) + db.ensure_conversation(conv_id, "web", "carol", None, Some("web")) .await .expect("ensure first"), "first ensure_conversation should create the row" @@ -761,7 +761,7 @@ mod tests { // Calling again with the same ID should not error. assert!( - db.ensure_conversation(conv_id, "web", "carol", None) + db.ensure_conversation(conv_id, "web", "carol", None, Some("web")) .await .expect("ensure second (idempotent)"), "second ensure_conversation should touch owned row" @@ -806,7 +806,7 @@ mod tests { tokio::time::sleep(std::time::Duration::from_millis(25)).await; assert!( - !db.ensure_conversation(conv_id, "web", "mallory", None) + !db.ensure_conversation(conv_id, "web", "mallory", None, None) .await .expect("foreign ensure should not error"), "foreign ensure_conversation should report not ensured" diff --git a/src/tools/builder/core.rs b/src/tools/builder/core.rs index 3cfc0cbaa57..5073b79d69e 100644 --- a/src/tools/builder/core.rs +++ b/src/tools/builder/core.rs @@ -48,6 +48,75 @@ use crate::tools::tool::{ }; use crate::tools::{ToolRegistry, prepare_tool_params}; +/// Deserialize `dependencies` from either a list of strings, a list of objects, +/// or a flat object map. LLMs often produce TOML-style inline tables +/// (`{"ureq": {"version": "2"}}`) instead of simple strings (`"ureq = \"2\""`); +/// this normalises all variants to `Vec<"name = \"version\"">` strings. +fn deserialize_dependencies<'de, D>(deserializer: D) -> Result, D::Error> +where + D: serde::Deserializer<'de>, +{ + use serde::de; + use serde_json::Value; + + let val = Value::deserialize(deserializer)?; + match val { + Value::Array(arr) => { + let mut out = Vec::with_capacity(arr.len()); + for item in arr { + match item { + Value::String(s) => out.push(s), + Value::Object(map) => { + for (name, spec) in map { + if let Some(dep) = flatten_dep(&name, &spec) { + out.push(dep); + } + } + } + other => { + return Err(de::Error::custom(format!( + "expected string or object in dependencies array, got {other}" + ))); + } + } + } + Ok(out) + } + Value::Object(map) => { + let out: Vec = map + .into_iter() + .filter_map(|(name, spec)| flatten_dep(&name, &spec)) + .collect(); + Ok(out) + } + Value::Null => Ok(Vec::new()), + other => Err(de::Error::custom(format!( + "expected array or object for dependencies, got {other}" + ))), + } +} + +/// Flatten a dependency entry like `("ureq", "2")` or `("serde", {"version":"1","features":["derive"]})` +/// into a TOML-compatible string like `ureq = "2"` or `serde = { version = "1", features = ["derive"] }`. +/// Returns `None` for values that cannot produce valid TOML (null, bool, array, number). +fn flatten_dep(name: &str, spec: &serde_json::Value) -> Option { + match spec { + serde_json::Value::String(version) => Some(format!("{name} = \"{version}\"")), + serde_json::Value::Object(map) => { + let parts: Vec = map.iter().map(|(k, v)| format!("{k} = {v}")).collect(); + Some(format!("{name} = {{ {} }}", parts.join(", "))) + } + _ => { + tracing::warn!( + name, + ?spec, + "Skipping dependency with unsupported value type" + ); + None + } + } +} + fn process_builder_tool_result( tool_name: &str, tool_call_id: &str, @@ -80,6 +149,7 @@ pub struct BuildRequirement { /// Expected output format. pub output_spec: Option, /// External dependencies needed. + #[serde(default, deserialize_with = "deserialize_dependencies")] pub dependencies: Vec, /// Security/capability requirements (for WASM tools). pub capabilities: Vec, @@ -1236,6 +1306,59 @@ mod tests { assert!(deserialized.capabilities.is_empty()); } + /// Regression test for #1640: LLM returns dependencies as inline tables + /// (`{"ureq": {"version": "2"}}`) instead of simple strings. + #[test] + fn test_build_requirement_deserialize_inline_table_deps() { + let json = r#"{ + "name": "gh-search", + "description": "Search GitHub repos", + "software_type": "wasm_tool", + "language": "rust", + "dependencies": [ + {"ureq": {"version": "2"}}, + {"serde": {"version": "1", "features": ["derive"]}} + ], + "capabilities": ["http"] + }"#; + let req: BuildRequirement = serde_json::from_str(json).unwrap(); + assert_eq!(req.dependencies.len(), 2); + assert!(req.dependencies[0].starts_with("ureq = ")); + assert!(req.dependencies[1].starts_with("serde = ")); + } + + /// Dependencies as a flat object map (`{"ureq": "2", "serde_json": "1"}`). + #[test] + fn test_build_requirement_deserialize_object_map_deps() { + let json = r#"{ + "name": "tool", + "description": "A tool", + "software_type": "wasm_tool", + "language": "rust", + "dependencies": {"ureq": "2", "serde_json": "1"}, + "capabilities": [] + }"#; + let req: BuildRequirement = serde_json::from_str(json).unwrap(); + assert_eq!(req.dependencies.len(), 2); + assert!(req.dependencies.iter().any(|d| d.contains("ureq"))); + assert!(req.dependencies.iter().any(|d| d.contains("serde_json"))); + } + + /// Null or missing dependencies should deserialize to empty vec. + #[test] + fn test_build_requirement_deserialize_null_deps() { + let json = r#"{ + "name": "tool", + "description": "A tool", + "software_type": "script", + "language": "bash", + "dependencies": null, + "capabilities": [] + }"#; + let req: BuildRequirement = serde_json::from_str(json).unwrap(); + assert!(req.dependencies.is_empty()); + } + #[test] fn test_builder_config_default_sensible_values() { let config = BuilderConfig::default(); diff --git a/src/tools/builtin/job.rs b/src/tools/builtin/job.rs index 4c711e6904a..940c7cfc693 100644 --- a/src/tools/builtin/job.rs +++ b/src/tools/builtin/job.rs @@ -21,7 +21,7 @@ use crate::context::{ContextManager, JobContext, JobState}; use crate::db::Database; use crate::history::SandboxJobRecord; use crate::orchestrator::auth::CredentialGrant; -use crate::orchestrator::job_manager::{ContainerJobManager, JobMode}; +use crate::orchestrator::job_manager::{ContainerJobManager, JobCreationParams, JobMode}; use crate::secrets::SecretsStore; use crate::tools::tool::{ApprovalRequirement, Tool, ToolError, ToolOutput, require_str}; use ironclaw_common::AppEvent; @@ -361,7 +361,7 @@ impl CreateJobTool { explicit_dir: Option, wait: bool, mode: JobMode, - credential_grants: Vec, + params: JobCreationParams, ctx: &JobContext, ) -> Result { let start = std::time::Instant::now(); @@ -376,7 +376,7 @@ impl CreateJobTool { let project_dir_str = project_dir.display().to_string(); // Serialize credential grants so restarts can reload them. - let credential_grants_json = match serde_json::to_string(&credential_grants) { + let credential_grants_json = match serde_json::to_string(¶ms.credential_grants) { Ok(json) => json, Err(e) => { tracing::warn!( @@ -431,7 +431,7 @@ impl CreateJobTool { // Create the container job with the pre-determined job_id. let _token = jm - .create_job(job_id, task, Some(project_dir), mode, credential_grants) + .create_job(job_id, task, Some(project_dir), mode, params) .await .map_err(|e| { self.update_status( @@ -849,6 +849,18 @@ impl Tool for CreateJobTool { secrets store (via 'ironclaw tool auth' or web UI). Example: \ {\"github_token\": \"GITHUB_TOKEN\", \"npm_token\": \"NPM_TOKEN\"}", "additionalProperties": { "type": "string" } + }, + "mcp_servers": { + "type": "array", + "items": { "type": "string" }, + "description": "Optional list of MCP server names to make available in the container. \ + If omitted, the full master config is mounted. If empty, no MCP servers \ + are available. Only effective when MCP_PER_JOB_ENABLED=true." + }, + "max_iterations": { + "type": "integer", + "description": "Maximum number of agent loop iterations for the worker. \ + Defaults to 50, capped at 500. Use lower values for simple tasks." } }, "required": ["title", "description"] @@ -909,10 +921,44 @@ impl Tool for CreateJobTool { // Parse and validate credential grants let credential_grants = self.parse_credentials(¶ms, &ctx.user_id).await?; + // Parse optional MCP server filter and iteration cap. + // Validate types: warn if present but wrong type so callers know why it was ignored. + let mcp_servers: Option> = match params.get("mcp_servers") { + Some(v) if v.is_array() => v.as_array().map(|arr| { + arr.iter() + .filter_map(|v| v.as_str().map(String::from)) + .collect() + }), + Some(_) => { + tracing::warn!("mcp_servers parameter is not an array — ignoring"); + None + } + None => None, + }; + let max_iterations: Option = match params.get("max_iterations") { + Some(v) if v.is_u64() || v.is_i64() => v.as_u64().map(|n| n.clamp(1, 500) as u32), + Some(_) => { + tracing::warn!("max_iterations parameter is not a number — ignoring"); + None + } + None => None, + }; + // Combine title and description into the task prompt for the sub-agent. let task = format!("{}\n\n{}", title, description); - self.execute_sandbox(&task, explicit_dir, wait, mode, credential_grants, ctx) - .await + self.execute_sandbox( + &task, + explicit_dir, + wait, + mode, + JobCreationParams { + credential_grants, + mcp_servers, + max_iterations, + }, + ctx, + ) + .await } else { self.execute_local(title, description, ctx).await } @@ -1568,7 +1614,7 @@ mod tests { None, false, JobMode::Worker, - vec![], + JobCreationParams::default(), &JobContext::default(), ) .await; diff --git a/src/tools/builtin/message.rs b/src/tools/builtin/message.rs index 33f1b74563b..1c1a5479cf5 100644 --- a/src/tools/builtin/message.rs +++ b/src/tools/builtin/message.rs @@ -205,11 +205,11 @@ impl Tool for MessageTool { }, "channel": { "type": "string", - "description": "Target channel (defaults to current channel if omitted)" + "description": "Transport/integration name: 'slack-relay', 'telegram', 'signal', 'gateway'. This is NOT a Slack channel ID — use target for that. Defaults to current channel if omitted." }, "target": { "type": "string", - "description": "Recipient: E.164 phone, group ID, chat ID (defaults to current sender/group if omitted)" + "description": "Recipient within the transport. Slack: channel ID (C0...), user ID (U0...), or #channel-name. Telegram: chat ID. Signal: E.164 phone or group ID. Defaults to current conversation target if omitted." }, "attachments": { "type": "array", diff --git a/src/tools/builtin/routine.rs b/src/tools/builtin/routine.rs index 0424add8f2d..676b5077aab 100644 --- a/src/tools/builtin/routine.rs +++ b/src/tools/builtin/routine.rs @@ -1178,8 +1178,22 @@ impl Tool for RoutineCreateTool { dedup_window: None, }, notify: NotifyConfig { - channel: normalized.delivery.channel.clone(), - user: normalized.delivery.user.clone(), + // Fall back to the current conversation's channel/target when + // the LLM omits delivery params, so routines created from + // e.g. a Slack channel know where to send results. + channel: normalized.delivery.channel.clone().or_else(|| { + ctx.metadata + .get("notify_channel") + .and_then(|v| v.as_str()) + .map(ToOwned::to_owned) + }), + user: normalized.delivery.user.clone().or_else(|| { + ctx.metadata + .get("notify_user") + .and_then(|v| v.as_str()) + .filter(|v| *v != "default") + .map(ToOwned::to_owned) + }), ..NotifyConfig::default() }, last_run_at: None, @@ -2614,6 +2628,28 @@ mod tests { ); } + /// Regression: routine_create must fall back to ctx.metadata for delivery + /// config when the LLM omits delivery.channel/user. This verifies the + /// parsing layer returns None so the execute path triggers the fallback. + #[test] + fn routine_create_omitted_delivery_enables_context_fallback() { + let params = serde_json::json!({ + "name": "ping-every-5", + "prompt": "Send Ping in this channel.", + "request": { "kind": "cron", "schedule": "*/5 * * * *" } + }); + + let parsed = parse_routine_create_request(¶ms).expect("parse"); + assert!( + parsed.delivery.channel.is_none(), + "omitted delivery.channel should be None so execute() falls back to ctx.metadata", + ); + assert!( + parsed.delivery.user.is_none(), + "omitted delivery.user should be None so execute() falls back to ctx.metadata", + ); + } + #[test] fn build_full_job_action_uses_live_owner_scope_defaults() { let execution = NormalizedExecutionRequest { diff --git a/src/tools/builtin/shell.rs b/src/tools/builtin/shell.rs index fa92cb37232..ed365be8a6b 100644 --- a/src/tools/builtin/shell.rs +++ b/src/tools/builtin/shell.rs @@ -855,7 +855,8 @@ impl Tool for ShellTool { }, "timeout": { "type": "integer", - "description": "Timeout in seconds (optional, default 120)" + "description": "Timeout in seconds (optional, default 120)", + "minimum": 1 } }, "required": ["command"] @@ -869,8 +870,41 @@ impl Tool for ShellTool { ) -> Result { let command = require_str(¶ms, "command")?; - let workdir = params.get("workdir").and_then(|v| v.as_str()); - let timeout = params.get("timeout").and_then(|v| v.as_u64()); + let workdir = match params.get("workdir") { + None => None, + Some(v) if v.is_null() => None, + Some(v) => { + let s = v.as_str().ok_or_else(|| { + ToolError::InvalidParameters("workdir must be a string".to_string()) + })?; + let trimmed = s.trim(); + if trimmed.is_empty() { + None + } else { + Some(trimmed) + } + } + }; + + let timeout = match params.get("timeout") { + None => None, + Some(v) if v.is_null() => None, + Some(v) => { + let n = v.as_u64().ok_or_else(|| { + ToolError::InvalidParameters( + "timeout must be a positive integer number of seconds".to_string(), + ) + })?; + + if n == 0 { + return Err(ToolError::InvalidParameters( + "timeout must be greater than 0".to_string(), + )); + } + + Some(n) + } + }; let start = std::time::Instant::now(); let (output, exit_code) = self @@ -952,17 +986,160 @@ fn truncate_for_error(s: &str) -> String { #[cfg(test)] mod tests { use super::*; + use tempfile::TempDir; + + async fn execute_shell( + tool: &ShellTool, + params: serde_json::Value, + ) -> Result { + tool.execute(params, &JobContext::default()).await + } + + fn assert_invalid_parameters(result: Result, expected_message: &str) { + match result { + Err(ToolError::InvalidParameters(message)) => { + assert_eq!(message, expected_message); + } + Err(other) => panic!("expected InvalidParameters, got {other:?}"), + Ok(output) => panic!("expected InvalidParameters, got success: {output:?}"), + } + } #[tokio::test] async fn test_echo_command() { let tool = ShellTool::new(); - let ctx = JobContext::default(); + let result = execute_shell(&tool, serde_json::json!({"command": "echo hello"})) + .await + .unwrap(); - let result = tool - .execute(serde_json::json!({"command": "echo hello"}), &ctx) + let output = result.result.get("output").unwrap().as_str().unwrap(); + assert!(output.contains("hello")); + assert_eq!(result.result.get("exit_code").unwrap().as_i64().unwrap(), 0); + } + + #[tokio::test] + async fn test_execute_treats_blank_workdir_as_none() { + let temp_dir = TempDir::new().unwrap(); + let tool = ShellTool::new().with_working_dir(temp_dir.path().to_path_buf()); + + for workdir in ["", " "] { + let result = execute_shell( + &tool, + serde_json::json!({ + "command": "pwd", + "workdir": workdir + }), + ) .await .unwrap(); + let output = result + .result + .get("output") + .unwrap() + .as_str() + .unwrap() + .trim(); + let output_path = PathBuf::from(output).canonicalize().unwrap(); + let expected_path = temp_dir.path().canonicalize().unwrap(); + assert_eq!(output_path, expected_path); + } + } + + #[tokio::test] + async fn test_execute_rejects_non_string_workdir() { + let tool = ShellTool::new(); + let result = execute_shell( + &tool, + serde_json::json!({ + "command": "pwd", + "workdir": 42 + }), + ) + .await; + + assert_invalid_parameters(result, "workdir must be a string"); + } + + #[tokio::test] + async fn test_execute_treats_missing_or_null_timeout_as_none() { + let tool = ShellTool::new(); + + for params in [ + serde_json::json!({"command": "echo hello"}), + serde_json::json!({"command": "echo hello", "timeout": null}), + ] { + let result = execute_shell(&tool, params).await.unwrap(); + let output = result.result.get("output").unwrap().as_str().unwrap(); + assert!(output.contains("hello")); + assert_eq!(result.result.get("exit_code").unwrap().as_i64().unwrap(), 0); + } + } + + #[tokio::test] + async fn test_execute_rejects_non_numeric_timeout_string() { + let tool = ShellTool::new(); + let result = execute_shell( + &tool, + serde_json::json!({ + "command": "echo hello", + "timeout": "abc" + }), + ) + .await; + + assert_invalid_parameters( + result, + "timeout must be a positive integer number of seconds", + ); + } + + #[tokio::test] + async fn test_execute_rejects_zero_timeout() { + let tool = ShellTool::new(); + let result = execute_shell( + &tool, + serde_json::json!({ + "command": "echo hello", + "timeout": 0 + }), + ) + .await; + + assert_invalid_parameters(result, "timeout must be greater than 0"); + } + + #[tokio::test] + async fn test_execute_rejects_float_timeout() { + let tool = ShellTool::new(); + let result = execute_shell( + &tool, + serde_json::json!({ + "command": "echo hello", + "timeout": 3.5 + }), + ) + .await; + + assert_invalid_parameters( + result, + "timeout must be a positive integer number of seconds", + ); + } + + #[tokio::test] + async fn test_execute_accepts_valid_timeout() { + let tool = ShellTool::new(); + let result = execute_shell( + &tool, + serde_json::json!({ + "command": "echo hello", + "timeout": 30 + }), + ) + .await + .unwrap(); + let output = result.result.get("output").unwrap().as_str().unwrap(); assert!(output.contains("hello")); assert_eq!(result.result.get("exit_code").unwrap().as_i64().unwrap(), 0); diff --git a/src/tools/registry.rs b/src/tools/registry.rs index 8c08633bbd2..ce6e3bf4c90 100644 --- a/src/tools/registry.rs +++ b/src/tools/registry.rs @@ -25,7 +25,7 @@ use crate::tools::builtin::{ ToolUpgradeTool, WriteFileTool, }; use crate::tools::rate_limiter::RateLimiter; -use crate::tools::tool::{ApprovalRequirement, Tool, ToolDomain}; +use crate::tools::tool::{ApprovalRequirement, Tool, ToolDiscoverySummary, ToolDomain}; use crate::tools::wasm::{ Capabilities, OAuthRefreshConfig, ResourceLimits, SharedCredentialRegistry, WasmError, WasmStorageError, WasmToolRuntime, WasmToolStore, WasmToolWrapper, @@ -674,6 +674,9 @@ impl ToolRegistry { if let Some(s) = reg.schema { wrapper = wrapper.with_schema(s); } + if let Some(summary) = reg.discovery_summary { + wrapper = wrapper.with_discovery_summary(summary); + } if let Some(store) = reg.secrets_store { wrapper = wrapper.with_secrets_store(store); } @@ -748,6 +751,7 @@ impl ToolRegistry { limits: None, description: Some(&tool_with_binary.tool.description), schema: Some(tool_with_binary.tool.parameters_schema.clone()), + discovery_summary: None, secrets_store: self.secrets_store.clone(), oauth_refresh: None, }) @@ -791,6 +795,8 @@ pub struct WasmToolRegistration<'a> { pub description: Option<&'a str>, /// Optional parameter schema override. pub schema: Option, + /// Optional curated discovery guidance for `tool_info(detail: "summary")`. + pub discovery_summary: Option, /// Secrets store for credential injection at request time. pub secrets_store: Option>, /// OAuth refresh configuration for auto-refreshing expired tokens. diff --git a/src/tools/wasm/capabilities_schema.rs b/src/tools/wasm/capabilities_schema.rs index 8ac7806ca58..7f3cbb08101 100644 --- a/src/tools/wasm/capabilities_schema.rs +++ b/src/tools/wasm/capabilities_schema.rs @@ -33,6 +33,7 @@ use std::time::Duration; use serde::{Deserialize, Serialize}; use crate::secrets::{CredentialLocation, CredentialMapping}; +use crate::tools::tool::ToolDiscoverySummary; use crate::tools::wasm::{ Capabilities, EndpointPattern, HttpCapability, RateLimitConfig, SecretsCapability, ToolInvokeCapability, WebhookCapability, WorkspaceCapability, @@ -47,6 +48,10 @@ pub struct CapabilitiesFile { #[serde(default)] pub description: Option, + /// Optional curated guidance surfaced by `tool_info(detail: "summary")`. + #[serde(default)] + pub discovery_summary: Option, + /// Extension version (semver). #[serde(default)] pub version: Option, @@ -154,6 +159,7 @@ impl CapabilitiesFile { if let Some(inner) = self.capabilities.take() { let inner = inner.resolve_nested_inner(depth + 1); self.description = self.description.or(inner.description); + self.discovery_summary = self.discovery_summary.or(inner.discovery_summary); self.http = self.http.or(inner.http); self.secrets = self.secrets.or(inner.secrets); self.tool_invoke = self.tool_invoke.or(inner.tool_invoke); @@ -1531,6 +1537,28 @@ mod tests { ); } + #[test] + fn test_discovery_summary_promoted_from_nested_capabilities() { + let json = r#"{ + "capabilities": { + "discovery_summary": { + "always_required": ["action"], + "notes": ["Use tool_info for full schema"] + } + } + }"#; + + let caps = CapabilitiesFile::from_json(json).unwrap(); + let summary = caps + .discovery_summary + .expect("discovery summary should be promoted"); + assert_eq!(summary.always_required, vec!["action".to_string()]); + assert_eq!( + summary.notes, + vec!["Use tool_info for full schema".to_string()] + ); + } + /// Regression test for issue #974: deeply nested capabilities wrappers /// must not cause stack overflow. resolve_nested should stop at /// MAX_NESTED_DEPTH and return gracefully. diff --git a/src/tools/wasm/loader.rs b/src/tools/wasm/loader.rs index 680abf939b9..8aa0ba472e5 100644 --- a/src/tools/wasm/loader.rs +++ b/src/tools/wasm/loader.rs @@ -128,48 +128,50 @@ impl WasmToolLoader { // capabilities file — it is auto-derived from the WASM module's // schema() export at prepare time (see WasmToolSchemas::compact_schema), // so no schema override is needed here. - let (capabilities, oauth_refresh, description) = if let Some(cap_path) = capabilities_path { - if cap_path.exists() { - let cap_bytes = fs::read(cap_path).await?; - let cap_file = CapabilitiesFile::from_bytes(&cap_bytes) - .map_err(|e| WasmLoadError::InvalidCapabilities(e.to_string()))?; - cap_file.validate(name); - - // Check WIT version compatibility - check_wit_version_compat( - name, - cap_file.wit_version.as_deref(), - crate::tools::wasm::WIT_TOOL_VERSION, - )?; - - let caps = cap_file.to_capabilities(); - let oauth = resolve_oauth_refresh_config(&cap_file); - let desc = cap_file.description.clone(); - if desc.is_none() { + let (capabilities, oauth_refresh, description, discovery_summary) = + if let Some(cap_path) = capabilities_path { + if cap_path.exists() { + let cap_bytes = fs::read(cap_path).await?; + let cap_file = CapabilitiesFile::from_bytes(&cap_bytes) + .map_err(|e| WasmLoadError::InvalidCapabilities(e.to_string()))?; + cap_file.validate(name); + + // Check WIT version compatibility + check_wit_version_compat( + name, + cap_file.wit_version.as_deref(), + crate::tools::wasm::WIT_TOOL_VERSION, + )?; + + let caps = cap_file.to_capabilities(); + let oauth = resolve_oauth_refresh_config(&cap_file); + let desc = cap_file.description.clone(); + let summary = cap_file.discovery_summary.clone(); + if desc.is_none() { + tracing::warn!( + tool = name, + path = %cap_path.display(), + "Capabilities file missing \"description\" field; \ + tool will use generic fallback description" + ); + } + (caps, oauth, desc, summary) + } else { tracing::warn!( tool = name, path = %cap_path.display(), - "Capabilities file missing \"description\" field; \ - tool will use generic fallback description" + "Capabilities file not found, using default (no permissions)" ); + (Capabilities::default(), None, None, None) } - (caps, oauth, desc) } else { tracing::warn!( tool = name, - path = %cap_path.display(), - "Capabilities file not found, using default (no permissions)" - ); - (Capabilities::default(), None, None) - } - } else { - tracing::warn!( - tool = name, - "No capabilities file for WASM tool; \ + "No capabilities file for WASM tool; \ tool will use generic fallback description" - ); - (Capabilities::default(), None, None) - }; + ); + (Capabilities::default(), None, None, None) + }; // Register the tool self.registry @@ -181,6 +183,7 @@ impl WasmToolLoader { limits: None, description: description.as_deref(), schema: None, + discovery_summary, secrets_store: self.secrets_store.clone(), oauth_refresh, }) diff --git a/src/tools/wasm/wrapper.rs b/src/tools/wasm/wrapper.rs index f8f69ba435a..ed2994fa42e 100644 --- a/src/tools/wasm/wrapper.rs +++ b/src/tools/wasm/wrapper.rs @@ -20,7 +20,7 @@ use crate::context::JobContext; use crate::llm::recording::{HttpExchangeRequest, HttpExchangeResponse, HttpInterceptor}; use crate::safety::LeakDetector; use crate::secrets::{DecryptedSecret, SecretsStore}; -use crate::tools::tool::{Tool, ToolError, ToolOutput}; +use crate::tools::tool::{Tool, ToolDiscoverySummary, ToolError, ToolOutput}; use crate::tools::wasm::capabilities::Capabilities; use crate::tools::wasm::credential_injector::{ InjectedCredentials, host_matches_pattern, inject_credential, @@ -585,6 +585,8 @@ pub struct WasmToolWrapper { description: String, /// Compact and discovery schemas for this tool. schemas: WasmToolSchemas, + /// Optional curated discovery guidance surfaced by `tool_info`. + discovery_summary: Option, /// Injected credentials for HTTP requests (e.g., OAuth tokens). /// Keys are placeholder names like "GOOGLE_ACCESS_TOKEN". credentials: HashMap, @@ -836,6 +838,7 @@ impl WasmToolWrapper { Self { description: prepared.description.clone(), schemas: WasmToolSchemas::new(prepared.schema.clone()), + discovery_summary: None, runtime, prepared, capabilities, @@ -882,6 +885,12 @@ impl WasmToolWrapper { self } + /// Override the curated discovery summary. + pub fn with_discovery_summary(mut self, summary: ToolDiscoverySummary) -> Self { + self.discovery_summary = Some(summary); + self + } + /// Set credentials for HTTP request placeholder injection. pub fn with_credentials(mut self, credentials: HashMap) -> Self { self.credentials = credentials; @@ -1098,6 +1107,10 @@ impl Tool for WasmToolWrapper { self.schemas.discovery() } + fn discovery_summary(&self) -> Option { + self.discovery_summary.clone() + } + /// Compose the tool schema for LLM function calling. /// /// When the advertised schema is permissive (no typed properties), appends @@ -1149,6 +1162,7 @@ impl Tool for WasmToolWrapper { let capabilities = self.capabilities.clone(); let description = self.description.clone(); let schemas = self.schemas.clone(); + let discovery_summary = self.discovery_summary.clone(); let credentials = self.credentials.clone(); // Execute in blocking task with timeout @@ -1159,6 +1173,7 @@ impl Tool for WasmToolWrapper { capabilities, description, schemas, + discovery_summary, credentials, secrets_store: None, // Not needed in blocking task oauth_refresh: None, // Already used above for pre-refresh @@ -3058,6 +3073,27 @@ mod tests { assert_eq!(wrapper.discovery_schema(), typed_schema); // safety: test-only assertion } + #[tokio::test] + async fn test_wrapper_returns_curated_discovery_summary() { + let runtime = Arc::new(WasmToolRuntime::new(WasmRuntimeConfig::for_testing()).unwrap()); // safety: test-only setup + let prepared = runtime + .prepare("github", b"\0asm\x0d\0\x01\0", None) + .await + .unwrap(); // safety: test-only setup + + let summary = crate::tools::tool::ToolDiscoverySummary { + always_required: vec!["action".into()], + notes: vec!["Use tool_info for the full schema".into()], + ..crate::tools::tool::ToolDiscoverySummary::default() + }; + + let wrapper = + super::WasmToolWrapper::new(Arc::clone(&runtime), prepared, Capabilities::default()) + .with_discovery_summary(summary.clone()); + + assert_eq!(wrapper.discovery_summary(), Some(summary)); + } + #[test] fn test_build_tool_usage_hint_detects_nullable_container_properties() { let schema = serde_json::json!({ diff --git a/src/worker/job.rs b/src/worker/job.rs index e64eef0139e..63923a23f89 100644 --- a/src/worker/job.rs +++ b/src/worker/job.rs @@ -316,7 +316,6 @@ Report when the job is complete or if you encounter issues you cannot resolve."# reasoning: &Reasoning, reason_ctx: &mut ReasoningContext, ) -> Result<(), Error> { - const MAX_WORKER_ITERATIONS: usize = 500; let max_iterations = self .context_manager() .get_context(self.job_id) @@ -324,7 +323,7 @@ Report when the job is complete or if you encounter issues you cannot resolve."# .ok() .and_then(|ctx| ctx.metadata.get("max_iterations").and_then(|v| v.as_u64())) .unwrap_or(50) as usize; - let max_iterations = max_iterations.min(MAX_WORKER_ITERATIONS); + let max_iterations = max_iterations.min(ironclaw_common::MAX_WORKER_ITERATIONS as usize); // Initial tool definitions for planning (will be refreshed in loop) reason_ctx.available_tools = self.tools().tool_definitions().await; diff --git a/tests/e2e/conftest.py b/tests/e2e/conftest.py index 7d223a472ee..a9e35af83a1 100644 --- a/tests/e2e/conftest.py +++ b/tests/e2e/conftest.py @@ -443,6 +443,200 @@ async def hosted_oauth_refresh_server( home_tmpdir.cleanup() +@pytest.fixture(scope="session") +async def loop_limited_server( + ironclaw_binary, + mock_llm_server, + wasm_tools_dir, +): + """Start an isolated ironclaw instance with a low tool-iteration limit.""" + reserved = _reserve_loopback_sockets(2) + db_tmpdir = tempfile.TemporaryDirectory(prefix="ironclaw-e2e-loop-limit-db-") + home_tmpdir = tempfile.TemporaryDirectory(prefix="ironclaw-e2e-loop-limit-home-") + + try: + gateway_port = reserved[0].getsockname()[1] + http_port = reserved[1].getsockname()[1] + for sock in reserved: + if sock.fileno() != -1: + sock.close() + + env = { + "PATH": os.environ.get("PATH", "/usr/bin:/bin"), + "HOME": home_tmpdir.name, + "IRONCLAW_BASE_DIR": os.path.join(home_tmpdir.name, ".ironclaw"), + "RUST_LOG": "ironclaw=info", + "RUST_BACKTRACE": "1", + "IRONCLAW_OWNER_ID": OWNER_SCOPE_ID, + "GATEWAY_ENABLED": "true", + "GATEWAY_HOST": "127.0.0.1", + "GATEWAY_PORT": str(gateway_port), + "GATEWAY_AUTH_TOKEN": AUTH_TOKEN, + "HTTP_HOST": "127.0.0.1", + "HTTP_PORT": str(http_port), + "HTTP_WEBHOOK_SECRET": HTTP_WEBHOOK_SECRET, + "CLI_ENABLED": "false", + "LLM_BACKEND": "openai_compatible", + "LLM_BASE_URL": mock_llm_server, + "LLM_MODEL": "mock-model", + "DATABASE_BACKEND": "libsql", + "LIBSQL_PATH": os.path.join(db_tmpdir.name, "loop-limited.db"), + "SANDBOX_ENABLED": "false", + "SKILLS_ENABLED": "true", + "ROUTINES_ENABLED": "true", + "HEARTBEAT_ENABLED": "false", + "EMBEDDING_ENABLED": "false", + "WASM_ENABLED": "true", + "WASM_TOOLS_DIR": wasm_tools_dir, + "WASM_CHANNELS_DIR": _WASM_CHANNELS_TMPDIR.name, + "ONBOARD_COMPLETED": "true", + "IRONCLAW_OAUTH_CALLBACK_URL": "https://oauth.test.example/oauth/callback", + "IRONCLAW_OAUTH_EXCHANGE_URL": mock_llm_server, + "AGENT_MAX_TOOL_ITERATIONS": "2", + } + _forward_coverage_env(env) + + proc = await asyncio.create_subprocess_exec( + ironclaw_binary, "--no-onboard", + stdin=asyncio.subprocess.DEVNULL, + stdout=asyncio.subprocess.PIPE, + stderr=asyncio.subprocess.PIPE, + env=env, + ) + startup_kill_attempted = False + base_url = f"http://127.0.0.1:{gateway_port}" + try: + await wait_for_ready(f"{base_url}/api/health", timeout=60) + yield base_url + except TimeoutError: + if proc.returncode is None: + startup_kill_attempted = True + await _stop_process(proc, timeout=2) + returncode = proc.returncode + stderr_bytes = b"" + if proc.stderr: + try: + stderr_bytes = await asyncio.wait_for(proc.stderr.read(8192), timeout=2) + except asyncio.TimeoutError: + pass + stderr_text = stderr_bytes.decode("utf-8", errors="replace") + pytest.fail( + f"loop-limited ironclaw server failed to start on port {gateway_port} " + f"(returncode={returncode}).\nstderr:\n{stderr_text}" + ) + finally: + if proc.returncode is None: + if startup_kill_attempted: + await _stop_process(proc, timeout=2) + else: + await _stop_process(proc, sig=signal.SIGINT, timeout=10) + if proc.returncode is None: + await _stop_process(proc, timeout=2) + finally: + for sock in reserved: + if sock.fileno() != -1: + sock.close() + db_tmpdir.cleanup() + home_tmpdir.cleanup() + + +@pytest.fixture(scope="session") +async def length_preserving_server( + ironclaw_binary, + mock_llm_server, + wasm_tools_dir, +): + """Start an isolated ironclaw instance using the NearAI provider path.""" + reserved = _reserve_loopback_sockets(2) + db_tmpdir = tempfile.TemporaryDirectory(prefix="ironclaw-e2e-length-db-") + home_tmpdir = tempfile.TemporaryDirectory(prefix="ironclaw-e2e-length-home-") + + try: + gateway_port = reserved[0].getsockname()[1] + http_port = reserved[1].getsockname()[1] + for sock in reserved: + if sock.fileno() != -1: + sock.close() + + env = { + "PATH": os.environ.get("PATH", "/usr/bin:/bin"), + "HOME": home_tmpdir.name, + "IRONCLAW_BASE_DIR": os.path.join(home_tmpdir.name, ".ironclaw"), + "RUST_LOG": "ironclaw=info", + "RUST_BACKTRACE": "1", + "IRONCLAW_OWNER_ID": OWNER_SCOPE_ID, + "GATEWAY_ENABLED": "true", + "GATEWAY_HOST": "127.0.0.1", + "GATEWAY_PORT": str(gateway_port), + "GATEWAY_AUTH_TOKEN": AUTH_TOKEN, + "HTTP_HOST": "127.0.0.1", + "HTTP_PORT": str(http_port), + "HTTP_WEBHOOK_SECRET": HTTP_WEBHOOK_SECRET, + "CLI_ENABLED": "false", + "LLM_BACKEND": "nearai", + "NEARAI_BASE_URL": mock_llm_server, + "NEARAI_MODEL": "mock-model", + "NEARAI_API_KEY": "mock-nearai-key", + "DATABASE_BACKEND": "libsql", + "LIBSQL_PATH": os.path.join(db_tmpdir.name, "length-preserving.db"), + "SANDBOX_ENABLED": "false", + "SKILLS_ENABLED": "true", + "ROUTINES_ENABLED": "true", + "HEARTBEAT_ENABLED": "false", + "EMBEDDING_ENABLED": "false", + "WASM_ENABLED": "true", + "WASM_TOOLS_DIR": wasm_tools_dir, + "WASM_CHANNELS_DIR": _WASM_CHANNELS_TMPDIR.name, + "ONBOARD_COMPLETED": "true", + "IRONCLAW_OAUTH_CALLBACK_URL": "https://oauth.test.example/oauth/callback", + "IRONCLAW_OAUTH_EXCHANGE_URL": mock_llm_server, + } + _forward_coverage_env(env) + + proc = await asyncio.create_subprocess_exec( + ironclaw_binary, "--no-onboard", + stdin=asyncio.subprocess.DEVNULL, + stdout=asyncio.subprocess.PIPE, + stderr=asyncio.subprocess.PIPE, + env=env, + ) + startup_kill_attempted = False + base_url = f"http://127.0.0.1:{gateway_port}" + try: + await wait_for_ready(f"{base_url}/api/health", timeout=60) + yield base_url + except TimeoutError: + if proc.returncode is None: + startup_kill_attempted = True + await _stop_process(proc, timeout=2) + returncode = proc.returncode + stderr_bytes = b"" + if proc.stderr: + try: + stderr_bytes = await asyncio.wait_for(proc.stderr.read(8192), timeout=2) + except asyncio.TimeoutError: + pass + stderr_text = stderr_bytes.decode("utf-8", errors="replace") + pytest.fail( + f"length-preserving ironclaw server failed to start on port {gateway_port} " + f"(returncode={returncode}).\nstderr:\n{stderr_text}" + ) + finally: + if proc.returncode is None: + if startup_kill_attempted: + await _stop_process(proc, timeout=2) + else: + await _stop_process(proc, sig=signal.SIGINT, timeout=10) + if proc.returncode is None: + await _stop_process(proc, timeout=2) + finally: + for sock in reserved: + if sock.fileno() != -1: + sock.close() + db_tmpdir.cleanup() + home_tmpdir.cleanup() + + @pytest.fixture(scope="session") async def extension_cleanup_server( ironclaw_binary, @@ -675,3 +869,25 @@ async def page(ironclaw_server, browser): await pg.wait_for_selector("#auth-screen", state="hidden", timeout=15000) yield pg await context.close() + + +@pytest.fixture +async def loop_limited_page(loop_limited_server, browser): + """Fresh Playwright page bound to the low-iteration gateway fixture.""" + context = await browser.new_context(viewport={"width": 1280, "height": 720}) + pg = await context.new_page() + await pg.goto(f"{loop_limited_server}/?token={AUTH_TOKEN}") + await pg.wait_for_selector("#auth-screen", state="hidden", timeout=15000) + yield pg + await context.close() + + +@pytest.fixture +async def length_preserving_page(length_preserving_server, browser): + """Fresh Playwright page bound to the length-preserving gateway fixture.""" + context = await browser.new_context(viewport={"width": 1280, "height": 720}) + pg = await context.new_page() + await pg.goto(f"{length_preserving_server}/?token={AUTH_TOKEN}") + await pg.wait_for_selector("#auth-screen", state="hidden", timeout=15000) + yield pg + await context.close() diff --git a/tests/e2e/helpers.py b/tests/e2e/helpers.py index 4cb7afebc98..e5eaab780ac 100644 --- a/tests/e2e/helpers.py +++ b/tests/e2e/helpers.py @@ -26,6 +26,7 @@ "chat_messages": "#chat-messages", "message_user": "#chat-messages .message.user", "message_assistant": "#chat-messages .message.assistant", + "message_system": "#chat-messages .message.system", # Skills "skill_search_input": "#skill-search-input", "skill_search_results": "#skill-search-results", @@ -182,6 +183,75 @@ async def api_post(base_url: str, path: str, **kwargs) -> httpx.Response: ) +async def send_chat_and_wait_for_terminal_message( + page, + message: str, + *, + timeout: int = 30000, +) -> dict[str, str]: + """Send a chat message and wait for the next terminal visible outcome. + + Returns a dict with: + - ``role``: ``assistant`` or ``system`` + - ``text``: rendered text of the newest terminal message + """ + chat_input = page.locator(SEL["chat_input"]) + await chat_input.wait_for(state="visible", timeout=5000) + + assistant_sel = SEL["message_assistant"] + system_sel = SEL["message_system"] + before_assistant = await page.locator(assistant_sel).count() + before_system = await page.locator(system_sel).count() + + await chat_input.fill(message) + await chat_input.press("Enter") + + handle = await page.wait_for_function( + """({ + assistantSelector, + systemSelector, + chatInputSelector, + assistantCount, + systemCount, + }) => { + const input = document.querySelector(chatInputSelector); + const systems = document.querySelectorAll(systemSelector); + if (systems.length > systemCount) { + const last = systems[systems.length - 1]; + const content = last.querySelector('.message-content'); + return { + role: 'system', + text: ((content && content.innerText) || last.innerText || '').trim(), + }; + } + + const assistants = document.querySelectorAll(assistantSelector); + if (assistants.length > assistantCount && input && !input.disabled) { + const last = assistants[assistants.length - 1]; + const content = last.querySelector('.message-content'); + const text = ((content && content.innerText) || last.innerText || '').trim(); + if (text.length > 0 && !last.hasAttribute('data-streaming')) { + return { + role: 'assistant', + text, + }; + } + } + + return null; + }""", + arg={ + "assistantSelector": assistant_sel, + "systemSelector": system_sel, + "chatInputSelector": SEL["chat_input"], + "assistantCount": before_assistant, + "systemCount": before_system, + }, + timeout=timeout, + ) + return await handle.json_value() + + def signed_http_webhook_headers(body: bytes) -> dict[str, str]: """Return headers for the owner-scoped HTTP webhook channel.""" digest = hmac.new( diff --git a/tests/e2e/mock_llm.py b/tests/e2e/mock_llm.py index 4119a1ab486..661e7e557a9 100644 --- a/tests/e2e/mock_llm.py +++ b/tests/e2e/mock_llm.py @@ -14,6 +14,7 @@ from aiohttp import web CANNED_RESPONSES = [ + (re.compile(r"empty routine response", re.IGNORECASE), ""), (re.compile(r"hello|hi|hey", re.IGNORECASE), "Hello! How can I help you today?"), (re.compile(r"2\s*\+\s*2|two plus two", re.IGNORECASE), "The answer is 4."), (re.compile(r"skill|install", re.IGNORECASE), "I can help you with skills management."), @@ -23,8 +24,21 @@ ] DEFAULT_RESPONSE = "I understand your request." +TOOL_FAILURE_TRIGGER = re.compile(r"issue 1780 tool failure", re.IGNORECASE) +TRUNCATED_TOOL_CALL_TRIGGER = re.compile( + r"issue 1780 truncated tool call", + re.IGNORECASE, +) +EMPTY_REPLY_TRIGGER = re.compile(r"issue 1780 empty reply", re.IGNORECASE) +LOOP_FOREVER_TRIGGER = re.compile(r"issue 1780 loop forever", re.IGNORECASE) + TOOL_CALL_PATTERNS = [ (re.compile(r"echo (.+)", re.IGNORECASE), "echo", lambda m: {"message": m.group(1)}), + ( + re.compile(r"loop until cap", re.IGNORECASE), + "echo", + lambda _: {"message": "loop-until-cap"}, + ), ( re.compile(r"make approval post (?P