paladin_memory/garrison/
in_memory_garrison.rs1use async_trait::async_trait;
20use paladin_core::platform::container::garrison::{
21 ConversationRole, EvictionStrategy, GarrisonConfig, GarrisonEntry,
22};
23use paladin_ports::output::garrison_port::{GarrisonError, GarrisonPort, GarrisonStats};
24use std::collections::VecDeque;
25use std::sync::RwLock;
26
27pub struct InMemoryGarrison {
59 entries: RwLock<VecDeque<GarrisonEntry>>,
60 config: GarrisonConfig,
61}
62
63impl InMemoryGarrison {
64 pub fn new(config: GarrisonConfig) -> Self {
80 Self {
81 entries: RwLock::new(VecDeque::new()),
82 config,
83 }
84 }
85
86 fn apply_windowing(&self, entries: &mut VecDeque<GarrisonEntry>) {
88 while entries.len() > self.config.max_entries {
90 self.evict_entry(entries);
91 }
92
93 if let Some(max_tokens) = self.config.max_tokens {
95 while self.calculate_total_tokens(entries) > max_tokens && !entries.is_empty() {
96 self.evict_entry(entries);
97 }
98 }
99 }
100
101 fn evict_entry(&self, entries: &mut VecDeque<GarrisonEntry>) {
103 match self.config.eviction_strategy {
104 EvictionStrategy::FIFO | EvictionStrategy::SlidingWindow => {
105 entries.pop_front();
106 }
107 EvictionStrategy::ImportanceBased => {
108 self.evict_importance_based(entries);
109 }
110 }
111 }
112
113 fn evict_importance_based(&self, entries: &mut VecDeque<GarrisonEntry>) {
115 let total_entries = entries.len();
116 if total_entries == 0 {
117 return;
118 }
119
120 let preserve_count = self.config.preserve_recent_count.min(total_entries);
121 let recent_start_idx = total_entries.saturating_sub(preserve_count);
122
123 for i in 0..recent_start_idx {
125 if entries[i].role != ConversationRole::System {
126 entries.remove(i);
127 return;
128 }
129 }
130
131 for i in recent_start_idx..total_entries {
133 if entries[i].role != ConversationRole::System {
134 entries.remove(i);
135 return;
136 }
137 }
138
139 entries.pop_front();
141 }
142
143 fn calculate_total_tokens(&self, entries: &VecDeque<GarrisonEntry>) -> u32 {
145 entries.iter().filter_map(|e| e.token_count).sum()
146 }
147
148 fn estimate_size_bytes(&self, entries: &VecDeque<GarrisonEntry>) -> usize {
150 entries.iter().map(|e| e.content.len()).sum()
151 }
152}
153
154#[async_trait]
155impl GarrisonPort for InMemoryGarrison {
156 async fn remember(&self, entry: GarrisonEntry) -> Result<(), GarrisonError> {
157 entry
159 .validate()
160 .map_err(|e| GarrisonError::Custom(format!("Invalid entry: {}", e)))?;
161
162 let mut entries = self
163 .entries
164 .write()
165 .map_err(|e| GarrisonError::StorageError(format!("Lock poisoned: {}", e)))?;
166
167 entries.push_back(entry);
168 self.apply_windowing(&mut entries);
169
170 Ok(())
171 }
172
173 async fn recall_recent(&self, limit: usize) -> Result<Vec<GarrisonEntry>, GarrisonError> {
174 let entries = self
175 .entries
176 .read()
177 .map_err(|e| GarrisonError::StorageError(format!("Lock poisoned: {}", e)))?;
178
179 let start = entries.len().saturating_sub(limit);
180 Ok(entries.range(start..).cloned().collect())
181 }
182
183 async fn search(&self, query: &str, limit: usize) -> Result<Vec<GarrisonEntry>, GarrisonError> {
184 let entries = self
185 .entries
186 .read()
187 .map_err(|e| GarrisonError::StorageError(format!("Lock poisoned: {}", e)))?;
188
189 let results: Vec<GarrisonEntry> = entries
190 .iter()
191 .filter(|e| e.content.contains(query))
192 .take(limit)
193 .cloned()
194 .collect();
195
196 Ok(results)
197 }
198
199 async fn forget_all(&self) -> Result<(), GarrisonError> {
200 let mut entries = self
201 .entries
202 .write()
203 .map_err(|e| GarrisonError::StorageError(format!("Lock poisoned: {}", e)))?;
204
205 entries.clear();
206 Ok(())
207 }
208
209 async fn stats(&self) -> Result<GarrisonStats, GarrisonError> {
210 let entries = self
211 .entries
212 .read()
213 .map_err(|e| GarrisonError::StorageError(format!("Lock poisoned: {}", e)))?;
214
215 Ok(GarrisonStats {
216 entry_count: entries.len(),
217 total_tokens: self.calculate_total_tokens(&entries),
218 size_bytes: Some(self.estimate_size_bytes(&entries)),
219 })
220 }
221}
222
223#[cfg(test)]
224mod tests {
225 use super::*;
226 use paladin_core::platform::container::garrison::ConversationRole;
227
228 #[tokio::test]
229 async fn test_remember_and_recall() {
230 let config = GarrisonConfig::default();
231 let garrison = InMemoryGarrison::new(config);
232
233 let entry = GarrisonEntry::new(ConversationRole::User, "Test message".to_string());
234 garrison.remember(entry).await.unwrap();
235
236 let recent = garrison.recall_recent(10).await.unwrap();
237 assert_eq!(recent.len(), 1);
238 assert_eq!(recent[0].content, "Test message");
239 }
240
241 #[tokio::test]
242 async fn test_windowing_by_count() {
243 let config = GarrisonConfig::new(3, None);
244 let garrison = InMemoryGarrison::new(config);
245
246 for i in 0..5 {
247 let entry = GarrisonEntry::new(ConversationRole::User, format!("Message {}", i));
248 garrison.remember(entry).await.unwrap();
249 }
250
251 let all = garrison.recall_recent(100).await.unwrap();
252 assert_eq!(all.len(), 3);
253 assert_eq!(all[0].content, "Message 2");
254 }
255
256 #[tokio::test]
257 async fn test_windowing_by_tokens() {
258 let config = GarrisonConfig::new(100, Some(50));
259 let garrison = InMemoryGarrison::new(config);
260
261 for i in 0..5 {
262 let entry = GarrisonEntry::with_token_count(
263 ConversationRole::User,
264 format!("Message {}", i),
265 20,
266 );
267 garrison.remember(entry).await.unwrap();
268 }
269
270 let all = garrison.recall_recent(100).await.unwrap();
271 assert!(all.len() <= 3);
273 }
274
275 #[tokio::test]
276 async fn test_search_functionality() {
277 let config = GarrisonConfig::default();
278 let garrison = InMemoryGarrison::new(config);
279
280 garrison
281 .remember(GarrisonEntry::new(
282 ConversationRole::User,
283 "Hello world".to_string(),
284 ))
285 .await
286 .unwrap();
287 garrison
288 .remember(GarrisonEntry::new(
289 ConversationRole::User,
290 "Goodbye world".to_string(),
291 ))
292 .await
293 .unwrap();
294 garrison
295 .remember(GarrisonEntry::new(
296 ConversationRole::User,
297 "Random message".to_string(),
298 ))
299 .await
300 .unwrap();
301
302 let results = garrison.search("world", 10).await.unwrap();
303 assert_eq!(results.len(), 2);
304 }
305
306 #[tokio::test]
307 async fn test_forget_all() {
308 let config = GarrisonConfig::default();
309 let garrison = InMemoryGarrison::new(config);
310
311 garrison
312 .remember(GarrisonEntry::new(
313 ConversationRole::User,
314 "Test".to_string(),
315 ))
316 .await
317 .unwrap();
318
319 garrison.forget_all().await.unwrap();
320
321 let recent = garrison.recall_recent(10).await.unwrap();
322 assert_eq!(recent.len(), 0);
323 }
324
325 #[tokio::test]
326 async fn test_stats() {
327 let config = GarrisonConfig::default();
328 let garrison = InMemoryGarrison::new(config);
329
330 garrison
331 .remember(GarrisonEntry::with_token_count(
332 ConversationRole::User,
333 "First".to_string(),
334 10,
335 ))
336 .await
337 .unwrap();
338 garrison
339 .remember(GarrisonEntry::with_token_count(
340 ConversationRole::Assistant,
341 "Second".to_string(),
342 20,
343 ))
344 .await
345 .unwrap();
346
347 let stats = garrison.stats().await.unwrap();
348 assert_eq!(stats.entry_count, 2);
349 assert_eq!(stats.total_tokens, 30);
350 assert!(stats.size_bytes.is_some());
351 }
352
353 #[tokio::test]
354 async fn test_importance_based_eviction() {
355 let config = GarrisonConfig::new(3, None)
356 .with_eviction_strategy(EvictionStrategy::ImportanceBased)
357 .with_preserve_recent(1);
358
359 let garrison = InMemoryGarrison::new(config);
360
361 garrison
362 .remember(GarrisonEntry::new(
363 ConversationRole::System,
364 "System prompt".to_string(),
365 ))
366 .await
367 .unwrap();
368 garrison
369 .remember(GarrisonEntry::new(
370 ConversationRole::User,
371 "User 1".to_string(),
372 ))
373 .await
374 .unwrap();
375 garrison
376 .remember(GarrisonEntry::new(
377 ConversationRole::User,
378 "User 2".to_string(),
379 ))
380 .await
381 .unwrap();
382 garrison
383 .remember(GarrisonEntry::new(
384 ConversationRole::User,
385 "User 3 - triggers eviction".to_string(),
386 ))
387 .await
388 .unwrap();
389
390 let all = garrison.recall_recent(100).await.unwrap();
391 assert_eq!(all.len(), 3);
392 assert!(all.iter().any(|e| e.role == ConversationRole::System));
394 }
395}