use agent_base::llm_trait::{ChatRequest, LlmProvider};
use agent_base::{AgentResult, ChatMessage, StreamChunk};
const SUMMARIZATION_PROMPT: &str = "\
You are performing a CONTEXT CHECKPOINT COMPACTION. \
Create a handoff summary for another LLM that will resume the task.
The original goal of this session was: {goal}
User messages have been preserved separately. \
Summarize ONLY the assistant responses and tool results below.
Include:
- What the assistant did (tools called, actions taken, results found)
- Key decisions made and important constraints discovered
- What remains to be done (clear next steps)
{lang}
Do NOT reproduce assistant text replies verbatim (poems, articles, code examples, etc.). \
Only describe what was done, not the content itself.
Be concise, structured, and focused on helping the next LLM seamlessly continue the work. \
Do not repeat work that has already been done. \
Output ONLY the summary text, no preamble, about {max_chars} characters max.
=== ASSISTANT AND TOOL RESPONSES ===
{transcript}";
const LANG_INSTRUCTION_CJK: &str =
"Respond in the same language as the conversation (CJK detected).";
const LANG_INSTRUCTION_DEFAULT: &str = "";
pub async fn summarize(
client: &dyn LlmProvider,
transcript: &str,
original_goal: &str,
max_chars: usize,
on_progress: Option<&(dyn Fn(usize) + Sync)>,
) -> AgentResult<String> {
if max_chars == 0 {
return Ok(String::new());
}
let lang = language_instruction(&format!("{original_goal}\n{transcript}"));
let prompt = build_prompt(original_goal, lang, max_chars, transcript);
let system = ChatMessage::system(
"You are a conversation summarizer for an AI agent that can call tools \
(browser, shell, search, etc.).",
);
let user = ChatMessage::user(prompt);
let request = ChatRequest::new(vec![system, user]);
let mut stream = client
.stream(request)
.await
.map_err(agent_base::AgentError::from)?;
let mut text = String::new();
while let Some(chunk) = stream.next().await {
match chunk.map_err(agent_base::AgentError::from)? {
StreamChunk::Text(t) => {
text.push_str(&t);
if let Some(cb) = on_progress {
cb(text.len());
}
}
StreamChunk::Stop { .. } => break,
_ => {}
}
}
Ok(truncate_summary_output(&text, max_chars))
}
fn build_prompt(goal: &str, lang: &str, max_chars: usize, transcript: &str) -> String {
let mut out = String::with_capacity(SUMMARIZATION_PROMPT.len() + goal.len() + transcript.len());
let mut chars = SUMMARIZATION_PROMPT.chars().peekable();
while let Some(c) = chars.next() {
if c == '{' {
let rest: String = chars.clone().take_while(|ch| *ch != '}').collect();
match rest.as_str() {
"goal" => {
out.push_str(goal);
for _ in 0..=rest.len() {
chars.next();
}
}
"lang" => {
out.push_str(lang);
for _ in 0..=rest.len() {
chars.next();
}
}
"max_chars" => {
out.push_str(&max_chars.to_string());
for _ in 0..=rest.len() {
chars.next();
}
}
"transcript" => {
out.push_str(transcript);
for _ in 0..=rest.len() {
chars.next();
}
}
_ => out.push(c),
}
} else {
out.push(c);
}
}
out
}
pub fn language_instruction(text: &str) -> &'static str {
let meaningful: Vec<char> = text.chars().filter(|c| !c.is_whitespace()).collect();
if meaningful.is_empty() {
return LANG_INSTRUCTION_DEFAULT;
}
let cjk_count = meaningful.iter().filter(|c| is_cjk(**c)).count();
if cjk_count * 5 >= meaningful.len() {
LANG_INSTRUCTION_CJK
} else {
LANG_INSTRUCTION_DEFAULT
}
}
fn is_cjk(c: char) -> bool {
matches!(c,
'\u{4E00}'..='\u{9FFF}' | '\u{3400}'..='\u{4DBF}' | '\u{F900}'..='\u{FAFF}' | '\u{3000}'..='\u{303F}' | '\u{FF00}'..='\u{FFEF}' | '\u{3040}'..='\u{309F}' | '\u{30A0}'..='\u{30FF}' | '\u{AC00}'..='\u{D7AF}' )
}
pub fn truncate_summary_output(text: &str, max_chars: usize) -> String {
let char_count = text.chars().count();
if char_count <= max_chars {
return text.to_string();
}
if max_chars == 0 {
return String::new();
}
let budget = max_chars.saturating_sub(1);
let front = (budget as f64 * 0.8) as usize;
let rear = budget.saturating_sub(front);
let front_s: String = text.chars().take(front).collect();
let rear_s: String = text
.chars()
.rev()
.take(rear)
.collect::<Vec<_>>()
.into_iter()
.rev()
.collect();
format!("{front_s}…{rear_s}")
}
#[cfg(test)]
mod tests {
use super::*;
use agent_base::llm_trait::response::FinishReason;
use agent_base::llm_trait::types::UsageInfo;
use agent_base::llm_trait::{
Capabilities, ChatRequest, ChatResponse, ChatStream, LlmError, LlmProvider, ProviderInfo,
};
struct PromptCapture {
captured: std::sync::Arc<std::sync::Mutex<Vec<String>>>,
response: String,
}
#[async_trait::async_trait]
impl LlmProvider for PromptCapture {
async fn stream(&self, request: ChatRequest) -> Result<ChatStream, LlmError> {
for msg in &request.messages {
if let ChatMessage::User { content, .. } = msg {
self.captured.lock().unwrap().push(content.clone());
}
}
let response = self.response.clone();
Ok(ChatStream::new(Box::pin(futures_util::stream::once(
async move { Ok(agent_base::StreamChunk::Text(response)) },
))))
}
async fn chat(&self, request: ChatRequest) -> Result<ChatResponse, LlmError> {
for msg in &request.messages {
if let ChatMessage::User { content, .. } = msg {
self.captured.lock().unwrap().push(content.clone());
}
}
Ok(ChatResponse {
content: self.response.clone(),
tool_calls: vec![],
usage: UsageInfo::default(),
finish_reason: FinishReason::Stop,
raw: None,
reasoning_content: None,
thinking_signature: None,
})
}
fn capabilities(&self) -> Capabilities {
Capabilities::default()
}
fn info(&self) -> ProviderInfo {
ProviderInfo {
name: "stub".to_string(),
model: "stub-model".to_string(),
version: None,
}
}
}
#[test]
fn test_language_instruction_cjk() {
assert_eq!(
language_instruction("用户问了关于日志分析的问题,发现了5次操作"),
LANG_INSTRUCTION_CJK
);
}
#[test]
fn test_language_instruction_english() {
assert_eq!(
language_instruction("The user asked about log analysis, found 5 operations"),
LANG_INSTRUCTION_DEFAULT
);
}
#[test]
fn test_language_instruction_mostly_latin_with_some_cjk() {
assert_eq!(
language_instruction("The user asked about 日志 analysis of the system"),
LANG_INSTRUCTION_DEFAULT
);
}
#[test]
fn test_language_instruction_mixed_heavy_cjk() {
assert_eq!(
language_instruction("分析日志时发现 operations 有5次 user asked 分析"),
LANG_INSTRUCTION_CJK
);
}
#[test]
fn test_language_instruction_empty() {
assert_eq!(language_instruction(""), LANG_INSTRUCTION_DEFAULT);
}
#[test]
fn test_language_instruction_whitespace_only() {
assert_eq!(language_instruction(" \n\t "), LANG_INSTRUCTION_DEFAULT);
}
#[test]
fn test_language_instruction_hangul() {
assert_eq!(
language_instruction("사용자가 로그 분석에 대해 물었습니다"),
LANG_INSTRUCTION_CJK
);
}
#[test]
fn test_truncate_short_text() {
assert_eq!(truncate_summary_output("short", 100), "short");
}
#[test]
fn test_truncate_long_text_preserves_ends() {
let text = "a".repeat(500) + "TAIL";
let result = truncate_summary_output(&text, 100);
assert!(result.chars().count() <= 100);
assert!(result.starts_with('a'));
assert!(result.contains("TAIL"));
assert!(result.contains('…'));
}
#[test]
fn test_truncate_exact_boundary() {
let text = "x".repeat(100);
assert_eq!(truncate_summary_output(&text, 100), text);
}
#[test]
fn test_truncate_zero() {
assert_eq!(truncate_summary_output("anything", 0), "");
}
#[test]
fn test_build_prompt_all_placeholders_filled() {
let prompt = build_prompt("fix the bug", LANG_INSTRUCTION_CJK, 5000, "user: hello");
assert!(prompt.contains("fix the bug"));
assert!(prompt.contains("CJK detected"));
assert!(prompt.contains("5000"));
assert!(prompt.contains("user: hello"));
assert!(!prompt.contains("{goal}"));
assert!(!prompt.contains("{lang}"));
assert!(!prompt.contains("{transcript}"));
assert!(!prompt.contains("{max_chars}"));
}
#[test]
fn test_build_prompt_goal_with_placeholder_literals_not_polluted() {
let goal = "按 {lang} 字段分组,max={max_chars}";
let prompt = build_prompt(goal, LANG_INSTRUCTION_CJK, 5000, "data");
assert!(
prompt.contains("按 {lang} 字段分组,max={max_chars}"),
"literal placeholders in goal must survive: {prompt}"
);
assert!(prompt.contains("CJK detected"));
assert!(prompt.contains("5000"));
}
#[test]
fn test_build_prompt_transcript_with_goal_placeholder_not_polluted() {
let transcript = "user: use {goal} as the key";
let prompt = build_prompt("real goal", LANG_INSTRUCTION_DEFAULT, 1000, transcript);
assert!(
prompt.contains("use {goal} as the key"),
"literal {{goal}} in transcript must survive: {prompt}"
);
assert!(prompt.contains("real goal"));
}
#[tokio::test]
async fn test_summarize_prompt_contains_goal_and_lang() {
let captured = std::sync::Arc::new(std::sync::Mutex::new(Vec::new()));
let client = std::sync::Arc::new(PromptCapture {
captured: captured.clone(),
response: "a summary".into(),
});
let _ = summarize(
client.as_ref(),
"tool output here",
"分析服务器日志中的延迟问题",
5000,
None,
)
.await
.unwrap();
let prompts = captured.lock().unwrap();
assert_eq!(prompts.len(), 1);
let prompt = &prompts[0];
assert!(
prompt.contains("分析服务器日志中的延迟问题"),
"goal missing"
);
assert!(prompt.contains("CJK detected"), "lang instruction missing");
assert!(prompt.contains("5000"), "max_chars missing");
assert!(prompt.contains("tool output here"), "transcript missing");
}
#[tokio::test]
async fn test_summarize_output_truncated() {
let long_response = "x".repeat(2000);
let captured = std::sync::Arc::new(std::sync::Mutex::new(Vec::new()));
let client = std::sync::Arc::new(PromptCapture {
captured: captured.clone(),
response: long_response,
});
let result = summarize(client.as_ref(), "t", "g", 100, None)
.await
.unwrap();
assert!(result.chars().count() <= 100);
}
#[tokio::test]
async fn test_summarize_max_chars_zero() {
let captured = std::sync::Arc::new(std::sync::Mutex::new(Vec::new()));
let client = std::sync::Arc::new(PromptCapture {
captured: captured.clone(),
response: "ignored".into(),
});
let result = summarize(client.as_ref(), "t", "g", 0, None).await.unwrap();
assert!(result.is_empty());
assert!(captured.lock().unwrap().is_empty());
}
#[tokio::test]
async fn test_summarize_returns_response_content() {
let expected_summary =
"User said hello and asked for a poem. Assistant provided a classical Chinese poem.";
let captured = std::sync::Arc::new(std::sync::Mutex::new(Vec::new()));
let client = std::sync::Arc::new(PromptCapture {
captured: captured.clone(),
response: expected_summary.into(),
});
let transcript = "[user] 你好\n[assistant] 你好!有什么我可以帮你的吗?\n[user] 来一首古诗";
let result = summarize(client.as_ref(), transcript, "你好", 5000, None)
.await
.unwrap();
assert_eq!(
result, expected_summary,
"summarize should return the LLM response"
);
let prompts = captured.lock().unwrap();
assert_eq!(prompts.len(), 1, "should have sent exactly one prompt");
let prompt = &prompts[0];
assert!(prompt.contains("你好"), "prompt should contain the goal");
assert!(
prompt.contains("来一首古诗"),
"prompt should contain the transcript"
);
assert!(
prompt.contains("CONTEXT CHECKPOINT COMPACTION"),
"prompt should contain the compaction instruction"
);
}
#[tokio::test]
#[ignore] async fn test_summarize_with_real_deepseek_api() {
let api_key = std::env::var("DEEPSEEK_API_KEY").unwrap_or_default();
if api_key.is_empty() {
eprintln!("Skipping test: DEEPSEEK_API_KEY not set");
return;
}
let _base_url = std::env::var("DEEPSEEK_BASE_URL")
.unwrap_or_else(|_| "https://api.deepseek.com".to_string());
return;
}
}