1use agent_base::llm_trait::{ChatRequest, LlmProvider};
9use agent_base::{AgentResult, ChatMessage, StreamChunk};
10
11const SUMMARIZATION_PROMPT: &str = "\
22You are performing a CONTEXT CHECKPOINT COMPACTION. \
23Create a handoff summary for another LLM that will resume the task.
24
25The original goal of this session was: {goal}
26
27User messages have been preserved separately. \
28Summarize ONLY the assistant responses and tool results below.
29
30Include:
31- What the assistant did (tools called, actions taken, results found)
32- Key decisions made and important constraints discovered
33- What remains to be done (clear next steps)
34{lang}
35Do NOT reproduce assistant text replies verbatim (poems, articles, code examples, etc.). \
36Only describe what was done, not the content itself.
37
38Be concise, structured, and focused on helping the next LLM seamlessly continue the work. \
39Do not repeat work that has already been done. \
40Output ONLY the summary text, no preamble, about {max_chars} characters max.
41
42=== ASSISTANT AND TOOL RESPONSES ===
43{transcript}";
44
45const LANG_INSTRUCTION_CJK: &str =
47 "Respond in the same language as the conversation (CJK detected).";
48
49const LANG_INSTRUCTION_DEFAULT: &str = "";
52
53pub async fn summarize(
66 client: &dyn LlmProvider,
67 transcript: &str,
68 original_goal: &str,
69 max_chars: usize,
70 on_progress: Option<&(dyn Fn(usize) + Sync)>,
71) -> AgentResult<String> {
72 if max_chars == 0 {
73 return Ok(String::new());
74 }
75
76 let lang = language_instruction(&format!("{original_goal}\n{transcript}"));
77 let prompt = build_prompt(original_goal, lang, max_chars, transcript);
78
79 let system = ChatMessage::system(
80 "You are a conversation summarizer for an AI agent that can call tools \
81 (browser, shell, search, etc.).",
82 );
83 let user = ChatMessage::user(prompt);
84
85 let request = ChatRequest::new(vec![system, user]);
86 let mut stream = client
87 .stream(request)
88 .await
89 .map_err(agent_base::AgentError::from)?;
90 let mut text = String::new();
91 while let Some(chunk) = stream.next().await {
92 match chunk.map_err(agent_base::AgentError::from)? {
93 StreamChunk::Text(t) => {
94 text.push_str(&t);
95 if let Some(cb) = on_progress {
96 cb(text.len());
97 }
98 }
99 StreamChunk::Stop { .. } => break,
100 _ => {}
101 }
102 }
103
104 Ok(truncate_summary_output(&text, max_chars))
105}
106
107fn build_prompt(goal: &str, lang: &str, max_chars: usize, transcript: &str) -> String {
115 let mut out = String::with_capacity(SUMMARIZATION_PROMPT.len() + goal.len() + transcript.len());
118 let mut chars = SUMMARIZATION_PROMPT.chars().peekable();
119 while let Some(c) = chars.next() {
120 if c == '{' {
121 let rest: String = chars.clone().take_while(|ch| *ch != '}').collect();
123 match rest.as_str() {
124 "goal" => {
125 out.push_str(goal);
126 for _ in 0..=rest.len() {
128 chars.next();
129 }
130 }
131 "lang" => {
132 out.push_str(lang);
133 for _ in 0..=rest.len() {
134 chars.next();
135 }
136 }
137 "max_chars" => {
138 out.push_str(&max_chars.to_string());
139 for _ in 0..=rest.len() {
140 chars.next();
141 }
142 }
143 "transcript" => {
144 out.push_str(transcript);
145 for _ in 0..=rest.len() {
146 chars.next();
147 }
148 }
149 _ => out.push(c),
150 }
151 } else {
152 out.push(c);
153 }
154 }
155 out
156}
157
158pub fn language_instruction(text: &str) -> &'static str {
166 let meaningful: Vec<char> = text.chars().filter(|c| !c.is_whitespace()).collect();
167 if meaningful.is_empty() {
168 return LANG_INSTRUCTION_DEFAULT;
169 }
170 let cjk_count = meaningful.iter().filter(|c| is_cjk(**c)).count();
171 if cjk_count * 5 >= meaningful.len() {
172 LANG_INSTRUCTION_CJK
173 } else {
174 LANG_INSTRUCTION_DEFAULT
175 }
176}
177
178fn is_cjk(c: char) -> bool {
180 matches!(c,
181 '\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}' )
190}
191
192pub fn truncate_summary_output(text: &str, max_chars: usize) -> String {
200 let char_count = text.chars().count();
201 if char_count <= max_chars {
202 return text.to_string();
203 }
204 if max_chars == 0 {
205 return String::new();
206 }
207 let budget = max_chars.saturating_sub(1);
209 let front = (budget as f64 * 0.8) as usize;
210 let rear = budget.saturating_sub(front);
211 let front_s: String = text.chars().take(front).collect();
212 let rear_s: String = text
213 .chars()
214 .rev()
215 .take(rear)
216 .collect::<Vec<_>>()
217 .into_iter()
218 .rev()
219 .collect();
220 format!("{front_s}…{rear_s}")
221}
222
223#[cfg(test)]
226mod tests {
227 use super::*;
228 use agent_base::llm_trait::response::FinishReason;
229 use agent_base::llm_trait::types::UsageInfo;
230 use agent_base::llm_trait::{
231 Capabilities, ChatRequest, ChatResponse, ChatStream, LlmError, LlmProvider, ProviderInfo,
232 };
233
234 struct PromptCapture {
238 captured: std::sync::Arc<std::sync::Mutex<Vec<String>>>,
239 response: String,
240 }
241
242 #[async_trait::async_trait]
243 impl LlmProvider for PromptCapture {
244 async fn stream(&self, request: ChatRequest) -> Result<ChatStream, LlmError> {
245 for msg in &request.messages {
247 if let ChatMessage::User { content, .. } = msg {
248 self.captured.lock().unwrap().push(content.clone());
249 }
250 }
251 let response = self.response.clone();
252 Ok(ChatStream::new(Box::pin(futures_util::stream::once(
253 async move { Ok(agent_base::StreamChunk::Text(response)) },
254 ))))
255 }
256
257 async fn chat(&self, request: ChatRequest) -> Result<ChatResponse, LlmError> {
258 for msg in &request.messages {
259 if let ChatMessage::User { content, .. } = msg {
260 self.captured.lock().unwrap().push(content.clone());
261 }
262 }
263 Ok(ChatResponse {
264 content: self.response.clone(),
265 tool_calls: vec![],
266 usage: UsageInfo::default(),
267 finish_reason: FinishReason::Stop,
268 raw: None,
269 reasoning_content: None,
270 thinking_signature: None,
271 })
272 }
273
274 fn capabilities(&self) -> Capabilities {
275 Capabilities::default()
276 }
277
278 fn info(&self) -> ProviderInfo {
279 ProviderInfo {
280 name: "stub".to_string(),
281 model: "stub-model".to_string(),
282 version: None,
283 }
284 }
285 }
286
287 #[test]
290 fn test_language_instruction_cjk() {
291 assert_eq!(
292 language_instruction("用户问了关于日志分析的问题,发现了5次操作"),
293 LANG_INSTRUCTION_CJK
294 );
295 }
296
297 #[test]
298 fn test_language_instruction_english() {
299 assert_eq!(
300 language_instruction("The user asked about log analysis, found 5 operations"),
301 LANG_INSTRUCTION_DEFAULT
302 );
303 }
304
305 #[test]
306 fn test_language_instruction_mostly_latin_with_some_cjk() {
307 assert_eq!(
308 language_instruction("The user asked about 日志 analysis of the system"),
309 LANG_INSTRUCTION_DEFAULT
310 );
311 }
312
313 #[test]
314 fn test_language_instruction_mixed_heavy_cjk() {
315 assert_eq!(
316 language_instruction("分析日志时发现 operations 有5次 user asked 分析"),
317 LANG_INSTRUCTION_CJK
318 );
319 }
320
321 #[test]
322 fn test_language_instruction_empty() {
323 assert_eq!(language_instruction(""), LANG_INSTRUCTION_DEFAULT);
324 }
325
326 #[test]
327 fn test_language_instruction_whitespace_only() {
328 assert_eq!(language_instruction(" \n\t "), LANG_INSTRUCTION_DEFAULT);
329 }
330
331 #[test]
332 fn test_language_instruction_hangul() {
333 assert_eq!(
334 language_instruction("사용자가 로그 분석에 대해 물었습니다"),
335 LANG_INSTRUCTION_CJK
336 );
337 }
338
339 #[test]
342 fn test_truncate_short_text() {
343 assert_eq!(truncate_summary_output("short", 100), "short");
344 }
345
346 #[test]
347 fn test_truncate_long_text_preserves_ends() {
348 let text = "a".repeat(500) + "TAIL";
349 let result = truncate_summary_output(&text, 100);
350 assert!(result.chars().count() <= 100);
351 assert!(result.starts_with('a'));
352 assert!(result.contains("TAIL"));
353 assert!(result.contains('…'));
354 }
355
356 #[test]
357 fn test_truncate_exact_boundary() {
358 let text = "x".repeat(100);
359 assert_eq!(truncate_summary_output(&text, 100), text);
360 }
361
362 #[test]
363 fn test_truncate_zero() {
364 assert_eq!(truncate_summary_output("anything", 0), "");
365 }
366
367 #[test]
370 fn test_build_prompt_all_placeholders_filled() {
371 let prompt = build_prompt("fix the bug", LANG_INSTRUCTION_CJK, 5000, "user: hello");
372 assert!(prompt.contains("fix the bug"));
373 assert!(prompt.contains("CJK detected"));
374 assert!(prompt.contains("5000"));
375 assert!(prompt.contains("user: hello"));
376 assert!(!prompt.contains("{goal}"));
378 assert!(!prompt.contains("{lang}"));
379 assert!(!prompt.contains("{transcript}"));
380 assert!(!prompt.contains("{max_chars}"));
381 }
382
383 #[test]
384 fn test_build_prompt_goal_with_placeholder_literals_not_polluted() {
385 let goal = "按 {lang} 字段分组,max={max_chars}";
387 let prompt = build_prompt(goal, LANG_INSTRUCTION_CJK, 5000, "data");
388 assert!(
389 prompt.contains("按 {lang} 字段分组,max={max_chars}"),
390 "literal placeholders in goal must survive: {prompt}"
391 );
392 assert!(prompt.contains("CJK detected"));
394 assert!(prompt.contains("5000"));
395 }
396
397 #[test]
398 fn test_build_prompt_transcript_with_goal_placeholder_not_polluted() {
399 let transcript = "user: use {goal} as the key";
400 let prompt = build_prompt("real goal", LANG_INSTRUCTION_DEFAULT, 1000, transcript);
401 assert!(
402 prompt.contains("use {goal} as the key"),
403 "literal {{goal}} in transcript must survive: {prompt}"
404 );
405 assert!(prompt.contains("real goal"));
406 }
407
408 #[tokio::test]
411 async fn test_summarize_prompt_contains_goal_and_lang() {
412 let captured = std::sync::Arc::new(std::sync::Mutex::new(Vec::new()));
413 let client = std::sync::Arc::new(PromptCapture {
414 captured: captured.clone(),
415 response: "a summary".into(),
416 });
417
418 let _ = summarize(
420 client.as_ref(),
421 "tool output here",
422 "分析服务器日志中的延迟问题",
423 5000,
424 None,
425 )
426 .await
427 .unwrap();
428
429 let prompts = captured.lock().unwrap();
430 assert_eq!(prompts.len(), 1);
431 let prompt = &prompts[0];
432 assert!(
433 prompt.contains("分析服务器日志中的延迟问题"),
434 "goal missing"
435 );
436 assert!(prompt.contains("CJK detected"), "lang instruction missing");
437 assert!(prompt.contains("5000"), "max_chars missing");
438 assert!(prompt.contains("tool output here"), "transcript missing");
439 }
440
441 #[tokio::test]
442 async fn test_summarize_output_truncated() {
443 let long_response = "x".repeat(2000);
444 let captured = std::sync::Arc::new(std::sync::Mutex::new(Vec::new()));
445 let client = std::sync::Arc::new(PromptCapture {
446 captured: captured.clone(),
447 response: long_response,
448 });
449
450 let result = summarize(client.as_ref(), "t", "g", 100, None)
451 .await
452 .unwrap();
453 assert!(result.chars().count() <= 100);
454 }
455
456 #[tokio::test]
457 async fn test_summarize_max_chars_zero() {
458 let captured = std::sync::Arc::new(std::sync::Mutex::new(Vec::new()));
459 let client = std::sync::Arc::new(PromptCapture {
460 captured: captured.clone(),
461 response: "ignored".into(),
462 });
463
464 let result = summarize(client.as_ref(), "t", "g", 0, None).await.unwrap();
465 assert!(result.is_empty());
466 assert!(captured.lock().unwrap().is_empty());
468 }
469
470 #[tokio::test]
471 async fn test_summarize_returns_response_content() {
472 let expected_summary =
474 "User said hello and asked for a poem. Assistant provided a classical Chinese poem.";
475 let captured = std::sync::Arc::new(std::sync::Mutex::new(Vec::new()));
476 let client = std::sync::Arc::new(PromptCapture {
477 captured: captured.clone(),
478 response: expected_summary.into(),
479 });
480
481 let transcript = "[user] 你好\n[assistant] 你好!有什么我可以帮你的吗?\n[user] 来一首古诗";
482 let result = summarize(client.as_ref(), transcript, "你好", 5000, None)
483 .await
484 .unwrap();
485
486 assert_eq!(
488 result, expected_summary,
489 "summarize should return the LLM response"
490 );
491
492 let prompts = captured.lock().unwrap();
494 assert_eq!(prompts.len(), 1, "should have sent exactly one prompt");
495 let prompt = &prompts[0];
496 assert!(prompt.contains("你好"), "prompt should contain the goal");
497 assert!(
498 prompt.contains("来一首古诗"),
499 "prompt should contain the transcript"
500 );
501 assert!(
502 prompt.contains("CONTEXT CHECKPOINT COMPACTION"),
503 "prompt should contain the compaction instruction"
504 );
505 }
506
507 #[tokio::test]
508 #[ignore] async fn test_summarize_with_real_deepseek_api() {
510 let api_key = std::env::var("DEEPSEEK_API_KEY").unwrap_or_default();
515 if api_key.is_empty() {
516 eprintln!("Skipping test: DEEPSEEK_API_KEY not set");
517 return;
518 }
519
520 let _base_url = std::env::var("DEEPSEEK_BASE_URL")
521 .unwrap_or_else(|_| "https://api.deepseek.com".to_string());
522
523 return; }
532}