ai-agents-memory 1.1.1

Memory implementations for AI Agents framework
Documentation
//! Summarizer trait and implementations for memory compression

use std::sync::Arc;

use async_trait::async_trait;

use ai_agents_core::{ChatMessage, LLMProvider, Result, Role};

use super::native::readable_projection;

/// Summarizes conversation messages for memory compression.
///
/// Built-in implementations: `LLMSummarizer` (uses an LLM to generate summaries)
/// and `NoopSummarizer` (concatenates messages, for testing).
/// Most users use `LLMSummarizer`, auto-configured from the YAML `summarizer_llm` field.
#[async_trait]
pub trait Summarizer: Send + Sync {
    /// Produce a summary from a batch of messages.
    async fn summarize(&self, messages: &[ChatMessage]) -> Result<String>;

    /// Maximum messages per summarization call. Returns 20 by default.
    fn max_batch_size(&self) -> usize {
        20
    }

    /// Combine multiple summaries into one. Joins with `\n\n` by default.
    async fn merge_summaries(&self, summaries: &[String]) -> Result<String> {
        Ok(summaries.join("\n\n"))
    }
}

pub struct LLMSummarizer {
    llm: Arc<dyn LLMProvider>,
    merge_llm: Arc<dyn LLMProvider>,
    prompt_template: String,
    merge_prompt_template: String,
    max_batch_size: usize,
}

impl LLMSummarizer {
    pub fn new(llm: Arc<dyn LLMProvider>) -> Self {
        Self {
            merge_llm: llm.clone(),
            llm,
            prompt_template: DEFAULT_SUMMARY_PROMPT.to_string(),
            merge_prompt_template: DEFAULT_MERGE_PROMPT.to_string(),
            max_batch_size: 20,
        }
    }

    /// Uses a separate provider for merging while keeping the one-provider constructor compatible.
    pub fn with_merge_llm(mut self, llm: Arc<dyn LLMProvider>) -> Self {
        self.merge_llm = llm;
        self
    }

    pub fn with_prompt(mut self, prompt: impl Into<String>) -> Self {
        self.prompt_template = prompt.into();
        self
    }

    pub fn with_merge_prompt(mut self, prompt: impl Into<String>) -> Self {
        self.merge_prompt_template = prompt.into();
        self
    }

    pub fn with_batch_size(mut self, size: usize) -> Self {
        self.max_batch_size = size.max(1);
        self
    }

    // Projects replay-bearing markers before building the auxiliary-model prompt.
    fn format_messages(&self, messages: &[ChatMessage]) -> Result<String> {
        Ok(readable_projection(messages)?
            .iter()
            .map(|m| format!("{}: {}", format_role(&m.role), m.content))
            .collect::<Vec<_>>()
            .join("\n"))
    }
}

fn format_role(role: &Role) -> &'static str {
    match role {
        Role::System => "System",
        Role::User => "User",
        Role::Assistant => "Assistant",
        Role::Tool => "Tool",
        Role::Function => "Function",
    }
}

#[async_trait]
impl Summarizer for LLMSummarizer {
    async fn summarize(&self, messages: &[ChatMessage]) -> Result<String> {
        if messages.is_empty() {
            return Ok(String::new());
        }

        let conversation = self.format_messages(messages)?;
        let prompt = self
            .prompt_template
            .replace("{conversation}", &conversation);

        let llm_messages = vec![ChatMessage::user(&prompt)];

        let response = self.llm.complete(&llm_messages, None).await?;
        Ok(response.content.trim().to_string())
    }

    fn max_batch_size(&self) -> usize {
        self.max_batch_size
    }

    async fn merge_summaries(&self, summaries: &[String]) -> Result<String> {
        if summaries.is_empty() {
            return Ok(String::new());
        }

        if summaries.len() == 1 {
            return Ok(summaries[0].clone());
        }

        let combined = summaries.join("\n---\n");
        let prompt = self.merge_prompt_template.replace("{summaries}", &combined);

        let llm_messages = vec![ChatMessage::user(&prompt)];

        let response = self.merge_llm.complete(&llm_messages, None).await?;
        Ok(response.content.trim().to_string())
    }
}

