281 lines
8.7 KiB
Rust
281 lines
8.7 KiB
Rust
//! 上下文自动压缩 —— 当对话历史过长时自动压缩。
|
||
|
||
use crate::llm::types::message::{ContentBlock, Message};
|
||
|
||
const AUTOCOMPACT_BUFFER_TOKENS: u32 = 13_000;
|
||
const RESERVED_OUTPUT_TOKENS: u32 = 20_000;
|
||
const MAX_CONSECUTIVE_FAILURES: u32 = 3;
|
||
const KEEP_RECENT: usize = 6;
|
||
|
||
/// 上下文压缩配置。
|
||
#[derive(Debug, Clone)]
|
||
pub struct CompactConfig {
|
||
/// 模型上下文窗口大小(token 数)。
|
||
pub context_window: u32,
|
||
/// 为输出预留的 token 数。
|
||
pub reserved_tokens: u32,
|
||
/// 微压缩保留的最近消息数。
|
||
pub keep_recent: usize,
|
||
}
|
||
|
||
impl Default for CompactConfig {
|
||
fn default() -> Self {
|
||
Self {
|
||
context_window: 128_000,
|
||
reserved_tokens: RESERVED_OUTPUT_TOKENS,
|
||
keep_recent: KEEP_RECENT,
|
||
}
|
||
}
|
||
}
|
||
|
||
impl CompactConfig {
|
||
/// 计算自动压缩触发的阈值。
|
||
pub fn threshold(&self) -> u32 {
|
||
self.context_window
|
||
.saturating_sub(self.reserved_tokens)
|
||
.saturating_sub(AUTOCOMPACT_BUFFER_TOKENS)
|
||
}
|
||
}
|
||
|
||
/// 压缩状态 —— 跟踪连续失败次数(断路器模式)。
|
||
#[derive(Debug, Clone)]
|
||
pub struct CompactState {
|
||
consecutive_failures: u32,
|
||
}
|
||
|
||
impl Default for CompactState {
|
||
fn default() -> Self {
|
||
Self::new()
|
||
}
|
||
}
|
||
|
||
impl CompactState {
|
||
/// 创建一个新的压缩状态。
|
||
pub fn new() -> Self {
|
||
Self {
|
||
consecutive_failures: 0,
|
||
}
|
||
}
|
||
|
||
/// 记录一次成功的压缩。
|
||
pub fn record_success(&mut self) {
|
||
self.consecutive_failures = 0;
|
||
}
|
||
|
||
/// 记录一次压缩失败。
|
||
///
|
||
/// 返回 `true` 表示已达断路器上限,不再尝试。
|
||
pub fn record_failure(&mut self) -> bool {
|
||
self.consecutive_failures += 1;
|
||
self.consecutive_failures >= MAX_CONSECUTIVE_FAILURES
|
||
}
|
||
}
|
||
|
||
/// 粗略估计消息列表的 token 数(基于字符数,4 字符 ≈ 1 token)。
|
||
pub fn estimate_message_tokens(messages: &[Message]) -> u32 {
|
||
messages.iter().map(estimate_single_message_tokens).sum()
|
||
}
|
||
|
||
fn estimate_single_message_tokens(msg: &Message) -> u32 {
|
||
let role_overhead: u32 = 4;
|
||
let content_tokens = match msg {
|
||
Message::System { content }
|
||
| Message::User { content }
|
||
| Message::Assistant { content }
|
||
| Message::ToolResult { content, .. } => estimate_content_blocks_tokens(content),
|
||
Message::UserImage { .. } => 50,
|
||
};
|
||
role_overhead + content_tokens
|
||
}
|
||
|
||
fn estimate_content_blocks_tokens(blocks: &[ContentBlock]) -> u32 {
|
||
blocks.iter().map(estimate_block_tokens).sum()
|
||
}
|
||
|
||
fn estimate_block_tokens(block: &ContentBlock) -> u32 {
|
||
match block {
|
||
ContentBlock::Text { text } => estimate_text_tokens(text),
|
||
ContentBlock::Thinking { text, .. } => estimate_text_tokens(text),
|
||
ContentBlock::ToolUse { input, .. } => estimate_text_tokens(&input.to_string()),
|
||
ContentBlock::ToolResult { content, .. } => estimate_content_blocks_tokens(content),
|
||
// ponytail: Image / Audio / File / Extension 在 IR 中固定估算。
|
||
// 无文本的视觉/音频 block 用兜底估算,避免 token 计数膨胀。
|
||
ContentBlock::Image { .. }
|
||
| ContentBlock::Audio { .. }
|
||
| ContentBlock::File { .. }
|
||
| ContentBlock::Extension { .. } => 50,
|
||
}
|
||
}
|
||
|
||
fn estimate_text_tokens(text: &str) -> u32 {
|
||
if text.is_empty() {
|
||
return 0;
|
||
}
|
||
let len = text.len() as u32;
|
||
(len * 4).div_ceil(3)
|
||
}
|
||
|
||
/// 判断是否需要触发自动压缩。
|
||
pub fn should_compact(messages: &[Message], config: &CompactConfig, state: &CompactState) -> bool {
|
||
if state.consecutive_failures >= MAX_CONSECUTIVE_FAILURES {
|
||
return false;
|
||
}
|
||
let tokens = estimate_message_tokens(messages);
|
||
tokens >= config.threshold()
|
||
}
|
||
|
||
/// 执行微压缩 —— 用 `[pruned]` 替换旧的 tool result 内容。
|
||
///
|
||
/// 这是最便宜的压缩方式,不需要 LLM 调用。
|
||
/// 保留最近的 `keep_recent` 条消息不变。
|
||
///
|
||
/// 返回释放的估算 token 数。
|
||
///
|
||
/// **审查 FIX-F**:仅压缩 `is_error: false` 的 `ToolResult` —— 错误结果包含对 LLM
|
||
/// 理解失败原因至关重要的诊断信息,压缩后 LLM 无法理解。
|
||
pub fn microcompact(messages: &mut [Message], keep_recent: usize) -> u32 {
|
||
if messages.len() <= keep_recent {
|
||
return 0;
|
||
}
|
||
|
||
let prune_start = messages.len() - keep_recent;
|
||
let mut freed_tokens: u32 = 0;
|
||
|
||
// 第一遍:计算可释放 token(仅非错误 ToolResult)
|
||
for msg in &messages[..prune_start] {
|
||
if matches!(
|
||
msg,
|
||
Message::ToolResult {
|
||
is_error: false,
|
||
..
|
||
}
|
||
) {
|
||
freed_tokens += estimate_single_message_tokens(msg);
|
||
}
|
||
}
|
||
|
||
// 第二遍:替换内容(仅非错误 ToolResult)
|
||
for msg in &mut messages[..prune_start] {
|
||
if let Message::ToolResult {
|
||
content,
|
||
is_error: false,
|
||
..
|
||
} = msg
|
||
{
|
||
*content = vec![ContentBlock::Text {
|
||
text: "[pruned]".to_string(),
|
||
}];
|
||
}
|
||
}
|
||
|
||
freed_tokens
|
||
}
|
||
|
||
#[cfg(test)]
|
||
mod tests {
|
||
use super::*;
|
||
|
||
fn user_msg(s: &str) -> Message {
|
||
Message::user_text(s)
|
||
}
|
||
|
||
#[test]
|
||
fn estimate_message_tokens_handles_all_variants() {
|
||
let messages = vec![
|
||
Message::System {
|
||
content: vec![ContentBlock::Text { text: "sys".into() }],
|
||
},
|
||
Message::user_text("hi"),
|
||
Message::assistant("ans"),
|
||
Message::user_image(
|
||
"b64",
|
||
"image/png",
|
||
crate::llm::types::shared::ImageDetail::Auto,
|
||
),
|
||
Message::tool_result("call_1", "tool res", false),
|
||
];
|
||
let tokens = estimate_message_tokens(&messages);
|
||
// 至少 5 条消息 × 4 role overhead = 20 + 文本/估算
|
||
assert!(tokens > 20);
|
||
}
|
||
|
||
#[test]
|
||
fn microcompact_replaces_old_tool_result_with_pruned() {
|
||
let mut messages = vec![
|
||
user_msg("hi"),
|
||
Message::tool_result("call_1", "raw result".repeat(50), false),
|
||
user_msg("again"),
|
||
user_msg("keep recent 1"),
|
||
user_msg("keep recent 2"),
|
||
];
|
||
let before_len = messages.len();
|
||
let freed = microcompact(&mut messages, 2);
|
||
assert!(freed > 0);
|
||
assert_eq!(messages.len(), before_len); // 只改内容,不删消息
|
||
// 索引 1 是被压缩的 ToolResult
|
||
if let Message::ToolResult {
|
||
content, is_error, ..
|
||
} = &messages[1]
|
||
{
|
||
assert_eq!(content.len(), 1);
|
||
assert!(matches!(&content[0], ContentBlock::Text { text } if text == "[pruned]"));
|
||
assert!(!is_error);
|
||
} else {
|
||
panic!("expected ToolResult at index 1");
|
||
}
|
||
}
|
||
|
||
/// FIX-F 验证:错误结果不被压缩,诊断信息完整保留。
|
||
#[test]
|
||
fn microcompact_preserves_error_tool_results() {
|
||
let mut messages = vec![
|
||
user_msg("hi"),
|
||
Message::tool_result("call_1", "important error info: backend down", true),
|
||
user_msg("keep recent 1"),
|
||
user_msg("keep recent 2"),
|
||
];
|
||
let before_len = messages.len();
|
||
let freed = microcompact(&mut messages, 2);
|
||
assert_eq!(freed, 0); // 错误 ToolResult 不计入
|
||
assert_eq!(messages.len(), before_len);
|
||
// 错误信息保留完整
|
||
if let Message::ToolResult {
|
||
content, is_error, ..
|
||
} = &messages[1]
|
||
{
|
||
assert!(is_error);
|
||
assert!(
|
||
matches!(&content[0], ContentBlock::Text { text } if text.contains("backend down"))
|
||
);
|
||
} else {
|
||
panic!("expected ToolResult at index 1");
|
||
}
|
||
}
|
||
|
||
#[test]
|
||
fn microcompact_keeps_recent_messages() {
|
||
let mut messages = vec![
|
||
Message::tool_result("c1", "old", false),
|
||
user_msg("m1"),
|
||
user_msg("m2"),
|
||
user_msg("recent"),
|
||
];
|
||
let freed = microcompact(&mut messages, 1);
|
||
// keep_recent=1 → 仅最近 1 条不动,其余 ToolResult 压缩
|
||
// 索引 0 (ToolResult) 被压缩,索引 1-3 保留
|
||
assert!(freed > 0);
|
||
if let Message::ToolResult { content, .. } = &messages[0] {
|
||
assert!(matches!(&content[0], ContentBlock::Text { text } if text == "[pruned]"));
|
||
}
|
||
assert!(matches!(messages[3], Message::User { .. }));
|
||
}
|
||
|
||
#[test]
|
||
fn should_compact_respects_threshold() {
|
||
let cfg = CompactConfig::default();
|
||
let state = CompactState::new();
|
||
let empty: Vec<Message> = vec![];
|
||
assert!(!should_compact(&empty, &cfg, &state));
|
||
}
|
||
}
|