Skip to main content

a3s_code_core/
memory.rs

1//! Memory and learning system for the agent.
2//!
3//! Core types (`MemoryStore`, `MemoryItem`, `MemoryType`, `RelevanceConfig`,
4//! `FileMemoryStore`, `InMemoryStore`) live in `a3s-memory`.
5//!
6//! This module owns `MemoryConfig`, `MemoryStats`, `AgentMemory` (three-tier
7//! session memory), and `MemoryContextProvider` (context injection bridge).
8
9use 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/// One durable-memory write as observed by a host integration.
18///
19/// `incoming` preserves the identity and per-turn metadata of this observation,
20/// while `stored` is the canonical item returned by the memory backend. They
21/// differ when a backend consolidates a duplicate into an existing item.
22#[derive(Debug, Clone)]
23pub struct MemoryObservation {
24    pub incoming: MemoryItem,
25    pub stored: MemoryItem,
26    pub merged: bool,
27}
28
29/// Host extension point invoked after a durable memory has been persisted.
30///
31/// Observer failures are logged but never roll back the memory write. This is
32/// intended for derived, auditable projections such as preference and workflow
33/// learning; the memory store remains the source of truth.
34#[async_trait::async_trait]
35pub trait MemoryObserver: Send + Sync {
36    async fn on_memory_stored(&self, observation: MemoryObservation) -> anyhow::Result<()>;
37}
38
39// ============================================================================
40// Configuration
41// ============================================================================
42
43/// Configuration for the agent memory system (three-tier: working/short-term/long-term)
44#[derive(Debug, Clone, Serialize, Deserialize)]
45#[serde(rename_all = "camelCase")]
46pub struct MemoryConfig {
47    /// Relevance scoring parameters
48    #[serde(default)]
49    pub relevance: RelevanceConfig,
50    /// Maximum short-term memory items (default: 100)
51    #[serde(default = "MemoryConfig::default_max_short_term")]
52    pub max_short_term: usize,
53    /// Maximum working memory items (default: 10)
54    #[serde(default = "MemoryConfig::default_max_working")]
55    pub max_working: usize,
56    /// Automatic pruning policy for long-term storage. `None` disables background pruning.
57    #[serde(default)]
58    pub prune_policy: Option<PrunePolicy>,
59    /// How often the background pruning task runs, in seconds (default: 3600).
60    #[serde(default = "MemoryConfig::default_prune_interval_secs")]
61    pub prune_interval_secs: u64,
62    /// Use an LLM after every completed, non-empty turn to judge whether the
63    /// turn contains durable memories and, when it does, distill them from the
64    /// transcript.
65    ///
66    /// Enabled by default when memory is configured. Semantic value decisions
67    /// belong to the LLM; the runtime does not use content-keyword gates.
68    #[serde(
69        default = "MemoryConfig::default_llm_extraction",
70        alias = "llm_extraction"
71    )]
72    pub llm_extraction: bool,
73    /// Maximum durable memories the LLM extractor may write per turn.
74    #[serde(default = "MemoryConfig::default_llm_extraction_max_items")]
75    pub llm_extraction_max_items: usize,
76    /// Maximum transcript characters passed into the LLM memory extractor.
77    #[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// ============================================================================
118// Memory Stats
119// ============================================================================
120
121/// Statistics for the three-tier agent memory system
122#[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// ============================================================================
130// Agent Memory (three-tier: working / short-term / long-term)
131// ============================================================================
132
133/// Three-tier agent memory: working, short-term (session), and long-term (persisted).
134#[derive(Clone)]
135pub struct AgentMemory {
136    /// Long-term memory store
137    pub(crate) store: Arc<dyn MemoryStore>,
138    /// Short-term memory (current session)
139    short_term: Arc<RwLock<VecDeque<MemoryItem>>>,
140    /// Working memory (active context)
141    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
164/// A FIFO ticket for one completed-turn extraction.
165///
166/// Registration happens before a background task is spawned, so session close
167/// can observe every accepted extraction even if the task has not been polled
168/// yet. Chaining each ticket to its predecessor preserves completed-turn order
169/// without blocking streaming callers.
170pub(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    /// Create a new agent memory system with default configuration
207    pub fn new(store: Arc<dyn MemoryStore>) -> Self {
208        Self::with_config(store, MemoryConfig::default())
209    }
210
211    /// Create a new agent memory system with custom configuration.
212    ///
213    /// If `config.prune_policy` is `Some`, a background Tokio task is spawned
214    /// that periodically calls `store.prune()` at the configured interval.
215    pub fn with_config(store: Arc<dyn MemoryStore>, config: MemoryConfig) -> Self {
216        Self::with_config_and_observers(store, config, Vec::new())
217    }
218
219    /// Create a memory system with host observers for successful durable
220    /// writes. Observers receive both the incoming observation and the
221    /// canonical stored item so duplicate consolidation remains auditable.
222    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; // skip the immediate first tick
236                        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    /// Store a memory in long-term storage and add to short-term
275    pub async fn remember(&self, item: MemoryItem) -> anyhow::Result<()> {
276        self.remember_item(item).await.map(|_| ())
277    }
278
279    /// Store a memory and return the normalized item that was sent to storage.
280    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    /// Remove a memory from long-term storage and session-local memory tiers.
313    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    /// Remember a successful pattern
321    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    /// Remember a successful pattern and return the stored memory item.
333    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    /// Remember a failure to avoid repeating
359    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    /// Remember a failed pattern and return the stored memory item.
371    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    /// Recall similar past experiences
397    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    /// Recall by tags
406    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    /// Get recent memories
415    pub async fn get_recent(&self, limit: usize) -> anyhow::Result<Vec<MemoryItem>> {
416        self.store.get_recent(limit).await
417    }
418
419    /// Add to working memory (auto-trims by relevance if over capacity)
420    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    /// Get working memory
436    pub async fn get_working(&self) -> Vec<MemoryItem> {
437        self.working.read().await.clone()
438    }
439
440    /// Clear working memory
441    pub async fn clear_working(&self) {
442        self.working.write().await.clear();
443    }
444
445    /// Get short-term memory
446    pub async fn get_short_term(&self) -> Vec<MemoryItem> {
447        self.short_term.read().await.iter().cloned().collect()
448    }
449
450    /// Clear short-term memory
451    pub async fn clear_short_term(&self) {
452        self.short_term.write().await.clear();
453    }
454
455    /// Get memory statistics
456    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    /// Get access to the underlying store
465    pub fn store(&self) -> &Arc<dyn MemoryStore> {
466        &self.store
467    }
468
469    /// Get working memory count
470    pub async fn working_count(&self) -> usize {
471        self.working.read().await.len()
472    }
473
474    /// Get short-term memory count
475    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    /// Wait until every extraction registered before this call has settled.
510    /// Returns `false` when the bounded close-time wait expires.
511    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
525// ============================================================================
526// Memory Context Provider
527// ============================================================================
528
529/// Context provider that surfaces past memories as agent context.
530pub 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        // Memory extraction is owned by the agent loop's LLM value judge.
666        // This provider only contributes recalled memories as prompt context.
667        Ok(())
668    }
669}
670
671// ============================================================================
672// Tests
673// ============================================================================
674
675#[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        // Insert one old low-importance item directly into the store.
929        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        // Calling prune on the underlying store via the public accessor works.
936        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}