1use async_trait::async_trait;
9
10use crate::agent::compaction::estimate_message_tokens;
11use crate::agent::prompt::PromptBlock;
12use crate::error::AgentError;
13use crate::model::message::{ContentBlock, Message, Role, ToolCall};
14use crate::model::provider::{CompletionRequest, ModelHint, ModelProvider};
15
16const HISTORY_SUMMARY_PROMPT: &str = include_str!("history_summary_prompt.md");
17
18fn content_str_len(value: &serde_json::Value) -> usize {
19 value.as_str().map(str::len).unwrap_or_else(|| value.to_string().len())
20}
21
22fn collapse_tool_pairs(messages: Vec<Message>) -> Vec<Message> {
26 let mut result: Vec<Message> = Vec::new();
27 let len = messages.len();
28 let mut i = 0;
29
30 while i < len {
31 let msg = &messages[i];
32
33 match msg.role {
34 Role::Assistant => {
35 let tool_calls: Vec<&ToolCall> = msg.tool_calls().collect();
36 if tool_calls.is_empty() {
37 result.push(msg.clone());
38 i += 1;
39 continue;
40 }
41
42 let assistant_text: String = msg
43 .blocks
44 .iter()
45 .filter_map(|b| match b {
46 ContentBlock::Text(t) => Some(t.text.as_str()),
47 _ => None,
48 })
49 .collect::<Vec<_>>()
50 .join("\n");
51
52 let mut collapsed = Vec::new();
53 if !assistant_text.is_empty() {
54 collapsed.push(assistant_text.to_string());
55 }
56
57 let mut j = i + 1;
58 while j < len && messages[j].role == Role::Tool {
59 for tr in messages[j].tool_results_iter() {
60 let content_len = content_str_len(&tr.content);
61 if tr.is_error {
62 collapsed.push(format!(
63 "Used tool `{}` — ERROR ({} chars)",
64 tr.name, content_len
65 ));
66 } else {
67 collapsed
68 .push(format!("Used tool `{}` — {} chars", tr.name, content_len));
69 }
70 }
71 j += 1;
72 }
73
74 let has_content = collapsed.iter().any(|s| !s.is_empty());
75 if !has_content {
76 result.push(msg.clone());
77 } else {
78 result.push(Message::user(collapsed.join("\n")));
79 }
80 i = j;
81 }
82 Role::Tool => {
83 let parts: Vec<String> = msg
84 .tool_results_iter()
85 .map(|tr| {
86 let content_len = content_str_len(&tr.content);
87 if tr.is_error {
88 format!("Used tool `{}` — ERROR ({} chars)", tr.name, content_len)
89 } else {
90 format!("Used tool `{}` — {} chars", tr.name, content_len)
91 }
92 })
93 .collect();
94 if !parts.is_empty() {
95 result.push(Message::user(parts.join("\n")));
96 }
97 i += 1;
98 }
99 _ => {
100 result.push(msg.clone());
101 i += 1;
102 }
103 }
104 }
105
106 result
107}
108
109#[async_trait]
111pub trait HistoryCompactionStrategy: Send + Sync + std::fmt::Debug {
112 async fn compact(
116 &self,
117 messages: &mut Vec<Message>,
118 provider: &dyn ModelProvider,
119 ) -> Result<bool, AgentError>;
120}
121
122#[derive(Debug)]
124pub struct SummaryHistoryCompaction {
125 pub keep_recent: usize,
127 pub max_summary_input_tokens: usize,
129 pub summary_output_tokens: usize,
131}
132
133#[async_trait]
134impl HistoryCompactionStrategy for SummaryHistoryCompaction {
135 async fn compact(
136 &self,
137 messages: &mut Vec<Message>,
138 provider: &dyn ModelProvider,
139 ) -> Result<bool, AgentError> {
140 let system_end = messages.iter().take_while(|m| m.role == Role::System).count();
141
142 let split_point = messages.len().saturating_sub(self.keep_recent);
143 let split_point = split_point.max(system_end);
144
145 if split_point <= system_end {
146 return Ok(false);
147 }
148
149 let old_messages = &messages[system_end..split_point];
150 let summary_text = self.summarize_pass(old_messages.to_vec(), provider).await?;
151
152 messages.splice(system_end..split_point, [Message::user(summary_text)]);
153
154 Ok(true)
155 }
156}
157
158impl SummaryHistoryCompaction {
159 async fn summarize_pass(
166 &self,
167 messages: Vec<Message>,
168 provider: &dyn ModelProvider,
169 ) -> Result<String, AgentError> {
170 let messages = collapse_tool_pairs(messages);
171 let prompt_tokens = provider.estimate_tokens(HISTORY_SUMMARY_PROMPT.trim());
172 let budget = self
173 .max_summary_input_tokens
174 .saturating_sub(prompt_tokens)
175 .saturating_sub(self.summary_output_tokens)
176 .max(1);
177
178 let mut selected = Vec::new();
179 let mut tokens_used = 0usize;
180
181 for msg in messages.into_iter().rev() {
182 let msg_tokens = estimate_message_tokens(std::slice::from_ref(&msg), provider);
183 if tokens_used + msg_tokens > budget && !selected.is_empty() {
184 break;
185 }
186 tokens_used += msg_tokens;
187 selected.push(msg);
188 }
189 selected.reverse();
190
191 if selected.is_empty() {
192 return Ok(String::new());
193 }
194
195 self.complete_summary(selected, provider).await
196 }
197
198 async fn complete_summary(
199 &self,
200 messages: Vec<Message>,
201 provider: &dyn ModelProvider,
202 ) -> Result<String, AgentError> {
203 let summary_request = CompletionRequest {
204 system_prompt_blocks: vec![PromptBlock::dynamic(
205 "history_summary",
206 HISTORY_SUMMARY_PROMPT.trim(),
207 )],
208 messages,
209 tools: vec![],
210 model_hint: Some(ModelHint::Summarization),
211 max_tokens: Some(self.summary_output_tokens.min(u32::MAX as usize) as u32),
212 };
213
214 let response = provider.complete(summary_request).await?;
215 Ok(response.message.text_content())
216 }
217}
218
219#[cfg(test)]
220mod tests {
221 use super::*;
222 use crate::model::message::{TextBlock, ToolCall, ToolResult};
223 use crate::model::mock::MockProvider;
224 use crate::model::provider::{
225 CompletionResponse, DeepSeekConfig, DeepSeekProvider, StopReason,
226 };
227
228 fn make_tool_call(id: &str, name: &str, args: serde_json::Value) -> ContentBlock {
229 ContentBlock::ToolCall(ToolCall { id: id.into(), name: name.into(), arguments: args })
230 }
231
232 fn make_tool_result(id: &str, name: &str, content: &str, is_error: bool) -> ContentBlock {
233 ContentBlock::ToolResult(ToolResult {
234 tool_call_id: id.into(),
235 name: name.into(),
236 content: serde_json::Value::String(content.into()),
237 is_error,
238 })
239 }
240
241 #[test]
242 fn collapse_tool_pairs_merges_assistant_tool_calls_with_results() {
243 let messages = vec![
244 Message {
245 role: Role::Assistant,
246 blocks: vec![
247 ContentBlock::Text(TextBlock { text: "Let me check the file.".into() }),
248 make_tool_call("call_1", "read", serde_json::json!({"path": "src/main.rs"})),
249 ],
250 },
251 Message {
252 role: Role::Tool,
253 blocks: vec![make_tool_result("call_1", "read", "file content here", false)],
254 },
255 ];
256
257 let collapsed = collapse_tool_pairs(messages);
258 assert_eq!(collapsed.len(), 1);
259 assert_eq!(collapsed[0].role, Role::User);
260 let text = collapsed[0].text_content();
261 assert!(text.contains("Let me check the file."));
262 assert!(text.contains("Used tool `read`"));
263 assert!(text.contains("17 chars"));
264 }
265
266 #[test]
267 fn collapse_tool_pairs_preserves_non_tool_messages() {
268 let messages = vec![
269 Message::user("Hello"),
270 Message::assistant("Hi there!"),
271 Message::user("Do something"),
272 ];
273
274 let collapsed = collapse_tool_pairs(messages);
275 assert_eq!(collapsed.len(), 3);
276 assert_eq!(collapsed[0].text_content(), "Hello");
277 assert_eq!(collapsed[1].text_content(), "Hi there!");
278 assert_eq!(collapsed[2].text_content(), "Do something");
279 }
280
281 #[test]
282 fn collapse_tool_pairs_handles_multiple_tool_calls_in_one_turn() {
283 let messages = vec![
284 Message {
285 role: Role::Assistant,
286 blocks: vec![
287 make_tool_call("call_1", "read", serde_json::json!({"path": "a.rs"})),
288 make_tool_call("call_2", "grep", serde_json::json!({"pattern": "TODO"})),
289 ],
290 },
291 Message {
292 role: Role::Tool,
293 blocks: vec![make_tool_result("call_1", "read", "aaa", false)],
294 },
295 Message {
296 role: Role::Tool,
297 blocks: vec![make_tool_result("call_2", "grep", "line 42: TODO", false)],
298 },
299 ];
300
301 let collapsed = collapse_tool_pairs(messages);
302 assert_eq!(collapsed.len(), 1);
303 let text = collapsed[0].text_content();
304 assert!(text.contains("Used tool `read`"));
305 assert!(text.contains("Used tool `grep`"));
306 }
307
308 #[test]
309 fn collapse_tool_pairs_handles_tool_error() {
310 let messages = vec![
311 Message {
312 role: Role::Assistant,
313 blocks: vec![
314 ContentBlock::Text(TextBlock { text: "Running test.".into() }),
315 make_tool_call("call_1", "bash", serde_json::json!({"command": "cargo test"})),
316 ],
317 },
318 Message {
319 role: Role::Tool,
320 blocks: vec![make_tool_result("call_1", "bash", "compilation failed", true)],
321 },
322 ];
323
324 let collapsed = collapse_tool_pairs(messages);
325 assert_eq!(collapsed.len(), 1);
326 let text = collapsed[0].text_content();
327 assert!(text.contains("Running test."));
328 assert!(text.contains("ERROR"));
329 assert!(text.contains("Used tool `bash`"));
330 }
331
332 #[test]
333 fn collapse_tool_pairs_does_not_collapse_lone_assistant_without_tool_results() {
334 let messages = vec![
335 Message {
336 role: Role::Assistant,
337 blocks: vec![make_tool_call("call_1", "read", serde_json::json!({"path": "x"}))],
338 },
339 Message::user("next message"),
340 ];
341
342 let collapsed = collapse_tool_pairs(messages);
343 assert_eq!(collapsed.len(), 2);
344 assert_eq!(collapsed[0].role, Role::Assistant);
345 assert_eq!(collapsed[1].role, Role::User);
346 }
347
348 const TEST_SUMMARY_FIXTURE: &str = include_str!(concat!(
349 env!("CARGO_MANIFEST_DIR"),
350 "/tests/fixtures/compaction/test_summary.txt"
351 ));
352 const TEST_SUMMARY_OUTPUT: &str = "test_summary_output.txt";
353
354 struct FakeProvider;
355
356 #[async_trait::async_trait]
357 impl ModelProvider for FakeProvider {
358 async fn complete(
359 &self,
360 _request: CompletionRequest,
361 ) -> Result<CompletionResponse, AgentError> {
362 Ok(CompletionResponse {
363 message: Message::assistant("summary text"),
364 stop_reason: StopReason::EndTurn,
365 usage: None,
366 model: None,
367 })
368 }
369 }
370
371 #[tokio::test]
372 async fn preserves_leading_system_prompt_without_summarizing_it() {
373 let compaction = SummaryHistoryCompaction {
374 keep_recent: 1,
375 max_summary_input_tokens: 120_000,
376 summary_output_tokens: 4_000,
377 };
378 let mut messages =
379 vec![Message::system("persona"), Message::user("first"), Message::user("second")];
380 let provider = MockProvider::new(vec![CompletionResponse {
381 message: Message::assistant("summary text"),
382 stop_reason: StopReason::EndTurn,
383 usage: None,
384 model: None,
385 }]);
386
387 let changed = compaction.compact(&mut messages, &provider).await.unwrap();
388
389 assert!(changed);
390 assert_eq!(messages[0].role, Role::System);
391 assert_eq!(messages[0].text_content(), "persona");
392 assert_eq!(messages[1].role, Role::User);
393 assert!(messages[1].text_content().contains("summary text"));
394
395 let requests = provider.requests.lock().await;
396 assert_eq!(requests.len(), 1);
397 assert_eq!(requests[0].messages, vec![Message::user("first")]);
398 }
399
400 fn real_deepseek_provider() -> Option<DeepSeekProvider> {
401 let _ = dotenvy::from_filename("src/provider/.env");
402 dotenvy::dotenv().ok();
403
404 let api_key = std::env::var("DEEPSEEK_TEST_KEY")
405 .or_else(|_| std::env::var("DEEPSEEK_API_KEY"))
406 .ok()?;
407 if api_key.is_empty() || api_key == "your_deepseek_api_key_here" {
408 return None;
409 }
410
411 Some(DeepSeekProvider::new(DeepSeekConfig {
412 api_key,
413 model: std::env::var("DEEPSEEK_TEST_MODEL").unwrap_or_else(|_| "deepseek-chat".into()),
414 base_url: std::env::var("DEEPSEEK_BASE_URL")
415 .unwrap_or_else(|_| "https://api.deepseek.com".into()),
416 }))
417 }
418
419 #[tokio::test]
420 async fn summarizes_test_summary_fixture_with_real_model() {
421 let provider = match real_deepseek_provider() {
422 Some(provider) => provider,
423 None => {
424 eprintln!("SKIP: DEEPSEEK_TEST_KEY or DEEPSEEK_API_KEY not set");
425 return;
426 }
427 };
428 assert!(!TEST_SUMMARY_FIXTURE.trim().is_empty());
429
430 let compaction = SummaryHistoryCompaction {
431 keep_recent: 1,
432 max_summary_input_tokens: 800_000,
433 summary_output_tokens: 4_000,
434 };
435 let recent_message = Message::assistant("Recent turn remains available verbatim.");
436 let mut messages = vec![
437 Message::system("persona"),
438 Message::user(TEST_SUMMARY_FIXTURE),
439 recent_message.clone(),
440 ];
441
442 let changed = compaction.compact(&mut messages, &provider).await.unwrap();
443
444 assert!(changed);
445 assert_eq!(messages[0], Message::system("persona"));
446 let summary_text = messages[1].text_content();
447 assert!(!summary_text.trim().is_empty());
448 assert_ne!(summary_text.trim(), TEST_SUMMARY_FIXTURE.trim());
449 let output_path = std::path::Path::new(env!("CARGO_MANIFEST_DIR"))
450 .join("src/compaction")
451 .join(TEST_SUMMARY_OUTPUT);
452 std::fs::write(&output_path, summary_text.trim()).unwrap_or_else(|err| {
453 panic!("failed to write summary output to {}: {err}", output_path.display())
454 });
455 assert_eq!(messages[2], recent_message);
456 }
457
458 #[tokio::test]
459 async fn skips_compaction_when_only_system_prompt_is_old() {
460 let compaction = SummaryHistoryCompaction {
461 keep_recent: 1,
462 max_summary_input_tokens: 120_000,
463 summary_output_tokens: 4_000,
464 };
465 let mut messages = vec![Message::system("long system prompt text"), Message::user("hi")];
466 let changed = compaction.compact(&mut messages, &FakeProvider).await.unwrap();
467 assert!(!changed);
468 }
469
470 #[tokio::test]
471 async fn summarizes_old_messages_in_one_pass() {
472 let compaction = SummaryHistoryCompaction {
473 keep_recent: 1,
474 max_summary_input_tokens: 120_000,
475 summary_output_tokens: 4_000,
476 };
477 let mut messages = vec![
478 Message::system("persona"),
479 Message::user("early chat"),
480 Message::user("some history"),
481 Message::user("recent"),
482 ];
483 let provider = MockProvider::new(vec![CompletionResponse {
484 message: Message::assistant("condensed summary"),
485 stop_reason: StopReason::EndTurn,
486 usage: None,
487 model: None,
488 }]);
489
490 let changed = compaction.compact(&mut messages, &provider).await.unwrap();
491
492 assert!(changed);
493 assert!(messages[1].text_content().contains("condensed summary"));
494
495 let requests = provider.requests.lock().await;
496 assert_eq!(requests.len(), 1);
497 assert_eq!(
498 requests[0].messages,
499 vec![Message::user("early chat"), Message::user("some history")]
500 );
501 }
502
503 #[tokio::test]
504 async fn drops_oldest_messages_when_exceeding_summary_input_budget() {
505 let prompt_tokens = FakeProvider.estimate_tokens(HISTORY_SUMMARY_PROMPT.trim());
506 let compaction = SummaryHistoryCompaction {
507 keep_recent: 1,
508 max_summary_input_tokens: prompt_tokens + 1 + 4_000,
509 summary_output_tokens: 4_000,
510 };
511 let mut messages = vec![
512 Message::system("persona"),
513 Message::user("very old message that is quite long and will be dropped"),
514 Message::user("more recent old message to keep"),
515 Message::user("recent"),
516 ];
517 let provider = MockProvider::new(vec![CompletionResponse {
518 message: Message::assistant("partial summary"),
519 stop_reason: StopReason::EndTurn,
520 usage: None,
521 model: None,
522 }]);
523
524 let changed = compaction.compact(&mut messages, &provider).await.unwrap();
525 assert!(changed);
526
527 let requests = provider.requests.lock().await;
528 assert_eq!(requests.len(), 1);
529 let summarized: Vec<String> =
530 requests[0].messages.iter().map(|m| m.text_content()).collect();
531 assert!(
532 summarized.iter().any(|s| s.contains("more recent old message")),
533 "should keep the more recent old message"
534 );
535 assert!(
536 !summarized
537 .iter()
538 .any(|s| s == "very old message that is quite long and will be dropped"),
539 "should drop the oldest message that doesn't fit budget"
540 );
541 }
542}