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 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}