From 5d1d504e1105778b13915f904a6dcde05cfda795 Mon Sep 17 00:00:00 2001 From: Zaki Date: Mon, 23 Mar 2026 06:40:40 -0700 Subject: [PATCH] fix(security): block cross-channel approval thread hijacking (#1485) Add source_channel to Thread and verify channel authorization before allowing approval messages to target threads by UUID. The web gateway channel is allowed as a trusted approval UI. Threads without source_channel (deserialized from older DB records) are permitted for backward compatibility. Closes #1485 Co-Authored-By: Claude Opus 4.6 (1M context) --- src/agent/agent_loop.rs | 17 +++- src/agent/compaction.rs | 10 +-- src/agent/dispatcher.rs | 4 +- src/agent/session.rs | 144 +++++++++++++++++++----------- src/agent/session_manager.rs | 34 +++++-- src/agent/thread_ops.rs | 14 +-- src/channels/web/handlers/chat.rs | 2 +- src/channels/web/server.rs | 2 +- 8 files changed, 146 insertions(+), 81 deletions(-) diff --git a/src/agent/agent_loop.rs b/src/agent/agent_loop.rs index 59f1f87d..310f9378 100644 --- a/src/agent/agent_loop.rs +++ b/src/agent/agent_loop.rs @@ -843,7 +843,7 @@ impl Agent { { use crate::agent::session::Thread; let mut sess = session.lock().await; - let thread = Thread::with_id(id, sess.id); + let thread = Thread::with_id(id, sess.id, None); sess.active_thread = Some(id); sess.threads.entry(id).or_insert(thread); } @@ -1148,7 +1148,20 @@ 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) { + let authorized = thread.source_channel.as_ref().is_none_or(|src| { + src == &message.channel || message.channel == "web" + }); + 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 30bb2b6c..c69f9608 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 9e639171..f678bedb 100644 --- a/src/agent/dispatcher.rs +++ b/src/agent/dispatcher.rs @@ -2299,7 +2299,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"); @@ -2412,7 +2412,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/session.rs b/src/agent/session.rs index 6c873e46..dd0ec726 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. @@ -235,7 +238,7 @@ pub const MAX_PENDING_MESSAGES: usize = 10; 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 +251,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 +269,7 @@ impl Thread { pending_approval: None, pending_auth: None, pending_messages: VecDeque::new(), + source_channel: source_channel.map(String::from), } } @@ -787,13 +792,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 +811,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 +834,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 +860,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 +879,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 +892,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 +907,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 +937,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 +951,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 +963,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 +982,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 +998,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 +1015,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 +1042,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 +1063,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 +1073,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 +1097,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 +1109,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 +1127,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 +1143,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 +1159,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 +1175,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 +1193,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 +1231,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 +1246,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 +1256,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 +1266,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 +1293,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 +1322,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 +1339,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 +1371,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 +1398,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 +1430,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 +1461,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 +1474,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 +1496,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 +1539,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 +1562,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 +1598,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 +1618,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 +1629,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 +1648,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 +1676,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 +1705,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()); @@ -1717,6 +1722,7 @@ mod tests { } #[test] +<<<<<<< HEAD fn test_record_tool_result_for_by_id() { let mut turn = Turn::new(0, "test"); turn.record_tool_call_with_reasoning( @@ -1811,4 +1817,34 @@ 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" + ); + } } diff --git a/src/agent/session_manager.rs b/src/agent/session_manager.rs index ae98b0b0..e7144bea 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); } @@ -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); } { @@ -1020,6 +1020,7 @@ mod tests { } #[tokio::test] +<<<<<<< HEAD async fn test_resolve_thread_with_pre_parsed_uuid_adopts_thread() { use crate::agent::session::Thread; @@ -1102,4 +1103,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 af0bd67f..a7f08919 100644 --- a/src/agent/thread_ops.rs +++ b/src/agent/thread_ops.rs @@ -141,7 +141,7 @@ impl Agent { 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, None); if !chat_messages.is_empty() { thread.restore_from_messages(chat_messages); } @@ -1781,7 +1781,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 +2117,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 +2185,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 +2211,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 +2241,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 +2268,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/web/handlers/chat.rs b/src/channels/web/handlers/chat.rs index d1580f5c..302f8051 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, diff --git a/src/channels/web/server.rs b/src/channels/web/server.rs index 8c9ecdde..030f9d5e 100644 --- a/src/channels/web/server.rs +++ b/src/channels/web/server.rs @@ -2014,7 +2014,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,