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

273 lines
11 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。
pub(crate) fn accumulate_tokens(old: Option<i64>, add: u32) -> Option<i64> {
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<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
)
}
/// 保存对话到数据库(按 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);
}
(
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("已截断"));
}
}