use std::sync::Arc;
use async_trait::async_trait;
use parking_lot::RwLock;
use ai_agents_core::{ChatMessage, MemorySnapshot, Result};
use super::Memory;
use super::native::NativeRetentionInspection;
pub struct InMemoryStore {
messages: Arc<RwLock<Vec<ChatMessage>>>,
max_messages: usize,
}
impl InMemoryStore {
pub fn new(max_messages: usize) -> Self {
Self {
messages: Arc::new(RwLock::new(Vec::new())),
max_messages,
}
}
pub fn max_messages(&self) -> usize {
self.max_messages
}
fn bounded_eviction_count(messages: &[ChatMessage], max_messages: usize) -> Result<usize> {
let inspection = NativeRetentionInspection::inspect(messages)?;
let required = messages.len().saturating_sub(max_messages);
if required == 0 {
return Ok(0);
}
inspection
.safe_prefix_len_between(required, messages.len())
.ok_or_else(|| {
ai_agents_core::AgentError::MemoryError(
"message limit cannot preserve the protected signed native exchange"
.to_string(),
)
})
}
}
impl Clone for InMemoryStore {
fn clone(&self) -> Self {
Self {
messages: Arc::clone(&self.messages),
max_messages: self.max_messages,
}
}
}
#[async_trait]
impl ai_agents_core::Memory for InMemoryStore {
async fn add_message(&self, message: ChatMessage) -> Result<()> {
let mut messages = self.messages.write();
let mut prospective = messages.clone();
prospective.push(message);
let evict_count = Self::bounded_eviction_count(&prospective, self.max_messages)?;
prospective.drain(..evict_count);
*messages = prospective;
Ok(())
}
async fn get_messages(&self, limit: Option<usize>) -> Result<Vec<ChatMessage>> {
let messages = self.messages.read();
match limit {
Some(n) => {
let start = messages.len().saturating_sub(n);
Ok(messages[start..].to_vec())
}
None => Ok(messages.clone()),
}
}
async fn clear(&self) -> Result<()> {
self.messages.write().clear();
Ok(())
}
fn len(&self) -> usize {
self.messages.read().len()
}
async fn restore(&self, snapshot: MemorySnapshot) -> Result<()> {
let mut prospective = snapshot.messages;
let evict_count = Self::bounded_eviction_count(&prospective, self.max_messages)?;
prospective.drain(..evict_count);
let mut messages = self.messages.write();
*messages = prospective;
Ok(())
}
async fn evict_oldest(&self, count: usize) -> Result<Vec<ChatMessage>> {
let mut messages = self.messages.write();
let requested = count.min(messages.len());
let inspection = NativeRetentionInspection::inspect(&messages)?;
let evict_count = if requested == 0 {
0
} else {
inspection
.safe_prefix_len_between(requested, messages.len())
.ok_or_else(|| {
ai_agents_core::AgentError::MemoryError(
"eviction would split the protected signed native exchange".to_string(),
)
})?
};
let evicted: Vec<ChatMessage> = messages.drain(..evict_count).collect();
Ok(evicted)
}
}
#[async_trait]
impl Memory for InMemoryStore {}
#[cfg(test)]
mod tests {
use super::*;
use ai_agents_core::{
Memory as CoreMemory, NativeCallBinding, NativeProviderState, NativeProviderTarget, Role,
ToolCall, encode_native_tool_call_markers, encode_native_tool_result_marker,
};
fn make_message(content: &str) -> ChatMessage {
ChatMessage {
role: Role::User,
content: content.to_string(),
name: None,
timestamp: None,
}
}
fn signed_exchange(exchange_id: &str) -> (ChatMessage, ChatMessage) {
let call = ToolCall {
id: format!("{exchange_id}-call"),
name: "lookup".to_string(),
arguments: serde_json::json!({"query":"fixture"}),
};
let state = NativeProviderState::new(
exchange_id,
"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(),
),
ChatMessage::function(
"lookup",
encode_native_tool_result_marker(&call, serde_json::json!({"ok":true})).unwrap(),
),
)
}
#[tokio::test]
async fn test_add_and_get_messages() {
let store = InMemoryStore::new(10);
store.add_message(make_message("hello")).await.unwrap();
store.add_message(make_message("world")).await.unwrap();
let messages = store.get_messages(None).await.unwrap();
assert_eq!(messages.len(), 2);
assert_eq!(messages[0].content, "hello");
assert_eq!(messages[1].content, "world");
}
#[tokio::test]
async fn test_max_messages_limit() {
let store = InMemoryStore::new(3);
for i in 0..5 {
store
.add_message(make_message(&format!("msg{}", i)))
.await
.unwrap();
}
let messages = store.get_messages(None).await.unwrap();
assert_eq!(messages.len(), 3);
assert_eq!(messages[0].content, "msg2");
assert_eq!(messages[1].content, "msg3");
assert_eq!(messages[2].content, "msg4");
}
#[tokio::test]
async fn test_get_messages_with_limit() {
let store = InMemoryStore::new(10);
for i in 0..5 {
store
.add_message(make_message(&format!("msg{}", i)))
.await
.unwrap();
}
let messages = store.get_messages(Some(2)).await.unwrap();
assert_eq!(messages.len(), 2);
assert_eq!(messages[0].content, "msg3");
assert_eq!(messages[1].content, "msg4");
}
#[tokio::test]
async fn test_clear() {
let store = InMemoryStore::new(10);
store.add_message(make_message("test")).await.unwrap();
assert!(!store.is_empty());
store.clear().await.unwrap();
assert!(store.is_empty());
}
#[tokio::test]
async fn test_clone_shares_state() {
let store1 = InMemoryStore::new(10);
let store2 = store1.clone();
store1
.add_message(make_message("from store1"))
.await
.unwrap();
let messages = store2.get_messages(None).await.unwrap();
assert_eq!(messages.len(), 1);
assert_eq!(messages[0].content, "from store1");
}
#[tokio::test]
async fn test_snapshot_restore() {
let store = InMemoryStore::new(10);
store.add_message(make_message("msg1")).await.unwrap();
store.add_message(make_message("msg2")).await.unwrap();
let snapshot = store.snapshot().await.unwrap();
assert_eq!(snapshot.messages.len(), 2);
store.clear().await.unwrap();
assert!(store.is_empty());
store.restore(snapshot).await.unwrap();
let messages = store.get_messages(None).await.unwrap();
assert_eq!(messages.len(), 2);
assert_eq!(messages[0].content, "msg1");
}
#[tokio::test]
async fn test_evict_oldest() {
let store = InMemoryStore::new(10);
for i in 0..5 {
store
.add_message(make_message(&format!("msg{}", i)))
.await
.unwrap();
}
let evicted = store.evict_oldest(2).await.unwrap();
assert_eq!(evicted.len(), 2);
assert_eq!(evicted[0].content, "msg0");
assert_eq!(evicted[1].content, "msg1");
let remaining = store.get_messages(None).await.unwrap();
assert_eq!(remaining.len(), 3);
assert_eq!(remaining[0].content, "msg2");
}
#[tokio::test]
async fn signed_add_rejects_limit_that_would_split_protected_turn_atomically() {
let store = InMemoryStore::new(2);
let (assistant, result) = signed_exchange("active-add");
store
.add_message(make_message("current user"))
.await
.unwrap();
store.add_message(assistant.clone()).await.unwrap();
let error = store.add_message(result).await.unwrap_err();
assert!(
error
.to_string()
.contains("protected signed native exchange")
);
let retained = store.get_messages(None).await.unwrap();
assert_eq!(retained.len(), 2);
assert_eq!(retained[0].content, "current user");
assert_eq!(retained[1].content, assistant.content);
}
#[tokio::test]
async fn signed_add_evicts_completed_past_turn_as_one_group() {
let store = InMemoryStore::new(3);
let (assistant, result) = signed_exchange("past-add");
store.add_message(make_message("old user")).await.unwrap();
store.add_message(assistant).await.unwrap();
store.add_message(result).await.unwrap();
store.add_message(make_message("new user")).await.unwrap();
let retained = store.get_messages(None).await.unwrap();
assert_eq!(retained.len(), 1);
assert_eq!(retained[0].content, "new user");
}
#[tokio::test]
async fn signed_restore_failure_leaves_existing_history_unchanged() {
let store = InMemoryStore::new(2);
store.add_message(make_message("existing")).await.unwrap();
let (assistant, result) = signed_exchange("restore-active");
let snapshot = MemorySnapshot::new(vec![make_message("restored user"), assistant, result]);
let error = store.restore(snapshot).await.unwrap_err();
assert!(
error
.to_string()
.contains("protected signed native exchange")
);
let retained = store.get_messages(None).await.unwrap();
assert_eq!(retained.len(), 1);
assert_eq!(retained[0].content, "existing");
}
#[tokio::test]
async fn signed_eviction_expands_to_complete_past_turn() {
let store = InMemoryStore::new(10);
let (assistant, result) = signed_exchange("past-evict");
store.add_message(make_message("old user")).await.unwrap();
store.add_message(assistant).await.unwrap();
store.add_message(result).await.unwrap();
store.add_message(make_message("new user")).await.unwrap();
let evicted = store.evict_oldest(1).await.unwrap();
assert_eq!(evicted.len(), 3);
let retained = store.get_messages(None).await.unwrap();
assert_eq!(retained.len(), 1);
assert_eq!(retained[0].content, "new user");
}
#[tokio::test]
async fn signed_eviction_rejects_protected_turn_without_mutation() {
let store = InMemoryStore::new(10);
let (assistant, result) = signed_exchange("active-evict");
store
.add_message(make_message("current user"))
.await
.unwrap();
store.add_message(assistant).await.unwrap();
store.add_message(result).await.unwrap();
let before = store.get_messages(None).await.unwrap();
let error = store.evict_oldest(1).await.unwrap_err();
assert!(
error
.to_string()
.contains("protected signed native exchange")
);
let after = store.get_messages(None).await.unwrap();
assert_eq!(
after
.iter()
.map(|message| &message.content)
.collect::<Vec<_>>(),
before
.iter()
.map(|message| &message.content)
.collect::<Vec<_>>()
);
}
}