Skip to main content

mnemo_amp/
store.rs

1//! `MemoryStore`-conformant surface over a [`MnemoEngine`].
2//!
3//! Maps each of the 5 AMP ops onto **real** engine calls. Two of them
4//! have no 1:1 engine primitive and are implemented as thin
5//! compositions, deliberately:
6//!
7//! - **`merge`** — mnemo's `MnemoEngine::merge` is a *branch-timeline*
8//!   merge (checkpoint thread/branch), not a memory-record merge. AMP
9//!   `merge` folds N memory records into one, so this adapter composes
10//!   `storage.get_memory` → `engine.remember` (the consolidated
11//!   record, `SourceType::Consolidation`) → `engine.forget`
12//!   (`ForgetStrategy::Consolidate` on the originals). No fictitious
13//!   engine method is assumed.
14//! - **`expire`** — there is no `engine.expire`; the lifecycle path is
15//!   `expires_at` + `run_ttl_sweep`. AMP `expire` sets `expires_at` on
16//!   the targets (immediately in the past when `ttl_seconds` is unset
17//!   or `0`, else `now + ttl`) via `storage.update_memory`, then runs
18//!   `run_ttl_sweep` for the immediate case so the records are
19//!   hard-deleted and a `MemoryExpired` audit event is emitted.
20
21use std::sync::Arc;
22
23use async_trait::async_trait;
24use uuid::Uuid;
25
26use mnemo_core::model::event::EventType;
27use mnemo_core::model::memory::{MemoryType, SourceType};
28use mnemo_core::query::MnemoEngine;
29use mnemo_core::query::forget::{ForgetRequest, ForgetStrategy};
30use mnemo_core::query::recall::RecallRequest;
31use mnemo_core::query::remember::RememberRequest;
32
33use crate::approval::{Approval, ApprovalHook, AutoApprove, WriteDiff};
34use crate::error::AmpError;
35use crate::wire::{AmpEnvelope, AmpHit, AmpMemoryType, AmpOp, AmpResult};
36
37/// Default recall depth (the conformance suite's recall@5).
38pub const DEFAULT_TOP_K: usize = 5;
39
40/// The AMP `MemoryStore` surface: 5 ops, each returning an
41/// [`AmpResult`]. Transport adapters (REST / MCP / a fan-out router)
42/// call these.
43#[async_trait]
44pub trait MemoryStore: Send + Sync {
45    async fn remember(&self, env: &AmpEnvelope) -> Result<AmpResult, AmpError>;
46    async fn recall(&self, env: &AmpEnvelope) -> Result<AmpResult, AmpError>;
47    async fn forget(&self, env: &AmpEnvelope) -> Result<AmpResult, AmpError>;
48    async fn merge(&self, env: &AmpEnvelope) -> Result<AmpResult, AmpError>;
49    async fn expire(&self, env: &AmpEnvelope) -> Result<AmpResult, AmpError>;
50
51    /// Dispatch an envelope to the matching op.
52    async fn dispatch(&self, env: &AmpEnvelope) -> Result<AmpResult, AmpError> {
53        match env.op {
54            AmpOp::Remember => self.remember(env).await,
55            AmpOp::Recall => self.recall(env).await,
56            AmpOp::Forget => self.forget(env).await,
57            AmpOp::Merge => self.merge(env).await,
58            AmpOp::Expire => self.expire(env).await,
59        }
60    }
61}
62
63/// AMP store backed by a shared [`MnemoEngine`], with an optional HITL
64/// approval gate on long-term writes.
65pub struct MnemoAmpStore {
66    engine: Arc<MnemoEngine>,
67    approval: Arc<dyn ApprovalHook>,
68}
69
70impl MnemoAmpStore {
71    /// Build a store with the default [`AutoApprove`] gate (no human in
72    /// the loop).
73    pub fn new(engine: Arc<MnemoEngine>) -> Self {
74        Self {
75            engine,
76            approval: Arc::new(AutoApprove),
77        }
78    }
79
80    /// Attach a HITL diff-and-approve hook consulted before every
81    /// long-term (`semantic` / `procedural`) write.
82    pub fn with_approval_hook(mut self, hook: Arc<dyn ApprovalHook>) -> Self {
83        self.approval = hook;
84        self
85    }
86
87    fn agent<'a>(&'a self, env: &'a AmpEnvelope) -> &'a str {
88        env.agent_id
89            .as_deref()
90            .unwrap_or(&self.engine.default_agent_id)
91    }
92
93    /// Run the approval gate for a long-term write and, on approval,
94    /// emit a `Decision` audit event through the existing hash chain.
95    /// Short-term tiers approve implicitly without an event.
96    async fn gate_long_term_write(
97        &self,
98        agent_id: &str,
99        diff: WriteDiff,
100    ) -> Result<bool, AmpError> {
101        if !diff.memory_type.is_long_term() {
102            return Ok(true);
103        }
104        match self.approval.review(&diff) {
105            Approval::Approve => {
106                let rendered = diff.render();
107                let event = mnemo_core::query::event_builder::build_event(
108                    &self.engine,
109                    agent_id,
110                    EventType::Decision,
111                    serde_json::json!({
112                        "amp_approval": "approved",
113                        "hook": self.approval.name(),
114                        "memory_type": diff.memory_type.as_str(),
115                        "diff": rendered,
116                    }),
117                    &rendered,
118                    None,
119                )
120                .await;
121                if let Err(e) = self.engine.storage.insert_event(&event).await {
122                    tracing::warn!(error = %e, "failed to insert AMP approval audit event");
123                }
124                Ok(true)
125            }
126            Approval::Reject(_) => Ok(false),
127        }
128    }
129}
130
131fn to_core_type(t: AmpMemoryType) -> MemoryType {
132    match t {
133        AmpMemoryType::Episodic => MemoryType::Episodic,
134        AmpMemoryType::Semantic => MemoryType::Semantic,
135        AmpMemoryType::Procedural => MemoryType::Procedural,
136        AmpMemoryType::Working => MemoryType::Working,
137    }
138}
139
140fn from_core_type(t: MemoryType) -> AmpMemoryType {
141    match t {
142        MemoryType::Episodic => AmpMemoryType::Episodic,
143        MemoryType::Semantic => AmpMemoryType::Semantic,
144        MemoryType::Procedural => AmpMemoryType::Procedural,
145        MemoryType::Working => AmpMemoryType::Working,
146    }
147}
148
149fn parse_ids(env: &AmpEnvelope) -> Result<Vec<Uuid>, AmpError> {
150    env.memory_ids
151        .iter()
152        .map(|s| {
153            Uuid::parse_str(s).map_err(|_| AmpError::Validation(format!("invalid memory id '{s}'")))
154        })
155        .collect()
156}
157
158#[async_trait]
159impl MemoryStore for MnemoAmpStore {
160    async fn remember(&self, env: &AmpEnvelope) -> Result<AmpResult, AmpError> {
161        let content = env
162            .content
163            .clone()
164            .ok_or_else(|| AmpError::Validation("remember requires `content`".into()))?;
165        let agent_id = self.agent(env).to_string();
166
167        // HITL gate on long-term writes BEFORE the record commits.
168        let diff = WriteDiff {
169            agent_id: agent_id.clone(),
170            memory_type: env.memory_type,
171            before: None,
172            after: content.clone(),
173            tags: env.tags.clone(),
174        };
175        let long_term = env.memory_type.is_long_term();
176        let approved = self.gate_long_term_write(&agent_id, diff).await?;
177        if !approved {
178            return Ok(AmpResult::rejected(
179                AmpOp::Remember,
180                "long-term write rejected by approval hook",
181            ));
182        }
183
184        let mut req = RememberRequest::new(content);
185        req.agent_id = Some(agent_id);
186        req.memory_type = Some(to_core_type(env.memory_type));
187        if !env.tags.is_empty() {
188            req.tags = Some(env.tags.clone());
189        }
190        req.metadata = env.metadata.clone();
191        if let Some(ttl) = env.ttl_seconds {
192            req.ttl_seconds = Some(ttl);
193        }
194
195        let resp = self.engine.remember(req).await?;
196        let mut out = AmpResult::ok(AmpOp::Remember);
197        out.ids = vec![resp.id.to_string()];
198        if long_term {
199            out.approved = Some(true);
200        }
201        Ok(out)
202    }
203
204    async fn recall(&self, env: &AmpEnvelope) -> Result<AmpResult, AmpError> {
205        let query = env
206            .query
207            .clone()
208            .ok_or_else(|| AmpError::Validation("recall requires `query`".into()))?;
209        let mut req = RecallRequest::new(query);
210        req.agent_id = Some(self.agent(env).to_string());
211        req.limit = Some(env.top_k.unwrap_or(DEFAULT_TOP_K));
212        req.memory_type = Some(to_core_type(env.memory_type));
213        if !env.tags.is_empty() {
214            req.tags = Some(env.tags.clone());
215        }
216        req.strategy = Some("auto".to_string());
217
218        let resp = self.engine.recall(req).await?;
219        let mut out = AmpResult::ok(AmpOp::Recall);
220        out.hits = resp
221            .memories
222            .into_iter()
223            .map(|m| AmpHit {
224                id: m.id.to_string(),
225                content: m.content,
226                memory_type: from_core_type(m.memory_type),
227                score: m.score,
228                tags: m.tags,
229            })
230            .collect();
231        Ok(out)
232    }
233
234    async fn forget(&self, env: &AmpEnvelope) -> Result<AmpResult, AmpError> {
235        let ids = parse_ids(env)?;
236        if ids.is_empty() {
237            return Err(AmpError::Validation("forget requires `memory_ids`".into()));
238        }
239        let mut req = ForgetRequest::new(ids);
240        req.agent_id = Some(self.agent(env).to_string());
241        req.strategy = Some(ForgetStrategy::SoftDelete);
242        let resp = self.engine.forget(req).await?;
243        let mut out = AmpResult::ok(AmpOp::Forget);
244        out.ids = resp.forgotten.iter().map(|id| id.to_string()).collect();
245        if !resp.errors.is_empty() {
246            out.detail = format!("{} id(s) failed", resp.errors.len());
247        }
248        Ok(out)
249    }
250
251    async fn merge(&self, env: &AmpEnvelope) -> Result<AmpResult, AmpError> {
252        // Thin composition over remember + forget — NOT engine.merge
253        // (which is a branch-timeline merge). Fold the source records'
254        // content into one consolidated record, then retire the
255        // originals with the Consolidate strategy so the chain stays
256        // auditable.
257        let ids = parse_ids(env)?;
258        if ids.len() < 2 {
259            return Err(AmpError::Validation(
260                "merge requires at least two `memory_ids`".into(),
261            ));
262        }
263        let agent_id = self.agent(env).to_string();
264
265        let mut contents = Vec::with_capacity(ids.len());
266        let mut tag_set: Vec<String> = env.tags.clone();
267        for id in &ids {
268            let rec = self
269                .engine
270                .storage
271                .get_memory(*id)
272                .await?
273                .ok_or_else(|| AmpError::NotFound(format!("memory {id} not found")))?;
274            contents.push(rec.content);
275            for t in rec.tags {
276                if !tag_set.contains(&t) {
277                    tag_set.push(t);
278                }
279            }
280        }
281        let merged_content = format!("[AMP merged from {}] {}", ids.len(), contents.join(" | "));
282
283        // HITL gate on the long-term merged write.
284        let diff = WriteDiff {
285            agent_id: agent_id.clone(),
286            memory_type: env.memory_type,
287            before: Some(contents_preview(&contents)),
288            after: merged_content.clone(),
289            tags: tag_set.clone(),
290        };
291        let approved = self.gate_long_term_write(&agent_id, diff).await?;
292        if !approved {
293            return Ok(AmpResult::rejected(
294                AmpOp::Merge,
295                "long-term merge rejected by approval hook",
296            ));
297        }
298
299        let mut req = RememberRequest::new(merged_content);
300        req.agent_id = Some(agent_id.clone());
301        req.memory_type = Some(to_core_type(env.memory_type));
302        req.source_type = Some(SourceType::Consolidation);
303        if !tag_set.is_empty() {
304            req.tags = Some(tag_set);
305        }
306        req.metadata = Some(serde_json::json!({
307            "amp_merged_from": ids.iter().map(|i| i.to_string()).collect::<Vec<_>>(),
308        }));
309        let merged = self.engine.remember(req).await?;
310
311        // Retire the originals.
312        let mut forget_req = ForgetRequest::new(ids);
313        forget_req.agent_id = Some(agent_id);
314        forget_req.strategy = Some(ForgetStrategy::Consolidate);
315        let forgotten = self.engine.forget(forget_req).await?;
316
317        let mut out = AmpResult::ok(AmpOp::Merge);
318        out.ids = vec![merged.id.to_string()];
319        if env.memory_type.is_long_term() {
320            out.approved = Some(true);
321        }
322        if !forgotten.errors.is_empty() {
323            out.detail = format!("{} original(s) failed to retire", forgotten.errors.len());
324        }
325        Ok(out)
326    }
327
328    async fn expire(&self, env: &AmpEnvelope) -> Result<AmpResult, AmpError> {
329        // Thin composition over expires_at + run_ttl_sweep — there is
330        // no engine.expire primitive. Immediate when ttl is unset/0.
331        let ids = parse_ids(env)?;
332        if ids.is_empty() {
333            return Err(AmpError::Validation("expire requires `memory_ids`".into()));
334        }
335        let ttl = env.ttl_seconds.unwrap_or(0);
336        let immediate = ttl == 0;
337        let now = chrono::Utc::now();
338        let expires_at = if immediate {
339            now - chrono::Duration::seconds(1)
340        } else {
341            now + chrono::Duration::seconds(ttl as i64)
342        };
343        let expires_str = expires_at.to_rfc3339();
344
345        let mut touched = Vec::new();
346        for id in &ids {
347            match self.engine.storage.get_memory(*id).await? {
348                Some(mut rec) => {
349                    rec.expires_at = Some(expires_str.clone());
350                    rec.updated_at = now.to_rfc3339();
351                    self.engine.storage.update_memory(&rec).await?;
352                    touched.push(id.to_string());
353                }
354                None => {
355                    return Err(AmpError::NotFound(format!("memory {id} not found")));
356                }
357            }
358        }
359
360        if immediate {
361            // Hard-delete the now-expired records + emit MemoryExpired
362            // audit events through the existing lifecycle path.
363            self.engine.run_ttl_sweep().await?;
364        }
365
366        let mut out = AmpResult::ok(AmpOp::Expire);
367        out.ids = touched;
368        if !immediate {
369            out.detail = format!("scheduled to expire at {expires_str}");
370        }
371        Ok(out)
372    }
373}
374
375/// A bounded, deterministic preview of the source contents folded by a
376/// merge, for the approval diff.
377fn contents_preview(contents: &[String]) -> String {
378    contents
379        .iter()
380        .map(|c| {
381            if c.len() > 80 {
382                format!("{}…", &c[..80])
383            } else {
384                c.clone()
385            }
386        })
387        .collect::<Vec<_>>()
388        .join(" | ")
389}