From 2094d6e30d6fe7330b423d0bbc00fbe759722c76 Mon Sep 17 00:00:00 2001 From: pikaxinge <68273313+pikaxinge@users.noreply.github.com> Date: Thu, 12 Mar 2026 07:52:31 +0800 Subject: [PATCH] fix(agent): block thread_id-based context pollution across users (#760) * fix(agent): prevent forged thread UUID context/write contamination * fix(agent): close thread_id race and reject forged UUID hydration * fix(ci): satisfy clippy and fmt checks after rebase --- src/agent/agent_loop.rs | 4 +- src/agent/thread_ops.rs | 190 ++++++++++++++++++++++++++---- src/channels/web/handlers/chat.rs | 10 +- src/channels/web/server.rs | 10 +- src/db/libsql/conversations.rs | 15 ++- src/db/mod.rs | 2 +- src/db/postgres.rs | 2 +- src/history/store.rs | 24 ++-- src/testing/mod.rs | 62 +++++++++- tests/e2e_thread_id_isolation.rs | 183 ++++++++++++++++++++++++++++ 10 files changed, 448 insertions(+), 54 deletions(-) create mode 100644 tests/e2e_thread_id_isolation.rs diff --git a/src/agent/agent_loop.rs b/src/agent/agent_loop.rs index d95f3e46..97c9b2b9 100644 --- a/src/agent/agent_loop.rs +++ b/src/agent/agent_loop.rs @@ -803,7 +803,9 @@ impl Agent { thread_id = %external_thread_id, "Hydrating thread from DB" ); - self.maybe_hydrate_thread(message, external_thread_id).await; + if let Some(rejection) = self.maybe_hydrate_thread(message, external_thread_id).await { + return Ok(Some(format!("Error: {}", rejection))); + } } // Resolve session and thread diff --git a/src/agent/thread_ops.rs b/src/agent/thread_ops.rs index f5d2c53e..566f5140 100644 --- a/src/agent/thread_ops.rs +++ b/src/agent/thread_ops.rs @@ -23,6 +23,14 @@ use crate::error::Error; use crate::llm::{ChatMessage, ToolCall}; use crate::tools::redact_params; +const FORGED_THREAD_ID_ERROR: &str = "Invalid or unauthorized thread ID."; + +fn requires_preexisting_uuid_thread(channel: &str) -> bool { + // Gateway-style channels send server-issued conversation UUIDs. + // Unknown UUIDs should be rejected instead of silently creating a new thread. + matches!(channel, "gateway" | "test") +} + impl Agent { /// Hydrate a historical thread from DB into memory if not already present. /// @@ -37,11 +45,11 @@ impl Agent { &self, message: &IncomingMessage, external_thread_id: &str, - ) { + ) -> Option { // Only hydrate UUID-shaped thread IDs (web gateway uses UUIDs) let thread_uuid = match Uuid::parse_str(external_thread_id) { Ok(id) => id, - Err(_) => return, + Err(_) => return None, }; // Check if already in memory @@ -52,7 +60,7 @@ impl Agent { { let sess = session.lock().await; if sess.threads.contains_key(&thread_uuid) { - return; + return None; } } @@ -61,6 +69,62 @@ impl Agent { let msg_count; if let Some(store) = self.store() { + // Never hydrate history from a conversation UUID that isn't owned + // by the current authenticated user. + let owned = match store + .conversation_belongs_to_user(thread_uuid, &message.user_id) + .await + { + Ok(v) => v, + Err(e) => { + tracing::warn!( + "Failed to verify conversation ownership for hydration {}: {}", + thread_uuid, + e + ); + if requires_preexisting_uuid_thread(&message.channel) { + return Some(FORGED_THREAD_ID_ERROR.to_string()); + } + return None; + } + }; + if !owned { + let exists = match store.get_conversation_metadata(thread_uuid).await { + Ok(Some(_)) => true, + Ok(None) => false, + Err(e) => { + tracing::warn!( + "Failed to inspect conversation metadata for hydration {}: {}", + thread_uuid, + e + ); + if requires_preexisting_uuid_thread(&message.channel) { + return Some(FORGED_THREAD_ID_ERROR.to_string()); + } + return None; + } + }; + + if requires_preexisting_uuid_thread(&message.channel) { + tracing::warn!( + user = %message.user_id, + channel = %message.channel, + thread_id = %thread_uuid, + exists, + "Rejected message for unavailable thread id" + ); + return Some(FORGED_THREAD_ID_ERROR.to_string()); + } + + tracing::warn!( + user = %message.user_id, + thread_id = %thread_uuid, + exists, + "Skipped hydration for thread id not owned by sender" + ); + return None; + } + let db_messages = store .list_conversation_messages(thread_uuid) .await @@ -104,6 +168,8 @@ impl Agent { thread_uuid, msg_count ); + + None } pub(super) async fn process_user_input( @@ -303,8 +369,13 @@ impl Agent { thread_id = %thread_id, "Persisting user message to DB" ); - self.persist_user_message(thread_id, &message.user_id, effective_content) - .await; + self.persist_user_message( + thread_id, + &message.channel, + &message.user_id, + effective_content, + ) + .await; tracing::debug!( message_id = %message.id, @@ -386,10 +457,21 @@ impl Agent { .await; // Persist tool calls then assistant response (user message already persisted at turn start) - self.persist_tool_calls(thread_id, &message.user_id, turn_number, &tool_calls) - .await; - self.persist_assistant_response(thread_id, &message.user_id, &response) - .await; + self.persist_tool_calls( + thread_id, + &message.channel, + &message.user_id, + turn_number, + &tool_calls, + ) + .await; + self.persist_assistant_response( + thread_id, + &message.channel, + &message.user_id, + &response, + ) + .await; Ok(SubmissionResult::response(response)) } @@ -423,6 +505,41 @@ impl Agent { } } + /// Ensure a thread UUID is writable for `(channel, user_id)`. + /// + /// Returns `false` for foreign/unowned conversation IDs or DB errors. + async fn ensure_writable_conversation( + &self, + store: &Arc, + thread_id: Uuid, + channel: &str, + user_id: &str, + ) -> bool { + match store + .ensure_conversation(thread_id, channel, user_id, None) + .await + { + Ok(true) => true, + Ok(false) => { + tracing::warn!( + user = %user_id, + channel = %channel, + thread_id = %thread_id, + "Rejected write for unavailable thread id" + ); + false + } + Err(e) => { + tracing::warn!( + "Failed to ensure writable conversation {}: {}", + thread_id, + e + ); + false + } + } + } + /// Persist the user message to the DB at turn start (before the agentic loop). /// /// This ensures the user message is durable even if the process crashes @@ -430,6 +547,7 @@ impl Agent { pub(super) async fn persist_user_message( &self, thread_id: Uuid, + channel: &str, user_id: &str, user_input: &str, ) { @@ -438,11 +556,10 @@ impl Agent { None => return, }; - if let Err(e) = store - .ensure_conversation(thread_id, "gateway", user_id, None) + if !self + .ensure_writable_conversation(&store, thread_id, channel, user_id) .await { - tracing::warn!("Failed to ensure conversation {}: {}", thread_id, e); return; } @@ -462,6 +579,7 @@ impl Agent { pub(super) async fn persist_assistant_response( &self, thread_id: Uuid, + channel: &str, user_id: &str, response: &str, ) { @@ -470,11 +588,10 @@ impl Agent { None => return, }; - if let Err(e) = store - .ensure_conversation(thread_id, "gateway", user_id, None) + if !self + .ensure_writable_conversation(&store, thread_id, channel, user_id) .await { - tracing::warn!("Failed to ensure conversation {}: {}", thread_id, e); return; } @@ -494,6 +611,7 @@ impl Agent { pub(super) async fn persist_tool_calls( &self, thread_id: Uuid, + channel: &str, user_id: &str, turn_number: usize, tool_calls: &[crate::agent::session::TurnToolCall], @@ -543,11 +661,10 @@ impl Agent { } }; - if let Err(e) = store - .ensure_conversation(thread_id, "gateway", user_id, None) + if !self + .ensure_writable_conversation(&store, thread_id, channel, user_id) .await { - tracing::warn!("Failed to ensure conversation {}: {}", thread_id, e); return; } @@ -1214,10 +1331,21 @@ impl Agent { .map(|t| (t.turn_number, t.tool_calls.clone())) .unwrap_or_default(); // User message already persisted at turn start; save tool calls then assistant response - self.persist_tool_calls(thread_id, &message.user_id, turn_number, &tool_calls) - .await; - self.persist_assistant_response(thread_id, &message.user_id, &response) - .await; + self.persist_tool_calls( + thread_id, + &message.channel, + &message.user_id, + turn_number, + &tool_calls, + ) + .await; + self.persist_assistant_response( + thread_id, + &message.channel, + &message.user_id, + &response, + ) + .await; let _ = self .channels .send_status( @@ -1270,8 +1398,13 @@ impl Agent { thread.clear_pending_approval(); thread.complete_turn(&rejection); // User message already persisted at turn start; save rejection response - self.persist_assistant_response(thread_id, &message.user_id, &rejection) - .await; + self.persist_assistant_response( + thread_id, + &message.channel, + &message.user_id, + &rejection, + ) + .await; } } @@ -1309,8 +1442,13 @@ impl Agent { thread.enter_auth_mode(ext_name.clone()); thread.complete_turn(&instructions); // User message already persisted at turn start; save auth instructions - self.persist_assistant_response(thread_id, &message.user_id, &instructions) - .await; + self.persist_assistant_response( + thread_id, + &message.channel, + &message.user_id, + &instructions, + ) + .await; } } let _ = self diff --git a/src/channels/web/handlers/chat.rs b/src/channels/web/handlers/chat.rs index 11054864..91c4533b 100644 --- a/src/channels/web/handlers/chat.rs +++ b/src/channels/web/handlers/chat.rs @@ -534,11 +534,17 @@ pub async fn chat_new_thread_handler( // Persist the empty conversation row with thread_type metadata synchronously // so that the subsequent loadThreads() call from the frontend sees it. if let Some(ref store) = state.store { - if let Err(e) = store + match store .ensure_conversation(thread_id, "gateway", &state.user_id, None) .await { - tracing::warn!("Failed to persist new thread: {}", e); + Ok(true) => {} + Ok(false) => tracing::warn!( + user = %state.user_id, + thread_id = %thread_id, + "Skipped persisting new thread due to ownership/channel conflict" + ), + Err(e) => tracing::warn!("Failed to persist new thread: {}", e), } let metadata_val = serde_json::json!("thread"); if let Err(e) = store diff --git a/src/channels/web/server.rs b/src/channels/web/server.rs index f9682db8..3beafc20 100644 --- a/src/channels/web/server.rs +++ b/src/channels/web/server.rs @@ -1448,11 +1448,17 @@ async fn chat_new_thread_handler( // Persist the empty conversation row with thread_type metadata synchronously // so that the subsequent loadThreads() call from the frontend sees it. if let Some(ref store) = state.store { - if let Err(e) = store + match store .ensure_conversation(thread_id, "gateway", &state.user_id, None) .await { - tracing::warn!("Failed to persist new thread: {}", e); + Ok(true) => {} + Ok(false) => tracing::warn!( + user = %state.user_id, + thread_id = %thread_id, + "Skipped persisting new thread due to ownership/channel conflict" + ), + Err(e) => tracing::warn!("Failed to persist new thread: {}", e), } let metadata_val = serde_json::json!("thread"); if let Err(e) = store diff --git a/src/db/libsql/conversations.rs b/src/db/libsql/conversations.rs index 2a7ef06c..911ee863 100644 --- a/src/db/libsql/conversations.rs +++ b/src/db/libsql/conversations.rs @@ -67,20 +67,23 @@ impl ConversationStore for LibSqlBackend { channel: &str, user_id: &str, thread_id: Option<&str>, - ) -> Result<(), DatabaseError> { + ) -> Result { let conn = self.connect().await?; let now = fmt_ts(&Utc::now()); - conn.execute( + 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) - ON CONFLICT (id) DO UPDATE SET last_activity = ?5 + 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], ) - .await - .map_err(|e| DatabaseError::Query(e.to_string()))?; - Ok(()) + .await + .map_err(|e| DatabaseError::Query(e.to_string()))?; + Ok(affected > 0) } async fn list_conversations_with_preview( diff --git a/src/db/mod.rs b/src/db/mod.rs index 9485d41c..4afd1db8 100644 --- a/src/db/mod.rs +++ b/src/db/mod.rs @@ -207,7 +207,7 @@ pub trait ConversationStore: Send + Sync { channel: &str, user_id: &str, thread_id: Option<&str>, - ) -> Result<(), DatabaseError>; + ) -> Result; async fn list_conversations_with_preview( &self, user_id: &str, diff --git a/src/db/postgres.rs b/src/db/postgres.rs index 9dd988bc..2cf6a65a 100644 --- a/src/db/postgres.rs +++ b/src/db/postgres.rs @@ -99,7 +99,7 @@ impl ConversationStore for PgBackend { channel: &str, user_id: &str, thread_id: Option<&str>, - ) -> Result<(), DatabaseError> { + ) -> Result { self.store .ensure_conversation(id, channel, user_id, thread_id) .await diff --git a/src/history/store.rs b/src/history/store.rs index e877cbbf..83f60d70 100644 --- a/src/history/store.rs +++ b/src/history/store.rs @@ -1407,25 +1407,31 @@ pub struct ConversationMessage { impl Store { /// Ensure a conversation row exists for a given UUID. /// - /// Idempotent: inserts on first call, bumps `last_activity` on subsequent calls. + /// Returns `true` when the row is inserted or refreshed for the same + /// `(channel, user_id)`. Returns `false` when the UUID already exists but + /// belongs to a different owner/channel. pub async fn ensure_conversation( &self, id: Uuid, channel: &str, user_id: &str, thread_id: Option<&str>, - ) -> Result<(), DatabaseError> { + ) -> Result { let conn = self.conn().await?; - conn.execute( - r#" + let affected = conn + .execute( + r#" INSERT INTO conversations (id, channel, user_id, thread_id) VALUES ($1, $2, $3, $4) - ON CONFLICT (id) DO UPDATE SET last_activity = NOW() + 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], - ) - .await?; - Ok(()) + &[&id, &channel, &user_id, &thread_id], + ) + .await?; + Ok(affected > 0) } /// List conversations with a title derived from the first user message. diff --git a/src/testing/mod.rs b/src/testing/mod.rs index 97612887..d2078a80 100644 --- a/src/testing/mod.rs +++ b/src/testing/mod.rs @@ -641,14 +641,20 @@ mod tests { let conv_id = uuid::Uuid::new_v4(); // ensure_conversation should create the row. - db.ensure_conversation(conv_id, "web", "carol", None) - .await - .expect("ensure first"); + assert!( + db.ensure_conversation(conv_id, "web", "carol", None) + .await + .expect("ensure first"), + "first ensure_conversation should create the row" + ); // Calling again with the same ID should not error. - db.ensure_conversation(conv_id, "web", "carol", None) - .await - .expect("ensure second (idempotent)"); + assert!( + db.ensure_conversation(conv_id, "web", "carol", None) + .await + .expect("ensure second (idempotent)"), + "second ensure_conversation should touch owned row" + ); // Should be able to add messages to it. let msg_id = db @@ -666,6 +672,50 @@ mod tests { assert_eq!(msgs[0].content, "test message"); } + #[cfg(feature = "libsql")] + #[tokio::test] + async fn test_ensure_conversation_foreign_conflict_does_not_touch_last_activity() { + let harness = TestHarnessBuilder::new().build().await; + let db = &harness.db; + + let conv_id = db + .create_conversation("web", "alice", None) + .await + .expect("create conversation"); + + let before = db + .list_conversations_all_channels("alice", 10) + .await + .expect("list conversations before foreign ensure") + .into_iter() + .find(|c| c.id == conv_id) + .expect("conversation must exist before foreign ensure") + .last_activity; + + tokio::time::sleep(std::time::Duration::from_millis(25)).await; + + assert!( + !db.ensure_conversation(conv_id, "web", "mallory", None) + .await + .expect("foreign ensure should not error"), + "foreign ensure_conversation should report not ensured" + ); + + let after = db + .list_conversations_all_channels("alice", 10) + .await + .expect("list conversations after foreign ensure") + .into_iter() + .find(|c| c.id == conv_id) + .expect("conversation must still exist after foreign ensure") + .last_activity; + + assert_eq!( + after, before, + "foreign ensure_conversation should not mutate last_activity" + ); + } + #[cfg(feature = "libsql")] #[tokio::test] async fn test_paginated_messages() { diff --git a/tests/e2e_thread_id_isolation.rs b/tests/e2e_thread_id_isolation.rs new file mode 100644 index 00000000..baec73c1 --- /dev/null +++ b/tests/e2e_thread_id_isolation.rs @@ -0,0 +1,183 @@ +//! E2E regression test: forged thread IDs must not cross user boundaries. +//! +//! Demonstrates that a client cannot provide another user's conversation UUID +//! and get that history hydrated into prompt context or written into. + +#[cfg(feature = "libsql")] +mod support; + +#[cfg(feature = "libsql")] +mod tests { + use std::time::Duration; + + use ironclaw::channels::{IncomingMessage, OutgoingResponse}; + use uuid::Uuid; + + use crate::support::test_rig::TestRigBuilder; + use crate::support::trace_llm::{LlmTrace, TraceResponse, TraceStep}; + + fn assert_safe_thread_rejection(response: &OutgoingResponse) { + let msg = response.content.to_lowercase(); + assert!( + msg.contains("thread") && (msg.contains("invalid") || msg.contains("unauthorized")), + "expected safe thread-id rejection response, got: {}", + response.content + ); + } + + #[tokio::test] + async fn forged_existing_foreign_thread_id_is_rejected_without_hydration_or_persistence() { + let trace = LlmTrace::single_turn( + "thread-id-isolation", + "attacker turn", + vec![TraceStep { + request_hint: None, + response: TraceResponse::Text { + content: "safe response".to_string(), + input_tokens: 12, + output_tokens: 4, + }, + expected_tool_results: Vec::new(), + }], + ); + + let rig = TestRigBuilder::new().with_trace(trace).build().await; + + let foreign_thread_id = Uuid::new_v4(); + let marker = format!("FOREIGN-MARKER-{}", Uuid::new_v4()); + let store = rig.database(); + assert!( + store + .ensure_conversation(foreign_thread_id, "gateway", "victim-user", None) + .await + .expect("failed to create victim conversation"), + "test setup failed: victim conversation was not created" + ); + store + .add_conversation_message( + foreign_thread_id, + "user", + &format!("victim-only secret marker: {marker}"), + ) + .await + .expect("failed to seed victim conversation message"); + + let before_messages = store + .list_conversation_messages(foreign_thread_id) + .await + .expect("failed to read victim conversation before forged send"); + assert!( + before_messages.iter().any(|m| m.content.contains(&marker)), + "test setup failed: victim marker message missing" + ); + let before_len = before_messages.len(); + + let forged = IncomingMessage::new("test", "test-user", "attacker turn") + .with_thread(foreign_thread_id.to_string()); + rig.send_incoming(forged).await; + let responses = rig.wait_for_responses(1, Duration::from_secs(20)).await; + assert_eq!( + responses.len(), + 1, + "expected one assistant response for forged-thread request" + ); + assert_safe_thread_rejection(&responses[0]); + + let captured = rig.captured_llm_requests(); + assert!( + captured.is_empty(), + "forged thread-id request should be rejected before any LLM call" + ); + let prompt_dump = captured + .iter() + .flat_map(|req| req.iter().map(|m| m.content.as_str())) + .collect::>() + .join("\n"); + assert!( + !prompt_dump.contains(&marker), + "forged thread_id leaked foreign marker into LLM prompt context: {prompt_dump}" + ); + + let after_messages = store + .list_conversation_messages(foreign_thread_id) + .await + .expect("failed to read victim conversation after forged send"); + assert_eq!( + after_messages.len(), + before_len, + "forged thread_id wrote new messages into victim conversation" + ); + assert!( + after_messages + .iter() + .all(|m| m.content != "attacker turn" && m.content != "safe response"), + "forged request content was persisted to victim conversation" + ); + + rig.shutdown(); + } + + #[tokio::test] + async fn forged_nonexistent_thread_id_is_rejected_and_followup_request_still_works() { + let trace = LlmTrace::single_turn( + "thread-id-isolation-nonexistent", + "real follow-up turn", + vec![TraceStep { + request_hint: None, + response: TraceResponse::Text { + content: "safe response".to_string(), + input_tokens: 12, + output_tokens: 4, + }, + expected_tool_results: Vec::new(), + }], + ); + + let rig = TestRigBuilder::new().with_trace(trace).build().await; + + let forged_thread_id = Uuid::new_v4(); + let store = rig.database(); + + let forged = IncomingMessage::new("test", "test-user", "attacker turn") + .with_thread(forged_thread_id.to_string()); + rig.send_incoming(forged).await; + let responses = rig.wait_for_responses(1, Duration::from_secs(20)).await; + assert_eq!( + responses.len(), + 1, + "expected one response for forged nonexistent-thread request" + ); + assert_safe_thread_rejection(&responses[0]); + assert!( + rig.captured_llm_requests().is_empty(), + "forged nonexistent thread-id request should be rejected before any LLM call" + ); + assert!( + store + .get_conversation_metadata(forged_thread_id) + .await + .expect("get metadata for forged thread id") + .is_none(), + "forged nonexistent thread id must not create a conversation row" + ); + + rig.send_message("real follow-up turn").await; + let responses = rig.wait_for_responses(2, Duration::from_secs(20)).await; + assert_eq!( + responses.len(), + 2, + "expected follow-up response after rejection" + ); + assert_eq!( + responses[1].content, "safe response", + "follow-up valid request should still be handled normally" + ); + assert_eq!( + rig.captured_llm_requests().len(), + 1, + "only follow-up request should reach LLM" + ); + + rig.shutdown(); + } +}