pub 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:

{conversation}

Summary:"#;

pub const DEFAULT_MERGE_PROMPT: &str = r#"Merge the following conversation summaries into a single coherent summary, preserving all important information:

{summaries}

Merged Summary:"#;

pub struct NoopSummarizer;

#[async_trait]
impl Summarizer for NoopSummarizer {
    async fn summarize(&self, messages: &[ChatMessage]) -> Result<String> {
        Ok(readable_projection(messages)?
            .iter()
            .map(|m| m.content.clone())
            .collect::<Vec<_>>()
            .join(" | "))
    }
}

#[cfg(test)]
mod tests {
    use super::*;
    use ai_agents_core::{FinishReason, LLMChunk, LLMConfig, LLMError, LLMFeature, LLMResponse};
    use parking_lot::Mutex;

    struct MockLLMProvider {
        responses: Mutex<Vec<String>>,
        requests: Mutex<Vec<Vec<ChatMessage>>>,
    }

    impl MockLLMProvider {
        fn new(responses: Vec<String>) -> Self {
            Self {
                responses: Mutex::new(responses),
                requests: Mutex::new(Vec::new()),
            }
        }

        fn requests(&self) -> Vec<Vec<ChatMessage>> {
            self.requests.lock().clone()
        }
    }

    #[async_trait]
    impl LLMProvider for MockLLMProvider {
        async fn complete(
            &self,
            messages: &[ChatMessage],
            _config: Option<&LLMConfig>,
        ) -> std::result::Result<LLMResponse, LLMError> {
            self.requests.lock().push(messages.to_vec());
            let response = self
                .responses
                .lock()
                .pop()
                .unwrap_or_else(|| "Summary of conversation".to_string());
            Ok(LLMResponse::new(response, FinishReason::Stop))
        }

        async fn complete_stream(
            &self,
            _messages: &[ChatMessage],
            _config: Option<&LLMConfig>,
        ) -> std::result::Result<
            Box<dyn futures::Stream<Item = std::result::Result<LLMChunk, LLMError>> + Unpin + Send>,
            LLMError,
        > {
            Err(LLMError::Other(
                "Streaming not supported in mock".to_string(),
            ))
        }

        fn provider_name(&self) -> &str {
            "mock"
        }

        fn supports(&self, _feature: LLMFeature) -> bool {
            true
        }
    }

    fn make_message(role: Role, content: &str) -> ChatMessage {
        ChatMessage {
            role,
            content: content.to_string(),
            name: None,
            timestamp: None,
        }
    }

    fn signed_assistant_message() -> ChatMessage {
        use ai_agents_core::{
            NativeCallBinding, NativeProviderState, NativeProviderTarget, ToolCall,
            encode_native_tool_call_markers,
        };

        let call = ToolCall {
            id: "summary-call".to_string(),
            name: "lookup".to_string(),
            arguments: serde_json::json!({"query":"fixture"}),
        };
        let state = NativeProviderState::new(
            "summary-exchange",
            "google",
            "generateContent",
            NativeProviderTarget::new("https://example.invalid/v1beta/", "fixture-model").unwrap(),
            serde_json::json!({
                "role":"model",
                "parts":[{
                    "functionCall":{"name":"lookup","args":{"query":"fixture"}},
                    "thoughtSignature":"fixture-signature"
                }]
            }),
            vec![NativeCallBinding::new(&call.id, 0).unwrap()],
        )
        .unwrap();
        ChatMessage::assistant(
            encode_native_tool_call_markers(std::slice::from_ref(&call), Some(&state)).unwrap(),
        )
    }

    #[tokio::test]
    async fn test_llm_summarizer_basic() {
        let provider = Arc::new(MockLLMProvider::new(vec!["Test summary".to_string()]));
        let summarizer = LLMSummarizer::new(provider);

        let messages = vec![
            make_message(Role::User, "Hello"),
            make_message(Role::Assistant, "Hi there!"),
        ];

        let summary = summarizer.summarize(&messages).await.unwrap();
        assert_eq!(summary, "Test summary");
    }

