feat: stream turns through Feishu cards

This commit is contained in:
xiaoxixi 2026-07-17 17:30:38 +08:00
parent c7ceb877a2
commit c6e022f6cb
6 changed files with 648 additions and 45 deletions

View File

@ -70,7 +70,9 @@
"allow_from": ["*"], "allow_from": ["*"],
"agent": "default", "agent": "default",
"media_dir": "~/.picobot/media/feishu", "media_dir": "~/.picobot/media/feishu",
"reaction_emoji": "Typing" "reaction_emoji": "Typing",
"live_updates": false,
"live_update_interval_ms": 500
} }
}, },
"memory": { "memory": {

View File

@ -6,15 +6,15 @@ use std::time::{Duration, Instant};
use async_trait::async_trait; use async_trait::async_trait;
use futures_util::{SinkExt, StreamExt}; use futures_util::{SinkExt, StreamExt};
use prost::{Message as ProstMessage, bytes::Bytes}; use prost::{Message as ProstMessage, bytes::Bytes};
use regex::Regex;
use serde::Deserialize; use serde::Deserialize;
use tokio::sync::{Mutex, RwLock}; use tokio::sync::{Mutex, RwLock};
use tokio::task::JoinHandle; use tokio::task::JoinHandle;
use tokio_util::sync::CancellationToken; use tokio_util::sync::CancellationToken;
use crate::bus::{MediaItem, MessageBus, OutboundMessage}; 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::config::FeishuChannelConfig;
use crate::session::{ToolStatus, TurnBlock, TurnSnapshot, TurnStatus};
const FEISHU_API_BASE: &str = "https://open.feishu.cn/open-apis"; const FEISHU_API_BASE: &str = "https://open.feishu.cn/open-apis";
const FEISHU_WS_BASE: &str = "https://open.feishu.cn"; 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 { impl FeishuChannel {
fn strip_thinking_tags(content: &str) -> String {
use std::sync::LazyLock;
static THINK_RE: LazyLock<Regex> =
LazyLock::new(|| Regex::new(r"(?s)<think>.*?</think>").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. /// Build a Card JSON 2.0 interactive card with a single markdown element.
fn build_card_content(markdown: &str) -> String { fn build_card_content(markdown: &str) -> String {
serde_json::json!({ serde_json::json!({
@ -1849,7 +1841,7 @@ impl FeishuChannel {
break; 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 search_region = &text[start..end];
let split_at = search_region let split_at = search_region
.rfind('\n') .rfind('\n')
@ -1886,7 +1878,7 @@ impl FeishuChannel {
receive_id: &str, receive_id: &str,
receive_id_type: &str, receive_id_type: &str,
card_content: &str, card_content: &str,
) -> Result<(), ChannelError> { ) -> Result<String, ChannelError> {
let token = self.get_tenant_access_token().await?; let token = self.get_tenant_access_token().await?;
let resp = self let resp = self
@ -1912,6 +1904,12 @@ impl FeishuChannel {
struct SendResp { struct SendResp {
code: i32, code: i32,
msg: String, msg: String,
data: Option<SendData>,
}
#[derive(Deserialize)]
struct SendData {
message_id: String,
} }
let send_resp: SendResp = resp.json().await.map_err(|e| { 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(()) Ok(())
} }
} }
#[async_trait]
trait FeishuTurnApi: Send {
async fn create_card(&mut self, markdown: &str) -> Result<String, ChannelError>;
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<String, String>,
}
#[async_trait]
impl FeishuTurnApi for FeishuTurnBackend {
async fn create_card(&mut self, markdown: &str) -> Result<String, ChannelError> {
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<dyn FeishuTurnApi>,
message_id: Option<String>,
cleaned_up: bool,
}
impl FeishuTurnSink {
fn new(api: Box<dyn FeishuTurnApi>) -> 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] #[async_trait]
impl Channel for FeishuChannel { impl Channel for FeishuChannel {
fn name(&self) -> &str { fn name(&self) -> &str {
@ -2043,11 +2275,37 @@ impl Channel for FeishuChannel {
self.running.try_read().map(|r| *r).unwrap_or(false) self.running.try_read().map(|r| *r).unwrap_or(false)
} }
async fn send(&self, msg: OutboundMessage) -> Result<(), ChannelError> { fn live_policy(&self) -> LivePolicy {
let msg = OutboundMessage { if self.config.live_updates {
content: Self::strip_thinking_tags(&msg.content), LivePolicy::Snapshot {
..msg 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<Box<dyn TurnSink>, 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_") { let receive_id = if msg.chat_id.starts_with("oc_") {
&msg.chat_id &msg.chat_id
} else { } else {
@ -2252,6 +2510,53 @@ impl Channel for FeishuChannel {
#[cfg(test)] #[cfg(test)]
mod tests { mod tests {
use super::*; use super::*;
use crate::agent::TurnEvent;
use crate::delivery::{PresentationPolicy, project_snapshot};
#[derive(Default)]
struct MockTurnState {
created: Vec<String>,
updated: Vec<(String, String)>,
cleanups: usize,
fail_updates: usize,
}
struct MockTurnApi {
state: Arc<Mutex<MockTurnState>>,
}
#[async_trait]
impl FeishuTurnApi for MockTurnApi {
async fn create_card(&mut self, markdown: &str) -> Result<String, ChannelError> {
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<Mutex<MockTurnState>>) -> FeishuTurnSink {
FeishuTurnSink::new(Box::new(MockTurnApi { state }))
}
fn test_channel() -> FeishuChannel { fn test_channel() -> FeishuChannel {
FeishuChannel::new( FeishuChannel::new(
@ -2263,12 +2568,169 @@ mod tests {
agent: String::new(), agent: String::new(),
media_dir: String::new(), media_dir: String::new(),
reaction_emoji: "THUMBSUP".to_string(), reaction_emoji: "THUMBSUP".to_string(),
live_updates: false,
live_update_interval_ms: 500,
}, },
Path::new("/tmp"), Path::new("/tmp"),
) )
.expect("test channel should be valid") .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] #[tokio::test]
async fn stop_aborts_connection_task_that_ignores_cancellation() { async fn stop_aborts_connection_task_that_ignores_cancellation() {
let channel = test_channel(); let channel = test_channel();

View File

@ -78,6 +78,11 @@ pub struct FeishuChannelConfig {
/// Emoji type for message reactions (e.g. "THUMBSUP", "OK", "EYES"). /// Emoji type for message reactions (e.g. "THUMBSUP", "OK", "EYES").
#[serde(default = "default_reaction_emoji")] #[serde(default = "default_reaction_emoji")]
pub reaction_emoji: String, 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<String> { fn default_allow_from() -> Vec<String> {
@ -95,6 +100,10 @@ fn default_reaction_emoji() -> String {
"Typing".to_string() "Typing".to_string()
} }
fn default_feishu_live_update_interval_ms() -> u64 {
500
}
#[derive(Debug, Clone, Deserialize, Serialize)] #[derive(Debug, Clone, Deserialize, Serialize)]
pub struct ProviderConfig { pub struct ProviderConfig {
#[serde(rename = "type")] #[serde(rename = "type")]

View File

@ -166,8 +166,8 @@ impl DeliveryCoordinator {
&self, &self,
supervisor: &TaskSupervisor, supervisor: &TaskSupervisor,
route: SinkRoute, route: SinkRoute,
snapshots: watch::Receiver<Arc<TurnSnapshot>>, mut snapshots: watch::Receiver<Arc<TurnSnapshot>>,
sink: Box<dyn TurnSink>, mut sink: Box<dyn TurnSink>,
) -> Result<oneshot::Receiver<Result<(), DeliveryError>>, DeliveryError> { ) -> Result<oneshot::Receiver<Result<(), DeliveryError>>, DeliveryError> {
let SinkRoute { let SinkRoute {
channel, channel,
@ -178,17 +178,28 @@ impl DeliveryCoordinator {
let (result_tx, result_rx) = oneshot::channel(); let (result_tx, result_rx) = oneshot::channel();
let coordinator = self.clone(); let coordinator = self.clone();
let task_name = format!("turn-delivery:{channel}:{chat_id}"); let task_name = format!("turn-delivery:{channel}:{chat_id}");
let spawned = supervisor.spawn(task_name, async move { let cancellation = supervisor.cancellation_token();
let result = coordinator let shutdown_snapshot = snapshots.clone();
.deliver( let spawned = supervisor.spawn_graceful(task_name, async move {
let mut delivery = Box::pin(coordinator.deliver_sink(
&channel, &channel,
&chat_id, &chat_id,
live_policy, live_policy,
presentation, presentation,
snapshots, &mut snapshots,
sink, &mut *sink,
) ));
.await; 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 { if let Err(error) = &result {
tracing::error!(channel, chat_id, error = %error, "Turn delivery failed"); tracing::error!(channel, chat_id, error = %error, "Turn delivery failed");
} }
@ -208,6 +219,26 @@ impl DeliveryCoordinator {
presentation: PresentationPolicy, presentation: PresentationPolicy,
mut snapshots: watch::Receiver<Arc<TurnSnapshot>>, mut snapshots: watch::Receiver<Arc<TurnSnapshot>>,
mut sink: Box<dyn TurnSink>, mut sink: Box<dyn TurnSink>,
) -> 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<Arc<TurnSnapshot>>,
sink: &mut dyn TurnSink,
) -> Result<(), DeliveryError> { ) -> Result<(), DeliveryError> {
let target_lock = self.write_locks.for_target(channel, chat_id); let target_lock = self.write_locks.for_target(channel, chat_id);
let min_interval = match live_policy { let min_interval = match live_policy {
@ -221,9 +252,7 @@ impl DeliveryCoordinator {
let snapshot = snapshots.borrow_and_update().clone(); let snapshot = snapshots.borrow_and_update().clone();
if snapshot.status != TurnStatus::Running { if snapshot.status != TurnStatus::Running {
let projected = project_snapshot(&snapshot, presentation); let projected = project_snapshot(&snapshot, presentation);
return self return self.deliver_terminal(&target_lock, sink, &projected).await;
.deliver_terminal(&target_lock, &mut *sink, &projected)
.await;
} }
if let Some(interval) = min_interval { if let Some(interval) = min_interval {
@ -234,7 +263,7 @@ impl DeliveryCoordinator {
let latest = snapshots.borrow_and_update().clone(); let latest = snapshots.borrow_and_update().clone();
if latest.status != TurnStatus::Running { if latest.status != TurnStatus::Running {
let projected = project_snapshot(&latest, presentation); 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, () = sleep_until(next_update_at) => break,
@ -244,9 +273,7 @@ impl DeliveryCoordinator {
let latest = snapshots.borrow_and_update().clone(); let latest = snapshots.borrow_and_update().clone();
if latest.status != TurnStatus::Running { if latest.status != TurnStatus::Running {
let projected = project_snapshot(&latest, presentation); let projected = project_snapshot(&latest, presentation);
return self return self.deliver_terminal(&target_lock, sink, &projected).await;
.deliver_terminal(&target_lock, &mut *sink, &projected)
.await;
} }
let projected = project_snapshot(&latest, presentation); let projected = project_snapshot(&latest, presentation);
let _guard = target_lock.lock().await; 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( async fn deliver_terminal(
&self, &self,
target_lock: &Arc<AsyncMutex<()>>, target_lock: &Arc<AsyncMutex<()>>,
@ -595,4 +638,39 @@ mod tests {
assert_eq!(channel.opened.load(Ordering::SeqCst), 1); assert_eq!(channel.opened.load(Ordering::SeqCst), 1);
assert_eq!(state.terminal.lock().await.len(), 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);
}
} }

View File

@ -288,6 +288,7 @@ impl GatewayState {
&inbound.chat_id, &inbound.chat_id,
&inbound.content, &inbound.content,
inbound.media, inbound.media,
inbound.forwarded_metadata.clone(),
).await { ).await {
Ok(crate::session::session::HandleResult::AgentResponse(content)) => { Ok(crate::session::session::HandleResult::AgentResponse(content)) => {
let outbound = crate::bus::OutboundMessage { let outbound = crate::bus::OutboundMessage {

View File

@ -25,6 +25,15 @@ fn outbound_session_metadata(session_id: &str) -> HashMap<String, String> {
HashMap::from([("_session_id".to_string(), session_id.to_string())]) HashMap::from([("_session_id".to_string(), session_id.to_string())])
} }
fn outbound_turn_metadata(
session_id: &str,
forwarded: &HashMap<String, String>,
) -> HashMap<String, String> {
let mut metadata = forwarded.clone();
metadata.insert("_session_id".to_string(), session_id.to_string());
metadata
}
tokio::task_local! { tokio::task_local! {
pub(super) static CURRENT_SOURCE_SESSION: Option<String>; pub(super) static CURRENT_SOURCE_SESSION: Option<String>;
} }
@ -117,6 +126,30 @@ mod cancelled_partial_tests {
use super::*; use super::*;
use crate::agent::TurnEvent; 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] #[test]
fn visible_partial_text_becomes_cancelled_persisted_message() { fn visible_partial_text_becomes_cancelled_persisted_message() {
let (controller, emitter, _) = TurnController::start("session", "message-id"); let (controller, emitter, _) = TurnController::start("session", "message-id");
@ -233,6 +266,7 @@ struct AgentTask {
chat_id: String, chat_id: String,
content: String, content: String,
media: Vec<MediaItem>, media: Vec<MediaItem>,
forwarded_metadata: HashMap<String, String>,
} }
#[derive(Clone)] #[derive(Clone)]
@ -2176,6 +2210,7 @@ impl SessionManager {
chat_id: &str, chat_id: &str,
content: &str, content: &str,
media: Vec<MediaItem>, media: Vec<MediaItem>,
forwarded_metadata: HashMap<String, String>,
) -> Result<HandleResult, AgentError> { ) -> Result<HandleResult, AgentError> {
let unified_id = self.resolve_dialog_id(channel, chat_id).await?; let unified_id = self.resolve_dialog_id(channel, chat_id).await?;
tracing::debug!(unified_id = %unified_id, "handle_message resolved unified_id"); tracing::debug!(unified_id = %unified_id, "handle_message resolved unified_id");
@ -2213,6 +2248,7 @@ impl SessionManager {
chat_id: chat_id.to_string(), chat_id: chat_id.to_string(),
content: content.to_string(), content: content.to_string(),
media, media,
forwarded_metadata,
}; };
let session_clone = session.clone(); let session_clone = session.clone();
let unified_str = unified_id.to_string(); let unified_str = unified_id.to_string();
@ -2351,6 +2387,7 @@ fn spawn_agent_worker(
'tasks: while let Some(task) = task_rx.recv().await { 'tasks: while let Some(task) = task_rx.recv().await {
let task_chan = task.channel.clone(); let task_chan = task.channel.clone();
let task_cid = task.chat_id.clone(); let task_cid = task.chat_id.clone();
let task_metadata = task.forwarded_metadata.clone();
let notification_session_id = unified_str.clone(); let notification_session_id = unified_str.clone();
let (notify_tx, mut notify_rx) = mpsc::unbounded_channel(); 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(), content: "Failed to save your message, please try again.".to_string(),
reply_to: None, reply_to: None,
media: vec![], media: vec![],
metadata: outbound_session_metadata(&unified_str), metadata: outbound_turn_metadata(&unified_str, &task_metadata),
delivery: None, delivery: None,
}; };
let _ = bus.publish_outbound(err_outbound).await; let _ = bus.publish_outbound(err_outbound).await;
@ -2434,7 +2471,7 @@ fn spawn_agent_worker(
.to_string(), .to_string(),
reply_to: None, reply_to: None,
media: vec![], media: vec![],
metadata: outbound_session_metadata(&unified_str), metadata: outbound_turn_metadata(&unified_str, &task_metadata),
delivery: None, delivery: None,
}; };
let _ = bus.publish_outbound(err_outbound).await; let _ = bus.publish_outbound(err_outbound).await;
@ -2561,7 +2598,7 @@ fn spawn_agent_worker(
chat_id: task_cid.clone(), chat_id: task_cid.clone(),
session_id: unified_str.clone(), session_id: unified_str.clone(),
reply_to: None, reply_to: None,
metadata: outbound_session_metadata(&unified_str), metadata: outbound_turn_metadata(&unified_str, &task_metadata),
}, },
turn_receiver, turn_receiver,
) )
@ -2604,6 +2641,7 @@ fn spawn_agent_worker(
let chan2 = task_chan.clone(); let chan2 = task_chan.clone();
let cid2 = task_cid.clone(); let cid2 = task_cid.clone();
let unified_str2 = unified_str.clone(); let unified_str2 = unified_str.clone();
let task_metadata2 = task_metadata.clone();
let turn_lifecycle = &turn_controller; let turn_lifecycle = &turn_controller;
let process_future = async move { let process_future = async move {
let response_session_id = unified_str2.clone(); let response_session_id = unified_str2.clone();
@ -2656,8 +2694,9 @@ fn spawn_agent_worker(
.to_string(), .to_string(),
reply_to: None, reply_to: None,
media: vec![], media: vec![],
metadata: outbound_session_metadata( metadata: outbound_turn_metadata(
&response_session_id, &response_session_id,
&task_metadata2,
), ),
delivery: None, delivery: None,
}; };
@ -2732,7 +2771,10 @@ fn spawn_agent_worker(
content: format!("Processing error: {}", e), content: format!("Processing error: {}", e),
reply_to: None, reply_to: None,
media: vec![], media: vec![],
metadata: outbound_session_metadata(&response_session_id), metadata: outbound_turn_metadata(
&response_session_id,
&task_metadata2,
),
delivery: None, delivery: None,
}; };
if !live_delivery_started { if !live_delivery_started {
@ -2756,7 +2798,10 @@ fn spawn_agent_worker(
content: format!("Processing error: {}", e), content: format!("Processing error: {}", e),
reply_to: None, reply_to: None,
media: vec![], media: vec![],
metadata: outbound_session_metadata(&response_session_id), metadata: outbound_turn_metadata(
&response_session_id,
&task_metadata2,
),
delivery: None, delivery: None,
}; };
if !live_delivery_started { if !live_delivery_started {
@ -2807,7 +2852,10 @@ fn spawn_agent_worker(
.to_string(), .to_string(),
reply_to: None, reply_to: None,
media: vec![], media: vec![],
metadata: outbound_session_metadata(&response_session_id), metadata: outbound_turn_metadata(
&response_session_id,
&task_metadata2,
),
delivery: None, delivery: None,
}; };
if !live_delivery_started { if !live_delivery_started {
@ -2827,7 +2875,10 @@ fn spawn_agent_worker(
content: response, content: response,
reply_to: None, reply_to: None,
media: vec![], media: vec![],
metadata: outbound_session_metadata(&response_session_id), metadata: outbound_turn_metadata(
&response_session_id,
&task_metadata2,
),
delivery: None, delivery: None,
}; };
let _ = bus2.publish_outbound(outbound).await; let _ = bus2.publish_outbound(outbound).await;