use crate::assembler::{Assembled, ResponseAssembler};
use crate::error::ProviderError;
use crate::ids::{EntryId, ModelId};
use crate::message::{ContentBlock, FinishReason, Message, Usage};
use crate::provider::{
CredentialProvider, GenerationOptions, ModelMessage, ModelProvider, ModelRequest,
ReasoningLevel,
};
use crate::store::SummaryRecord;
use futures::StreamExt;
use serde::{Deserialize, Serialize};
use std::collections::hash_map::DefaultHasher;
use std::hash::Hasher;
use tokio_util::sync::CancellationToken;
#[derive(Clone, Copy, Debug, PartialEq, Serialize, Deserialize)]
pub struct CompactionConfig {
pub trigger_ratio: f64,
pub keep_recent_turns: usize,
pub max_attempts: u32,
}
impl Default for CompactionConfig {
fn default() -> Self {
Self {
trigger_ratio: 0.8,
keep_recent_turns: 4,
max_attempts: 3,
}
}
}
pub struct CompactionContext {
pub current_model: ModelId,
}
pub trait CompactionModelSelector: Send + Sync {
fn select_model(&self, current: &ModelId) -> ModelId;
}
pub struct CurrentModelSelector;
impl CompactionModelSelector for CurrentModelSelector {
fn select_model(&self, current: &ModelId) -> ModelId {
current.clone()
}
}
pub fn select_compaction_range(
messages: &[Message],
keep_recent_turns: usize,
) -> Option<(usize, usize)> {
let user_indices: Vec<usize> = messages
.iter()
.enumerate()
.filter_map(|(i, m)| matches!(m, Message::User { .. }).then_some(i))
.collect();
if keep_recent_turns == 0 || user_indices.len() < keep_recent_turns + 1 {
return None;
}
let cover_len = *user_indices
.get(user_indices.len() - keep_recent_turns)
.expect("index guarded above");
let retain_from = cover_len;
(cover_len > 0).then_some((cover_len, retain_from))
}
pub fn apply_summary(messages: &[Message], summary_text: &str, retain_from: usize) -> Vec<Message> {
let mut out = Vec::with_capacity(1 + messages.len().saturating_sub(retain_from));
out.push(Message::Summary {
text: summary_text.to_string(),
});
out.extend(messages.iter().skip(retain_from).cloned());
out
}
#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)]
pub struct SummaryPayload {
pub text: String,
pub covered_message_count: usize,
pub retain_from: usize,
pub source_hash: String,
pub usage: Usage,
}
pub fn source_hash(covered: &[Message]) -> String {
let mut hasher = DefaultHasher::new();
for message in covered {
if let Ok(json) = serde_json::to_string(message) {
hasher.write(json.as_bytes());
}
}
format!("{:016x}", hasher.finish())
}
const SUMMARY_SYSTEM_PROMPT: &str = "你是助手,请将以下对话历史压缩为保留关键信息(目标、已做决定、结果、未决事项)的摘要,直接输出摘要正文。";
pub async fn generate_summary(
provider: &dyn ModelProvider,
credentials: &dyn CredentialProvider,
cancel: CancellationToken,
summary_model: &ModelId,
covered: &[Message],
) -> Result<String, ProviderError> {
let request = ModelRequest {
model: summary_model.clone(),
system_prompt: SUMMARY_SYSTEM_PROMPT.to_string(),
messages: covered.iter().map(project_for_summary).collect(),
tools: Vec::new(),
reasoning: ReasoningLevel::Off,
generation: GenerationOptions::default(),
provider_options: serde_json::Value::Null,
};
let stream = tokio::select! {
biased;
_ = cancel.cancelled() => return Err(ProviderError::Cancelled),
stream = provider.stream(request, credentials, cancel.clone()) => stream?,
};
let mut assembler = ResponseAssembler::new();
let mut stream = stream;
while let Some(item) = stream.next().await {
match item {
Ok(event) => {
if let Err(error) = assembler.push(event) {
return Err(ProviderError::Protocol(format!("{error:?}")));
}
}
Err(error) => return Err(error),
}
}
match assembler.finalize() {
Assembled::Complete {
blocks,
finish_reason: FinishReason::Stop,
..
} => {
let text = blocks
.iter()
.filter_map(|b| match b {
ContentBlock::Text { text } => Some(text.as_str()),
_ => None,
})
.collect::<Vec<_>>()
.join("");
if text.trim().is_empty() {
Err(ProviderError::Protocol(
"summary truncated or invalid".into(),
))
} else {
Ok(text)
}
}
_ => Err(ProviderError::Protocol(
"summary truncated or invalid".into(),
)),
}
}
fn project_for_summary(message: &Message) -> ModelMessage {
match message {
Message::User { blocks } => ModelMessage::User {
blocks: blocks
.iter()
.filter(|b| matches!(b, ContentBlock::Text { .. }))
.cloned()
.collect(),
},
Message::Assistant { blocks, .. } => {
let texts: Vec<ContentBlock> = blocks
.iter()
.filter_map(|b| match b {
ContentBlock::Text { text } => Some(ContentBlock::Text { text: text.clone() }),
_ => None,
})
.collect();
ModelMessage::Assistant {
blocks: if texts.is_empty() {
vec![ContentBlock::Text {
text: "(调用了工具)".into(),
}]
} else {
texts
},
}
}
Message::ToolResult { results } => ModelMessage::User {
blocks: vec![ContentBlock::Text {
text: format!(
"(工具结果) {}",
results
.iter()
.map(|r| r.text.as_str())
.collect::<Vec<_>>()
.join("\n")
),
}],
},
Message::Summary { text } => ModelMessage::Summary { text: text.clone() },
}
}
pub fn latest_valid_summary<'a>(
entry_ids: &[EntryId],
summaries: &'a [SummaryRecord],
) -> Option<(usize, &'a SummaryRecord)> {
summaries.iter().enumerate().rev().find_map(|(_i, record)| {
let position = entry_ids
.iter()
.position(|id| *id == record.covered_until_entry)?;
Some((position, record))
})
}
pub fn apply_latest_valid_summary(
messages: &[Message],
entry_ids: &[EntryId],
summaries: &[SummaryRecord],
) -> Vec<Message> {
match latest_valid_summary(entry_ids, summaries) {
Some((position, record)) => {
let mut out = Vec::with_capacity(1 + messages.len().saturating_sub(position + 1));
out.push(Message::Summary {
text: record.text.clone(),
});
out.extend(messages.iter().skip(position + 1).cloned());
out
}
None => messages.to_vec(),
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::ids::ToolCallId;
use crate::ids::{BranchId, ProviderId};
use crate::message::{FinishReason, ToolResultPayload};
use crate::provider::ProviderEvent;
use crate::testing::{FakeCredentialProvider, FakeProvider};
fn user(text: &str) -> Message {
Message::User {
blocks: vec![ContentBlock::Text { text: text.into() }],
}
}
fn assistant_text(text: &str) -> Message {
Message::Assistant {
blocks: vec![ContentBlock::Text { text: text.into() }],
finish_reason: FinishReason::Stop,
truncated: false,
}
}
fn assistant_tool_call(id: &str) -> Message {
Message::Assistant {
blocks: vec![ContentBlock::ToolCall {
id: ToolCallId::from(id),
name: "t".into(),
arguments: serde_json::json!({}),
}],
finish_reason: FinishReason::Stop,
truncated: false,
}
}
fn results(ids: &[&str]) -> Message {
Message::ToolResult {
results: ids
.iter()
.map(|id| ToolResultPayload {
call_id: ToolCallId::from(*id),
is_error: false,
text: "ok".into(),
})
.collect(),
}
}
fn alternating(turns: usize) -> Vec<Message> {
let mut out = Vec::new();
for i in 0..turns {
out.push(user(&format!("q{i}")));
out.push(assistant_text(&format!("a{i}")));
}
out
}
fn caps() -> crate::provider::ModelCapabilities {
crate::provider::ModelCapabilities {
context_tokens: 200_000,
max_output_tokens: 8_192,
supports_tools: true,
supports_images: true,
supports_reasoning: true,
}
}
fn text_events(text: &str, finish: FinishReason) -> Vec<ProviderEvent> {
vec![
ProviderEvent::ResponseStarted,
ProviderEvent::TextDelta {
block: 0,
text: text.into(),
},
ProviderEvent::ResponseCompleted {
finish_reason: finish,
},
]
}
#[test]
fn config_defaults_match_spec() {
let config = CompactionConfig::default();
assert_eq!(config.trigger_ratio, 0.8);
assert_eq!(config.keep_recent_turns, 4);
assert_eq!(config.max_attempts, 3);
let json = serde_json::to_string(&config).unwrap();
assert_eq!(
serde_json::from_str::<CompactionConfig>(&json).unwrap(),
config
);
}
#[test]
fn range_six_turns_keep_four_covers_two() {
let history = alternating(6);
let (cover, retain) = select_compaction_range(&history, 4).expect("compactable");
assert_eq!(cover, 4);
assert_eq!(retain, 4);
assert!(matches!(history[cover], Message::User { .. }));
}
#[test]
fn range_five_turns_keep_four_covers_one_turn() {
let history = alternating(5);
let (cover, _) = select_compaction_range(&history, 4).expect("compactable");
assert_eq!(cover, 2); }
#[test]
fn range_four_turns_keep_four_is_none() {
assert!(select_compaction_range(&alternating(4), 4).is_none());
}
#[test]
fn range_keep_more_than_turns_is_none() {
assert!(select_compaction_range(&alternating(2), 5).is_none());
}
#[test]
fn range_boundary_never_splits_tool_pairs() {
let history = vec![
user("q0"),
assistant_tool_call("c1"),
results(&["c1"]),
user("q1"),
assistant_tool_call("c2"),
results(&["c2"]),
user("q2"),
assistant_tool_call("c3"),
results(&["c3"]),
user("q3"),
assistant_text("a3"),
];
for keep in 1..=3 {
if let Some((cover, _)) = select_compaction_range(&history, keep) {
assert!(
matches!(history[cover], Message::User { .. }),
"boundary at {cover} must be a User message"
);
let covered = &history[..cover];
assert!(
crate::chain::project_and_validate(covered).is_ok(),
"covered prefix must project as a valid chain"
);
}
}
let (cover, _) = select_compaction_range(&history, 2).expect("compactable");
assert_eq!(cover, 6); }
#[test]
fn apply_summary_shapes_working_context() {
let history = alternating(3);
let out = apply_summary(&history, "摘要", 4);
assert_eq!(out.len(), 1 + (6 - 4));
assert!(matches!(&out[0], Message::Summary { text } if text == "摘要"));
assert_eq!(out[1], history[4]);
}
#[tokio::test]
async fn generate_summary_returns_joined_text() {
let provider = FakeProvider::new(caps());
provider.push_events(text_events("目标:测试", FinishReason::Stop));
let creds = FakeCredentialProvider::default();
let text = generate_summary(
&provider,
&creds,
CancellationToken::new(),
&ModelId::from("sum-model"),
&alternating(3),
)
.await
.expect("summary ok");
assert_eq!(text, "目标:测试");
let request = &provider.requests()[0];
assert_eq!(request.model, ModelId::from("sum-model"));
assert!(request.tools.is_empty());
assert_eq!(request.reasoning, ReasoningLevel::Off);
assert!(request.system_prompt.contains("摘要"));
assert_eq!(request.messages.len(), 6);
}
#[tokio::test]
async fn generate_summary_length_finish_is_protocol_error() {
let provider = FakeProvider::new(caps());
provider.push_events(text_events("截断", FinishReason::Length));
let err = generate_summary(
&provider,
&FakeCredentialProvider::default(),
CancellationToken::new(),
&ModelId::from("m"),
&alternating(2),
)
.await
.expect_err("length must fail");
assert!(matches!(err, ProviderError::Protocol(_)));
}
#[tokio::test]
async fn generate_summary_passes_provider_error_through() {
let provider = FakeProvider::new(caps());
provider.push_error(ProviderError::Network("down".into()));
let err = generate_summary(
&provider,
&FakeCredentialProvider::default(),
CancellationToken::new(),
&ModelId::from("m"),
&alternating(2),
)
.await
.expect_err("provider error passes through");
assert!(matches!(err, ProviderError::Network(_)));
}
#[tokio::test]
async fn summary_projection_flattens_tools() {
let provider = FakeProvider::new(caps());
provider.push_events(text_events("s", FinishReason::Stop));
let covered = vec![user("q"), assistant_tool_call("c1"), results(&["c1"])];
generate_summary(
&provider,
&FakeCredentialProvider::default(),
CancellationToken::new(),
&ModelId::from("m"),
&covered,
)
.await
.unwrap();
let messages = &provider.requests()[0].messages;
assert_eq!(messages.len(), 3);
assert!(matches!(&messages[1], ModelMessage::Assistant { blocks }
if matches!(&blocks[0], ContentBlock::Text { text } if text == "(调用了工具)")));
assert!(matches!(&messages[2], ModelMessage::User { blocks }
if matches!(&blocks[0], ContentBlock::Text { text } if text.starts_with("(工具结果)"))));
}
fn summary_record(entry: &str, text: &str) -> SummaryRecord {
SummaryRecord {
summary_id: format!("summary-{entry}"),
branch_id: BranchId::from("b1"),
covered_until_entry: EntryId::from(entry),
text: text.to_string(),
source_hash: "h".to_string(),
provider: ProviderId::from("fake"),
model: ModelId::from("m"),
prompt_version: String::new(),
usage: Usage::default(),
}
}
#[test]
fn latest_valid_summary_prefers_latest_applicable() {
let entry_ids: Vec<EntryId> = (0..4).map(|i| EntryId::from(format!("e{i}"))).collect();
let messages: Vec<Message> = entry_ids
.iter()
.enumerate()
.map(|(i, _)| {
if i % 2 == 0 {
user(&format!("q{i}"))
} else {
assistant_text("a")
}
})
.collect();
let summaries = vec![
summary_record("e1", "旧摘要"),
summary_record("e9", "不在链上"),
summary_record("e2", "新摘要"),
];
let out = apply_latest_valid_summary(&messages, &entry_ids, &summaries);
assert_eq!(out.len(), 1 + (4 - 3));
assert!(matches!(&out[0], Message::Summary { text } if text == "新摘要"));
assert_eq!(out[1], messages[3]);
}
#[test]
fn latest_valid_summary_none_when_covered_entry_missing() {
let entry_ids = vec![EntryId::from("e0"), EntryId::from("e1")];
let messages = vec![user("q0"), assistant_text("a0")];
let summaries = vec![summary_record("e9", "陈旧")];
let out = apply_latest_valid_summary(&messages, &entry_ids, &summaries);
assert_eq!(out, messages);
}
#[test]
fn source_hash_is_stable_and_sensitive() {
let covered = alternating(2);
assert_eq!(source_hash(&covered), source_hash(&covered));
let mut other = covered.clone();
other[0] = user("different");
assert_ne!(source_hash(&covered), source_hash(&other));
}
#[test]
fn current_model_selector_is_identity() {
let model = ModelId::from("m1");
assert_eq!(CurrentModelSelector.select_model(&model), model);
let _context = CompactionContext {
current_model: model.clone(),
};
}
}