Skip to main content

ai_agents_memory/
in_memory.rs

1use std::sync::Arc;
2
3use async_trait::async_trait;
4use parking_lot::RwLock;
5
6use ai_agents_core::{ChatMessage, MemorySnapshot, Result};
7
8use super::Memory;
9use super::native::NativeRetentionInspection;
10
11pub struct InMemoryStore {
12    messages: Arc<RwLock<Vec<ChatMessage>>>,
13    max_messages: usize,
14}
15
16impl InMemoryStore {
17    pub fn new(max_messages: usize) -> Self {
18        Self {
19            messages: Arc::new(RwLock::new(Vec::new())),
20            max_messages,
21        }
22    }
23
24    pub fn max_messages(&self) -> usize {
25        self.max_messages
26    }
27
28    // Computes the complete prefix that can be removed before mutating the store.
29    fn bounded_eviction_count(messages: &[ChatMessage], max_messages: usize) -> Result<usize> {
30        let inspection = NativeRetentionInspection::inspect(messages)?;
31        let required = messages.len().saturating_sub(max_messages);
32        if required == 0 {
33            return Ok(0);
34        }
35        inspection
36            .safe_prefix_len_between(required, messages.len())
37            .ok_or_else(|| {
38                ai_agents_core::AgentError::MemoryError(
39                    "message limit cannot preserve the protected signed native exchange"
40                        .to_string(),
41                )
42            })
43    }
44}
45
46impl Clone for InMemoryStore {
47    fn clone(&self) -> Self {
48        Self {
49            messages: Arc::clone(&self.messages),
50            max_messages: self.max_messages,
51        }
52    }
53}
54
55#[async_trait]
56impl ai_agents_core::Memory for InMemoryStore {
57    async fn add_message(&self, message: ChatMessage) -> Result<()> {
58        let mut messages = self.messages.write();
59        let mut prospective = messages.clone();
60        prospective.push(message);
61        let evict_count = Self::bounded_eviction_count(&prospective, self.max_messages)?;
62        prospective.drain(..evict_count);
63        *messages = prospective;
64
65        Ok(())
66    }
67
68    async fn get_messages(&self, limit: Option<usize>) -> Result<Vec<ChatMessage>> {
69        let messages = self.messages.read();
70        match limit {
71            Some(n) => {
72                let start = messages.len().saturating_sub(n);
73                Ok(messages[start..].to_vec())
74            }
75            None => Ok(messages.clone()),
76        }
77    }
78
79    async fn clear(&self) -> Result<()> {
80        self.messages.write().clear();
81        Ok(())
82    }
83
84    fn len(&self) -> usize {
85        self.messages.read().len()
86    }
87
88    async fn restore(&self, snapshot: MemorySnapshot) -> Result<()> {
89        let mut prospective = snapshot.messages;
90        let evict_count = Self::bounded_eviction_count(&prospective, self.max_messages)?;
91        prospective.drain(..evict_count);
92        let mut messages = self.messages.write();
93        *messages = prospective;
94        Ok(())
95    }
96
97    async fn evict_oldest(&self, count: usize) -> Result<Vec<ChatMessage>> {
98        let mut messages = self.messages.write();
99        let requested = count.min(messages.len());
100        let inspection = NativeRetentionInspection::inspect(&messages)?;
101        let evict_count = if requested == 0 {
102            0
103        } else {
104            inspection
105                .safe_prefix_len_between(requested, messages.len())
106                .ok_or_else(|| {
107                    ai_agents_core::AgentError::MemoryError(
108                        "eviction would split the protected signed native exchange".to_string(),
109                    )
110                })?
111        };
112        let evicted: Vec<ChatMessage> = messages.drain(..evict_count).collect();
113        Ok(evicted)
114    }
115}
116
117#[async_trait]
118impl Memory for InMemoryStore {}
119
120#[cfg(test)]
121mod tests {
122    use super::*;
123    use ai_agents_core::{
124        Memory as CoreMemory, NativeCallBinding, NativeProviderState, NativeProviderTarget, Role,
125        ToolCall, encode_native_tool_call_markers, encode_native_tool_result_marker,
126    };
127
128    fn make_message(content: &str) -> ChatMessage {
129        ChatMessage {
130            role: Role::User,
131            content: content.to_string(),
132            name: None,
133            timestamp: None,
134        }
135    }
136
137    fn signed_exchange(exchange_id: &str) -> (ChatMessage, ChatMessage) {
138        let call = ToolCall {
139            id: format!("{exchange_id}-call"),
140            name: "lookup".to_string(),
141            arguments: serde_json::json!({"query":"fixture"}),
142        };
143        let state = NativeProviderState::new(
144            exchange_id,
145            "google",
146            "generateContent",
147            NativeProviderTarget::new("https://example.invalid/v1beta/", "fixture-model").unwrap(),
148            serde_json::json!({
149                "role":"model",
150                "parts":[{
151                    "functionCall":{"name":"lookup","args":{"query":"fixture"}},
152                    "thoughtSignature":"fixture-signature"
153                }]
154            }),
155            vec![NativeCallBinding::new(&call.id, 0).unwrap()],
156        )
157        .unwrap();
158        (
159            ChatMessage::assistant(
160                encode_native_tool_call_markers(std::slice::from_ref(&call), Some(&state)).unwrap(),
161            ),
162            ChatMessage::function(
163                "lookup",
164                encode_native_tool_result_marker(&call, serde_json::json!({"ok":true})).unwrap(),
165            ),
166        )
167    }
168
169    #[tokio::test]
170    async fn test_add_and_get_messages() {
171        let store = InMemoryStore::new(10);
172
173        store.add_message(make_message("hello")).await.unwrap();
174        store.add_message(make_message("world")).await.unwrap();
175
176        let messages = store.get_messages(None).await.unwrap();
177        assert_eq!(messages.len(), 2);
178        assert_eq!(messages[0].content, "hello");
179        assert_eq!(messages[1].content, "world");
180    }
181
182    #[tokio::test]
183    async fn test_max_messages_limit() {
184        let store = InMemoryStore::new(3);
185
186        for i in 0..5 {
187            store
188                .add_message(make_message(&format!("msg{}", i)))
189                .await
190                .unwrap();
191        }
192
193        let messages = store.get_messages(None).await.unwrap();
194        assert_eq!(messages.len(), 3);
195        assert_eq!(messages[0].content, "msg2");
196        assert_eq!(messages[1].content, "msg3");
197        assert_eq!(messages[2].content, "msg4");
198    }
199
200    #[tokio::test]
201    async fn test_get_messages_with_limit() {
202        let store = InMemoryStore::new(10);
203
204        for i in 0..5 {
205            store
206                .add_message(make_message(&format!("msg{}", i)))
207                .await
208                .unwrap();
209        }
210
211        let messages = store.get_messages(Some(2)).await.unwrap();
212        assert_eq!(messages.len(), 2);
213        assert_eq!(messages[0].content, "msg3");
214        assert_eq!(messages[1].content, "msg4");
215    }
216
217    #[tokio::test]
218    async fn test_clear() {
219        let store = InMemoryStore::new(10);
220
221        store.add_message(make_message("test")).await.unwrap();
222        assert!(!store.is_empty());
223
224        store.clear().await.unwrap();
225        assert!(store.is_empty());
226    }
227
228    #[tokio::test]
229    async fn test_clone_shares_state() {
230        let store1 = InMemoryStore::new(10);
231        let store2 = store1.clone();
232
233        store1
234            .add_message(make_message("from store1"))
235            .await
236            .unwrap();
237
238        let messages = store2.get_messages(None).await.unwrap();
239        assert_eq!(messages.len(), 1);
240        assert_eq!(messages[0].content, "from store1");
241    }
242
243    #[tokio::test]
244    async fn test_snapshot_restore() {
245        let store = InMemoryStore::new(10);
246        store.add_message(make_message("msg1")).await.unwrap();
247        store.add_message(make_message("msg2")).await.unwrap();
248
249        let snapshot = store.snapshot().await.unwrap();
250        assert_eq!(snapshot.messages.len(), 2);
251
252        store.clear().await.unwrap();
253        assert!(store.is_empty());
254
255        store.restore(snapshot).await.unwrap();
256        let messages = store.get_messages(None).await.unwrap();
257        assert_eq!(messages.len(), 2);
258        assert_eq!(messages[0].content, "msg1");
259    }
260
261    #[tokio::test]
262    async fn test_evict_oldest() {
263        let store = InMemoryStore::new(10);
264        for i in 0..5 {
265            store
266                .add_message(make_message(&format!("msg{}", i)))
267                .await
268                .unwrap();
269        }
270
271        let evicted = store.evict_oldest(2).await.unwrap();
272        assert_eq!(evicted.len(), 2);
273        assert_eq!(evicted[0].content, "msg0");
274        assert_eq!(evicted[1].content, "msg1");
275
276        let remaining = store.get_messages(None).await.unwrap();
277        assert_eq!(remaining.len(), 3);
278        assert_eq!(remaining[0].content, "msg2");
279    }
280
281    #[tokio::test]
282    async fn signed_add_rejects_limit_that_would_split_protected_turn_atomically() {
283        let store = InMemoryStore::new(2);
284        let (assistant, result) = signed_exchange("active-add");
285        store
286            .add_message(make_message("current user"))
287            .await
288            .unwrap();
289        store.add_message(assistant.clone()).await.unwrap();
290
291        let error = store.add_message(result).await.unwrap_err();
292
293        assert!(
294            error
295                .to_string()
296                .contains("protected signed native exchange")
297        );
298        let retained = store.get_messages(None).await.unwrap();
299        assert_eq!(retained.len(), 2);
300        assert_eq!(retained[0].content, "current user");
301        assert_eq!(retained[1].content, assistant.content);
302    }
303
304    #[tokio::test]
305    async fn signed_add_evicts_completed_past_turn_as_one_group() {
306        let store = InMemoryStore::new(3);
307        let (assistant, result) = signed_exchange("past-add");
308        store.add_message(make_message("old user")).await.unwrap();
309        store.add_message(assistant).await.unwrap();
310        store.add_message(result).await.unwrap();
311
312        store.add_message(make_message("new user")).await.unwrap();
313
314        let retained = store.get_messages(None).await.unwrap();
315        assert_eq!(retained.len(), 1);
316        assert_eq!(retained[0].content, "new user");
317    }
318
319    #[tokio::test]
320    async fn signed_restore_failure_leaves_existing_history_unchanged() {
321        let store = InMemoryStore::new(2);
322        store.add_message(make_message("existing")).await.unwrap();
323        let (assistant, result) = signed_exchange("restore-active");
324        let snapshot = MemorySnapshot::new(vec![make_message("restored user"), assistant, result]);
325
326        let error = store.restore(snapshot).await.unwrap_err();
327
328        assert!(
329            error
330                .to_string()
331                .contains("protected signed native exchange")
332        );
333        let retained = store.get_messages(None).await.unwrap();
334        assert_eq!(retained.len(), 1);
335        assert_eq!(retained[0].content, "existing");
336    }
337
338    #[tokio::test]
339    async fn signed_eviction_expands_to_complete_past_turn() {
340        let store = InMemoryStore::new(10);
341        let (assistant, result) = signed_exchange("past-evict");
342        store.add_message(make_message("old user")).await.unwrap();
343        store.add_message(assistant).await.unwrap();
344        store.add_message(result).await.unwrap();
345        store.add_message(make_message("new user")).await.unwrap();
346
347        let evicted = store.evict_oldest(1).await.unwrap();
348
349        assert_eq!(evicted.len(), 3);
350        let retained = store.get_messages(None).await.unwrap();
351        assert_eq!(retained.len(), 1);
352        assert_eq!(retained[0].content, "new user");
353    }
354
355    #[tokio::test]
356    async fn signed_eviction_rejects_protected_turn_without_mutation() {
357        let store = InMemoryStore::new(10);
358        let (assistant, result) = signed_exchange("active-evict");
359        store
360            .add_message(make_message("current user"))
361            .await
362            .unwrap();
363        store.add_message(assistant).await.unwrap();
364        store.add_message(result).await.unwrap();
365
366        let before = store.get_messages(None).await.unwrap();
367        let error = store.evict_oldest(1).await.unwrap_err();
368
369        assert!(
370            error
371                .to_string()
372                .contains("protected signed native exchange")
373        );
374        let after = store.get_messages(None).await.unwrap();
375        assert_eq!(
376            after
377                .iter()
378                .map(|message| &message.content)
379                .collect::<Vec<_>>(),
380            before
381                .iter()
382                .map(|message| &message.content)
383                .collect::<Vec<_>>()
384        );
385    }
386}