agent-infra-sdk 0.2.0

Unified Rust SDK for Gateway-backed and local Agent Infra APIs
use agent_context_contract::ConversationRecord;
use serde_json::Value;

/// Configuration for context compression.
#[derive(Debug, Clone)]
pub struct ContextCompressionConfig {
    /// Maximum number of messages to keep in the compressed context.
    pub max_messages: usize,
    /// Maximum total characters across all messages in the compressed context.
    pub max_chars: usize,
    /// Whether to summarize older messages.
    pub summarize_older: bool,
}

impl Default for ContextCompressionConfig {
    fn default() -> Self {
        Self {
            max_messages: 100,
            max_chars: 100_000,
            summarize_older: true,
        }
    }
}

/// Compress a context window (list of messages) to fit within the configured limits.
///
/// Strategy:
/// 1. If messages fit within limits, return all.
/// 2. Keep the most recent messages, dropping oldest first.
/// 3. If summarize_older is enabled, prepend a summary message.
pub fn compress_context(
    messages: &[ConversationRecord],
    config: &ContextCompressionConfig,
) -> Vec<ConversationRecord> {
    if messages.len() <= config.max_messages {
        let total_chars: usize = messages.iter().map(|m| estimate_chars(&m.content)).sum();
        if total_chars <= config.max_chars {
            return messages.to_vec();
        }
    }

    // Keep the most recent messages
    let keep_count = config.max_messages.min(messages.len());
    let mut kept: Vec<ConversationRecord> = messages[messages.len() - keep_count..].to_vec();

    // Trim by character count from the front
    let mut total_chars: usize = kept.iter().map(|m| estimate_chars(&m.content)).sum();
    while total_chars > config.max_chars && kept.len() > 1 {
        let removed = kept.remove(0);
        total_chars = total_chars.saturating_sub(estimate_chars(&removed.content));
    }

    // Prepend a summary if we dropped messages
    if config.summarize_older && kept.len() < messages.len() {
        let dropped_count = messages.len() - kept.len();
        let summary = ConversationRecord {
            id: format!("__compressed_{}", dropped_count),
            scope: messages
                .first()
                .map(|m| m.scope.clone())
                .unwrap_or_default(),
            role: "system".to_string(),
            content: Value::String(format!(
                "[Context compressed: {} earlier messages were summarized to fit the context window]",
                dropped_count
            )),
            name: None,
            tool_call_id: None,
            sequence: 0,
            created_at_ms: 0,
            metadata: Value::Null,
        };
        kept.insert(0, summary);
    }

    kept
}

fn estimate_chars(content: &Value) -> usize {
    match content {
        Value::String(s) => s.len(),
        Value::Array(parts) => parts
            .iter()
            .map(|p| {
                p.get("text")
                    .and_then(Value::as_str)
                    .map(|s| s.len())
                    .unwrap_or(0)
            })
            .sum(),
        _ => content.to_string().len(),
    }
}

#[cfg(test)]
mod tests {
    use super::*;
    use agent_context_contract::ConversationScope;

    fn make_msg(id: &str, text: &str) -> ConversationRecord {
        ConversationRecord {
            id: id.to_string(),
            scope: ConversationScope::conversation("conv-1"),
            role: "user".to_string(),
            content: Value::String(text.to_string()),
            name: None,
            tool_call_id: None,
            sequence: 0,
            created_at_ms: 0,
            metadata: Value::Null,
        }
    }

    #[test]
    fn compress_under_limits_returns_all() {
        let msgs: Vec<_> = (0..5)
            .map(|i| make_msg(&format!("m{i}"), "short"))
            .collect();
        let result = compress_context(&msgs, &ContextCompressionConfig::default());
        assert_eq!(result.len(), 5);
    }

    #[test]
    fn compress_over_message_limit_drops_oldest() {
        let msgs: Vec<_> = (0..200)
            .map(|i| make_msg(&format!("m{i}"), "msg"))
            .collect();
        let config = ContextCompressionConfig {
            max_messages: 50,
            max_chars: 1_000_000,
            summarize_older: false,
        };
        let result = compress_context(&msgs, &config);
        assert!(result.len() <= 50);
    }

    #[test]
    fn compress_adds_summary_when_dropping() {
        let msgs: Vec<_> = (0..200)
            .map(|i| make_msg(&format!("m{i}"), "msg"))
            .collect();
        let config = ContextCompressionConfig {
            max_messages: 50,
            max_chars: 1_000_000,
            summarize_older: true,
        };
        let result = compress_context(&msgs, &config);
        assert_eq!(result[0].role, "system");
        assert!(result[0].content.as_str().unwrap().contains("compressed"));
    }
}