    #[tokio::test]
    async fn test_llm_summarizer_empty_messages() {
        let provider = Arc::new(MockLLMProvider::new(vec![]));
        let summarizer = LLMSummarizer::new(provider);

        let summary = summarizer.summarize(&[]).await.unwrap();
        assert!(summary.is_empty());
    }

    #[tokio::test]
    async fn test_llm_summarizer_custom_prompt() {
        let provider = Arc::new(MockLLMProvider::new(vec!["Custom summary".to_string()]));
        let summarizer = LLMSummarizer::new(provider).with_prompt("Custom prompt: {conversation}");

        let messages = vec![make_message(Role::User, "Test")];
        let summary = summarizer.summarize(&messages).await.unwrap();
        assert_eq!(summary, "Custom summary");
    }

    #[tokio::test]
    async fn llm_summarizer_projects_native_provider_state_before_prompting() {
        let provider = Arc::new(MockLLMProvider::new(vec!["Projected summary".to_string()]));
        let summarizer = LLMSummarizer::new(provider.clone());

        let summary = summarizer
            .summarize(&[signed_assistant_message()])
            .await
            .unwrap();

        assert_eq!(summary, "Projected summary");
        let requests = provider.requests();
        let prompt = &requests[0][0].content;
        assert!(prompt.contains("native_tool_calls"));
        assert!(!prompt.contains("fixture-signature"));
        assert!(!prompt.contains("_ai_agents_provider_state"));
    }

    #[tokio::test]
    async fn noop_summarizer_projects_native_provider_state() {
        let summary = NoopSummarizer
            .summarize(&[signed_assistant_message()])
            .await
            .unwrap();

        assert!(summary.contains("native_tool_calls"));
        assert!(!summary.contains("fixture-signature"));
        assert!(!summary.contains("_ai_agents_provider_state"));
    }

    #[tokio::test]
    async fn summarizer_does_not_promote_user_marker_text_to_native_history() {
        let user_marker_text = signed_assistant_message().content;

        let summary = NoopSummarizer
            .summarize(&[ChatMessage::user(user_marker_text)])
            .await
            .unwrap();

        assert!(summary.contains("fixture-signature"));
        assert!(summary.contains("_ai_agents_provider_state"));
    }

    #[tokio::test]
    async fn test_merge_summaries() {
        let provider = Arc::new(MockLLMProvider::new(vec!["Merged summary".to_string()]));
        let summarizer = LLMSummarizer::new(provider);

        let summaries = vec!["Summary 1".to_string(), "Summary 2".to_string()];
        let merged = summarizer.merge_summaries(&summaries).await.unwrap();
        assert_eq!(merged, "Merged summary");
    }

    #[tokio::test]
    async fn test_merge_single_summary() {
        let provider = Arc::new(MockLLMProvider::new(vec![]));
        let summarizer = LLMSummarizer::new(provider);

        let summaries = vec!["Only summary".to_string()];
        let merged = summarizer.merge_summaries(&summaries).await.unwrap();
        assert_eq!(merged, "Only summary");
    }

    #[tokio::test]
    async fn test_noop_summarizer() {
        let summarizer = NoopSummarizer;

        let messages = vec![
            make_message(Role::User, "Hello"),
            make_message(Role::Assistant, "Hi"),
        ];

        let summary = summarizer.summarize(&messages).await.unwrap();
        assert!(summary.contains("Hello"));
        assert!(summary.contains("Hi"));
    }

    #[test]
    fn test_max_batch_size() {
        let provider = Arc::new(MockLLMProvider::new(vec![]));
        let summarizer = LLMSummarizer::new(provider).with_batch_size(10);
        assert_eq!(summarizer.max_batch_size(), 10);
    }

    #[test]
    fn test_format_role() {
        assert_eq!(format_role(&Role::User), "User");
        assert_eq!(format_role(&Role::Assistant), "Assistant");
        assert_eq!(format_role(&Role::System), "System");
        assert_eq!(format_role(&Role::Tool), "Tool");
        assert_eq!(format_role(&Role::Function), "Function");
    }
}