//! 对话持久化 + Token 累加器 use std::sync::Arc; use tokio::sync::Mutex; use df_storage::crud::AiConversationRepo; use df_storage::db::Database; use crate::commands::now_millis; use super::AiSession; /// Token 用量累加器(agent loop 生命周期内各轮叠加) /// /// 纯结构 + 方法:抽自 run_agentic_loop 的 `total_prompt`/`total_completion` 双计数器, /// 保证多轮累加、None 起始、跨 loop 实例叠加语义一致且可单测。 #[derive(Debug, Clone, Default)] pub(crate) struct TokenAccumulator { prompt: u32, completion: u32, } impl TokenAccumulator { /// 叠加一轮用量(round_usage 为本轮流式末 chunk 的累计用量) /// /// saturating_add:恶意/异常 provider 返回巨大值时,避免 u32 += 溢出回绕打乱后续 budget 判定。 pub(crate) fn add(&mut self, prompt: u32, completion: u32) { self.prompt = self.prompt.saturating_add(prompt); self.completion = self.completion.saturating_add(completion); } pub(crate) fn prompt(&self) -> u32 { self.prompt } pub(crate) fn completion(&self) -> u32 { self.completion } pub(crate) fn total(&self) -> u32 { self.prompt.saturating_add(self.completion) } } /// 把单轮增量叠加到 DB 的 Option 字段(读旧值+增量,跨 loop 实例防覆盖) /// /// 纯函数:抽自 save_conversation 的 token 累加逻辑,None 起始当作 0。 pub(crate) fn accumulate_tokens(old: Option, add: u32) -> Option { Some(old.unwrap_or(0) + add as i64) } /// 持久化截断阈值:超过此长度的消息 content 落库前截断头尾各保 HEAD/TAIL 字符。 /// /// 防 read_file 1MB 洞 / list_directory 大体量结果落库后每轮重发累积致 token 暴增 /// (Sprint 19 实测单对话 in=115万 / 消息体 1.6MB)。仅作用于持久化视图,不污染内存真相源。 pub(crate) const TRUNCATE_THRESHOLD: usize = 50_000; const TRUNCATE_HEAD: usize = 20_000; const TRUNCATE_TAIL: usize = 20_000; /// 落库前对超长 content 做截断(保留头尾各 ~20KB + 中段标注省略字符数)。 /// 50KB 阈值以下原样返回(零开销);按字符而非字节切避免 UTF-8 切坏中文。 pub(crate) fn truncate_for_persist(content: &str) -> String { let chars: Vec = content.chars().collect(); if chars.len() <= TRUNCATE_THRESHOLD { return content.to_string(); } let head: String = chars.iter().take(TRUNCATE_HEAD).collect(); let tail: String = chars[chars.len() - TRUNCATE_TAIL..].iter().collect(); let omitted = chars.len() - TRUNCATE_HEAD - TRUNCATE_TAIL; format!( "{}\n\n[...省略 {} 字符(已截断,完整内容仅在内存态可读)...]\n\n{}", head, omitted, tail ) } /// 保存对话到数据库(按 conv_id 写库,不受 active_conversation_id 切换影响) /// /// 写 messages + updated_at + 累加 token 用量 + 首次落库的 model;标题由 ensure_conversation_title 单独生成。 /// token 走累加模式:upsert 读旧值叠加,保证审批暂停→恢复跨 loop 实例不覆盖丢失。 /// model 仅首次落库写入 + 旧记录缺值时补填(不覆盖历史已存值,兼容本次改造前的老对话)。 pub(crate) async fn save_conversation( session_arc: &Arc>, db: &Arc, conv_id: &str, usage: Option<&df_ai::provider::TokenUsage>, model: Option<&str>, ) { // 取 messages + 懒创建首次落库所需的 provider_id/created_at // 工具结果(content)超 50KB 时截断头尾各 ~20KB + 中段标注,防大体量结果(read_file 1MB 洞 / // list_directory 13782 项)落库后每轮重发累积致 token 暴增。仅影响持久化视图,不污染 // 内存真相源(ContextManager)——build_for_request 仍读全量 messages。 let (messages_json, provider_id, created_at) = { let session = session_arc.lock().await; let mut msgs = session.messages.all_messages_clone(); for m in &mut msgs { m.content = truncate_for_persist(&m.content); } ( serde_json::to_string(&msgs).unwrap_or_else(|_| "[]".to_string()), session.active_provider_id.clone(), session.active_conv_created_at.clone(), ) }; let conv_repo = AiConversationRepo::new(db); match conv_repo.get_by_id(conv_id).await { Ok(Some(mut rec)) => { // 已落库:更新 messages + updated_at;token 累加(读旧值+新值,跨 loop 实例防覆盖) rec.messages = messages_json; rec.updated_at = now_millis(); if let Some(u) = usage { rec.prompt_tokens = accumulate_tokens(rec.prompt_tokens, u.prompt_tokens); rec.completion_tokens = accumulate_tokens(rec.completion_tokens, u.completion_tokens); } // model: 旧记录缺值时补填(不覆盖已有);models: 去重追加用过的所有 model(JSON 数组) if let Some(m) = model { if rec.model.is_none() { rec.model = Some(m.to_string()); } let mut list: Vec = rec.models .as_deref() .and_then(|s| serde_json::from_str(s).ok()) .unwrap_or_default(); if !list.iter().any(|x| x == m) { list.push(m.to_string()); } rec.models = Some(serde_json::to_string(&list).unwrap_or_else(|_| "[]".to_string())); } if let Err(e) = conv_repo.update_full(&rec).await { tracing::warn!("更新对话失败 {conv_id}: {e}"); } } Ok(None) => { // 懒创建首次落库(此为空对话不落库的落库点:走到这里 messages 必非空) let now = now_millis(); let rec = df_storage::models::AiConversationRecord { id: conv_id.to_string(), title: None, messages: messages_json, provider_id, model: model.map(|m| m.to_string()), models: model.map(|m| serde_json::to_string(&[m]).unwrap_or_else(|_| "[]".to_string())), archived: false, pinned: false, prompt_tokens: usage.map(|u| u.prompt_tokens as i64), completion_tokens: usage.map(|u| u.completion_tokens as i64), created_at: created_at.unwrap_or_else(|| now.clone()), updated_at: now, }; if let Err(e) = conv_repo.insert(rec).await { tracing::warn!("落库对话失败 {conv_id}: {e}"); } } Err(e) => tracing::warn!("读取对话 {conv_id} 失败: {e}"), } } #[cfg(test)] mod tests { use super::*; // ---------- TokenAccumulator + accumulate_tokens ---------- #[test] fn accumulator_starts_zero() { let acc = TokenAccumulator::default(); assert_eq!(acc.prompt(), 0); assert_eq!(acc.completion(), 0); assert_eq!(acc.total(), 0); } #[test] fn accumulator_single_add() { let mut acc = TokenAccumulator::default(); acc.add(100, 50); assert_eq!(acc.prompt(), 100); assert_eq!(acc.completion(), 50); assert_eq!(acc.total(), 150); } #[test] fn accumulator_multi_round_accumulation() { // 多轮累加(模拟 agent loop 多次迭代) let mut acc = TokenAccumulator::default(); acc.add(100, 20); // 轮1 acc.add(200, 40); // 轮2 acc.add(50, 10); // 轮3 assert_eq!(acc.prompt(), 350); assert_eq!(acc.completion(), 70); assert_eq!(acc.total(), 420); } #[test] fn accumulator_add_zero_is_noop() { let mut acc = TokenAccumulator::default(); acc.add(10, 5); acc.add(0, 0); assert_eq!(acc.total(), 15); } #[test] fn accumulate_tokens_from_none() { // 新记录(None 起始)落库 assert_eq!(accumulate_tokens(None, 100), Some(100)); assert_eq!(accumulate_tokens(None, 0), Some(0)); } #[test] fn accumulate_tokens_adds_to_existing() { // 跨 loop 实例叠加:旧值 + 新增不覆盖 assert_eq!(accumulate_tokens(Some(500), 100), Some(600)); assert_eq!(accumulate_tokens(Some(0), 42), Some(42)); } #[test] fn accumulate_tokens_multi_round_db_simulation() { // 模拟 save_conversation 多次落库累加(审批暂停→恢复跨 loop) let mut field: Option = None; field = accumulate_tokens(field, 100); // 首次 field = accumulate_tokens(field, 200); // 二次 field = accumulate_tokens(field, 50); // 三次 assert_eq!(field, Some(350)); } #[test] fn accumulator_and_db_accumulate_are_consistent() { // loop 内 TokenAccumulator 与落库 accumulate_tokens 总量语义一致 let mut acc = TokenAccumulator::default(); let mut db_prompt: Option = None; let mut db_completion: Option = None; for (p, c) in [(100u32, 20u32), (200, 40), (50, 10)] { acc.add(p, c); db_prompt = accumulate_tokens(db_prompt, p); db_completion = accumulate_tokens(db_completion, c); } assert_eq!(acc.prompt() as i64, db_prompt.unwrap()); assert_eq!(acc.completion() as i64, db_completion.unwrap()); } // ---------- truncate_for_persist ---------- #[test] fn truncate_short_content_unchanged() { // 阈值以下原样返回 assert_eq!(truncate_for_persist("hello"), "hello"); assert_eq!(truncate_for_persist(""), ""); let near_limit: String = "a".repeat(TRUNCATE_THRESHOLD); assert_eq!(truncate_for_persist(&near_limit).len(), TRUNCATE_THRESHOLD); } #[test] fn truncate_long_content_keeps_head_and_tail() { // 超阈值:保留头尾各 TRUNCATE_HEAD/TAIL 字符 + 中段标注 let long: String = "x".repeat(TRUNCATE_THRESHOLD + 1000); let result = truncate_for_persist(&long); // 头尾各 20k 字符应在结果中 assert!(result.starts_with(&"x".repeat(TRUNCATE_HEAD))); assert!(result.ends_with(&"x".repeat(TRUNCATE_TAIL))); // 中段标注存在 + 标注省略字符数(中段 = 总长 - 头 - 尾 = 51000 - 20000 - 20000 = 11000) assert!(result.contains("已截断")); assert!(result.contains("省略 11000 字符")); // 结果总长应远小于原长(20k 头 + 20k 尾 + 标注) assert!(result.chars().count() < TRUNCATE_THRESHOLD + 1000); } #[test] fn truncate_preserves_utf8_chinese() { // 按字符切不切坏 UTF-8 中文 let chinese: String = "中".repeat(TRUNCATE_THRESHOLD + 500); let result = truncate_for_persist(&chinese); assert!(result.starts_with('中')); assert!(result.ends_with('中')); assert!(result.contains("已截断")); } }