Files
DevFlow/src-tauri/src/commands/ai/conversation.rs

387 lines
16 KiB
Rust
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
//! 对话持久化 + 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<i64> 字段(读旧值+增量,跨 loop 实例防覆盖)
///
/// 纯函数:抽自 save_conversation 的 token 累加逻辑,None 起始当作 0。
/// saturating_add:长期对话累积接近 i64::MAX 时不再翻负,封顶在 i64::MAX(统计语义安全,
/// 溢出回绕成负值会污染前端用量展示与计费/限额判定)。
pub(crate) fn accumulate_tokens(old: Option<i64>, add: u32) -> Option<i64> {
Some(old.unwrap_or(0).saturating_add(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<char> = 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
)
}
/// 落库前对 `ChatMessage.parts` 做截断(F-260614-05 Phase 2a)。
///
/// - `Text` 片:复用 `truncate_for_persist` 头尾截断语义(单 Text 片 > 50KB 才截)。
/// - `Image` 片:base64 通常已 50KB+,落库前替换为占位 `Text` 片
/// `<image: base64 已省略, 共 N 字节>`,避免大体量图把对话 JSON 撑爆
/// (一张 1MB PNG 的 base64 ≈ 1.3MB 字符串)。原 Image 片仅在内存真相源(ContextManager)
/// 保留,重发时仍带图——对齐现有「持久化视图不污染内存真相源」约定。
/// - `Image` url 模式(无 base64):url 本身短,原样保留。
///
/// 返回 None 表示 parts 无需保留(全空或全部被替换且原 parts 仅含图);调用方据此清空 parts。
pub(crate) fn truncate_parts_for_persist(parts: &[df_ai::provider::ContentPart]) -> Option<Vec<df_ai::provider::ContentPart>> {
use df_ai::provider::ContentPart;
let mut out: Vec<ContentPart> = Vec::with_capacity(parts.len());
let mut changed = false;
for p in parts {
match p {
ContentPart::Text { text } => {
let truncated = truncate_for_persist(text);
if truncated.len() != text.len() {
changed = true;
}
out.push(ContentPart::Text { text: truncated });
}
ContentPart::Image { url, base64, media_type, alt } => {
// base64 模式:体量大,替换占位 Text 片
if let Some(b) = base64 {
let base64_len = b.len(); // 纯 base64 字符串长度(不含 data: 前缀)
let placeholder = format!("<image: base64 已省略, 共约 {} 字符>", base64_len);
out.push(ContentPart::Text { text: placeholder });
changed = true;
} else {
// url 模式:url 短,原样保留(含 media_type/alt)
out.push(ContentPart::Image {
url: url.clone(),
base64: None,
media_type: media_type.clone(),
alt: alt.clone(),
});
}
}
}
}
if out.is_empty() {
None
} else if changed {
Some(out)
} else {
// 无变化:返回克隆的原 parts(保持引用语义一致)
Some(parts.to_vec())
}
}
/// 保存对话到数据库(按 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<Mutex<AiSession>>,
db: &Arc<Database>,
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);
// F-260614-05 Phase 2a: parts(Image base64) 同样截断(替换占位 Text 片),
// 防大体量图把对话 JSON 撑爆。仅作用于持久化副本,不污染内存真相源。
if let Some(parts) = m.parts.as_ref() {
m.parts = truncate_parts_for_persist(parts);
}
}
(
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<String> = 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<i64> = 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<i64> = None;
let mut db_completion: Option<i64> = 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("已截断"));
}
// ---------- F-260614-05 Phase 2a truncate_parts_for_persist ----------
#[test]
fn truncate_parts_image_base64_replaced_with_placeholder() {
use df_ai::provider::ContentPart;
let parts = vec![
ContentPart::text("前缀文本"),
ContentPart::image_base64("image/png", &"A".repeat(60_000)),
ContentPart::text("后缀"),
];
let out = truncate_parts_for_persist(&parts).expect("应有结果");
// Image base64 被替换为占位 Text 片
assert!(out.iter().all(|p| !matches!(p, ContentPart::Image { base64: Some(_), .. })));
let placeholder = out.iter().find_map(|p| match p {
ContentPart::Text { text } if text.contains("base64 已省略") => Some(text.clone()),
_ => None,
}).expect("应含占位 Text 片");
assert!(placeholder.contains("60000"));
// 非 base64 的 Text 片保留原值
assert!(out.iter().any(|p| matches!(p, ContentPart::Text { text } if text == "前缀文本")));
assert!(out.iter().any(|p| matches!(p, ContentPart::Text { text } if text == "后缀")));
}
#[test]
fn truncate_parts_image_url_preserved() {
use df_ai::provider::ContentPart;
// url 模式:url 短,原样保留(含 media_type)
let parts = vec![ContentPart::Image {
url: Some("https://x/a.png".into()),
base64: None,
media_type: None,
alt: None,
}];
let out = truncate_parts_for_persist(&parts).expect("应有结果");
assert!(matches!(out[0], ContentPart::Image { ref url, .. } if url.as_deref() == Some("https://x/a.png")));
}
#[test]
fn truncate_parts_long_text_part_truncated() {
use df_ai::provider::ContentPart;
// 单 Text 片超阈值 → 头尾截断
let long: String = "x".repeat(TRUNCATE_THRESHOLD + 1000);
let parts = vec![ContentPart::Text { text: long }];
let out = truncate_parts_for_persist(&parts).expect("应有结果");
match &out[0] {
ContentPart::Text { text } => assert!(text.contains("已截断"), "超长 Text 片应被截断"),
_ => panic!("应仍是 Text 片"),
}
}
#[test]
fn truncate_parts_empty_returns_none() {
assert_eq!(truncate_parts_for_persist(&[]), None);
}
}