mirror of
https://github.com/outbackdingo/optimclaw.git
synced 2026-08-26 15:40:18 +00:00
Merge remote-tracking branch 'origin/staging' into feat/multi-tenant-isolation-phases-2-4
This commit is contained in:
@@ -1091,10 +1091,11 @@ impl Agent {
|
||||
} else {
|
||||
drop(sess);
|
||||
self.session_manager
|
||||
.resolve_thread(
|
||||
.resolve_thread_with_parsed_uuid(
|
||||
&message.user_id,
|
||||
&message.channel,
|
||||
message.conversation_scope(),
|
||||
approval_thread_uuid,
|
||||
)
|
||||
.await
|
||||
}
|
||||
@@ -1175,9 +1176,9 @@ impl Agent {
|
||||
&& let Submission::UserInput { ref content } = submission
|
||||
&& let Some(engine) = self.routine_engine().await
|
||||
{
|
||||
let fired = engine
|
||||
.check_event_triggers(&message.user_id, &message.channel, content)
|
||||
.await;
|
||||
// Use post-hook content so that BeforeInbound hooks that rewrite
|
||||
// input are respected by event trigger matching.
|
||||
let fired = engine.check_event_triggers(message, content).await;
|
||||
if fired > 0 {
|
||||
tracing::debug!(
|
||||
channel = %message.channel,
|
||||
|
||||
+186
-19
@@ -24,7 +24,7 @@ use crate::agent::Scheduler;
|
||||
use crate::agent::routine::{
|
||||
NotifyConfig, Routine, RoutineAction, RoutineRun, RunStatus, Trigger, next_cron_fire,
|
||||
};
|
||||
use crate::channels::OutgoingResponse;
|
||||
use crate::channels::{IncomingMessage, OutgoingResponse};
|
||||
use crate::config::RoutineConfig;
|
||||
use crate::context::{JobContext, JobState};
|
||||
use crate::db::Database;
|
||||
@@ -56,6 +56,40 @@ pub enum SandboxReadiness {
|
||||
DockerUnavailable,
|
||||
}
|
||||
|
||||
/// Check whether an event-triggered routine's user/channel filters match an
|
||||
/// incoming message.
|
||||
///
|
||||
/// Returns `true` if:
|
||||
/// - The routine has an `Event` trigger (non-Event routines always return `false`)
|
||||
/// - The routine's `user_id` matches the message's user scope
|
||||
/// - The routine's channel filter (if any) matches the message channel
|
||||
/// case-insensitively
|
||||
///
|
||||
/// This is a pure function extracted from `check_event_triggers` so the
|
||||
/// filter logic can be unit-tested without async infrastructure.
|
||||
pub(crate) fn routine_matches_message(routine: &Routine, message: &IncomingMessage) -> bool {
|
||||
// Only Event-triggered routines can match incoming messages.
|
||||
if !matches!(routine.trigger, Trigger::Event { .. }) {
|
||||
return false;
|
||||
}
|
||||
|
||||
// User ownership filter — only fire routines scoped to this user.
|
||||
if routine.user_id != message.user_id {
|
||||
return false;
|
||||
}
|
||||
|
||||
// Channel filter (case-insensitive, matching emit_system_event behavior)
|
||||
if let Trigger::Event {
|
||||
channel: Some(ch), ..
|
||||
} = &routine.trigger
|
||||
&& !ch.eq_ignore_ascii_case(&message.channel)
|
||||
{
|
||||
return false;
|
||||
}
|
||||
|
||||
true
|
||||
}
|
||||
|
||||
/// The routine execution engine.
|
||||
pub struct RoutineEngine {
|
||||
config: RoutineConfig,
|
||||
@@ -167,10 +201,7 @@ impl RoutineEngine {
|
||||
}
|
||||
|
||||
/// Check incoming message against event triggers. Returns number of routines fired.
|
||||
///
|
||||
/// Accepts only the three fields needed for matching (user scope, channel,
|
||||
/// message content) so callers never need to clone a full `IncomingMessage`.
|
||||
pub async fn check_event_triggers(&self, user_id: &str, channel: &str, content: &str) -> usize {
|
||||
pub async fn check_event_triggers(&self, message: &IncomingMessage, content: &str) -> usize {
|
||||
let cache = self.event_cache.read().await;
|
||||
|
||||
// Early return if there are no message matchers at all.
|
||||
@@ -208,16 +239,24 @@ impl RoutineEngine {
|
||||
EventMatcher::System { .. } => continue,
|
||||
};
|
||||
|
||||
if routine.user_id != user_id {
|
||||
continue;
|
||||
}
|
||||
|
||||
// Channel filter
|
||||
if let Trigger::Event {
|
||||
channel: Some(ch), ..
|
||||
} = &routine.trigger
|
||||
&& ch != channel
|
||||
{
|
||||
// User ownership + channel filter (extracted for testability).
|
||||
if !routine_matches_message(routine, message) {
|
||||
// User mismatch is expected for multi-user setups — keep at
|
||||
// trace to avoid one log per routine per inbound message.
|
||||
if routine.user_id != message.user_id {
|
||||
tracing::trace!(
|
||||
routine = %routine.name,
|
||||
routine_user = %routine.user_id,
|
||||
message_user = %message.user_id,
|
||||
"Skipped: user scope mismatch"
|
||||
);
|
||||
} else {
|
||||
tracing::debug!(
|
||||
routine = %routine.name,
|
||||
channel = %message.channel,
|
||||
"Skipped: channel mismatch"
|
||||
);
|
||||
}
|
||||
continue;
|
||||
}
|
||||
|
||||
@@ -228,14 +267,14 @@ impl RoutineEngine {
|
||||
|
||||
// Cooldown check
|
||||
if !self.check_cooldown(routine) {
|
||||
tracing::trace!(routine = %routine.name, "Skipped: cooldown active");
|
||||
tracing::debug!(routine = %routine.name, "Skipped: cooldown active");
|
||||
continue;
|
||||
}
|
||||
|
||||
// Concurrent run check (using batch-loaded counts)
|
||||
let running_count = concurrent_counts.get(&routine.id).copied().unwrap_or(0);
|
||||
if running_count >= routine.guardrails.max_concurrent as i64 {
|
||||
tracing::trace!(routine = %routine.name, "Skipped: max concurrent reached");
|
||||
tracing::debug!(routine = %routine.name, "Skipped: max concurrent reached");
|
||||
continue;
|
||||
}
|
||||
|
||||
@@ -1790,6 +1829,13 @@ pub fn spawn_cron_ticker(
|
||||
engine.check_cron_triggers().await;
|
||||
|
||||
let mut ticker = tokio::time::interval(interval);
|
||||
ticker.set_missed_tick_behavior(tokio::time::MissedTickBehavior::Skip);
|
||||
// Periodic event cache refresh so web/CLI mutations are picked up
|
||||
// without requiring tool-path code to call refresh_event_cache().
|
||||
// Uses wall-clock elapsed time so the refresh cadence is stable
|
||||
// regardless of the cron tick interval configuration.
|
||||
let refresh_interval = Duration::from_secs(60);
|
||||
let mut last_refresh = tokio::time::Instant::now();
|
||||
|
||||
loop {
|
||||
ticker.tick().await;
|
||||
@@ -1797,7 +1843,11 @@ pub fn spawn_cron_ticker(
|
||||
// never races with FullJobWatcher instances from this process.
|
||||
engine.sync_dispatched_runs().await;
|
||||
engine.check_cron_triggers().await;
|
||||
engine.sync_dispatched_runs().await;
|
||||
|
||||
if last_refresh.elapsed() >= refresh_interval {
|
||||
engine.refresh_event_cache().await;
|
||||
last_refresh = tokio::time::Instant::now();
|
||||
}
|
||||
}
|
||||
})
|
||||
}
|
||||
@@ -1863,7 +1913,13 @@ fn strip_html_tags(s: &str) -> String {
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use crate::agent::routine::{NotifyConfig, RunStatus};
|
||||
use chrono::Utc;
|
||||
use uuid::Uuid;
|
||||
|
||||
use crate::agent::routine::{
|
||||
NotifyConfig, Routine, RoutineAction, RoutineGuardrails, RunStatus, Trigger,
|
||||
};
|
||||
use crate::channels::IncomingMessage;
|
||||
use crate::config::RoutineConfig;
|
||||
|
||||
#[test]
|
||||
@@ -2061,6 +2117,117 @@ mod tests {
|
||||
}
|
||||
}
|
||||
|
||||
/// Helper to build a test routine with the given user_id and trigger.
|
||||
fn make_routine(user_id: &str, trigger: Trigger) -> Routine {
|
||||
Routine {
|
||||
id: Uuid::new_v4(),
|
||||
name: "test".to_string(),
|
||||
description: String::new(),
|
||||
user_id: user_id.to_string(),
|
||||
enabled: true,
|
||||
trigger,
|
||||
action: RoutineAction::Lightweight {
|
||||
prompt: String::new(),
|
||||
context_paths: vec![],
|
||||
max_tokens: 1000,
|
||||
use_tools: false,
|
||||
max_tool_rounds: 0,
|
||||
},
|
||||
guardrails: RoutineGuardrails::default(),
|
||||
notify: Default::default(),
|
||||
last_run_at: None,
|
||||
next_fire_at: None,
|
||||
run_count: 0,
|
||||
consecutive_failures: 0,
|
||||
state: serde_json::Value::Null,
|
||||
created_at: Utc::now(),
|
||||
updated_at: Utc::now(),
|
||||
}
|
||||
}
|
||||
|
||||
/// Helper to build a test IncomingMessage.
|
||||
fn make_message(user_id: &str, channel: &str, content: &str) -> IncomingMessage {
|
||||
IncomingMessage {
|
||||
id: Uuid::new_v4(),
|
||||
channel: channel.to_string(),
|
||||
user_id: user_id.to_string(),
|
||||
owner_id: user_id.to_string(),
|
||||
sender_id: user_id.to_string(),
|
||||
user_name: None,
|
||||
content: content.to_string(),
|
||||
thread_id: None,
|
||||
conversation_scope_id: None,
|
||||
received_at: Utc::now(),
|
||||
metadata: serde_json::Value::Null,
|
||||
timezone: None,
|
||||
attachments: vec![],
|
||||
is_internal: false,
|
||||
}
|
||||
}
|
||||
|
||||
/// Regression test for issue #1051: event triggers used case-sensitive
|
||||
/// channel comparison, so "Telegram" != "telegram" caused silent mismatch.
|
||||
/// Tests the actual `routine_matches_message` function used in `check_event_triggers`.
|
||||
#[test]
|
||||
fn test_channel_filter_is_case_insensitive() {
|
||||
let routine = make_routine(
|
||||
"user1",
|
||||
Trigger::Event {
|
||||
pattern: ".*".to_string(),
|
||||
channel: Some("Telegram".to_string()),
|
||||
},
|
||||
);
|
||||
let msg = make_message("user1", "telegram", "hello");
|
||||
|
||||
// Case-insensitive channel match must succeed
|
||||
assert!(super::routine_matches_message(&routine, &msg));
|
||||
|
||||
// Exact case must also work
|
||||
let msg_exact = make_message("user1", "Telegram", "hello");
|
||||
assert!(super::routine_matches_message(&routine, &msg_exact));
|
||||
|
||||
// Different channel must not match
|
||||
let msg_wrong = make_message("user1", "discord", "hello");
|
||||
assert!(!super::routine_matches_message(&routine, &msg_wrong));
|
||||
}
|
||||
|
||||
/// Regression test for issue #1051: event triggers did not filter by
|
||||
/// user_id, so routines from user A could fire on messages from user B.
|
||||
/// Tests the actual `routine_matches_message` function used in `check_event_triggers`.
|
||||
#[test]
|
||||
fn test_event_trigger_requires_user_match() {
|
||||
let routine = make_routine(
|
||||
"alice",
|
||||
Trigger::Event {
|
||||
pattern: ".*".to_string(),
|
||||
channel: None,
|
||||
},
|
||||
);
|
||||
|
||||
// Different user must not match
|
||||
let msg_bob = make_message("bob", "telegram", "hello");
|
||||
assert!(!super::routine_matches_message(&routine, &msg_bob));
|
||||
|
||||
// Same user must match
|
||||
let msg_alice = make_message("alice", "telegram", "hello");
|
||||
assert!(super::routine_matches_message(&routine, &msg_alice));
|
||||
}
|
||||
|
||||
/// When no channel filter is set, any channel should match (given user matches).
|
||||
#[test]
|
||||
fn test_no_channel_filter_matches_any_channel() {
|
||||
let routine = make_routine(
|
||||
"user1",
|
||||
Trigger::Event {
|
||||
pattern: ".*".to_string(),
|
||||
channel: None,
|
||||
},
|
||||
);
|
||||
|
||||
let msg = make_message("user1", "whatever_channel", "hello");
|
||||
assert!(super::routine_matches_message(&routine, &msg));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_routine_tool_denylist_blocks_self_management_tools() {
|
||||
let denylisted = vec
|
||||
/// with `parsed_uuid: None`.
|
||||
pub async fn resolve_thread(
|
||||
&self,
|
||||
user_id: &str,
|
||||
channel: &str,
|
||||
external_thread_id: Option<&str>,
|
||||
) -> (Arc<Mutex<Session>>, Uuid) {
|
||||
self.resolve_thread_with_parsed_uuid(user_id, channel, external_thread_id, None)
|
||||
.await
|
||||
}
|
||||
|
||||
/// Like [`resolve_thread`](Self::resolve_thread), but accepts a pre-parsed
|
||||
/// UUID to skip redundant parsing when the caller has already validated
|
||||
/// the external thread ID as a UUID (e.g. the approval routing path).
|
||||
///
|
||||
/// Uses a single read-lock acquisition for both the key lookup and the UUID
|
||||
/// adoption check to reduce contention under concurrent approval load.
|
||||
pub async fn resolve_thread_with_parsed_uuid(
|
||||
&self,
|
||||
user_id: &str,
|
||||
channel: &str,
|
||||
external_thread_id: Option<&str>,
|
||||
parsed_uuid: Option<Uuid>,
|
||||
) -> (Arc<Mutex<Session>>, Uuid) {
|
||||
let session = self.get_or_create_session(user_id).await;
|
||||
|
||||
@@ -116,51 +135,65 @@ impl SessionManager {
|
||||
external_thread_id: external_thread_id.map(String::from),
|
||||
};
|
||||
|
||||
// Check if we have a mapping
|
||||
{
|
||||
// Use pre-parsed UUID if available, otherwise parse from string.
|
||||
let ext_uuid = parsed_uuid
|
||||
.or_else(|| external_thread_id.and_then(|ext_tid| Uuid::parse_str(ext_tid).ok()));
|
||||
|
||||
// Validate that parsed_uuid (if provided) is consistent with external_thread_id.
|
||||
#[cfg(debug_assertions)]
|
||||
if let (Some(parsed), Some(ext_tid)) = (&parsed_uuid, external_thread_id) {
|
||||
debug_assert_eq!(
|
||||
Uuid::parse_str(ext_tid).ok().as_ref(),
|
||||
Some(parsed),
|
||||
"parsed_uuid must be the parsed form of external_thread_id"
|
||||
);
|
||||
}
|
||||
|
||||
// Single read lock for both the key lookup and UUID adoption check
|
||||
let adoptable_uuid = {
|
||||
let thread_map = self.thread_map.read().await;
|
||||
|
||||
// Fast path: exact key match
|
||||
if let Some(&thread_id) = thread_map.get(&key) {
|
||||
// Verify thread still exists in session
|
||||
let sess = session.lock().await;
|
||||
if sess.threads.contains_key(&thread_id) {
|
||||
return (Arc::clone(&session), thread_id);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// Check if external_thread_id is itself a known thread UUID that
|
||||
// exists in the session but was never registered in the thread_map
|
||||
// (e.g. created by chat_new_thread_handler or hydrated from DB).
|
||||
// We only adopt it if no thread_map entry maps to this UUID —
|
||||
// otherwise it belongs to a different channel scope.
|
||||
if let Some(ext_tid) = external_thread_id
|
||||
&& let Ok(ext_uuid) = Uuid::parse_str(ext_tid)
|
||||
{
|
||||
let thread_map = self.thread_map.read().await;
|
||||
let mapped_elsewhere = thread_map.values().any(|&v| v == ext_uuid);
|
||||
drop(thread_map);
|
||||
// UUID adoption check (still under the same read lock).
|
||||
// If external_thread_id is a valid UUID not mapped elsewhere,
|
||||
// it may be a thread created by chat_new_thread_handler or
|
||||
// hydrated from DB that we can adopt.
|
||||
// Only attempt adoption when external_thread_id is Some, preserving
|
||||
// the invariant that None external_thread_id never triggers adoption.
|
||||
if external_thread_id.is_some() {
|
||||
ext_uuid.filter(|&uuid| !thread_map.values().any(|&v| v == uuid))
|
||||
} else {
|
||||
None
|
||||
}
|
||||
}; // Single read lock dropped here
|
||||
|
||||
if !mapped_elsewhere {
|
||||
let sess = session.lock().await;
|
||||
if sess.threads.contains_key(&ext_uuid) {
|
||||
drop(sess);
|
||||
// If we found an adoptable UUID, verify it exists in session and acquire write lock
|
||||
if let Some(ext_uuid) = adoptable_uuid {
|
||||
let sess = session.lock().await;
|
||||
if sess.threads.contains_key(&ext_uuid) {
|
||||
drop(sess);
|
||||
|
||||
let mut thread_map = self.thread_map.write().await;
|
||||
// Re-check after acquiring write lock to prevent race condition
|
||||
// where another task mapped this UUID between our read and write.
|
||||
if !thread_map.values().any(|&v| v == ext_uuid) {
|
||||
thread_map.insert(key, ext_uuid);
|
||||
drop(thread_map);
|
||||
// Ensure undo manager exists
|
||||
let mut undo_managers = self.undo_managers.write().await;
|
||||
undo_managers
|
||||
.entry(ext_uuid)
|
||||
.or_insert_with(|| Arc::new(Mutex::new(UndoManager::new())));
|
||||
return (session, ext_uuid);
|
||||
}
|
||||
// If it was mapped elsewhere while we were unlocked, fall through
|
||||
// to create a new thread, preserving channel isolation.
|
||||
let mut thread_map = self.thread_map.write().await;
|
||||
// Re-check after acquiring write lock to prevent race condition
|
||||
// where another task mapped this UUID between our read and write.
|
||||
if !thread_map.values().any(|&v| v == ext_uuid) {
|
||||
thread_map.insert(key, ext_uuid);
|
||||
drop(thread_map);
|
||||
// Ensure undo manager exists
|
||||
let mut undo_managers = self.undo_managers.write().await;
|
||||
undo_managers
|
||||
.entry(ext_uuid)
|
||||
.or_insert_with(|| Arc::new(Mutex::new(UndoManager::new())));
|
||||
return (session, ext_uuid);
|
||||
}
|
||||
// If mapped elsewhere while unlocked, fall through to create new thread
|
||||
}
|
||||
}
|
||||
|
||||
@@ -909,6 +942,44 @@ mod tests {
|
||||
}
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn test_resolve_thread_consolidates_read_path() {
|
||||
// Verify that resolve_thread still correctly handles:
|
||||
// 1. Fast path: key exists in thread_map
|
||||
// 2. UUID adoption: external_thread_id is a UUID in session but not in map
|
||||
// 3. New thread: neither path matches
|
||||
use crate::agent::session::Thread;
|
||||
|
||||
let manager = SessionManager::new();
|
||||
|
||||
// Case 1: Normal resolution creates thread and maps it
|
||||
let (session1, tid1) = manager
|
||||
.resolve_thread("user1", "chan1", Some("ext-1"))
|
||||
.await;
|
||||
// Resolving again with same key should return same thread (fast path)
|
||||
let (_, tid1_again) = manager
|
||||
.resolve_thread("user1", "chan1", Some("ext-1"))
|
||||
.await;
|
||||
assert_eq!(tid1, tid1_again);
|
||||
|
||||
// Case 2: UUID adoption - insert a thread directly into session
|
||||
let adopted_id = Uuid::new_v4();
|
||||
{
|
||||
let mut sess = session1.lock().await;
|
||||
let thread = Thread::with_id(adopted_id, sess.id);
|
||||
sess.threads.insert(adopted_id, thread);
|
||||
}
|
||||
// Resolve with the UUID as external_thread_id -- should adopt it
|
||||
let (_, resolved) = manager
|
||||
.resolve_thread("user1", "chan1", Some(&adopted_id.to_string()))
|
||||
.await;
|
||||
assert_eq!(resolved, adopted_id);
|
||||
|
||||
// Case 3: Different channel gets different thread
|
||||
let (_, tid2) = manager.resolve_thread("user1", "chan2", None).await;
|
||||
assert_ne!(tid1, tid2);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn test_resolve_thread_finds_existing_session_thread_by_uuid() {
|
||||
use crate::agent::session::{Session, Thread};
|
||||
@@ -947,4 +1018,88 @@ mod tests {
|
||||
"should have exactly 1 thread, not a duplicate"
|
||||
);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn test_resolve_thread_with_pre_parsed_uuid_adopts_thread() {
|
||||
use crate::agent::session::Thread;
|
||||
|
||||
let manager = SessionManager::new();
|
||||
let (session, _) = manager.resolve_thread("user1", "chan1", None).await;
|
||||
|
||||
// Manually insert a thread with a known UUID
|
||||
let known_id = Uuid::new_v4();
|
||||
{
|
||||
let mut sess = session.lock().await;
|
||||
let thread = Thread::with_id(known_id, sess.id);
|
||||
sess.threads.insert(known_id, thread);
|
||||
}
|
||||
|
||||
// Resolve with pre-parsed UUID -- should adopt it without re-parsing
|
||||
let (_, resolved) = manager
|
||||
.resolve_thread_with_parsed_uuid(
|
||||
"user1",
|
||||
"chan1",
|
||||
Some(&known_id.to_string()),
|
||||
Some(known_id),
|
||||
)
|
||||
.await;
|
||||
assert_eq!(resolved, known_id);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn test_resolve_thread_with_parsed_uuid_none_delegates_to_parse() {
|
||||
use crate::agent::session::Thread;
|
||||
|
||||
let manager = SessionManager::new();
|
||||
let (session, _) = manager.resolve_thread("user2", "chan2", None).await;
|
||||
|
||||
// Insert a thread with a known UUID
|
||||
let known_id = Uuid::new_v4();
|
||||
{
|
||||
let mut sess = session.lock().await;
|
||||
let thread = Thread::with_id(known_id, sess.id);
|
||||
sess.threads.insert(known_id, thread);
|
||||
}
|
||||
|
||||
// Resolve with parsed_uuid=None but a valid UUID string -- should
|
||||
// fall back to parsing the string and still adopt the thread
|
||||
let (_, resolved) = manager
|
||||
.resolve_thread_with_parsed_uuid("user2", "chan2", Some(&known_id.to_string()), None)
|
||||
.await;
|
||||
assert_eq!(resolved, known_id);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn test_resolve_thread_with_none_external_thread_id_does_not_adopt() {
|
||||
use crate::agent::session::Thread;
|
||||
|
||||
let manager = SessionManager::new();
|
||||
let (session, default_tid) = manager.resolve_thread("user3", "chan3", None).await;
|
||||
|
||||
// Manually insert a thread with a known UUID (simulating a thread
|
||||
// created by chat_new_thread_handler)
|
||||
let known_id = Uuid::new_v4();
|
||||
{
|
||||
let mut sess = session.lock().await;
|
||||
let thread = Thread::with_id(known_id, sess.id);
|
||||
sess.threads.insert(known_id, thread);
|
||||
}
|
||||
|
||||
// Resolve with external_thread_id=None but parsed_uuid=Some.
|
||||
// This should NOT adopt the UUID — the old code prevented adoption
|
||||
// when external_thread_id was None, and we preserve that invariant.
|
||||
let (_, resolved) = manager
|
||||
.resolve_thread_with_parsed_uuid("user3", "chan3", None, Some(known_id))
|
||||
.await;
|
||||
|
||||
// Should return the existing default thread, not the injected UUID
|
||||
assert_eq!(
|
||||
resolved, default_tid,
|
||||
"should return existing default thread when external_thread_id is None"
|
||||
);
|
||||
assert_ne!(
|
||||
resolved, known_id,
|
||||
"should NOT adopt UUID when external_thread_id is None"
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
@@ -114,7 +114,7 @@ pub async fn routines_detail_handler(
|
||||
trigger_type: run.trigger_type.clone(),
|
||||
started_at: run.started_at.to_rfc3339(),
|
||||
completed_at: run.completed_at.map(|dt| dt.to_rfc3339()),
|
||||
status: format!("{:?}", run.status),
|
||||
status: run.status.to_string(),
|
||||
result_summary: run.result_summary.clone(),
|
||||
tokens_used: run.tokens_used,
|
||||
job_id: run.job_id,
|
||||
@@ -324,7 +324,7 @@ pub async fn routines_runs_handler(
|
||||
trigger_type: run.trigger_type.clone(),
|
||||
started_at: run.started_at.to_rfc3339(),
|
||||
completed_at: run.completed_at.map(|dt| dt.to_rfc3339()),
|
||||
status: format!("{:?}", run.status),
|
||||
status: run.status.to_string(),
|
||||
result_summary: run.result_summary.clone(),
|
||||
tokens_used: run.tokens_used,
|
||||
job_id: run.job_id,
|
||||
|
||||
@@ -2572,7 +2572,7 @@ async fn routines_runs_handler(
|
||||
trigger_type: run.trigger_type.clone(),
|
||||
started_at: run.started_at.to_rfc3339(),
|
||||
completed_at: run.completed_at.map(|dt| dt.to_rfc3339()),
|
||||
status: format!("{:?}", run.status),
|
||||
status: run.status.to_string(),
|
||||
result_summary: run.result_summary.clone(),
|
||||
tokens_used: run.tokens_used,
|
||||
job_id: run.job_id,
|
||||
|
||||
@@ -4265,9 +4265,9 @@ function renderRoutineDetail(routine) {
|
||||
+ '<th>Trigger</th><th>Started</th><th>Completed</th><th>Status</th><th>Summary</th><th>Tokens</th>'
|
||||
+ '</tr></thead><tbody>';
|
||||
for (const run of routine.recent_runs) {
|
||||
const runStatusClass = run.status === 'Ok' ? 'completed'
|
||||
: run.status === 'Failed' ? 'failed'
|
||||
: run.status === 'Attention' ? 'stuck'
|
||||
const runStatusClass = run.status === 'ok' ? 'completed'
|
||||
: run.status === 'failed' ? 'failed'
|
||||
: run.status === 'attention' ? 'stuck'
|
||||
: 'in_progress';
|
||||
html += '<tr>'
|
||||
+ '<td>' + escapeHtml(run.trigger_type) + '</td>'
|
||||
|
||||
+19
-8
@@ -10,7 +10,7 @@ use clap::Subcommand;
|
||||
use uuid::Uuid;
|
||||
|
||||
use crate::agent::routine::{
|
||||
NotifyConfig, Routine, RoutineAction, RoutineGuardrails, Trigger, next_cron_fire,
|
||||
NotifyConfig, Routine, RoutineAction, RoutineGuardrails, RunStatus, Trigger, next_cron_fire,
|
||||
};
|
||||
use crate::db::Database;
|
||||
|
||||
@@ -251,15 +251,26 @@ async fn list(
|
||||
);
|
||||
println!("{}", "-".repeat(130));
|
||||
|
||||
// Fetch last-run status for all routines in a single batch query
|
||||
let routine_ids: Vec<Uuid> = filtered.iter().map(|r| r.id).collect();
|
||||
let last_run_results = db
|
||||
.batch_get_last_run_status(&routine_ids)
|
||||
.await
|
||||
.unwrap_or_default();
|
||||
|
||||
for r in &filtered {
|
||||
let status = if r.enabled {
|
||||
if r.consecutive_failures > 0 {
|
||||
format!("err({})", r.consecutive_failures)
|
||||
} else {
|
||||
"active".to_string()
|
||||
}
|
||||
} else {
|
||||
let last_run_status = last_run_results.get(&r.id).copied();
|
||||
|
||||
let status = if !r.enabled {
|
||||
"disabled".to_string()
|
||||
} else if last_run_status == Some(RunStatus::Running) {
|
||||
"running".to_string()
|
||||
} else if r.consecutive_failures > 0 {
|
||||
format!("err({})", r.consecutive_failures)
|
||||
} else if last_run_status == Some(RunStatus::Attention) {
|
||||
"attention".to_string()
|
||||
} else {
|
||||
"active".to_string()
|
||||
};
|
||||
|
||||
let next_fire = r
|
||||
|
||||
@@ -462,6 +462,56 @@ impl RoutineStore for LibSqlBackend {
|
||||
Ok(counts)
|
||||
}
|
||||
|
||||
async fn batch_get_last_run_status(
|
||||
&self,
|
||||
routine_ids: &[Uuid],
|
||||
) -> Result<HashMap<Uuid, RunStatus>, DatabaseError> {
|
||||
if routine_ids.is_empty() {
|
||||
return Ok(HashMap::new());
|
||||
}
|
||||
|
||||
let conn = self.connect().await?;
|
||||
|
||||
// SQLite doesn't support ANY($1), so we query all latest runs and filter in memory.
|
||||
// Uses a subquery to pick only the most recent run per routine.
|
||||
let mut rows = conn
|
||||
.query(
|
||||
"SELECT routine_id, status FROM routine_runs r1
|
||||
WHERE started_at = (
|
||||
SELECT MAX(started_at) FROM routine_runs r2
|
||||
WHERE r2.routine_id = r1.routine_id
|
||||
)
|
||||
GROUP BY routine_id",
|
||||
params![],
|
||||
)
|
||||
.await
|
||||
.map_err(|e| {
|
||||
DatabaseError::Query(format!("Failed to batch get last run status: {}", e))
|
||||
})?;
|
||||
|
||||
let routine_id_set: HashSet<Uuid> = routine_ids.iter().copied().collect();
|
||||
let mut statuses = HashMap::new();
|
||||
|
||||
while let Some(row) = rows
|
||||
.next()
|
||||
.await
|
||||
.map_err(|e| DatabaseError::Query(e.to_string()))?
|
||||
{
|
||||
let id_str: String = get_text(&row, 0);
|
||||
let id = Uuid::parse_str(&id_str)
|
||||
.map_err(|e| DatabaseError::Query(format!("Invalid routine UUID: {}", e)))?;
|
||||
|
||||
if routine_id_set.contains(&id) {
|
||||
let status_str: String = get_text(&row, 1);
|
||||
if let std::result::Result::Ok(status) = status_str.parse::<RunStatus>() {
|
||||
statuses.insert(id, status);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
Ok(statuses)
|
||||
}
|
||||
|
||||
async fn link_routine_run_to_job(
|
||||
&self,
|
||||
run_id: Uuid,
|
||||
|
||||
@@ -528,6 +528,15 @@ pub trait RoutineStore: Send + Sync {
|
||||
&self,
|
||||
routine_ids: &[Uuid],
|
||||
) -> Result<HashMap<Uuid, i64>, DatabaseError>;
|
||||
|
||||
/// Fetch the last run status for multiple routines in a single query.
|
||||
/// Returns a map from routine_id to its most recent RunStatus.
|
||||
/// Routines with no runs are omitted from the result.
|
||||
async fn batch_get_last_run_status(
|
||||
&self,
|
||||
routine_ids: &[Uuid],
|
||||
) -> Result<HashMap<Uuid, RunStatus>, DatabaseError>;
|
||||
|
||||
async fn link_routine_run_to_job(
|
||||
&self,
|
||||
run_id: Uuid,
|
||||
|
||||
@@ -510,6 +510,14 @@ impl RoutineStore for PgBackend {
|
||||
.await
|
||||
}
|
||||
|
||||
async fn batch_get_last_run_status(
|
||||
&self,
|
||||
routine_ids: &[Uuid],
|
||||
) -> Result<std::collections::HashMap<Uuid, crate::agent::routine::RunStatus>, DatabaseError>
|
||||
{
|
||||
self.store.batch_get_last_run_status(routine_ids).await
|
||||
}
|
||||
|
||||
async fn link_routine_run_to_job(
|
||||
&self,
|
||||
run_id: Uuid,
|
||||
|
||||
+49
-53
@@ -891,24 +891,27 @@ impl ExtensionManager {
|
||||
*self.relay_channel_manager.write().await = Some(channel_manager);
|
||||
}
|
||||
|
||||
/// Check if a channel name corresponds to a relay extension (has stored stream token
|
||||
/// Check if a channel name corresponds to a relay extension (has stored team_id
|
||||
/// or is tracked in the installed relay extensions set).
|
||||
pub async fn is_relay_channel(&self, name: &str, user_id: &str) -> bool {
|
||||
// Check in-memory installed set first (supports no-store mode)
|
||||
if self.installed_relay_extensions.read().await.contains(name) {
|
||||
return true;
|
||||
}
|
||||
// Then check for stored stream token
|
||||
self.secrets
|
||||
.exists(user_id, &format!("relay:{}:stream_token", name))
|
||||
.await
|
||||
.unwrap_or(false)
|
||||
// Check for stored team_id (persisted across restarts by the OAuth callback)
|
||||
if let Some(ref store) = self.store {
|
||||
let key = format!("relay:{}:team_id", name);
|
||||
if let Ok(Some(v)) = store.get_setting(user_id, &key).await {
|
||||
return v.as_str().is_some_and(|s| !s.is_empty());
|
||||
}
|
||||
}
|
||||
false
|
||||
}
|
||||
|
||||
/// Restore persisted relay channels after startup.
|
||||
///
|
||||
/// Loads the persisted active channel list, filters to relay types (those with
|
||||
/// a stored stream token), and activates each via `activate_stored_relay()`.
|
||||
/// a stored team_id setting), and activates each via `activate_stored_relay()`.
|
||||
/// Skips channels that are already active.
|
||||
///
|
||||
/// Call this only after `set_relay_channel_manager()` or `set_channel_runtime()`.
|
||||
@@ -1428,9 +1431,11 @@ impl ExtensionManager {
|
||||
if kind_filter.is_none() || kind_filter == Some(ExtensionKind::ChannelRelay) {
|
||||
let installed = self.installed_relay_extensions.read().await;
|
||||
let active_names = self.active_channel_names.read().await;
|
||||
let errors = self.activation_errors.read().await;
|
||||
for name in installed.iter() {
|
||||
let active = active_names.contains(name);
|
||||
let has_token = self.is_relay_channel(name, user_id).await;
|
||||
let authenticated = self.is_relay_channel(name, user_id).await;
|
||||
let activation_error = errors.get(name).cloned();
|
||||
let registry_entry = self
|
||||
.registry
|
||||
.get_with_kind(name, Some(ExtensionKind::ChannelRelay))
|
||||
@@ -1443,13 +1448,13 @@ impl ExtensionManager {
|
||||
display_name,
|
||||
description,
|
||||
url: None,
|
||||
authenticated: has_token,
|
||||
authenticated,
|
||||
active,
|
||||
tools: Vec::new(),
|
||||
needs_setup: false,
|
||||
has_auth: true,
|
||||
installed: true,
|
||||
activation_error: None,
|
||||
activation_error,
|
||||
version: None,
|
||||
});
|
||||
}
|
||||
@@ -1626,7 +1631,22 @@ impl ExtensionManager {
|
||||
self.persist_active_channels(user_id).await;
|
||||
self.activation_errors.write().await.remove(name);
|
||||
|
||||
// Remove stored stream token
|
||||
// Remove stored team_id setting and clean up secrets
|
||||
if let Some(ref store) = self.store
|
||||
&& let Err(e) = store
|
||||
.delete_setting(user_id, &format!("relay:{}:team_id", name))
|
||||
.await
|
||||
{
|
||||
tracing::warn!(error = %e, name, "Failed to delete relay team_id setting on removal");
|
||||
}
|
||||
if let Err(e) = self
|
||||
.secrets
|
||||
.delete(user_id, &format!("relay:{}:oauth_state", name))
|
||||
.await
|
||||
{
|
||||
tracing::warn!(error = %e, name, "Failed to delete relay oauth_state secret on removal");
|
||||
}
|
||||
// Clean up legacy stream_token secret from pre-webhook installs
|
||||
let _ = self
|
||||
.secrets
|
||||
.delete(user_id, &format!("relay:{}:stream_token", name))
|
||||
@@ -4181,13 +4201,13 @@ impl ExtensionManager {
|
||||
///
|
||||
/// For Slack: initiates OAuth flow (redirect-based).
|
||||
/// For Telegram: accepts a bot token, registers it with channel-relay,
|
||||
/// and stores the returned stream token.
|
||||
/// and stores the team_id setting.
|
||||
async fn auth_channel_relay(
|
||||
&self,
|
||||
name: &str,
|
||||
user_id: &str,
|
||||
) -> Result<AuthResult, ExtensionError> {
|
||||
// Check if already authenticated (stream token exists)
|
||||
// Check if already authenticated (team_id setting exists)
|
||||
if self.is_relay_channel(name, user_id).await {
|
||||
return Ok(AuthResult::authenticated(name, ExtensionKind::ChannelRelay));
|
||||
}
|
||||
@@ -4233,19 +4253,9 @@ impl ExtensionManager {
|
||||
name: &str,
|
||||
user_id: &str,
|
||||
) -> Result<ActivateResult, ExtensionError> {
|
||||
let token_key = format!("relay:{}:stream_token", name);
|
||||
let team_id_key = format!("relay:{}:team_id", name);
|
||||
|
||||
// Check if we have a stream token
|
||||
// Verify auth: stream token must exist (even though we don't use it in this constructor path)
|
||||
let _stream_token = match self.secrets.get_decrypted(user_id, &token_key).await {
|
||||
Ok(secret) => secret.expose().to_string(),
|
||||
Err(_) => {
|
||||
return Err(ExtensionError::AuthRequired);
|
||||
}
|
||||
};
|
||||
|
||||
// Get team_id from settings
|
||||
// Get team_id from settings (stored by the OAuth callback)
|
||||
let team_id = if let Some(ref store) = self.store {
|
||||
store
|
||||
.get_setting(user_id, &team_id_key)
|
||||
@@ -4258,6 +4268,10 @@ impl ExtensionManager {
|
||||
String::new()
|
||||
};
|
||||
|
||||
if team_id.is_empty() {
|
||||
return Err(ExtensionError::AuthRequired);
|
||||
}
|
||||
|
||||
// Use relay config captured at startup
|
||||
let relay_config = self.relay_config()?;
|
||||
|
||||
@@ -4367,11 +4381,11 @@ impl ExtensionManager {
|
||||
return Ok(ExtensionKind::WasmChannel);
|
||||
}
|
||||
|
||||
// Check channel-relay extensions (installed in memory or has stored token)
|
||||
// Check channel-relay extensions (installed in memory or has stored team_id)
|
||||
if self.installed_relay_extensions.read().await.contains(name) {
|
||||
return Ok(ExtensionKind::ChannelRelay);
|
||||
}
|
||||
// Also check if there's a stored stream token (persisted across restarts)
|
||||
// Also check if there's a stored team_id setting (persisted across restarts)
|
||||
if self.is_relay_channel(name, user_id).await {
|
||||
return Ok(ExtensionKind::ChannelRelay);
|
||||
}
|
||||
@@ -4999,11 +5013,7 @@ impl ExtensionManager {
|
||||
names.insert(server.token_secret_name());
|
||||
(names, Vec::new())
|
||||
}
|
||||
ExtensionKind::ChannelRelay => {
|
||||
let mut names = std::collections::HashSet::new();
|
||||
names.insert(format!("relay:{}:stream_token", name));
|
||||
(names, Vec::new())
|
||||
}
|
||||
ExtensionKind::ChannelRelay => (std::collections::HashSet::new(), Vec::new()),
|
||||
};
|
||||
|
||||
let allowed_fields: std::collections::HashSet<String> =
|
||||
@@ -5434,7 +5444,9 @@ impl ExtensionManager {
|
||||
.map_err(|e| ExtensionError::NotInstalled(e.to_string()))?;
|
||||
server.token_secret_name()
|
||||
}
|
||||
ExtensionKind::ChannelRelay => format!("relay:{}:stream_token", name),
|
||||
ExtensionKind::ChannelRelay => {
|
||||
return Err(ExtensionError::AuthRequired);
|
||||
}
|
||||
};
|
||||
|
||||
let mut secrets = std::collections::HashMap::new();
|
||||
@@ -7043,7 +7055,7 @@ mod tests {
|
||||
let dir = tempfile::tempdir().expect("temp dir");
|
||||
let mgr = make_test_manager(None, dir.path().to_path_buf());
|
||||
|
||||
// No token stored → not a relay channel
|
||||
// No store configured, no team_id → not a relay channel
|
||||
assert!(!mgr.is_relay_channel("slack-relay", "test").await);
|
||||
}
|
||||
|
||||
@@ -7862,19 +7874,13 @@ mod tests {
|
||||
.await
|
||||
.insert("test-relay".to_string());
|
||||
|
||||
// configure() should dispatch to activate_channel_relay(), not
|
||||
// activate_wasm_channel(). Both will fail (no runtime configured),
|
||||
// but the error should be about relay config, not WASM channels.
|
||||
let mut secrets = std::collections::HashMap::new();
|
||||
secrets.insert(
|
||||
"relay:test-relay:stream_token".to_string(),
|
||||
"tok".to_string(),
|
||||
);
|
||||
|
||||
// configure() with empty secrets should dispatch to
|
||||
// activate_channel_relay(), not activate_wasm_channel(). Relay auth
|
||||
// is OAuth-only so there are no manual secrets to pass.
|
||||
let result = mgr
|
||||
.configure(
|
||||
"test-relay",
|
||||
&secrets,
|
||||
&std::collections::HashMap::new(),
|
||||
&std::collections::HashMap::new(),
|
||||
"test",
|
||||
)
|
||||
@@ -7886,7 +7892,6 @@ mod tests {
|
||||
);
|
||||
|
||||
let result = result.unwrap();
|
||||
// Activation will fail (no relay config), but secrets should still be stored
|
||||
assert!(
|
||||
!result.activated,
|
||||
"activation should fail without relay config"
|
||||
@@ -7896,15 +7901,6 @@ mod tests {
|
||||
"error should not mention WASM — got: {}",
|
||||
result.message
|
||||
);
|
||||
|
||||
// Verify the secret was stored
|
||||
assert!(
|
||||
mgr.secrets
|
||||
.exists("test", "relay:test-relay:stream_token")
|
||||
.await
|
||||
.unwrap_or(false),
|
||||
"configure should have stored the relay stream token"
|
||||
);
|
||||
}
|
||||
#[test]
|
||||
fn test_validation_failed_is_distinct_error_variant() {
|
||||
|
||||
@@ -1403,6 +1403,40 @@ impl Store {
|
||||
Ok(counts)
|
||||
}
|
||||
|
||||
/// Batch-load the most recent run status for multiple routines in a single query.
|
||||
/// Uses a window function to pick only the latest run per routine.
|
||||
#[cfg(feature = "postgres")]
|
||||
pub async fn batch_get_last_run_status(
|
||||
&self,
|
||||
routine_ids: &[Uuid],
|
||||
) -> Result<HashMap<Uuid, RunStatus>, DatabaseError> {
|
||||
if routine_ids.is_empty() {
|
||||
return Ok(HashMap::new());
|
||||
}
|
||||
|
||||
let conn = self.conn().await?;
|
||||
let rows = conn
|
||||
.query(
|
||||
"SELECT DISTINCT ON (routine_id) routine_id, status
|
||||
FROM routine_runs
|
||||
WHERE routine_id = ANY($1)
|
||||
ORDER BY routine_id, started_at DESC",
|
||||
&[&routine_ids],
|
||||
)
|
||||
.await?;
|
||||
|
||||
let mut statuses = HashMap::new();
|
||||
for row in rows {
|
||||
let id: Uuid = row.get("routine_id");
|
||||
let status_str: String = row.get("status");
|
||||
if let std::result::Result::Ok(status) = status_str.parse::<RunStatus>() {
|
||||
statuses.insert(id, status);
|
||||
}
|
||||
}
|
||||
|
||||
Ok(statuses)
|
||||
}
|
||||
|
||||
/// Link a routine run to a dispatched job.
|
||||
pub async fn link_routine_run_to_job(
|
||||
&self,
|
||||
|
||||
@@ -111,10 +111,23 @@ impl Tunnel for CloudflareTunnel {
|
||||
}
|
||||
}
|
||||
|
||||
// Drain stderr in the background to prevent SIGPIPE/buffer stalls.
|
||||
tokio::spawn(async move { while let Ok(Some(_)) = reader.next_line().await {} });
|
||||
if let Ok(mut guard) = self.url.write() {
|
||||
*guard = Some(public_url.clone());
|
||||
}
|
||||
|
||||
// Drain stdout silently.
|
||||
// We took ownership of cloudflared's stderr pipe above to parse the URL.
|
||||
// cloudflared continues writing logs for its entire lifetime. If we drop
|
||||
// the reader, the pipe closes and cloudflared gets SIGPIPE on its next
|
||||
// write. We can't just store the reader without reading — the OS pipe
|
||||
// buffer fills up and cloudflared blocks. So we drain it in a background
|
||||
// task. The task exits naturally when cloudflared is killed (EOF).
|
||||
let drain_handle = tokio::spawn(async move {
|
||||
while let Ok(Some(line)) = reader.next_line().await {
|
||||
tracing::trace!("cloudflared: {line}");
|
||||
}
|
||||
});
|
||||
|
||||
// Drain stdout silently to prevent SIGPIPE/buffer stalls.
|
||||
if let Some(stdout) = stdout {
|
||||
tokio::spawn(async move {
|
||||
let mut out_reader = tokio::io::BufReader::new(stdout).lines();
|
||||
@@ -122,12 +135,11 @@ impl Tunnel for CloudflareTunnel {
|
||||
});
|
||||
}
|
||||
|
||||
if let Ok(mut guard) = self.url.write() {
|
||||
*guard = Some(public_url.clone());
|
||||
}
|
||||
|
||||
let mut guard = self.proc.lock().await;
|
||||
*guard = Some(TunnelProcess { child });
|
||||
*guard = Some(TunnelProcess {
|
||||
child,
|
||||
_pipe_drain: Some(drain_handle),
|
||||
});
|
||||
|
||||
Ok(public_url)
|
||||
}
|
||||
|
||||
+18
-5
@@ -73,6 +73,7 @@ impl Tunnel for CustomTunnel {
|
||||
let stderr = child.stderr.take();
|
||||
|
||||
let mut public_url = format!("http://{local_host}:{local_port}");
|
||||
let mut drain_handle: Option<tokio::task::JoinHandle<()>> = None;
|
||||
|
||||
if self.url_pattern.is_some()
|
||||
&& let Some(stdout) = stdout
|
||||
@@ -103,17 +104,26 @@ impl Tunnel for CustomTunnel {
|
||||
Err(_) => {}
|
||||
}
|
||||
}
|
||||
// Drain remaining stdout to prevent SIGPIPE/buffer stalls.
|
||||
tokio::spawn(async move { while let Ok(Some(_)) = reader.next_line().await {} });
|
||||
// We took ownership of the process's stdout pipe above to parse the
|
||||
// URL. The process may continue writing to stdout for its lifetime.
|
||||
// If we drop the reader, the pipe closes and the process gets SIGPIPE.
|
||||
// We can't just store the reader without reading — the OS pipe buffer
|
||||
// fills up and the process blocks. So we drain it in a background task.
|
||||
// The task exits naturally when the process is killed (EOF).
|
||||
drain_handle = Some(tokio::spawn(async move {
|
||||
while let Ok(Some(line)) = reader.next_line().await {
|
||||
tracing::trace!("custom-tunnel: {line}");
|
||||
}
|
||||
}));
|
||||
} else if let Some(stdout) = stdout {
|
||||
// No url_pattern: still drain stdout to prevent pipe stalls.
|
||||
// No url_pattern: still drain stdout to prevent SIGPIPE/buffer stalls.
|
||||
tokio::spawn(async move {
|
||||
let mut reader = tokio::io::BufReader::new(stdout).lines();
|
||||
while let Ok(Some(_)) = reader.next_line().await {}
|
||||
});
|
||||
}
|
||||
|
||||
// Drain stderr silently.
|
||||
// Drain stderr to prevent SIGPIPE/buffer stalls.
|
||||
if let Some(stderr) = stderr {
|
||||
tokio::spawn(async move {
|
||||
let mut reader = tokio::io::BufReader::new(stderr).lines();
|
||||
@@ -126,7 +136,10 @@ impl Tunnel for CustomTunnel {
|
||||
}
|
||||
|
||||
let mut guard = self.proc.lock().await;
|
||||
*guard = Some(TunnelProcess { child });
|
||||
*guard = Some(TunnelProcess {
|
||||
child,
|
||||
_pipe_drain: drain_handle,
|
||||
});
|
||||
|
||||
Ok(public_url)
|
||||
}
|
||||
|
||||
+126
-16
@@ -66,6 +66,10 @@ pub trait Tunnel: Send + Sync {
|
||||
/// Wraps a spawned tunnel child process.
|
||||
pub(crate) struct TunnelProcess {
|
||||
pub child: tokio::process::Child,
|
||||
/// Background task that drains the process's output pipe (stdout or stderr).
|
||||
/// Must stay alive or the process dies (SIGPIPE from closed pipe) or hangs
|
||||
/// (OS pipe buffer fills up, blocking the process's writes).
|
||||
pub _pipe_drain: Option<tokio::task::JoinHandle<()>>,
|
||||
}
|
||||
|
||||
pub(crate) type SharedProcess = Arc<Mutex<Option<TunnelProcess>>>;
|
||||
@@ -182,6 +186,22 @@ pub fn create_tunnel(config: &TunnelProviderConfig) -> Result<Option<Box<dyn Tun
|
||||
|
||||
// ── Managed tunnel startup ───────────────────────────────────────
|
||||
|
||||
/// Determine which local address the tunnel should forward traffic to.
|
||||
///
|
||||
/// Prefers the webhook server (`HTTP_PORT`) since that's where webhook routes
|
||||
/// (Telegram, etc.) are served. Falls back to the gateway port if configured,
|
||||
/// otherwise defaults to 0.0.0.0:8080 (the same fallback the webhook server
|
||||
/// uses in main.rs when no HTTP config is present).
|
||||
fn resolve_tunnel_target(channels: &crate::config::ChannelsConfig) -> (&str, u16) {
|
||||
if let Some(ref http) = channels.http {
|
||||
return (http.host.as_str(), http.port);
|
||||
}
|
||||
if let Some(ref gw) = channels.gateway {
|
||||
return (gw.host.as_str(), gw.port);
|
||||
}
|
||||
("0.0.0.0", 8080)
|
||||
}
|
||||
|
||||
/// Start a managed tunnel if configured and no static URL is already set.
|
||||
///
|
||||
/// Returns the (potentially mutated) config with `tunnel.public_url` set,
|
||||
@@ -201,28 +221,17 @@ pub async fn start_managed_tunnel(
|
||||
return (config, None);
|
||||
};
|
||||
|
||||
let gateway_port = config
|
||||
.channels
|
||||
.gateway
|
||||
.as_ref()
|
||||
.map(|g| g.port)
|
||||
.unwrap_or(3000);
|
||||
let gateway_host = config
|
||||
.channels
|
||||
.gateway
|
||||
.as_ref()
|
||||
.map(|g| g.host.as_str())
|
||||
.unwrap_or("127.0.0.1");
|
||||
let (tunnel_host, tunnel_port) = resolve_tunnel_target(&config.channels);
|
||||
|
||||
match create_tunnel(provider_config) {
|
||||
Ok(Some(tunnel)) => {
|
||||
tracing::debug!(
|
||||
"Starting {} tunnel on {}:{}...",
|
||||
tunnel.name(),
|
||||
gateway_host,
|
||||
gateway_port
|
||||
tunnel_host,
|
||||
tunnel_port
|
||||
);
|
||||
match tunnel.start(gateway_host, gateway_port).await {
|
||||
match tunnel.start(tunnel_host, tunnel_port).await {
|
||||
Ok(url) => {
|
||||
tracing::debug!("Tunnel started: {}", url);
|
||||
config.tunnel.public_url = Some(url);
|
||||
@@ -383,10 +392,111 @@ mod tests {
|
||||
|
||||
{
|
||||
let mut guard = proc.lock().await;
|
||||
*guard = Some(TunnelProcess { child });
|
||||
*guard = Some(TunnelProcess {
|
||||
child,
|
||||
_pipe_drain: None,
|
||||
});
|
||||
}
|
||||
|
||||
kill_shared(&proc).await.unwrap();
|
||||
assert!(proc.lock().await.is_none());
|
||||
}
|
||||
|
||||
// ── Port selection regression tests ──────────────────────────────
|
||||
|
||||
fn base_channels() -> crate::config::ChannelsConfig {
|
||||
crate::config::ChannelsConfig {
|
||||
cli: crate::config::CliConfig { enabled: false },
|
||||
http: None,
|
||||
gateway: None,
|
||||
signal: None,
|
||||
wasm_channels_dir: std::env::temp_dir().join("ironclaw-test-channels"),
|
||||
wasm_channels_enabled: false,
|
||||
wasm_channel_owner_ids: std::collections::HashMap::new(),
|
||||
}
|
||||
}
|
||||
|
||||
fn channels_with_http(host: &str, port: u16) -> crate::config::ChannelsConfig {
|
||||
let mut c = base_channels();
|
||||
c.http = Some(crate::config::HttpConfig {
|
||||
host: host.to_string(),
|
||||
port,
|
||||
webhook_secret: None,
|
||||
user_id: "test".to_string(),
|
||||
});
|
||||
c.gateway = Some(crate::config::GatewayConfig {
|
||||
host: "127.0.0.1".to_string(),
|
||||
port: 3000,
|
||||
auth_token: None,
|
||||
user_id: "test".to_string(),
|
||||
workspace_read_scopes: Vec::new(),
|
||||
memory_layers: Vec::new(),
|
||||
user_tokens: None,
|
||||
});
|
||||
c
|
||||
}
|
||||
|
||||
fn channels_gateway_only(host: &str, port: u16) -> crate::config::ChannelsConfig {
|
||||
let mut c = base_channels();
|
||||
c.gateway = Some(crate::config::GatewayConfig {
|
||||
host: host.to_string(),
|
||||
port,
|
||||
auth_token: None,
|
||||
user_id: "test".to_string(),
|
||||
workspace_read_scopes: Vec::new(),
|
||||
memory_layers: Vec::new(),
|
||||
user_tokens: None,
|
||||
});
|
||||
c
|
||||
}
|
||||
|
||||
fn channels_neither() -> crate::config::ChannelsConfig {
|
||||
base_channels()
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn tunnel_target_prefers_http_port() {
|
||||
let channels = channels_with_http("0.0.0.0", 8080);
|
||||
let (host, port) = resolve_tunnel_target(&channels);
|
||||
assert_eq!(host, "0.0.0.0"); // safety: test-only
|
||||
assert_eq!(port, 8080); // safety: test-only
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn tunnel_target_falls_back_to_gateway() {
|
||||
let channels = channels_gateway_only("10.0.0.1", 4000);
|
||||
let (host, port) = resolve_tunnel_target(&channels);
|
||||
assert_eq!(host, "10.0.0.1"); // safety: test-only
|
||||
assert_eq!(port, 4000); // safety: test-only
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn tunnel_target_defaults_to_webhook_fallback() {
|
||||
let channels = channels_neither();
|
||||
let (host, port) = resolve_tunnel_target(&channels);
|
||||
// Matches the webhook server's hardcoded fallback in main.rs
|
||||
assert_eq!(host, "0.0.0.0"); // safety: test-only
|
||||
assert_eq!(port, 8080); // safety: test-only
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn tunnel_target_http_takes_priority_over_gateway() {
|
||||
let channels = channels_with_http("192.168.1.1", 9090);
|
||||
let (host, port) = resolve_tunnel_target(&channels);
|
||||
// Should use HTTP config, not gateway's 127.0.0.1:3000
|
||||
assert_eq!(host, "192.168.1.1"); // safety: test-only
|
||||
assert_eq!(port, 9090); // safety: test-only
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn tunnel_target_no_http_no_gateway_matches_webhook_fallback() {
|
||||
// When HTTP_PORT is not set and gateway is not configured (e.g. WASM
|
||||
// channels exist but no explicit HTTP config), the webhook server in
|
||||
// main.rs binds to 0.0.0.0:8080 as a hardcoded fallback. The tunnel
|
||||
// must target the same address so webhook traffic reaches the right
|
||||
// server.
|
||||
let channels = channels_neither();
|
||||
let (host, port) = resolve_tunnel_target(&channels);
|
||||
assert_eq!((host, port), ("0.0.0.0", 8080)); // safety: test-only
|
||||
}
|
||||
}
|
||||
|
||||
+21
-10
@@ -110,12 +110,24 @@ impl Tunnel for NgrokTunnel {
|
||||
}
|
||||
}
|
||||
|
||||
// Drain stdout silently — ngrok only emits low-level connection events
|
||||
// to stdout; the pipe must be consumed to prevent SIGPIPE/buffer stalls.
|
||||
tokio::spawn(async move { while let Ok(Some(_)) = reader.next_line().await {} });
|
||||
if let Ok(mut guard) = self.url.write() {
|
||||
*guard = Some(public_url.clone());
|
||||
}
|
||||
|
||||
// Drain stderr silently — with --log stdout all meaningful output goes
|
||||
// to stdout; stderr only needs to be consumed to prevent pipe stalls.
|
||||
// We took ownership of ngrok's stdout pipe above to parse the URL.
|
||||
// ngrok continues writing logs to stdout for its entire lifetime.
|
||||
// If we drop the reader, the pipe closes and ngrok gets SIGPIPE on
|
||||
// its next write → process dies. We can't just store the reader
|
||||
// without reading — the OS pipe buffer (~64KB) fills up and ngrok
|
||||
// blocks. So we drain it in a background task. The task exits
|
||||
// naturally when ngrok is killed (EOF on the pipe).
|
||||
let drain_handle = tokio::spawn(async move {
|
||||
while let Ok(Some(line)) = reader.next_line().await {
|
||||
tracing::trace!("ngrok: {line}");
|
||||
}
|
||||
});
|
||||
|
||||
// Drain stderr silently to prevent SIGPIPE/buffer stalls.
|
||||
if let Some(stderr) = stderr {
|
||||
tokio::spawn(async move {
|
||||
let mut err_reader = tokio::io::BufReader::new(stderr).lines();
|
||||
@@ -123,12 +135,11 @@ impl Tunnel for NgrokTunnel {
|
||||
});
|
||||
}
|
||||
|
||||
if let Ok(mut guard) = self.url.write() {
|
||||
*guard = Some(public_url.clone());
|
||||
}
|
||||
|
||||
let mut guard = self.proc.lock().await;
|
||||
*guard = Some(TunnelProcess { child });
|
||||
*guard = Some(TunnelProcess {
|
||||
child,
|
||||
_pipe_drain: Some(drain_handle),
|
||||
});
|
||||
|
||||
Ok(public_url)
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user