1use std::sync::Arc;
4
5use async_trait::async_trait;
6
7use ai_agents_core::{ChatMessage, LLMProvider, Result, Role};
8
9use super::native::readable_projection;
10
11#[async_trait]
17pub trait Summarizer: Send + Sync {
18 async fn summarize(&self, messages: &[ChatMessage]) -> Result<String>;
20
21 fn max_batch_size(&self) -> usize {
23 20
24 }
25
26 async fn merge_summaries(&self, summaries: &[String]) -> Result<String> {
28 Ok(summaries.join("\n\n"))
29 }
30}
31
32pub struct LLMSummarizer {
33 llm: Arc<dyn LLMProvider>,
34 prompt_template: String,
35 merge_prompt_template: String,
36 max_batch_size: usize,
37}
38
39impl LLMSummarizer {
40 pub fn new(llm: Arc<dyn LLMProvider>) -> Self {
41 Self {
42 llm,
43 prompt_template: DEFAULT_SUMMARY_PROMPT.to_string(),
44 merge_prompt_template: DEFAULT_MERGE_PROMPT.to_string(),
45 max_batch_size: 20,
46 }
47 }
48
49 pub fn with_prompt(mut self, prompt: impl Into<String>) -> Self {
50 self.prompt_template = prompt.into();
51 self
52 }
53
54 pub fn with_merge_prompt(mut self, prompt: impl Into<String>) -> Self {
55 self.merge_prompt_template = prompt.into();
56 self
57 }
58
59 pub fn with_batch_size(mut self, size: usize) -> Self {
60 self.max_batch_size = size.max(1);
61 self
62 }
63
64 fn format_messages(&self, messages: &[ChatMessage]) -> Result<String> {
66 Ok(readable_projection(messages)?
67 .iter()
68 .map(|m| format!("{}: {}", format_role(&m.role), m.content))
69 .collect::<Vec<_>>()
70 .join("\n"))
71 }
72}
73
74fn format_role(role: &Role) -> &'static str {
75 match role {
76 Role::System => "System",
77 Role::User => "User",
78 Role::Assistant => "Assistant",
79 Role::Tool => "Tool",
80 Role::Function => "Function",
81 }
82}
83
84#[async_trait]
85impl Summarizer for LLMSummarizer {
86 async fn summarize(&self, messages: &[ChatMessage]) -> Result<String> {
87 if messages.is_empty() {
88 return Ok(String::new());
89 }
90
91 let conversation = self.format_messages(messages)?;
92 let prompt = self
93 .prompt_template
94 .replace("{conversation}", &conversation);
95
96 let llm_messages = vec![ChatMessage::user(&prompt)];
97
98 let response = self.llm.complete(&llm_messages, None).await?;
99 Ok(response.content.trim().to_string())
100 }
101
102 fn max_batch_size(&self) -> usize {
103 self.max_batch_size
104 }
105
106 async fn merge_summaries(&self, summaries: &[String]) -> Result<String> {
107 if summaries.is_empty() {
108 return Ok(String::new());
109 }
110
111 if summaries.len() == 1 {
112 return Ok(summaries[0].clone());
113 }
114
115 let combined = summaries.join("\n---\n");
116 let prompt = self.merge_prompt_template.replace("{summaries}", &combined);
117
118 let llm_messages = vec![ChatMessage::user(&prompt)];
119
120 let response = self.llm.complete(&llm_messages, None).await?;
121 Ok(response.content.trim().to_string())
122 }
123}
124
125pub const DEFAULT_SUMMARY_PROMPT: &str = r#"Summarize the following conversation concisely, preserving key information, decisions, and context that would be important for continuing the conversation:
126
127{conversation}
128
129Summary:"#;
130
131pub const DEFAULT_MERGE_PROMPT: &str = r#"Merge the following conversation summaries into a single coherent summary, preserving all important information:
132
133{summaries}
134
135Merged Summary:"#;
136
137pub struct NoopSummarizer;
138
139#[async_trait]
140impl Summarizer for NoopSummarizer {
141 async fn summarize(&self, messages: &[ChatMessage]) -> Result<String> {
142 Ok(readable_projection(messages)?
143 .iter()
144 .map(|m| m.content.clone())
145 .collect::<Vec<_>>()
146 .join(" | "))
147 }
148}
149
150#[cfg(test)]
151mod tests {
152 use super::*;
153 use ai_agents_core::{FinishReason, LLMChunk, LLMConfig, LLMError, LLMFeature, LLMResponse};
154 use parking_lot::Mutex;
155
156 struct MockLLMProvider {
157 responses: Mutex<Vec<String>>,
158 requests: Mutex<Vec<Vec<ChatMessage>>>,
159 }
160
161 impl MockLLMProvider {
162 fn new(responses: Vec<String>) -> Self {
163 Self {
164 responses: Mutex::new(responses),
165 requests: Mutex::new(Vec::new()),
166 }
167 }
168
169 fn requests(&self) -> Vec<Vec<ChatMessage>> {
170 self.requests.lock().clone()
171 }
172 }
173
174 #[async_trait]
175 impl LLMProvider for MockLLMProvider {
176 async fn complete(
177 &self,
178 messages: &[ChatMessage],
179 _config: Option<&LLMConfig>,
180 ) -> std::result::Result<LLMResponse, LLMError> {
181 self.requests.lock().push(messages.to_vec());
182 let response = self
183 .responses
184 .lock()
185 .pop()
186 .unwrap_or_else(|| "Summary of conversation".to_string());
187 Ok(LLMResponse::new(response, FinishReason::Stop))
188 }
189
190 async fn complete_stream(
191 &self,
192 _messages: &[ChatMessage],
193 _config: Option<&LLMConfig>,
194 ) -> std::result::Result<
195 Box<dyn futures::Stream<Item = std::result::Result<LLMChunk, LLMError>> + Unpin + Send>,
196 LLMError,
197 > {
198 Err(LLMError::Other(
199 "Streaming not supported in mock".to_string(),
200 ))
201 }
202
203 fn provider_name(&self) -> &str {
204 "mock"
205 }
206
207 fn supports(&self, _feature: LLMFeature) -> bool {
208 true
209 }
210 }
211
212 fn make_message(role: Role, content: &str) -> ChatMessage {
213 ChatMessage {
214 role,
215 content: content.to_string(),
216 name: None,
217 timestamp: None,
218 }
219 }
220
221 fn signed_assistant_message() -> ChatMessage {
222 use ai_agents_core::{
223 NativeCallBinding, NativeProviderState, NativeProviderTarget, ToolCall,
224 encode_native_tool_call_markers,
225 };
226
227 let call = ToolCall {
228 id: "summary-call".to_string(),
229 name: "lookup".to_string(),
230 arguments: serde_json::json!({"query":"fixture"}),
231 };
232 let state = NativeProviderState::new(
233 "summary-exchange",
234 "google",
235 "generateContent",
236 NativeProviderTarget::new("https://example.invalid/v1beta/", "fixture-model").unwrap(),
237 serde_json::json!({
238 "role":"model",
239 "parts":[{
240 "functionCall":{"name":"lookup","args":{"query":"fixture"}},
241 "thoughtSignature":"fixture-signature"
242 }]
243 }),
244 vec![NativeCallBinding::new(&call.id, 0).unwrap()],
245 )
246 .unwrap();
247 ChatMessage::assistant(
248 encode_native_tool_call_markers(std::slice::from_ref(&call), Some(&state)).unwrap(),
249 )
250 }
251
252 #[tokio::test]
253 async fn test_llm_summarizer_basic() {
254 let provider = Arc::new(MockLLMProvider::new(vec!["Test summary".to_string()]));
255 let summarizer = LLMSummarizer::new(provider);
256
257 let messages = vec![
258 make_message(Role::User, "Hello"),
259 make_message(Role::Assistant, "Hi there!"),
260 ];
261
262 let summary = summarizer.summarize(&messages).await.unwrap();
263 assert_eq!(summary, "Test summary");
264 }
265
266 #[tokio::test]
267 async fn test_llm_summarizer_empty_messages() {
268 let provider = Arc::new(MockLLMProvider::new(vec![]));
269 let summarizer = LLMSummarizer::new(provider);
270
271 let summary = summarizer.summarize(&[]).await.unwrap();
272 assert!(summary.is_empty());
273 }
274
275 #[tokio::test]
276 async fn test_llm_summarizer_custom_prompt() {
277 let provider = Arc::new(MockLLMProvider::new(vec!["Custom summary".to_string()]));
278 let summarizer = LLMSummarizer::new(provider).with_prompt("Custom prompt: {conversation}");
279
280 let messages = vec![make_message(Role::User, "Test")];
281 let summary = summarizer.summarize(&messages).await.unwrap();
282 assert_eq!(summary, "Custom summary");
283 }
284
285 #[tokio::test]
286 async fn llm_summarizer_projects_native_provider_state_before_prompting() {
287 let provider = Arc::new(MockLLMProvider::new(vec!["Projected summary".to_string()]));
288 let summarizer = LLMSummarizer::new(provider.clone());
289
290 let summary = summarizer
291 .summarize(&[signed_assistant_message()])
292 .await
293 .unwrap();
294
295 assert_eq!(summary, "Projected summary");
296 let requests = provider.requests();
297 let prompt = &requests[0][0].content;
298 assert!(prompt.contains("native_tool_calls"));
299 assert!(!prompt.contains("fixture-signature"));
300 assert!(!prompt.contains("_ai_agents_provider_state"));
301 }
302
303 #[tokio::test]
304 async fn noop_summarizer_projects_native_provider_state() {
305 let summary = NoopSummarizer
306 .summarize(&[signed_assistant_message()])
307 .await
308 .unwrap();
309
310 assert!(summary.contains("native_tool_calls"));
311 assert!(!summary.contains("fixture-signature"));
312 assert!(!summary.contains("_ai_agents_provider_state"));
313 }
314
315 #[tokio::test]
316 async fn summarizer_does_not_promote_user_marker_text_to_native_history() {
317 let user_marker_text = signed_assistant_message().content;
318
319 let summary = NoopSummarizer
320 .summarize(&[ChatMessage::user(user_marker_text)])
321 .await
322 .unwrap();
323
324 assert!(summary.contains("fixture-signature"));
325 assert!(summary.contains("_ai_agents_provider_state"));
326 }
327
328 #[tokio::test]
329 async fn test_merge_summaries() {
330 let provider = Arc::new(MockLLMProvider::new(vec!["Merged summary".to_string()]));
331 let summarizer = LLMSummarizer::new(provider);
332
333 let summaries = vec!["Summary 1".to_string(), "Summary 2".to_string()];
334 let merged = summarizer.merge_summaries(&summaries).await.unwrap();
335 assert_eq!(merged, "Merged summary");
336 }
337
338 #[tokio::test]
339 async fn test_merge_single_summary() {
340 let provider = Arc::new(MockLLMProvider::new(vec![]));
341 let summarizer = LLMSummarizer::new(provider);
342
343 let summaries = vec!["Only summary".to_string()];
344 let merged = summarizer.merge_summaries(&summaries).await.unwrap();
345 assert_eq!(merged, "Only summary");
346 }
347
348 #[tokio::test]
349 async fn test_noop_summarizer() {
350 let summarizer = NoopSummarizer;
351
352 let messages = vec![
353 make_message(Role::User, "Hello"),
354 make_message(Role::Assistant, "Hi"),
355 ];
356
357 let summary = summarizer.summarize(&messages).await.unwrap();
358 assert!(summary.contains("Hello"));
359 assert!(summary.contains("Hi"));
360 }
361
362 #[test]
363 fn test_max_batch_size() {
364 let provider = Arc::new(MockLLMProvider::new(vec![]));
365 let summarizer = LLMSummarizer::new(provider).with_batch_size(10);
366 assert_eq!(summarizer.max_batch_size(), 10);
367 }
368
369 #[test]
370 fn test_format_role() {
371 assert_eq!(format_role(&Role::User), "User");
372 assert_eq!(format_role(&Role::Assistant), "Assistant");
373 assert_eq!(format_role(&Role::System), "System");
374 assert_eq!(format_role(&Role::Tool), "Tool");
375 assert_eq!(format_role(&Role::Function), "Function");
376 }
377}