273 lines
11 KiB
Rust
273 lines
11 KiB
Rust
//! 对话持久化 + 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("已截断"));
|
||
}
|
||
}
|