Skip to main content

agent_works/compression/
summarizer.rs

1//! LLM summarisation for context compression.
2//!
3//! Provides the [`summarize`] function that calls an LLM to produce a compact
4//! handoff summary of an older conversation block.  The summary is designed for
5//! *another* LLM to pick up seamlessly — it preserves the original goal, key
6//! decisions, tool findings, and clear next steps.
7
8use agent_base::llm_trait::{ChatRequest, LlmProvider};
9use agent_base::{AgentResult, ChatMessage, StreamChunk};
10
11/// Handoff-style summarisation prompt template.
12///
13/// Framed as a "context checkpoint compaction" — tells the LLM this is a
14/// handoff to another model, not a self-summary.  Placeholders `{goal}`,
15/// `{lang}`, `{max_chars}`, and `{transcript}` are filled by [`build_prompt`].
16///
17/// User messages are preserved verbatim (handled by the compactor), so the
18/// summarizer only receives assistant and tool responses.  The prompt
19/// explicitly instructs the LLM to describe *what happened* without
20/// reproducing assistant text verbatim (poems, code, articles, etc.).
21const 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
45/// Language instruction injected when CJK content is detected.
46const LANG_INSTRUCTION_CJK: &str =
47    "Respond in the same language as the conversation (CJK detected).";
48
49/// No extra language instruction for predominantly Latin text — English is the
50/// default and the prompt itself is English.
51const LANG_INSTRUCTION_DEFAULT: &str = "";
52
53// ── Public API ───────────────────────────────────────────────────────────────
54
55/// Summarise a transcript block via the LLM.
56///
57/// * `client` — the LLM client to call.
58/// * `transcript` — serialised older conversation (output of `serialize_block`).
59/// * `original_goal` — first user message, truncated to avoid blowing the prompt.
60/// * `max_chars` — target max length for the summary.
61/// * `on_progress` — optional callback invoked with cumulative character count
62///   as the LLM streams in. Useful for showing "generating summary... X chars".
63///
64/// Returns the summary text, truncated to `max_chars` if the LLM over-shoots.
65pub 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
107// ── Prompt construction ──────────────────────────────────────────────────────
108
109/// Build the summarisation prompt with all placeholders filled.
110///
111/// Uses single-pass character replacement to prevent placeholder pollution:
112/// if `original_goal` contains `{lang}` etc. as literal text, it is preserved
113/// verbatim (unlike a `.replace()` chain which would corrupt it).
114fn build_prompt(goal: &str, lang: &str, max_chars: usize, transcript: &str) -> String {
115    // Single-pass replacement: walk the template, emit literal chars or fill
116    // placeholders.  Content of goal/transcript is never re-scanned.
117    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            // Try to match a placeholder.
122            let rest: String = chars.clone().take_while(|ch| *ch != '}').collect();
123            match rest.as_str() {
124                "goal" => {
125                    out.push_str(goal);
126                    // Skip past the closing '}'.
127                    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
158// ── Language detection ───────────────────────────────────────────────────────
159
160/// Detect the dominant script of a text and return the appropriate language
161/// instruction for the summarisation prompt.
162///
163/// Returns `LANG_INSTRUCTION_CJK` when ≥ 20 % of non-whitespace characters
164/// are CJK, `LANG_INSTRUCTION_DEFAULT` otherwise.
165pub 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
178/// Returns `true` if `c` is a CJK ideograph, kana, hangul, or punctuation.
179fn is_cjk(c: char) -> bool {
180    matches!(c,
181        '\u{4E00}'..='\u{9FFF}'   // CJK Unified Ideographs
182        | '\u{3400}'..='\u{4DBF}' // CJK Unified Ideographs Extension A
183        | '\u{F900}'..='\u{FAFF}' // CJK Compatibility Ideographs
184        | '\u{3000}'..='\u{303F}' // CJK Symbols and Punctuation
185        | '\u{FF00}'..='\u{FFEF}' // Fullwidth Forms
186        | '\u{3040}'..='\u{309F}' // Hiragana
187        | '\u{30A0}'..='\u{30FF}' // Katakana
188        | '\u{AC00}'..='\u{D7AF}' // Hangul Syllables
189    )
190}
191
192// ── Output truncation ────────────────────────────────────────────────────────
193
194/// Truncate a summary to `max_chars` (front 80 % + rear 20 %).
195///
196/// Preserves the beginning (which typically contains the most important
197/// context) and the end (recent conclusions / next steps), dropping the
198/// middle when the LLM over-shoots.
199pub 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    // Reserve 1 char for the '…' separator.
208    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// ── Tests ────────────────────────────────────────────────────────────────────
224
225#[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    // ── Test helpers ──────────────────────────────────────────────────────
235
236    /// Mock that captures the prompt sent to the LLM.
237    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            // Capture the user message content.
246            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    // ── language_instruction ──────────────────────────────────────────────
288
289    #[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    // ── truncate_summary_output ───────────────────────────────────────────
340
341    #[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    // ── build_prompt ──────────────────────────────────────────────────────
368
369    #[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        // No unfilled placeholders.
377        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        // goal contains literal {lang} and {max_chars} — must NOT be replaced.
386        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        // The actual {lang} and {max_chars} placeholders should still be filled.
393        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    // ── summarize (mock-client integration) ───────────────────────────────
409
410    #[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        // CJK goal → should inject CJK instruction.
419        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        // Should not even call the LLM.
467        assert!(captured.lock().unwrap().is_empty());
468    }
469
470    #[tokio::test]
471    async fn test_summarize_returns_response_content() {
472        // Test that summarize() returns the LLM's response content.
473        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        // The result should be the mock response.
487        assert_eq!(
488            result, expected_summary,
489            "summarize should return the LLM response"
490        );
491
492        // Verify the prompt was sent correctly.
493        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] // Requires real API key: DEEPSEEK_API_KEY
509    async fn test_summarize_with_real_deepseek_api() {
510        // This test calls the real DeepSeek API to verify the summarization works.
511        // It's skipped by default because it requires an API key.
512        // Run with: cargo test test_summarize_with_real_deepseek_api -- --nocapture
513
514        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        // Create a real DeepSeek provider.
524        // TODO: Use llm-unified factory to create provider.
525        // let client: Arc<dyn agent_base::llm_trait::LlmProvider> = ...;
526        return; // TODO: update to new LlmProvider API
527
528        // The rest of this test needs to be updated for the new LlmProvider API.
529        // It previously used agent_base::llm::adapt(OpenAiClient::new(...)).
530        // TODO: Use llm-unified factory to create a real provider for integration testing.
531    }
532}