From 954bfd1d757b765bd94422fa36a8e0de214aa112 Mon Sep 17 00:00:00 2001 From: xiaoxixi Date: Tue, 14 Jul 2026 11:49:02 +0800 Subject: [PATCH] fix(messaging): make persistence and delivery explicit --- src/bus/dispatcher.rs | 81 +++++++++-- src/bus/message.rs | 9 ++ src/bus/mod.rs | 23 ++++ src/channels/cli_chat.rs | 1 + src/gateway/mod.rs | 2 + src/session/session.rs | 288 ++++++++++++++++++++------------------- 6 files changed, 259 insertions(+), 145 deletions(-) diff --git a/src/bus/dispatcher.rs b/src/bus/dispatcher.rs index e2db5fc..afe0d7a 100644 --- a/src/bus/dispatcher.rs +++ b/src/bus/dispatcher.rs @@ -56,10 +56,14 @@ impl OutboundDispatcher { if sender.as_ref().is_none_or(mpsc::Sender::is_closed) { let Some(channel) = self.channel_manager.get_channel(&msg.channel).await else { tracing::warn!(channel = %msg.channel, "No channel found for message"); + msg.complete_delivery(Err(format!("channel not found: {}", msg.channel))); continue; }; let (new_sender, receiver) = mpsc::channel(LANE_CAPACITY); - self.spawn_lane(channel, receiver, msg.channel.clone(), msg.chat_id.clone()); + if !self.spawn_lane(channel, receiver, msg.channel.clone(), msg.chat_id.clone()) { + msg.complete_delivery(Err("dispatcher is shutting down".to_string())); + continue; + } lanes.insert(lane_key.clone(), new_sender.clone()); sender = Some(new_sender); } @@ -77,6 +81,7 @@ impl OutboundDispatcher { capacity = LANE_CAPACITY, "Outbound lane full; rejecting message instead of blocking other destinations" ); + msg.complete_delivery(Err("outbound lane is full".to_string())); } Err(mpsc::error::TrySendError::Closed(msg)) => { // The lane may have expired between the closed check and @@ -84,12 +89,24 @@ impl OutboundDispatcher { lanes.remove(&lane_key); let Some(channel) = self.channel_manager.get_channel(&msg.channel).await else { tracing::warn!(channel = %msg.channel, "No channel found for message"); + msg.complete_delivery(Err(format!("channel not found: {}", msg.channel))); continue; }; let (new_sender, receiver) = mpsc::channel(LANE_CAPACITY); - self.spawn_lane(channel, receiver, msg.channel.clone(), msg.chat_id.clone()); - if new_sender.try_send(msg).is_ok() { - lanes.insert(lane_key, new_sender); + if !self.spawn_lane(channel, receiver, msg.channel.clone(), msg.chat_id.clone()) + { + msg.complete_delivery(Err("dispatcher is shutting down".to_string())); + continue; + } + match new_sender.try_send(msg) { + Ok(()) => { + lanes.insert(lane_key, new_sender); + } + Err(error) => { + error.into_inner().complete_delivery(Err( + "outbound lane could not be restarted during shutdown".to_string(), + )); + } } } } @@ -102,7 +119,7 @@ impl OutboundDispatcher { mut receiver: mpsc::Receiver, channel_name: String, chat_id: String, - ) { + ) -> bool { self.task_supervisor.spawn( format!("outbound-lane:{channel_name}:{chat_id}"), async move { @@ -111,7 +128,8 @@ impl OutboundDispatcher { Ok(Some(msg)) => msg, Ok(None) | Err(_) => break, }; - if let Err(error) = Self::send_with_retry(&*channel, msg).await { + let result = Self::send_with_retry(&*channel, &msg).await; + if let Err(error) = &result { tracing::error!( channel = %channel_name, chat_id = %chat_id, @@ -119,14 +137,15 @@ impl OutboundDispatcher { "Failed to send message after retries" ); } + msg.complete_delivery(result.map_err(|error| error.to_string())); } }, - ); + ) } async fn send_with_retry( channel: &dyn Channel, - msg: OutboundMessage, + msg: &OutboundMessage, ) -> Result<(), ChannelError> { const DELAYS: &[u64] = &[1, 2, 4]; @@ -201,6 +220,7 @@ mod tests { reply_to: None, media: vec![], metadata: HashMap::new(), + delivery: None, } } @@ -247,4 +267,49 @@ mod tests { task.abort(); supervisor.shutdown(Duration::from_secs(1)).await; } + + #[tokio::test] + async fn confirmed_delivery_reports_missing_channel() { + let bus = MessageBus::new(8); + let manager = ChannelManager::with_bus( + Arc::new(crate::channels::CliChatChannel::new()), + bus.clone(), + ); + let supervisor = TaskSupervisor::new(); + let dispatcher = OutboundDispatcher::new(bus.clone(), manager, supervisor.clone()); + let task = tokio::spawn(async move { dispatcher.run().await }); + + let mut message = outbound("missing", "not delivered"); + message.channel = "missing".to_string(); + let error = bus.deliver_outbound(message).await.unwrap_err(); + + assert!(matches!(error, crate::bus::BusError::DeliveryFailed(_))); + task.abort(); + supervisor.shutdown(Duration::from_secs(1)).await; + } + + #[tokio::test] + async fn confirmed_delivery_waits_for_channel_send() { + let bus = MessageBus::new(8); + let manager = ChannelManager::with_bus( + Arc::new(crate::channels::CliChatChannel::new()), + bus.clone(), + ); + let channel = Arc::new(RecordingChannel { + sent: Mutex::new(Vec::new()), + notify: Notify::new(), + }); + manager.register_channel("recording", channel.clone()).await; + let supervisor = TaskSupervisor::new(); + let dispatcher = OutboundDispatcher::new(bus.clone(), manager, supervisor.clone()); + let task = tokio::spawn(async move { dispatcher.run().await }); + + bus.deliver_outbound(outbound("confirmed", "delivered")) + .await + .unwrap(); + + assert_eq!(channel.sent.lock().await.as_slice(), &["delivered"]); + task.abort(); + supervisor.shutdown(Duration::from_secs(1)).await; + } } diff --git a/src/bus/message.rs b/src/bus/message.rs index 6a60c46..9d8dd57 100644 --- a/src/bus/message.rs +++ b/src/bus/message.rs @@ -276,6 +276,15 @@ pub struct OutboundMessage { pub reply_to: Option, pub media: Vec, pub metadata: HashMap, + pub(crate) delivery: Option>>>, +} + +impl OutboundMessage { + pub(crate) fn complete_delivery(&self, result: Result<(), String>) { + if let Some(delivery) = &self.delivery { + delivery.send_replace(Some(result)); + } + } } // ============================================================================ diff --git a/src/bus/mod.rs b/src/bus/mod.rs index f22829e..bde9eb4 100644 --- a/src/bus/mod.rs +++ b/src/bus/mod.rs @@ -68,6 +68,25 @@ impl MessageBus { .map_err(|_| BusError::Closed) } + /// Publish an outbound message and wait for the dispatcher to report the + /// actual channel delivery result. + pub async fn deliver_outbound(&self, mut msg: OutboundMessage) -> Result<(), BusError> { + let (delivery_tx, mut delivery_rx) = tokio::sync::watch::channel(None); + msg.delivery = Some(delivery_tx); + self.publish_outbound(msg).await?; + + tokio::time::timeout(std::time::Duration::from_secs(120), async { + loop { + delivery_rx.changed().await.map_err(|_| BusError::Closed)?; + if let Some(result) = delivery_rx.borrow().clone() { + return result.map_err(BusError::DeliveryFailed); + } + } + }) + .await + .map_err(|_| BusError::DeliveryTimedOut)? + } + /// Consume an outbound message (Dispatcher -> Bus) pub async fn consume_outbound(&self) -> Option { self.outbound_rx.lock().await.recv().await @@ -95,12 +114,16 @@ impl MessageBus { #[derive(Debug)] pub enum BusError { Closed, + DeliveryFailed(String), + DeliveryTimedOut, } impl std::fmt::Display for BusError { fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { match self { BusError::Closed => write!(f, "Bus channel closed"), + BusError::DeliveryFailed(error) => write!(f, "Outbound delivery failed: {error}"), + BusError::DeliveryTimedOut => write!(f, "Outbound delivery confirmation timed out"), } } } diff --git a/src/channels/cli_chat.rs b/src/channels/cli_chat.rs index da34e7b..b55fcda 100644 --- a/src/channels/cli_chat.rs +++ b/src/channels/cli_chat.rs @@ -650,6 +650,7 @@ mod tests { reply_to: None, media: Vec::new(), metadata: Default::default(), + delivery: None, }) .await .unwrap(); diff --git a/src/gateway/mod.rs b/src/gateway/mod.rs index ee1613c..6587926 100644 --- a/src/gateway/mod.rs +++ b/src/gateway/mod.rs @@ -226,6 +226,7 @@ impl GatewayState { reply_to: None, media: vec![], metadata: inbound.forwarded_metadata, + delivery: None, }; if let Err(e) = bus.publish_outbound(outbound).await { tracing::error!(error = %e, "Failed to publish outbound"); @@ -239,6 +240,7 @@ impl GatewayState { reply_to: None, media: vec![], metadata: inbound.forwarded_metadata, + delivery: None, }; if let Err(e) = bus.publish_outbound(outbound).await { tracing::error!(error = %e, "Failed to publish outbound"); diff --git a/src/session/session.rs b/src/session/session.rs index c97d934..447a6c4 100644 --- a/src/session/session.rs +++ b/src/session/session.rs @@ -98,6 +98,9 @@ pub struct Session { /// not overwrite a session that was changed by a command such as /clear or /// /delete while the slow work was in flight. state_version: u64, + /// Serializes durable mutations while allowing the session state mutex to + /// be released during SQLite I/O. + persistence_lock: Arc>, } /// A task to be processed by the per-session agent worker @@ -171,6 +174,7 @@ impl Session { current_cancel: None, worker_generation: 0, state_version: 0, + persistence_lock: Arc::new(Mutex::new(())), }) } @@ -346,6 +350,7 @@ impl Session { current_cancel: None, worker_generation: 0, state_version: 0, + persistence_lock: Arc::new(Mutex::new(())), }) } @@ -354,22 +359,6 @@ impl Session { self.id.to_string() } - /// 添加消息到历史并持久化到 Storage - /// 如果 `persist` 为 false,只更新内存(用于 compaction 场景) - pub async fn add_message( - &mut self, - message: ChatMessage, - persist: bool, - ) -> Result<(), StorageError> { - let message_id = message.id.clone(); - let snapshot = self.add_message_in_memory(message, persist); - if let Err(error) = persist_added_message(snapshot).await { - self.rollback_message_suffix(std::slice::from_ref(&message_id)); - return Err(error); - } - Ok(()) - } - fn add_message_in_memory( &mut self, message: ChatMessage, @@ -480,30 +469,6 @@ impl Session { &self.messages } - /// 清除历史消息 - pub fn clear_history(&mut self) { - let len = self.messages.len(); - self.messages.clear(); - self.seq_counter = 1; - self.total_message_count = 0; - self.message_count = 0; - self.state_version = self.state_version.wrapping_add(1); - #[cfg(debug_assertions)] - tracing::debug!(session_id = %self.id, previous_len = len, "Chat history cleared"); - } - - /// 重置对话上下文 - pub fn reset_context(&mut self) { - let len = self.messages.len(); - self.messages.clear(); - self.seq_counter = 1; - self.total_message_count = 0; - self.message_count = 0; - self.state_version = self.state_version.wrapping_add(1); - #[cfg(debug_assertions)] - tracing::debug!(session_id = %self.id, previous_len = len, "Chat context reset in memory"); - } - pub fn create_user_message(&self, content: &str, media_refs: Vec) -> ChatMessage { if media_refs.is_empty() { ChatMessage::user(content) @@ -535,14 +500,6 @@ impl Session { message } - /// 将 session 元数据写回 Storage - pub async fn persist_session_meta(&self) -> Result<(), StorageError> { - if let Some((storage, meta)) = self.session_meta_snapshot() { - storage.upsert_session(&meta).await?; - } - Ok(()) - } - fn session_meta_snapshot( &self, ) -> Option<(StdArc, crate::storage::session::SessionMeta)> { @@ -1079,6 +1036,7 @@ impl SessionManager { reply_to: None, media: vec![], metadata: std::collections::HashMap::new(), + delivery: None, }; let _ = sm_bus.publish_outbound(outbound).await; } @@ -1701,17 +1659,28 @@ impl SessionManager { ) -> Result<(), AgentError> { // Update in-memory session let session = self.get_or_create_session(session_id).await?; - let mut session_guard = session.lock().await; - session_guard.title = title.to_string(); - session_guard - .persist_session_meta() - .await - .map_err(|e| AgentError::Other(format!("failed to rename dialog: {}", e)))?; + let persistence_lock = { session.lock().await.persistence_lock.clone() }; + let _persistence_guard = persistence_lock.lock().await; + let meta_snapshot = { + let mut session_guard = session.lock().await; + session_guard.title = title.to_string(); + session_guard.state_version = session_guard.state_version.wrapping_add(1); + session_guard.session_meta_snapshot() + }; + if let Some((storage, meta)) = meta_snapshot { + storage + .upsert_session(&meta) + .await + .map_err(|e| AgentError::Other(format!("failed to rename dialog: {}", e)))?; + } Ok(()) } pub async fn delete_dialog(&self, session_id: &UnifiedSessionId) -> Result<(), AgentError> { let session_id_str = session_id.to_string(); + let session = self.get_or_create_session(session_id).await?; + let persistence_lock = { session.lock().await.persistence_lock.clone() }; + let _persistence_guard = persistence_lock.lock().await; // Soft delete from Storage self.storage @@ -1730,6 +1699,9 @@ impl SessionManager { pub async fn archive_dialog(&self, session_id: &UnifiedSessionId) -> Result<(), AgentError> { let session_id_str = session_id.to_string(); + let session = self.get_or_create_session(session_id).await?; + let persistence_lock = { session.lock().await.persistence_lock.clone() }; + let _persistence_guard = persistence_lock.lock().await; self.storage .archive_session(&session_id_str) .await @@ -1842,22 +1814,18 @@ impl SessionManager { ) -> Result<(), AgentError> { let unified_id = self.resolve_dialog_id(channel, chat_id).await?; let session = self.get_or_create_session(&unified_id).await?; - { - let mut guard = session.lock().await; - let source = MessageSource { - kind: SourceKind::SystemNotification, - from_channel: None, - from_session: None, - from_user_id: None, - system_name: Some(system_name.to_string()), - task_id: task_id.map(|s| s.to_string()), - }; - let msg = ChatMessage::assistant_with_source(content, source); - guard - .add_message(msg, true) - .await - .map_err(|e| AgentError::Other(format!("persist error: {}", e)))?; - } + let source = MessageSource { + kind: SourceKind::SystemNotification, + from_channel: None, + from_session: None, + from_user_id: None, + system_name: Some(system_name.to_string()), + task_id: task_id.map(|s| s.to_string()), + }; + let msg = ChatMessage::assistant_with_source(content, source); + append_persisted_messages(&session, vec![msg]) + .await + .map_err(|e| AgentError::Other(format!("persist error: {}", e)))?; let outbound = OutboundMessage { channel: channel.to_string(), @@ -1866,9 +1834,10 @@ impl SessionManager { reply_to: None, media: vec![], metadata: HashMap::new(), + delivery: None, }; self.bus - .publish_outbound(outbound) + .deliver_outbound(outbound) .await .map_err(|e| AgentError::Other(format!("bus publish error: {}", e)))?; @@ -2014,6 +1983,8 @@ async fn maybe_generate_title_outside_lock(session: Arc>) -> Resu .map_err(|e| AgentError::Other(format!("LLM call failed: {}", e)))?; let title = response.content.trim().to_string(); + let persistence_lock = { session.lock().await.persistence_lock.clone() }; + let _persistence_guard = persistence_lock.lock().await; let meta_snapshot = { let mut guard = session.lock().await; if guard.apply_generated_title(title) { @@ -2033,22 +2004,6 @@ async fn maybe_generate_title_outside_lock(session: Arc>) -> Resu Ok(()) } -async fn persist_added_message( - snapshot: Option, -) -> Result<(), StorageError> { - let Some((storage, session_id, msg_meta, session_meta)) = snapshot else { - return Ok(()); - }; - - storage - .persist_message_batch_with_retry( - &session_id, - std::slice::from_ref(&msg_meta), - &session_meta, - ) - .await -} - async fn persist_added_messages( snapshots: Vec>, ) -> Result<(), StorageError> { @@ -2081,6 +2036,32 @@ async fn persist_added_messages( .await } +async fn append_persisted_messages( + session: &Arc>, + messages: Vec, +) -> Result<(), StorageError> { + if messages.is_empty() { + return Ok(()); + } + + let persistence_lock = { session.lock().await.persistence_lock.clone() }; + let _persistence_guard = persistence_lock.lock().await; + let message_ids: Vec<_> = messages.iter().map(|message| message.id.clone()).collect(); + let snapshots = { + let mut guard = session.lock().await; + messages + .into_iter() + .map(|message| guard.add_message_in_memory(message, true)) + .collect() + }; + + if let Err(error) = persist_added_messages(snapshots).await { + session.lock().await.rollback_message_suffix(&message_ids); + return Err(error); + } + Ok(()) +} + fn spawn_agent_worker( mut task_rx: mpsc::Receiver, session: Arc>, @@ -2122,6 +2103,7 @@ fn spawn_agent_worker( reply_to: None, media: vec![], metadata, + delivery: None, }; let _ = bus.publish_outbound(outbound).await; } @@ -2134,6 +2116,30 @@ fn spawn_agent_worker( // /stop and other commands are not blocked behind slow I/O or // LLM-backed compaction. let skills_prompt = skills_loader.build_skills_prompt(); + let user_message = { + let guard = session.lock().await; + if guard.worker_generation != worker_gen { + return; + } + let media_refs: Vec = + task.media.iter().map(MediaItem::to_media_ref).collect(); + guard.create_user_message(&task.content, media_refs) + }; + if let Err(e) = append_persisted_messages(&session, vec![user_message]).await { + tracing::error!(error = %e, "Failed to persist user message"); + let err_outbound = OutboundMessage { + channel: task_chan.clone(), + chat_id: task_cid.clone(), + content: "Failed to save your message, please try again.".to_string(), + reply_to: None, + media: vec![], + metadata: HashMap::new(), + delivery: None, + }; + let _ = bus.publish_outbound(err_outbound).await; + continue 'tasks; + } + let (agent, history_raw, mut compressor, base_version, cancel_rx) = { let mut guard = session.lock().await; @@ -2141,31 +2147,6 @@ fn spawn_agent_worker( return; // stale worker } - let media_refs: Vec = - task.media.iter().map(|m| m.to_media_ref()).collect(); - let user_message = guard.create_user_message(&task.content, media_refs); - let user_message_id = user_message.id.clone(); - let user_persist = guard.add_message_in_memory(user_message, true); - if let Err(e) = persist_added_message(user_persist).await { - guard.rollback_message_suffix(std::slice::from_ref(&user_message_id)); - drop(guard); - tracing::error!(error = %e, "Failed to persist user message"); - let err_outbound = OutboundMessage { - channel: task_chan.clone(), - chat_id: task_cid.clone(), - content: "Failed to save your message, please try again." - .to_string(), - reply_to: None, - media: vec![], - metadata: HashMap::new(), - }; - let _ = bus.publish_outbound(err_outbound).await; - continue 'tasks; - } - if guard.worker_generation != worker_gen { - return; - } - let history_raw = guard.get_history().to_vec(); let agent = match guard.create_agent_with_notify(notify_tx) { @@ -2180,6 +2161,7 @@ fn spawn_agent_worker( reply_to: None, media: vec![], metadata: HashMap::new(), + delivery: None, }; let _ = bus.publish_outbound(err_outbound).await; continue 'tasks; @@ -2332,6 +2314,7 @@ fn spawn_agent_worker( reply_to: None, media: vec![], metadata: HashMap::new(), + delivery: None, }; let _ = bus2.publish_outbound(err_outbound).await; return; @@ -2390,6 +2373,7 @@ fn spawn_agent_worker( reply_to: None, media: vec![], metadata: HashMap::new(), + delivery: None, }; let _ = bus2.publish_outbound(err_outbound).await; return; @@ -2405,30 +2389,25 @@ fn spawn_agent_worker( reply_to: None, media: vec![], metadata: HashMap::new(), + delivery: None, }; let _ = bus2.publish_outbound(err_outbound).await; return; } }; - let response = { - let mut guard = session2.lock().await; - let mut persist_snapshots = Vec::new(); - let mut message_ids = Vec::new(); - for msg in result.emitted_messages { - message_ids.push(msg.id.clone()); - persist_snapshots.push(guard.add_message_in_memory(msg, true)); - } - let sent_count = guard.messages.len(); - guard.compressor.set_last_api_info(sent_count, result.total_tokens); - if let Err(e) = persist_added_messages(persist_snapshots).await { - guard.rollback_message_suffix(&message_ids); + let response_content = result.final_response.content; + let total_tokens = result.total_tokens; + let response = + if let Err(e) = append_persisted_messages(&session2, result.emitted_messages).await { tracing::error!(error = %e, "Failed to atomically persist agent turn"); None } else { - Some(result.final_response.content) - } - }; + let mut guard = session2.lock().await; + let sent_count = guard.messages.len(); + guard.compressor.set_last_api_info(sent_count, total_tokens); + Some(response_content) + }; let Some(response) = response else { let err_outbound = OutboundMessage { @@ -2439,6 +2418,7 @@ fn spawn_agent_worker( reply_to: None, media: vec![], metadata: HashMap::new(), + delivery: None, }; let _ = bus2.publish_outbound(err_outbound).await; return; @@ -2455,6 +2435,7 @@ fn spawn_agent_worker( reply_to: None, media: vec![], metadata: HashMap::new(), + delivery: None, }; let _ = bus2.publish_outbound(outbound).await; }; @@ -2533,6 +2514,8 @@ impl SessionManager { unified_id: &UnifiedSessionId, ) -> Result<(), AgentError> { let session = self.get_or_create_session(unified_id).await?; + let persistence_lock = { session.lock().await.persistence_lock.clone() }; + let _persistence_guard = persistence_lock.lock().await; let (storage, session_id, meta_snapshot) = { let mut session_guard = session.lock().await; // Clear in-memory @@ -2618,15 +2601,11 @@ impl OutboundMessenger for SessionManager { format!("[message from {}] \n{}", origin, content) }; - // Write source-tagged assistant message to target session history - { - let mut guard = session.lock().await; - let msg = ChatMessage::assistant_with_source(marked_content.clone(), source); - guard - .add_message(msg, true) - .await - .map_err(|e| e.to_string())?; - } + // Write the same text and media delivered to the target into history. + let msg = outbound_history_message(marked_content.clone(), source, &media); + append_persisted_messages(&session, vec![msg]) + .await + .map_err(|e| e.to_string())?; // Restore active dialog if source and target share channel:chat_id but differ in dialog_id if let Some(ref origin_id) = origin_id { @@ -2653,9 +2632,10 @@ impl OutboundMessenger for SessionManager { reply_to: None, media, metadata: HashMap::new(), + delivery: None, }; self.bus - .publish_outbound(outbound) + .deliver_outbound(outbound) .await .map_err(|e| e.to_string())?; @@ -2663,6 +2643,16 @@ impl OutboundMessenger for SessionManager { } } +fn outbound_history_message( + content: impl Into, + source: MessageSource, + media: &[MediaItem], +) -> ChatMessage { + let mut message = ChatMessage::assistant_with_source(content, source); + message.media_refs = media.iter().map(MediaItem::to_media_ref).collect(); + message +} + fn format_task_notification( task_id: &str, status: &crate::agent::TaskStatus, @@ -2680,3 +2670,27 @@ fn format_task_notification( crate::agent::TaskStatus::TimedOut => format!("📋 后台任务超时\n\n任务 ID: {}", task_id), } } + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn outbound_history_preserves_delivered_media() { + let source = MessageSource { + kind: SourceKind::CrossChannel, + from_channel: Some("cli_chat".into()), + from_session: Some("cli_chat:source:dialog".into()), + from_user_id: None, + system_name: None, + task_id: None, + }; + let media = vec![MediaItem::new("/tmp/report.pdf", "file")]; + + let message = outbound_history_message("report", source, &media); + + assert_eq!(message.media_refs.len(), 1); + assert_eq!(message.media_refs[0].path, "/tmp/report.pdf"); + assert_eq!(message.media_refs[0].media_type, "file"); + } +}