1use a3s_memory::{MemoryItem, MemoryStore, MemoryType, PrunePolicy, RelevanceConfig};
10use chrono::{DateTime, Utc};
11use serde::{Deserialize, Serialize};
12use std::collections::VecDeque;
13use std::sync::atomic::{AtomicUsize, Ordering};
14use std::sync::Arc;
15use tokio::sync::{oneshot, Notify, RwLock};
16
17#[derive(Debug, Clone)]
23pub struct MemoryObservation {
24 pub incoming: MemoryItem,
25 pub stored: MemoryItem,
26 pub merged: bool,
27}
28
29#[async_trait::async_trait]
35pub trait MemoryObserver: Send + Sync {
36 async fn on_memory_stored(&self, observation: MemoryObservation) -> anyhow::Result<()>;
37}
38
39#[derive(Debug, Clone, Serialize, Deserialize)]
45#[serde(rename_all = "camelCase")]
46pub struct MemoryConfig {
47 #[serde(default)]
49 pub relevance: RelevanceConfig,
50 #[serde(default = "MemoryConfig::default_max_short_term")]
52 pub max_short_term: usize,
53 #[serde(default = "MemoryConfig::default_max_working")]
55 pub max_working: usize,
56 #[serde(default)]
58 pub prune_policy: Option<PrunePolicy>,
59 #[serde(default = "MemoryConfig::default_prune_interval_secs")]
61 pub prune_interval_secs: u64,
62 #[serde(
69 default = "MemoryConfig::default_llm_extraction",
70 alias = "llm_extraction"
71 )]
72 pub llm_extraction: bool,
73 #[serde(default = "MemoryConfig::default_llm_extraction_max_items")]
75 pub llm_extraction_max_items: usize,
76 #[serde(default = "MemoryConfig::default_llm_extraction_max_input_chars")]
78 pub llm_extraction_max_input_chars: usize,
79}
80
81impl MemoryConfig {
82 fn default_max_short_term() -> usize {
83 100
84 }
85 fn default_max_working() -> usize {
86 10
87 }
88 fn default_prune_interval_secs() -> u64 {
89 3600
90 }
91 fn default_llm_extraction() -> bool {
92 true
93 }
94 fn default_llm_extraction_max_items() -> usize {
95 5
96 }
97 fn default_llm_extraction_max_input_chars() -> usize {
98 8_000
99 }
100}
101
102impl Default for MemoryConfig {
103 fn default() -> Self {
104 Self {
105 relevance: RelevanceConfig::default(),
106 max_short_term: 100,
107 max_working: 10,
108 prune_policy: None,
109 prune_interval_secs: 3600,
110 llm_extraction: true,
111 llm_extraction_max_items: 5,
112 llm_extraction_max_input_chars: 8_000,
113 }
114 }
115}
116
117#[derive(Debug, Clone, Serialize, Deserialize)]
123pub struct MemoryStats {
124 pub long_term_count: usize,
125 pub short_term_count: usize,
126 pub working_count: usize,
127}
128
129#[derive(Clone)]
135pub struct AgentMemory {
136 pub(crate) store: Arc<dyn MemoryStore>,
138 short_term: Arc<RwLock<VecDeque<MemoryItem>>>,
140 working: Arc<RwLock<Vec<MemoryItem>>>,
142 pub(crate) max_short_term: usize,
143 pub(crate) max_working: usize,
144 pub(crate) relevance_config: RelevanceConfig,
145 pub(crate) llm_extraction: bool,
146 pub(crate) llm_extraction_max_items: usize,
147 pub(crate) llm_extraction_max_input_chars: usize,
148 extraction_queue: Arc<MemoryExtractionQueue>,
149 observers: Arc<Vec<Arc<dyn MemoryObserver>>>,
150}
151
152#[derive(Default)]
153struct MemoryExtractionQueue {
154 state: std::sync::Mutex<MemoryExtractionQueueState>,
155 pending: AtomicUsize,
156 idle: Notify,
157}
158
159#[derive(Default)]
160struct MemoryExtractionQueueState {
161 tail: Option<oneshot::Receiver<()>>,
162}
163
164pub(crate) struct MemoryExtractionTicket {
171 predecessor: Option<oneshot::Receiver<()>>,
172 completion: Option<oneshot::Sender<()>>,
173 queue: Arc<MemoryExtractionQueue>,
174}
175
176impl MemoryExtractionTicket {
177 pub(crate) async fn wait_for_turn(&mut self) {
178 if let Some(predecessor) = self.predecessor.take() {
179 let _ = predecessor.await;
180 }
181 }
182}
183
184impl Drop for MemoryExtractionTicket {
185 fn drop(&mut self) {
186 if let Some(completion) = self.completion.take() {
187 let _ = completion.send(());
188 }
189 if self.queue.pending.fetch_sub(1, Ordering::AcqRel) == 1 {
190 self.queue.idle.notify_waiters();
191 }
192 }
193}
194
195impl std::fmt::Debug for AgentMemory {
196 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
197 f.debug_struct("AgentMemory")
198 .field("max_short_term", &self.max_short_term)
199 .field("max_working", &self.max_working)
200 .field("observers", &self.observers.len())
201 .finish()
202 }
203}
204
205impl AgentMemory {
206 pub fn new(store: Arc<dyn MemoryStore>) -> Self {
208 Self::with_config(store, MemoryConfig::default())
209 }
210
211 pub fn with_config(store: Arc<dyn MemoryStore>, config: MemoryConfig) -> Self {
216 Self::with_config_and_observers(store, config, Vec::new())
217 }
218
219 pub fn with_config_and_observers(
223 store: Arc<dyn MemoryStore>,
224 config: MemoryConfig,
225 observers: Vec<Arc<dyn MemoryObserver>>,
226 ) -> Self {
227 if let Some(policy) = config.prune_policy.clone() {
228 let store_for_task = Arc::clone(&store);
229 let interval_secs = config.prune_interval_secs;
230 match tokio::runtime::Handle::try_current() {
231 Ok(handle) => {
232 handle.spawn(async move {
233 let mut ticker =
234 tokio::time::interval(std::time::Duration::from_secs(interval_secs));
235 ticker.tick().await; loop {
237 ticker.tick().await;
238 if let Err(e) = store_for_task.prune(&policy).await {
239 tracing::warn!("memory prune failed: {e}");
240 }
241 }
242 });
243 }
244 Err(_) => {
245 tracing::warn!(
246 "memory prune policy configured but no async runtime is available"
247 );
248 }
249 }
250 }
251
252 Self {
253 store,
254 short_term: Arc::new(RwLock::new(VecDeque::new())),
255 working: Arc::new(RwLock::new(Vec::new())),
256 max_short_term: config.max_short_term,
257 max_working: config.max_working,
258 relevance_config: config.relevance,
259 llm_extraction: config.llm_extraction,
260 llm_extraction_max_items: config.llm_extraction_max_items,
261 llm_extraction_max_input_chars: config.llm_extraction_max_input_chars,
262 extraction_queue: Arc::new(MemoryExtractionQueue::default()),
263 observers: Arc::new(observers),
264 }
265 }
266
267 pub(crate) fn score(&self, item: &MemoryItem, now: DateTime<Utc>) -> f32 {
268 let age_days = (now - item.timestamp).num_seconds() as f32 / 86400.0;
269 let decay = (-age_days / self.relevance_config.decay_days).exp();
270 item.importance * self.relevance_config.importance_weight
271 + decay * self.relevance_config.recency_weight
272 }
273
274 pub async fn remember(&self, item: MemoryItem) -> anyhow::Result<()> {
276 self.remember_item(item).await.map(|_| ())
277 }
278
279 pub async fn remember_item(&self, item: MemoryItem) -> anyhow::Result<MemoryItem> {
281 let incoming = item.clone();
282 let item = self.store.store_and_return(item).await?;
283 let mut short_term = self.short_term.write().await;
284 if let Some(existing) = short_term
285 .iter_mut()
286 .find(|existing| existing.id == item.id)
287 {
288 *existing = item.clone();
289 } else {
290 short_term.push_back(item.clone());
291 }
292 if short_term.len() > self.max_short_term {
293 short_term.pop_front();
294 }
295 drop(short_term);
296
297 if !self.observers.is_empty() {
298 let observation = MemoryObservation {
299 merged: item.id != incoming.id,
300 incoming,
301 stored: item.clone(),
302 };
303 for observer in self.observers.iter() {
304 if let Err(error) = observer.on_memory_stored(observation.clone()).await {
305 tracing::warn!(%error, "memory observer failed after persistence");
306 }
307 }
308 }
309 Ok(item)
310 }
311
312 pub async fn forget(&self, id: &str) -> anyhow::Result<()> {
314 self.store.delete(id).await?;
315 self.short_term.write().await.retain(|item| item.id != id);
316 self.working.write().await.retain(|item| item.id != id);
317 Ok(())
318 }
319
320 pub async fn remember_success(
322 &self,
323 prompt: &str,
324 tools_used: &[String],
325 result: &str,
326 ) -> anyhow::Result<()> {
327 self.remember_success_item(prompt, tools_used, result)
328 .await
329 .map(|_| ())
330 }
331
332 pub async fn remember_success_item(
334 &self,
335 prompt: &str,
336 tools_used: &[String],
337 result: &str,
338 ) -> anyhow::Result<MemoryItem> {
339 let content = format!(
340 "Success: {}\nTools: {}\nResult: {}",
341 prompt,
342 tools_used.join(", "),
343 result
344 );
345 let mut item = MemoryItem::new(content)
346 .with_importance(0.8)
347 .with_tag("success")
348 .with_tag("pattern")
349 .with_type(MemoryType::Procedural)
350 .with_metadata("prompt", prompt)
351 .with_metadata("tools", tools_used.join(","));
352 for tool in tools_used {
353 item = item.with_tag(tool.clone());
354 }
355 self.remember_item(item).await
356 }
357
358 pub async fn remember_failure(
360 &self,
361 prompt: &str,
362 error: &str,
363 attempted_tools: &[String],
364 ) -> anyhow::Result<()> {
365 self.remember_failure_item(prompt, error, attempted_tools)
366 .await
367 .map(|_| ())
368 }
369
370 pub async fn remember_failure_item(
372 &self,
373 prompt: &str,
374 error: &str,
375 attempted_tools: &[String],
376 ) -> anyhow::Result<MemoryItem> {
377 let content = format!(
378 "Failure: {}\nError: {}\nAttempted tools: {}",
379 prompt,
380 error,
381 attempted_tools.join(", ")
382 );
383 let mut item = MemoryItem::new(content)
384 .with_importance(0.9)
385 .with_tag("failure")
386 .with_tag("avoid")
387 .with_type(MemoryType::Episodic)
388 .with_metadata("prompt", prompt)
389 .with_metadata("error", error);
390 for tool in attempted_tools {
391 item = item.with_tag(tool.clone());
392 }
393 self.remember_item(item).await
394 }
395
396 pub async fn recall_similar(
398 &self,
399 prompt: &str,
400 limit: usize,
401 ) -> anyhow::Result<Vec<MemoryItem>> {
402 self.store.search(prompt, limit).await
403 }
404
405 pub async fn recall_by_tags(
407 &self,
408 tags: &[String],
409 limit: usize,
410 ) -> anyhow::Result<Vec<MemoryItem>> {
411 self.store.search_by_tags(tags, limit).await
412 }
413
414 pub async fn get_recent(&self, limit: usize) -> anyhow::Result<Vec<MemoryItem>> {
416 self.store.get_recent(limit).await
417 }
418
419 pub async fn add_to_working(&self, item: MemoryItem) -> anyhow::Result<()> {
421 let mut working = self.working.write().await;
422 working.push(item);
423 if working.len() > self.max_working {
424 let now = Utc::now();
425 working.sort_by(|a, b| {
426 self.score(b, now)
427 .partial_cmp(&self.score(a, now))
428 .unwrap_or(std::cmp::Ordering::Equal)
429 });
430 working.truncate(self.max_working);
431 }
432 Ok(())
433 }
434
435 pub async fn get_working(&self) -> Vec<MemoryItem> {
437 self.working.read().await.clone()
438 }
439
440 pub async fn clear_working(&self) {
442 self.working.write().await.clear();
443 }
444
445 pub async fn get_short_term(&self) -> Vec<MemoryItem> {
447 self.short_term.read().await.iter().cloned().collect()
448 }
449
450 pub async fn clear_short_term(&self) {
452 self.short_term.write().await.clear();
453 }
454
455 pub async fn stats(&self) -> anyhow::Result<MemoryStats> {
457 Ok(MemoryStats {
458 long_term_count: self.store.count().await?,
459 short_term_count: self.short_term.read().await.len(),
460 working_count: self.working.read().await.len(),
461 })
462 }
463
464 pub fn store(&self) -> &Arc<dyn MemoryStore> {
466 &self.store
467 }
468
469 pub async fn working_count(&self) -> usize {
471 self.working.read().await.len()
472 }
473
474 pub async fn short_term_count(&self) -> usize {
476 self.short_term.read().await.len()
477 }
478
479 pub(crate) fn llm_extraction_enabled(&self) -> bool {
480 self.llm_extraction
481 }
482
483 pub(crate) fn llm_extraction_max_items(&self) -> usize {
484 self.llm_extraction_max_items
485 }
486
487 pub(crate) fn llm_extraction_max_input_chars(&self) -> usize {
488 self.llm_extraction_max_input_chars
489 }
490
491 pub(crate) fn enqueue_llm_extraction(&self) -> MemoryExtractionTicket {
492 let (completion, receiver) = oneshot::channel();
493 let predecessor = {
494 let mut state = self
495 .extraction_queue
496 .state
497 .lock()
498 .unwrap_or_else(std::sync::PoisonError::into_inner);
499 state.tail.replace(receiver)
500 };
501 self.extraction_queue.pending.fetch_add(1, Ordering::AcqRel);
502 MemoryExtractionTicket {
503 predecessor,
504 completion: Some(completion),
505 queue: Arc::clone(&self.extraction_queue),
506 }
507 }
508
509 pub(crate) async fn drain_llm_extractions(&self, timeout: std::time::Duration) -> bool {
512 let wait_until_idle = async {
513 loop {
514 let notified = self.extraction_queue.idle.notified();
515 if self.extraction_queue.pending.load(Ordering::Acquire) == 0 {
516 return;
517 }
518 notified.await;
519 }
520 };
521 tokio::time::timeout(timeout, wait_until_idle).await.is_ok()
522 }
523}
524
525pub struct MemoryContextProvider {
531 memory: AgentMemory,
532}
533
534impl MemoryContextProvider {
535 pub fn new(memory: AgentMemory) -> Self {
536 Self { memory }
537 }
538}
539
540pub(crate) fn memory_items_to_context_result(
541 provider: impl Into<String>,
542 items: Vec<MemoryItem>,
543) -> crate::context::ContextResult {
544 let mut result = crate::context::ContextResult::new(provider);
545 let total = items.len().max(1);
546 for (index, item) in items.into_iter().enumerate() {
547 let supersedes = relation_ids(&item, "supersedes");
548 let conflicts_with = relation_ids(&item, "conflicts_with");
549 let content = memory_context_content(&item, &supersedes, &conflicts_with);
550 let token_count = (content.len() / 4).max(1);
551 let recall_rank_score = 1.0 - (index as f32 / total as f32);
552 let relevance = (item.relevance_score() * 0.35 + recall_rank_score * 0.65).clamp(0.0, 1.0);
553 let context_item = crate::context::ContextItem::new(
554 &item.id,
555 crate::context::ContextType::Memory,
556 content,
557 )
558 .with_relevance(relevance)
559 .with_token_count(token_count)
560 .with_source(format!("memory://{}", item.id))
561 .with_metadata("memory_id", serde_json::json!(item.id))
562 .with_metadata(
563 "memory_type",
564 serde_json::json!(memory_type_label(item.memory_type)),
565 )
566 .with_metadata("tags", serde_json::json!(item.tags))
567 .with_metadata("importance", serde_json::json!(item.importance))
568 .with_provenance("long_term_memory")
569 .with_priority(0.35)
570 .with_trust(0.7)
571 .with_freshness(0.5);
572 let context_item = add_relation_metadata(context_item, "supersedes", supersedes);
573 let context_item = add_relation_metadata(context_item, "conflicts_with", conflicts_with);
574 result.add_item(context_item);
575 }
576 result
577}
578
579fn relation_ids(item: &MemoryItem, key: &str) -> Vec<String> {
580 item.metadata
581 .get(key)
582 .map(|value| {
583 value
584 .split(',')
585 .map(str::trim)
586 .filter(|id| !id.is_empty())
587 .map(ToOwned::to_owned)
588 .collect()
589 })
590 .unwrap_or_default()
591}
592
593fn memory_context_content(
594 item: &MemoryItem,
595 supersedes: &[String],
596 conflicts_with: &[String],
597) -> String {
598 let mut content = item.content.clone();
599 if supersedes.is_empty() && conflicts_with.is_empty() {
600 return content;
601 }
602
603 content.push_str("\n\nMemory relations:");
604 if !supersedes.is_empty() {
605 content.push_str("\n- supersedes: ");
606 content.push_str(&relation_sources(supersedes));
607 }
608 if !conflicts_with.is_empty() {
609 content.push_str("\n- conflicts_with: ");
610 content.push_str(&relation_sources(conflicts_with));
611 }
612 content
613}
614
615fn relation_sources(ids: &[String]) -> String {
616 ids.iter()
617 .map(|id| format!("memory://{id}"))
618 .collect::<Vec<_>>()
619 .join(", ")
620}
621
622fn add_relation_metadata(
623 item: crate::context::ContextItem,
624 key: &str,
625 ids: Vec<String>,
626) -> crate::context::ContextItem {
627 if ids.is_empty() {
628 item
629 } else {
630 item.with_metadata(key, serde_json::json!(ids))
631 }
632}
633
634fn memory_type_label(memory_type: MemoryType) -> &'static str {
635 match memory_type {
636 MemoryType::Episodic => "episodic",
637 MemoryType::Semantic => "semantic",
638 MemoryType::Procedural => "procedural",
639 MemoryType::Working => "working",
640 }
641}
642
643#[async_trait::async_trait]
644impl crate::context::ContextProvider for MemoryContextProvider {
645 fn name(&self) -> &str {
646 "memory"
647 }
648
649 async fn query(
650 &self,
651 query: &crate::context::ContextQuery,
652 ) -> anyhow::Result<crate::context::ContextResult> {
653 let limit = query.max_results.min(5);
654 let items = self.memory.recall_similar(&query.query, limit).await?;
655
656 Ok(memory_items_to_context_result("memory", items))
657 }
658
659 async fn on_turn_complete(
660 &self,
661 _session_id: &str,
662 _prompt: &str,
663 _response: &str,
664 ) -> anyhow::Result<()> {
665 Ok(())
668 }
669}
670
671#[cfg(test)]
676mod tests {
677 use super::*;
678 use crate::context::ContextProvider;
679 use a3s_memory::InMemoryStore;
680 use std::sync::{Arc, Mutex};
681
682 #[derive(Default)]
683 struct RecordingObserver {
684 observations: Mutex<Vec<MemoryObservation>>,
685 fail: bool,
686 }
687
688 #[async_trait::async_trait]
689 impl MemoryObserver for RecordingObserver {
690 async fn on_memory_stored(&self, observation: MemoryObservation) -> anyhow::Result<()> {
691 self.observations.lock().unwrap().push(observation);
692 if self.fail {
693 anyhow::bail!("observer projection failed");
694 }
695 Ok(())
696 }
697 }
698
699 #[tokio::test]
700 async fn test_agent_memory_remember_and_recall() {
701 let memory = AgentMemory::new(Arc::new(InMemoryStore::new()));
702 memory
703 .remember_success("create file", &["write".to_string()], "ok")
704 .await
705 .unwrap();
706 memory
707 .remember_failure("delete file", "denied", &["bash".to_string()])
708 .await
709 .unwrap();
710
711 let results = memory.recall_similar("create", 10).await.unwrap();
712 assert!(!results.is_empty());
713
714 let stats = memory.stats().await.unwrap();
715 assert_eq!(stats.long_term_count, 2);
716 assert_eq!(stats.short_term_count, 2);
717 }
718
719 #[tokio::test]
720 async fn test_agent_memory_forget_removes_all_tiers() {
721 let memory = AgentMemory::new(Arc::new(InMemoryStore::new()));
722 let item = memory
723 .remember_item(MemoryItem::new("superseded memory"))
724 .await
725 .unwrap();
726 memory.add_to_working(item.clone()).await.unwrap();
727
728 memory.forget(&item.id).await.unwrap();
729
730 assert_eq!(memory.stats().await.unwrap().long_term_count, 0);
731 assert!(memory.get_short_term().await.is_empty());
732 assert!(memory.get_working().await.is_empty());
733 }
734
735 #[tokio::test]
736 async fn test_agent_memory_uses_canonical_store_item_for_duplicates() {
737 let memory = AgentMemory::new(Arc::new(InMemoryStore::new()));
738 let first = memory
739 .remember_item(
740 MemoryItem::new("Run focused memory extraction tests after parser changes.")
741 .with_importance(0.3)
742 .with_tag("memory"),
743 )
744 .await
745 .unwrap();
746
747 let duplicate = memory
748 .remember_item(
749 MemoryItem::new(" run focused MEMORY extraction tests after parser changes. ")
750 .with_importance(0.9)
751 .with_tag("tests"),
752 )
753 .await
754 .unwrap();
755
756 assert_eq!(duplicate.id, first.id);
757 assert_eq!(memory.stats().await.unwrap().long_term_count, 1);
758 let short_term = memory.get_short_term().await;
759 assert_eq!(short_term.len(), 1);
760 assert_eq!(short_term[0].id, first.id);
761 assert_eq!(short_term[0].importance, 0.9);
762 assert!(short_term[0].tags.contains(&"memory".to_string()));
763 assert!(short_term[0].tags.contains(&"tests".to_string()));
764 }
765
766 #[tokio::test]
767 async fn test_memory_observer_receives_incoming_and_canonical_duplicate() {
768 let observer = Arc::new(RecordingObserver::default());
769 let memory = AgentMemory::with_config_and_observers(
770 Arc::new(InMemoryStore::new()),
771 MemoryConfig::default(),
772 vec![observer.clone()],
773 );
774 let first = memory
775 .remember_item(
776 MemoryItem::new("Run focused observer tests after memory persistence changes.")
777 .with_importance(0.8)
778 .with_metadata("session_id", "session-one"),
779 )
780 .await
781 .unwrap();
782 let duplicate_input =
783 MemoryItem::new(" run focused OBSERVER tests after memory persistence changes. ")
784 .with_importance(0.95)
785 .with_metadata("session_id", "session-two");
786 let duplicate_input_id = duplicate_input.id.clone();
787 let duplicate = memory.remember_item(duplicate_input).await.unwrap();
788
789 let observations = observer.observations.lock().unwrap();
790 assert_eq!(observations.len(), 2);
791 assert!(!observations[0].merged);
792 assert_eq!(observations[0].incoming.id, observations[0].stored.id);
793 assert!(observations[1].merged);
794 assert_eq!(observations[1].incoming.id, duplicate_input_id);
795 assert_eq!(observations[1].stored.id, first.id);
796 assert_eq!(observations[1].stored.id, duplicate.id);
797 assert_ne!(observations[1].incoming.id, observations[1].stored.id);
798 assert_eq!(
799 observations[1]
800 .incoming
801 .metadata
802 .get("session_id")
803 .map(String::as_str),
804 Some("session-two")
805 );
806 }
807
808 #[tokio::test]
809 async fn test_memory_observer_failure_does_not_roll_back_persistence() {
810 let store = Arc::new(InMemoryStore::new());
811 let observer = Arc::new(RecordingObserver {
812 observations: Mutex::new(Vec::new()),
813 fail: true,
814 });
815 let memory = AgentMemory::with_config_and_observers(
816 store.clone(),
817 MemoryConfig::default(),
818 vec![observer.clone()],
819 );
820
821 let stored = memory
822 .remember_item(MemoryItem::new(
823 "Persist even if a derived projection fails.",
824 ))
825 .await
826 .expect("observer errors must not fail the durable memory write");
827
828 assert_eq!(store.count().await.unwrap(), 1);
829 assert_eq!(memory.short_term_count().await, 1);
830 assert_eq!(memory.get_short_term().await[0].id, stored.id);
831 assert_eq!(observer.observations.lock().unwrap().len(), 1);
832 }
833
834 #[tokio::test]
835 async fn test_agent_memory_working() {
836 let memory = AgentMemory::new(Arc::new(InMemoryStore::new()));
837 memory
838 .add_to_working(MemoryItem::new("task").with_type(MemoryType::Working))
839 .await
840 .unwrap();
841 assert_eq!(memory.working_count().await, 1);
842 memory.clear_working().await;
843 assert_eq!(memory.working_count().await, 0);
844 }
845
846 #[tokio::test]
847 async fn test_agent_memory_working_overflow_trims() {
848 let memory = AgentMemory {
849 store: Arc::new(InMemoryStore::new()),
850 short_term: Arc::new(RwLock::new(VecDeque::new())),
851 working: Arc::new(RwLock::new(Vec::new())),
852 max_short_term: 100,
853 max_working: 3,
854 relevance_config: RelevanceConfig::default(),
855 llm_extraction: false,
856 llm_extraction_max_items: 5,
857 llm_extraction_max_input_chars: 8_000,
858 extraction_queue: Arc::new(MemoryExtractionQueue::default()),
859 observers: Arc::new(Vec::new()),
860 };
861 for i in 0..5 {
862 memory
863 .add_to_working(
864 MemoryItem::new(format!("task {i}")).with_importance(i as f32 * 0.2),
865 )
866 .await
867 .unwrap();
868 }
869 assert_eq!(memory.get_working().await.len(), 3);
870 }
871
872 #[tokio::test]
873 async fn test_agent_memory_recall_by_tags() {
874 let memory = AgentMemory::new(Arc::new(InMemoryStore::new()));
875 memory
876 .remember_success("create file", &["write".to_string()], "ok")
877 .await
878 .unwrap();
879 memory
880 .remember_failure("delete file", "denied", &["bash".to_string()])
881 .await
882 .unwrap();
883
884 let successes = memory
885 .recall_by_tags(&["success".to_string()], 10)
886 .await
887 .unwrap();
888 assert_eq!(successes.len(), 1);
889 let failures = memory
890 .recall_by_tags(&["failure".to_string()], 10)
891 .await
892 .unwrap();
893 assert_eq!(failures.len(), 1);
894 }
895
896 #[tokio::test]
897 async fn test_agent_memory_short_term_trim() {
898 let store = Arc::new(InMemoryStore::new());
899 let memory = AgentMemory {
900 store,
901 short_term: Arc::new(RwLock::new(VecDeque::new())),
902 working: Arc::new(RwLock::new(Vec::new())),
903 max_short_term: 3,
904 max_working: 10,
905 relevance_config: RelevanceConfig::default(),
906 llm_extraction: false,
907 llm_extraction_max_items: 5,
908 llm_extraction_max_input_chars: 8_000,
909 extraction_queue: Arc::new(MemoryExtractionQueue::default()),
910 observers: Arc::new(Vec::new()),
911 };
912 for i in 0..5 {
913 memory
914 .remember(MemoryItem::new(format!("item {i}")))
915 .await
916 .unwrap();
917 }
918 assert_eq!(memory.short_term_count().await, 3);
919 }
920
921 #[tokio::test]
922 async fn test_agent_memory_prune_delegates() {
923 use a3s_memory::PrunePolicy;
924
925 let store = Arc::new(InMemoryStore::new());
926 let memory = AgentMemory::new(store.clone());
927
928 let mut old_item = a3s_memory::MemoryItem::new("stale").with_importance(0.2);
930 old_item.timestamp = chrono::Utc::now() - chrono::Duration::days(100);
931 store.store(old_item).await.unwrap();
932
933 assert_eq!(store.count().await.unwrap(), 1);
934
935 let policy = PrunePolicy {
937 max_age_days: 90,
938 min_importance_to_keep: 0.5,
939 max_items: 0,
940 };
941 let deleted = memory.store().prune(&policy).await.unwrap();
942 assert_eq!(deleted, 1);
943 assert_eq!(store.count().await.unwrap(), 0);
944 }
945
946 #[test]
947 fn test_agent_memory_score_uses_config() {
948 let config = MemoryConfig {
949 relevance: RelevanceConfig {
950 decay_days: 7.0,
951 importance_weight: 0.9,
952 recency_weight: 0.1,
953 },
954 ..Default::default()
955 };
956 let memory = AgentMemory::with_config(Arc::new(InMemoryStore::new()), config);
957 let item = MemoryItem::new("Test").with_importance(1.0);
958 let score = memory.score(&item, Utc::now());
959 assert!(score > 0.95, "Score was {score}");
960 }
961
962 #[test]
963 fn test_memory_config_partial_deserialize_keeps_llm_extraction_enabled() {
964 let config: MemoryConfig = serde_json::from_str(r#"{"maxShortTerm": 12}"#).unwrap();
965 assert!(config.llm_extraction);
966 assert_eq!(config.max_short_term, 12);
967 }
968
969 #[test]
970 fn test_memory_config_allows_explicit_llm_extraction_disable() {
971 let config: MemoryConfig =
972 serde_json::from_str(r#"{"llmExtraction": false, "maxShortTerm": 12}"#).unwrap();
973 assert!(!config.llm_extraction);
974 assert_eq!(config.max_short_term, 12);
975 }
976
977 #[test]
978 fn test_memory_context_result_includes_relation_context() {
979 let item = MemoryItem::new("Use the file memory store for local sessions.")
980 .with_type(MemoryType::Procedural)
981 .with_tag("consolidated")
982 .with_tag("conflict")
983 .with_metadata("supersedes", "old-preference, old-workflow")
984 .with_metadata("conflicts_with", "legacy-default");
985
986 let result = memory_items_to_context_result("memory", vec![item.clone()]);
987
988 assert_eq!(result.items.len(), 1);
989 let context_item = &result.items[0];
990 assert!(context_item
991 .content
992 .contains("Use the file memory store for local sessions."));
993 assert!(context_item.content.contains("Memory relations:"));
994 assert!(context_item
995 .content
996 .contains("supersedes: memory://old-preference, memory://old-workflow"));
997 assert!(context_item
998 .content
999 .contains("conflicts_with: memory://legacy-default"));
1000 assert_eq!(
1001 context_item.metadata.get("memory_id"),
1002 Some(&serde_json::json!(item.id))
1003 );
1004 assert_eq!(
1005 context_item.metadata.get("memory_type"),
1006 Some(&serde_json::json!("procedural"))
1007 );
1008 assert_eq!(
1009 context_item.metadata.get("tags"),
1010 Some(&serde_json::json!(["consolidated", "conflict"]))
1011 );
1012 assert_eq!(
1013 context_item.metadata.get("supersedes"),
1014 Some(&serde_json::json!(["old-preference", "old-workflow"]))
1015 );
1016 assert_eq!(
1017 context_item.metadata.get("conflicts_with"),
1018 Some(&serde_json::json!(["legacy-default"]))
1019 );
1020 assert_eq!(
1021 context_item.token_count,
1022 (context_item.content.len() / 4).max(1)
1023 );
1024 }
1025
1026 #[test]
1027 fn test_memory_context_relevance_preserves_recall_order() {
1028 let top_match =
1029 MemoryItem::new("Run focused memory extraction tests after parser changes.")
1030 .with_importance(0.2)
1031 .with_type(MemoryType::Procedural);
1032 let generic_high_importance = MemoryItem::new("Remember general memory behavior.")
1033 .with_importance(1.0)
1034 .with_type(MemoryType::Semantic);
1035
1036 let result =
1037 memory_items_to_context_result("memory", vec![top_match, generic_high_importance]);
1038
1039 assert_eq!(result.items.len(), 2);
1040 assert!(
1041 result.items[0].relevance > result.items[1].relevance,
1042 "search recall order should remain a strong memory context ranking signal"
1043 );
1044 }
1045
1046 #[tokio::test]
1047 async fn test_memory_context_provider_does_not_mechanically_store_turns() {
1048 let memory = AgentMemory::new(Arc::new(InMemoryStore::new()));
1049 let provider = MemoryContextProvider::new(memory.clone());
1050
1051 provider
1052 .on_turn_complete("session-1", "remember nothing", "ok")
1053 .await
1054 .unwrap();
1055
1056 assert_eq!(memory.stats().await.unwrap().long_term_count, 0);
1057 }
1058}