From c6e022f6cb4cace408ae1710a598f4a6ed440fd1 Mon Sep 17 00:00:00 2001 From: xiaoxixi Date: Fri, 17 Jul 2026 17:30:38 +0800 Subject: [PATCH] feat: stream turns through Feishu cards --- resources/templates/config.example.json | 4 +- src/channels/feishu.rs | 494 +++++++++++++++++++++++- src/config/mod.rs | 9 + src/delivery/coordinator.rs | 118 +++++- src/gateway/mod.rs | 1 + src/session/session.rs | 67 +++- 6 files changed, 648 insertions(+), 45 deletions(-) diff --git a/resources/templates/config.example.json b/resources/templates/config.example.json index fb18857..bedc9ad 100644 --- a/resources/templates/config.example.json +++ b/resources/templates/config.example.json @@ -70,7 +70,9 @@ "allow_from": ["*"], "agent": "default", "media_dir": "~/.picobot/media/feishu", - "reaction_emoji": "Typing" + "reaction_emoji": "Typing", + "live_updates": false, + "live_update_interval_ms": 500 } }, "memory": { diff --git a/src/channels/feishu.rs b/src/channels/feishu.rs index 9f6b354..30d1d02 100644 --- a/src/channels/feishu.rs +++ b/src/channels/feishu.rs @@ -6,15 +6,15 @@ use std::time::{Duration, Instant}; use async_trait::async_trait; use futures_util::{SinkExt, StreamExt}; use prost::{Message as ProstMessage, bytes::Bytes}; -use regex::Regex; use serde::Deserialize; use tokio::sync::{Mutex, RwLock}; use tokio::task::JoinHandle; use tokio_util::sync::CancellationToken; use crate::bus::{MediaItem, MessageBus, OutboundMessage}; -use crate::channels::base::{Channel, ChannelError}; +use crate::channels::base::{Channel, ChannelError, LivePolicy, TurnSink, TurnTarget}; use crate::config::FeishuChannelConfig; +use crate::session::{ToolStatus, TurnBlock, TurnSnapshot, TurnStatus}; const FEISHU_API_BASE: &str = "https://open.feishu.cn/open-apis"; const FEISHU_WS_BASE: &str = "https://open.feishu.cn"; @@ -1807,14 +1807,6 @@ async fn resolve_unique_path(dir: &Path, filename: &str) -> std::path::PathBuf { } impl FeishuChannel { - fn strip_thinking_tags(content: &str) -> String { - use std::sync::LazyLock; - static THINK_RE: LazyLock = - LazyLock::new(|| Regex::new(r"(?s).*?").unwrap()); - let stripped = THINK_RE.replace_all(content, ""); - stripped.trim().to_string() - } - /// Build a Card JSON 2.0 interactive card with a single markdown element. fn build_card_content(markdown: &str) -> String { serde_json::json!({ @@ -1849,7 +1841,7 @@ impl FeishuChannel { break; } - let end = start + Self::CARD_MARKDOWN_MAX_BYTES; + let end = text.floor_char_boundary(start + Self::CARD_MARKDOWN_MAX_BYTES); let search_region = &text[start..end]; let split_at = search_region .rfind('\n') @@ -1886,7 +1878,7 @@ impl FeishuChannel { receive_id: &str, receive_id_type: &str, card_content: &str, - ) -> Result<(), ChannelError> { + ) -> Result { let token = self.get_tenant_access_token().await?; let resp = self @@ -1912,6 +1904,12 @@ impl FeishuChannel { struct SendResp { code: i32, msg: String, + data: Option, + } + + #[derive(Deserialize)] + struct SendData { + message_id: String, } let send_resp: SendResp = resp.json().await.map_err(|e| { @@ -1925,10 +1923,244 @@ impl FeishuChannel { ))); } + send_resp + .data + .map(|data| data.message_id) + .filter(|message_id| !message_id.is_empty()) + .ok_or_else(|| ChannelError::Other("Feishu send response has no message_id".into())) + } + + async fn update_interactive_card( + &self, + message_id: &str, + card_content: &str, + ) -> Result<(), ChannelError> { + let token = self.get_tenant_access_token().await?; + let card: serde_json::Value = serde_json::from_str(card_content) + .map_err(|error| ChannelError::Other(format!("Invalid card JSON: {error}")))?; + let response = self + .http_client + .patch(format!("{}/im/v1/messages/{}", FEISHU_API_BASE, message_id)) + .header("Content-Type", "application/json") + .header("Authorization", format!("Bearer {}", token)) + .json(&serde_json::json!({ "card": card })) + .send() + .await + .map_err(|error| { + ChannelError::ConnectionError(format!("Update card HTTP error: {error}")) + })?; + + #[derive(Deserialize)] + struct UpdateResp { + code: i32, + msg: String, + } + let result: UpdateResp = response.json().await.map_err(|error| { + ChannelError::Other(format!("Parse update card response error: {error}")) + })?; + if result.code != 0 { + return Err(ChannelError::Other(format!( + "Update card failed: code={} msg={}", + result.code, result.msg + ))); + } Ok(()) } } +#[async_trait] +trait FeishuTurnApi: Send { + async fn create_card(&mut self, markdown: &str) -> Result; + async fn update_card(&mut self, message_id: &str, markdown: &str) -> Result<(), ChannelError>; + async fn cleanup(&mut self); +} + +struct FeishuTurnBackend { + channel: FeishuChannel, + receive_id: String, + receive_id_type: &'static str, + metadata: HashMap, +} + +#[async_trait] +impl FeishuTurnApi for FeishuTurnBackend { + async fn create_card(&mut self, markdown: &str) -> Result { + let card = FeishuChannel::build_card_content(markdown); + self.channel + .send_interactive_card(&self.receive_id, self.receive_id_type, &card) + .await + } + + async fn update_card(&mut self, message_id: &str, markdown: &str) -> Result<(), ChannelError> { + let card = FeishuChannel::build_card_content(markdown); + self.channel + .update_interactive_card(message_id, &card) + .await + } + + async fn cleanup(&mut self) { + self.channel + .remove_reaction_from_metadata(&self.metadata) + .await; + } +} + +struct FeishuTurnSink { + api: Box, + message_id: Option, + cleaned_up: bool, +} + +impl FeishuTurnSink { + fn new(api: Box) -> Self { + Self { + api, + message_id: None, + cleaned_up: false, + } + } + + async fn cleanup(&mut self) { + if !self.cleaned_up { + self.api.cleanup().await; + self.cleaned_up = true; + } + } + + async fn send_chunks(&mut self, chunks: &[String]) -> Result<(), ChannelError> { + for chunk in chunks { + self.api.create_card(chunk).await?; + } + Ok(()) + } + + async fn finish_snapshot(&mut self, snapshot: &TurnSnapshot) -> Result<(), ChannelError> { + let markdown = render_feishu_turn(snapshot); + let chunks = if markdown.is_empty() { + Vec::new() + } else { + FeishuChannel::split_markdown_chunks(&markdown) + }; + + let result = if chunks.is_empty() { + Ok(()) + } else if let Some(message_id) = self.message_id.clone() { + match self.api.update_card(&message_id, &chunks[0]).await { + Ok(()) => self.send_chunks(&chunks[1..]).await, + Err(error) => { + tracing::warn!(error = %error, "Final Feishu card update failed; sending complete fallback"); + self.send_chunks(&chunks).await + } + } + } else { + self.send_chunks(&chunks).await + }; + self.cleanup().await; + result + } +} + +#[async_trait] +impl TurnSink for FeishuTurnSink { + async fn update(&mut self, snapshot: &TurnSnapshot) -> Result<(), ChannelError> { + let markdown = render_feishu_turn(snapshot); + if markdown.is_empty() { + return Ok(()); + } + let live_markdown = truncate_feishu_live_markdown(&markdown); + if let Some(message_id) = self.message_id.clone() { + self.api.update_card(&message_id, &live_markdown).await + } else { + self.message_id = Some(self.api.create_card(&live_markdown).await?); + Ok(()) + } + } + + async fn finish(&mut self, snapshot: &TurnSnapshot) -> Result<(), ChannelError> { + self.finish_snapshot(snapshot).await + } + + async fn abort(&mut self, snapshot: &TurnSnapshot) -> Result<(), ChannelError> { + self.finish_snapshot(snapshot).await + } +} + +fn render_feishu_turn(snapshot: &TurnSnapshot) -> String { + let mut sections = Vec::new(); + for block in &snapshot.blocks { + match block { + TurnBlock::Reasoning { text, .. } if !text.trim().is_empty() => { + sections.push(format!( + "> **思考过程**\n> {}", + text.trim().replace('\n', "\n> ") + )); + } + TurnBlock::Assistant { text, .. } if !text.trim().is_empty() => { + sections.push(text.trim().to_string()); + } + TurnBlock::Tool { + name, + status, + preview, + .. + } => { + let status = match status { + ToolStatus::Running => "执行中", + ToolStatus::Completed => "已完成", + ToolStatus::Failed => "失败", + }; + let mut section = format!("> 🔧 **{name}** · {status}"); + if let Some(preview) = preview.as_deref().filter(|value| !value.trim().is_empty()) { + section.push_str("\n> "); + section.push_str(&preview.trim().replace('\n', "\n> ")); + } + sections.push(section); + } + _ => {} + } + } + + if sections.is_empty() { + if snapshot.status == TurnStatus::Failed { + sections.push(format!( + "⚠️ 回复失败:{}", + snapshot.error.as_deref().unwrap_or("未知错误") + )); + } else if snapshot.status == TurnStatus::Cancelled { + sections.push("已停止生成。".to_string()); + } else { + return String::new(); + } + } + + let status = match snapshot.status { + TurnStatus::Running => Some(match snapshot.phase { + crate::session::TurnPhase::Queued => "排队中", + crate::session::TurnPhase::Reasoning => "思考中", + crate::session::TurnPhase::Responding => "生成中", + crate::session::TurnPhase::Acting => "调用工具中", + crate::session::TurnPhase::Finalizing => "收尾中", + }), + TurnStatus::Cancelled => Some("已停止"), + TurnStatus::Failed => Some("失败"), + TurnStatus::Completed => None, + }; + if let Some(status) = status { + sections.push(format!("_{status}_")); + } + sections.join("\n\n") +} + +fn truncate_feishu_live_markdown(markdown: &str) -> String { + if markdown.len() <= FeishuChannel::CARD_MARKDOWN_MAX_BYTES { + return markdown.to_string(); + } + const SUFFIX: &str = "\n\n_内容仍在生成,已暂时截断…_"; + let limit = FeishuChannel::CARD_MARKDOWN_MAX_BYTES.saturating_sub(SUFFIX.len()); + let boundary = markdown.floor_char_boundary(limit); + format!("{}{SUFFIX}", &markdown[..boundary]) +} + #[async_trait] impl Channel for FeishuChannel { fn name(&self) -> &str { @@ -2043,11 +2275,37 @@ impl Channel for FeishuChannel { self.running.try_read().map(|r| *r).unwrap_or(false) } - async fn send(&self, msg: OutboundMessage) -> Result<(), ChannelError> { - let msg = OutboundMessage { - content: Self::strip_thinking_tags(&msg.content), - ..msg + fn live_policy(&self) -> LivePolicy { + if self.config.live_updates { + LivePolicy::Snapshot { + min_interval: Duration::from_millis( + self.config.live_update_interval_ms.clamp(250, 5_000), + ), + } + } else { + LivePolicy::FinalOnly + } + } + + fn presentation_policy(&self) -> crate::delivery::PresentationPolicy { + crate::delivery::PresentationPolicy::external(self.config.live_updates) + } + + async fn open_turn(&self, target: TurnTarget) -> Result, ChannelError> { + let (receive_id, receive_id_type) = if target.chat_id.starts_with("oc_") { + (target.chat_id, "chat_id") + } else { + (target.reply_to.unwrap_or(target.chat_id), "open_id") }; + Ok(Box::new(FeishuTurnSink::new(Box::new(FeishuTurnBackend { + channel: self.clone(), + receive_id, + receive_id_type, + metadata: target.metadata, + })))) + } + + async fn send(&self, msg: OutboundMessage) -> Result<(), ChannelError> { let receive_id = if msg.chat_id.starts_with("oc_") { &msg.chat_id } else { @@ -2252,6 +2510,53 @@ impl Channel for FeishuChannel { #[cfg(test)] mod tests { use super::*; + use crate::agent::TurnEvent; + use crate::delivery::{PresentationPolicy, project_snapshot}; + + #[derive(Default)] + struct MockTurnState { + created: Vec, + updated: Vec<(String, String)>, + cleanups: usize, + fail_updates: usize, + } + + struct MockTurnApi { + state: Arc>, + } + + #[async_trait] + impl FeishuTurnApi for MockTurnApi { + async fn create_card(&mut self, markdown: &str) -> Result { + let mut state = self.state.lock().await; + state.created.push(markdown.to_string()); + Ok(format!("card-{}", state.created.len())) + } + + async fn update_card( + &mut self, + message_id: &str, + markdown: &str, + ) -> Result<(), ChannelError> { + let mut state = self.state.lock().await; + if state.fail_updates > 0 { + state.fail_updates -= 1; + return Err(ChannelError::Other("card can no longer be edited".into())); + } + state + .updated + .push((message_id.to_string(), markdown.to_string())); + Ok(()) + } + + async fn cleanup(&mut self) { + self.state.lock().await.cleanups += 1; + } + } + + fn mock_sink(state: Arc>) -> FeishuTurnSink { + FeishuTurnSink::new(Box::new(MockTurnApi { state })) + } fn test_channel() -> FeishuChannel { FeishuChannel::new( @@ -2263,12 +2568,169 @@ mod tests { agent: String::new(), media_dir: String::new(), reaction_emoji: "THUMBSUP".to_string(), + live_updates: false, + live_update_interval_ms: 500, }, Path::new("/tmp"), ) .expect("test channel should be valid") } + #[tokio::test] + async fn turn_sink_creates_once_updates_same_card_and_cleans_up_at_finish() { + let state = Arc::new(Mutex::new(MockTurnState::default())); + let mut sink = mock_sink(state.clone()); + let (controller, emitter, _) = + crate::session::TurnController::start("feishu:chat:dialog", "message"); + emitter + .emit(TurnEvent::TextDelta { + iteration: 0, + delta: "hello".into(), + }) + .unwrap(); + sink.update(&project_snapshot( + &controller.snapshot(), + PresentationPolicy::external(true), + )) + .await + .unwrap(); + emitter + .emit(TurnEvent::TextDelta { + iteration: 0, + delta: " world".into(), + }) + .unwrap(); + sink.update(&project_snapshot( + &controller.snapshot(), + PresentationPolicy::external(true), + )) + .await + .unwrap(); + controller.complete(None); + sink.finish(&project_snapshot( + &controller.snapshot(), + PresentationPolicy::external(true), + )) + .await + .unwrap(); + + let state = state.lock().await; + assert_eq!(state.created.len(), 1); + assert_eq!(state.updated.len(), 2); + assert!(state.updated.iter().all(|(id, _)| id == "card-1")); + assert!(state.updated.last().unwrap().1.contains("hello world")); + assert_eq!(state.cleanups, 1); + } + + #[tokio::test] + async fn final_update_failure_sends_complete_fallback_and_cleanup_is_idempotent() { + let state = Arc::new(Mutex::new(MockTurnState::default())); + let mut sink = mock_sink(state.clone()); + let (controller, emitter, _) = + crate::session::TurnController::start("feishu:chat:dialog", "message"); + emitter + .emit(TurnEvent::TextDelta { + iteration: 0, + delta: "partial".into(), + }) + .unwrap(); + sink.update(&controller.snapshot()).await.unwrap(); + state.lock().await.fail_updates = 1; + emitter + .emit(TurnEvent::TextDelta { + iteration: 0, + delta: " final".into(), + }) + .unwrap(); + controller.complete(None); + sink.finish(&controller.snapshot()).await.unwrap(); + sink.finish(&controller.snapshot()).await.unwrap(); + + let state = state.lock().await; + assert_eq!(state.created.len(), 2); + assert!(state.created[1].contains("partial final")); + assert_eq!(state.cleanups, 1); + } + + #[tokio::test] + async fn final_only_sink_sends_no_fragments_and_abort_without_text_is_visible() { + let state = Arc::new(Mutex::new(MockTurnState::default())); + let mut sink = mock_sink(state.clone()); + let (controller, _emitter, _) = + crate::session::TurnController::start("feishu:chat:dialog", "message"); + controller.fail("provider unavailable"); + + sink.abort(&controller.snapshot()).await.unwrap(); + + let state = state.lock().await; + assert_eq!(state.created.len(), 1); + assert!(state.created[0].contains("provider unavailable")); + assert_eq!(state.updated.len(), 0); + assert_eq!(state.cleanups, 1); + } + + #[test] + fn external_projection_removes_reasoning_before_feishu_rendering() { + let (controller, emitter, _) = + crate::session::TurnController::start("feishu:chat:dialog", "message"); + emitter + .emit(TurnEvent::ReasoningDelta { + iteration: 0, + delta: "private".into(), + }) + .unwrap(); + emitter + .emit(TurnEvent::TextDelta { + iteration: 0, + delta: "public".into(), + }) + .unwrap(); + let projected = + project_snapshot(&controller.snapshot(), PresentationPolicy::external(true)); + + let markdown = render_feishu_turn(&projected); + assert!(markdown.contains("public")); + assert!(!markdown.contains("private")); + } + + #[test] + fn live_card_truncation_preserves_utf8_and_payload_limit() { + let markdown = "你".repeat(FeishuChannel::CARD_MARKDOWN_MAX_BYTES); + let truncated = truncate_feishu_live_markdown(&markdown); + + assert!(truncated.len() <= FeishuChannel::CARD_MARKDOWN_MAX_BYTES); + assert!(truncated.ends_with("_内容仍在生成,已暂时截断…_")); + } + + #[test] + fn final_card_chunking_preserves_long_utf8_content() { + let markdown = "你".repeat(FeishuChannel::CARD_MARKDOWN_MAX_BYTES); + let chunks = FeishuChannel::split_markdown_chunks(&markdown); + + assert!(chunks.len() > 1); + assert!( + chunks + .iter() + .all(|chunk| chunk.len() <= FeishuChannel::CARD_MARKDOWN_MAX_BYTES) + ); + assert_eq!(chunks.concat(), markdown); + } + + #[test] + fn live_policy_uses_configured_bounded_interval() { + let mut channel = test_channel(); + assert_eq!(channel.live_policy(), LivePolicy::FinalOnly); + channel.config.live_updates = true; + channel.config.live_update_interval_ms = 10; + + assert_eq!( + channel.live_policy(), + LivePolicy::Snapshot { + min_interval: Duration::from_millis(250) + } + ); + } + #[tokio::test] async fn stop_aborts_connection_task_that_ignores_cancellation() { let channel = test_channel(); diff --git a/src/config/mod.rs b/src/config/mod.rs index f05b366..a3e0e37 100644 --- a/src/config/mod.rs +++ b/src/config/mod.rs @@ -78,6 +78,11 @@ pub struct FeishuChannelConfig { /// Emoji type for message reactions (e.g. "THUMBSUP", "OK", "EYES"). #[serde(default = "default_reaction_emoji")] pub reaction_emoji: String, + /// Edit one card with latest Turn snapshots instead of sending only the final result. + #[serde(default)] + pub live_updates: bool, + #[serde(default = "default_feishu_live_update_interval_ms")] + pub live_update_interval_ms: u64, } fn default_allow_from() -> Vec { @@ -95,6 +100,10 @@ fn default_reaction_emoji() -> String { "Typing".to_string() } +fn default_feishu_live_update_interval_ms() -> u64 { + 500 +} + #[derive(Debug, Clone, Deserialize, Serialize)] pub struct ProviderConfig { #[serde(rename = "type")] diff --git a/src/delivery/coordinator.rs b/src/delivery/coordinator.rs index 86b2749..789f123 100644 --- a/src/delivery/coordinator.rs +++ b/src/delivery/coordinator.rs @@ -166,8 +166,8 @@ impl DeliveryCoordinator { &self, supervisor: &TaskSupervisor, route: SinkRoute, - snapshots: watch::Receiver>, - sink: Box, + mut snapshots: watch::Receiver>, + mut sink: Box, ) -> Result>, DeliveryError> { let SinkRoute { channel, @@ -178,17 +178,28 @@ impl DeliveryCoordinator { let (result_tx, result_rx) = oneshot::channel(); let coordinator = self.clone(); let task_name = format!("turn-delivery:{channel}:{chat_id}"); - let spawned = supervisor.spawn(task_name, async move { - let result = coordinator - .deliver( - &channel, - &chat_id, - live_policy, - presentation, - snapshots, - sink, - ) - .await; + let cancellation = supervisor.cancellation_token(); + let shutdown_snapshot = snapshots.clone(); + let spawned = supervisor.spawn_graceful(task_name, async move { + let mut delivery = Box::pin(coordinator.deliver_sink( + &channel, + &chat_id, + live_policy, + presentation, + &mut snapshots, + &mut *sink, + )); + let result = tokio::select! { + result = &mut delivery => result, + () = cancellation.cancelled() => { + drop(delivery); + let snapshot = shutdown_snapshot.borrow().clone(); + let projected = project_snapshot(&snapshot, presentation); + coordinator + .abort_for_shutdown(&channel, &chat_id, &mut *sink, &projected) + .await + } + }; if let Err(error) = &result { tracing::error!(channel, chat_id, error = %error, "Turn delivery failed"); } @@ -208,6 +219,26 @@ impl DeliveryCoordinator { presentation: PresentationPolicy, mut snapshots: watch::Receiver>, mut sink: Box, + ) -> Result<(), DeliveryError> { + self.deliver_sink( + channel, + chat_id, + live_policy, + presentation, + &mut snapshots, + &mut *sink, + ) + .await + } + + async fn deliver_sink( + &self, + channel: &str, + chat_id: &str, + live_policy: LivePolicy, + presentation: PresentationPolicy, + snapshots: &mut watch::Receiver>, + sink: &mut dyn TurnSink, ) -> Result<(), DeliveryError> { let target_lock = self.write_locks.for_target(channel, chat_id); let min_interval = match live_policy { @@ -221,9 +252,7 @@ impl DeliveryCoordinator { let snapshot = snapshots.borrow_and_update().clone(); if snapshot.status != TurnStatus::Running { let projected = project_snapshot(&snapshot, presentation); - return self - .deliver_terminal(&target_lock, &mut *sink, &projected) - .await; + return self.deliver_terminal(&target_lock, sink, &projected).await; } if let Some(interval) = min_interval { @@ -234,7 +263,7 @@ impl DeliveryCoordinator { let latest = snapshots.borrow_and_update().clone(); if latest.status != TurnStatus::Running { let projected = project_snapshot(&latest, presentation); - return self.deliver_terminal(&target_lock, &mut *sink, &projected).await; + return self.deliver_terminal(&target_lock, sink, &projected).await; } } () = sleep_until(next_update_at) => break, @@ -244,9 +273,7 @@ impl DeliveryCoordinator { let latest = snapshots.borrow_and_update().clone(); if latest.status != TurnStatus::Running { let projected = project_snapshot(&latest, presentation); - return self - .deliver_terminal(&target_lock, &mut *sink, &projected) - .await; + return self.deliver_terminal(&target_lock, sink, &projected).await; } let projected = project_snapshot(&latest, presentation); let _guard = target_lock.lock().await; @@ -272,6 +299,22 @@ impl DeliveryCoordinator { } } + async fn abort_for_shutdown( + &self, + channel: &str, + chat_id: &str, + sink: &mut dyn TurnSink, + snapshot: &TurnSnapshot, + ) -> Result<(), DeliveryError> { + let target_lock = self.write_locks.for_target(channel, chat_id); + let _guard = target_lock.lock().await; + match timeout(self.sink_call_timeout, sink.abort(snapshot)).await { + Ok(Ok(())) => Ok(()), + Ok(Err(error)) => Err(DeliveryError::FinalFailed(error)), + Err(_) => Err(DeliveryError::FinalTimedOut), + } + } + async fn deliver_terminal( &self, target_lock: &Arc>, @@ -595,4 +638,39 @@ mod tests { assert_eq!(channel.opened.load(Ordering::SeqCst), 1); assert_eq!(state.terminal.lock().await.len(), 1); } + + #[tokio::test] + async fn supervisor_shutdown_aborts_sink_and_waits_for_cleanup() { + let state = Arc::new(SinkState::default()); + let (_controller, emitter, receiver) = TurnController::start("session", "message"); + emitter + .emit(TurnEvent::TextDelta { + iteration: 0, + delta: "partial".into(), + }) + .unwrap(); + let supervisor = TaskSupervisor::new(); + let coordinator = DeliveryCoordinator::for_test(Duration::from_secs(1), []); + let result = coordinator + .spawn_sink( + &supervisor, + SinkRoute { + channel: "channel".into(), + chat_id: "chat".into(), + live_policy: LivePolicy::FinalOnly, + presentation: PresentationPolicy::unattended(), + }, + receiver, + sink(state.clone()), + ) + .unwrap(); + tokio::task::yield_now().await; + + supervisor.shutdown(Duration::from_secs(1)).await; + + assert!(result.await.unwrap().is_ok()); + let terminal = state.terminal.lock().await; + assert_eq!(terminal.len(), 1); + assert_eq!(terminal[0].status, TurnStatus::Running); + } } diff --git a/src/gateway/mod.rs b/src/gateway/mod.rs index c7e0f10..37eef9b 100644 --- a/src/gateway/mod.rs +++ b/src/gateway/mod.rs @@ -288,6 +288,7 @@ impl GatewayState { &inbound.chat_id, &inbound.content, inbound.media, + inbound.forwarded_metadata.clone(), ).await { Ok(crate::session::session::HandleResult::AgentResponse(content)) => { let outbound = crate::bus::OutboundMessage { diff --git a/src/session/session.rs b/src/session/session.rs index 3ec08d7..11bac93 100644 --- a/src/session/session.rs +++ b/src/session/session.rs @@ -25,6 +25,15 @@ fn outbound_session_metadata(session_id: &str) -> HashMap { HashMap::from([("_session_id".to_string(), session_id.to_string())]) } +fn outbound_turn_metadata( + session_id: &str, + forwarded: &HashMap, +) -> HashMap { + let mut metadata = forwarded.clone(); + metadata.insert("_session_id".to_string(), session_id.to_string()); + metadata +} + tokio::task_local! { pub(super) static CURRENT_SOURCE_SESSION: Option; } @@ -117,6 +126,30 @@ mod cancelled_partial_tests { use super::*; use crate::agent::TurnEvent; + #[test] + fn turn_metadata_preserves_channel_cleanup_fields() { + let forwarded = HashMap::from([ + ("feishu.message_id".to_string(), "message-1".to_string()), + ("feishu.reaction_id".to_string(), "reaction-1".to_string()), + ("_session_id".to_string(), "stale".to_string()), + ]); + + let metadata = outbound_turn_metadata("session-1", &forwarded); + + assert_eq!( + metadata.get("_session_id").map(String::as_str), + Some("session-1") + ); + assert_eq!( + metadata.get("feishu.message_id").map(String::as_str), + Some("message-1") + ); + assert_eq!( + metadata.get("feishu.reaction_id").map(String::as_str), + Some("reaction-1") + ); + } + #[test] fn visible_partial_text_becomes_cancelled_persisted_message() { let (controller, emitter, _) = TurnController::start("session", "message-id"); @@ -233,6 +266,7 @@ struct AgentTask { chat_id: String, content: String, media: Vec, + forwarded_metadata: HashMap, } #[derive(Clone)] @@ -2176,6 +2210,7 @@ impl SessionManager { chat_id: &str, content: &str, media: Vec, + forwarded_metadata: HashMap, ) -> Result { let unified_id = self.resolve_dialog_id(channel, chat_id).await?; tracing::debug!(unified_id = %unified_id, "handle_message resolved unified_id"); @@ -2213,6 +2248,7 @@ impl SessionManager { chat_id: chat_id.to_string(), content: content.to_string(), media, + forwarded_metadata, }; let session_clone = session.clone(); let unified_str = unified_id.to_string(); @@ -2351,6 +2387,7 @@ fn spawn_agent_worker( 'tasks: while let Some(task) = task_rx.recv().await { let task_chan = task.channel.clone(); let task_cid = task.chat_id.clone(); + let task_metadata = task.forwarded_metadata.clone(); let notification_session_id = unified_str.clone(); let (notify_tx, mut notify_rx) = mpsc::unbounded_channel(); @@ -2407,7 +2444,7 @@ fn spawn_agent_worker( content: "Failed to save your message, please try again.".to_string(), reply_to: None, media: vec![], - metadata: outbound_session_metadata(&unified_str), + metadata: outbound_turn_metadata(&unified_str, &task_metadata), delivery: None, }; let _ = bus.publish_outbound(err_outbound).await; @@ -2434,7 +2471,7 @@ fn spawn_agent_worker( .to_string(), reply_to: None, media: vec![], - metadata: outbound_session_metadata(&unified_str), + metadata: outbound_turn_metadata(&unified_str, &task_metadata), delivery: None, }; let _ = bus.publish_outbound(err_outbound).await; @@ -2561,7 +2598,7 @@ fn spawn_agent_worker( chat_id: task_cid.clone(), session_id: unified_str.clone(), reply_to: None, - metadata: outbound_session_metadata(&unified_str), + metadata: outbound_turn_metadata(&unified_str, &task_metadata), }, turn_receiver, ) @@ -2604,6 +2641,7 @@ fn spawn_agent_worker( let chan2 = task_chan.clone(); let cid2 = task_cid.clone(); let unified_str2 = unified_str.clone(); + let task_metadata2 = task_metadata.clone(); let turn_lifecycle = &turn_controller; let process_future = async move { let response_session_id = unified_str2.clone(); @@ -2656,8 +2694,9 @@ fn spawn_agent_worker( .to_string(), reply_to: None, media: vec![], - metadata: outbound_session_metadata( + metadata: outbound_turn_metadata( &response_session_id, + &task_metadata2, ), delivery: None, }; @@ -2732,7 +2771,10 @@ fn spawn_agent_worker( content: format!("Processing error: {}", e), reply_to: None, media: vec![], - metadata: outbound_session_metadata(&response_session_id), + metadata: outbound_turn_metadata( + &response_session_id, + &task_metadata2, + ), delivery: None, }; if !live_delivery_started { @@ -2756,7 +2798,10 @@ fn spawn_agent_worker( content: format!("Processing error: {}", e), reply_to: None, media: vec![], - metadata: outbound_session_metadata(&response_session_id), + metadata: outbound_turn_metadata( + &response_session_id, + &task_metadata2, + ), delivery: None, }; if !live_delivery_started { @@ -2807,7 +2852,10 @@ fn spawn_agent_worker( .to_string(), reply_to: None, media: vec![], - metadata: outbound_session_metadata(&response_session_id), + metadata: outbound_turn_metadata( + &response_session_id, + &task_metadata2, + ), delivery: None, }; if !live_delivery_started { @@ -2827,7 +2875,10 @@ fn spawn_agent_worker( content: response, reply_to: None, media: vec![], - metadata: outbound_session_metadata(&response_session_id), + metadata: outbound_turn_metadata( + &response_session_id, + &task_metadata2, + ), delivery: None, }; let _ = bus2.publish_outbound(outbound).await;