Skip to main content

remem/memory/
edge.rs

1use anyhow::{Context, Result};
2use rusqlite::{params, Connection, OptionalExtension};
3use serde::Serialize;
4
5#[derive(Debug, Clone, Copy, PartialEq, Eq)]
6pub enum MemoryEdgeType {
7    Supersedes,
8    Duplicates,
9    Conflicts,
10    DerivedFrom,
11    MergedInto,
12    SplitFrom,
13}
14
15impl MemoryEdgeType {
16    pub const fn as_str(self) -> &'static str {
17        match self {
18            Self::Supersedes => "supersedes",
19            Self::Duplicates => "duplicates",
20            Self::Conflicts => "conflicts",
21            Self::DerivedFrom => "derived_from",
22            Self::MergedInto => "merged_into",
23            Self::SplitFrom => "split_from",
24        }
25    }
26}
27
28#[derive(Debug, Clone, Copy, Default)]
29pub struct MemoryEdgeWriteContext<'a> {
30    pub state_key_id: Option<i64>,
31    pub source_candidate_id: Option<i64>,
32    pub evidence_event_ids: &'a [i64],
33    pub source_operation_id: Option<i64>,
34    pub confidence: Option<f64>,
35    pub reason: Option<&'a str>,
36}
37
38#[derive(Debug, Clone, PartialEq)]
39pub struct MemoryEdgeInput<'a> {
40    pub edge_type: MemoryEdgeType,
41    pub from_memory_id: Option<i64>,
42    pub to_memory_id: Option<i64>,
43    pub state_key_id: Option<i64>,
44    pub source_candidate_id: Option<i64>,
45    pub evidence_event_ids: &'a [i64],
46    pub source_operation_id: Option<i64>,
47    pub confidence: Option<f64>,
48    pub reason: Option<&'a str>,
49}
50
51pub fn insert_memory_edge(conn: &Connection, input: &MemoryEdgeInput<'_>) -> Result<i64> {
52    let now = chrono::Utc::now().timestamp();
53    let evidence_event_ids = if input.evidence_event_ids.is_empty() {
54        None
55    } else {
56        Some(
57            serde_json::to_string(input.evidence_event_ids)
58                .context("serialize memory edge evidence event ids")?,
59        )
60    };
61    conn.execute(
62        "INSERT INTO memory_edges
63         (edge_type, from_memory_id, to_memory_id, state_key_id, source_candidate_id,
64          evidence_event_ids, source_operation_id, confidence, reason, created_at_epoch)
65         VALUES (?1, ?2, ?3, ?4, ?5, ?6, ?7, ?8, ?9, ?10)",
66        params![
67            input.edge_type.as_str(),
68            input.from_memory_id,
69            input.to_memory_id,
70            input.state_key_id,
71            input.source_candidate_id,
72            evidence_event_ids.as_deref(),
73            input.source_operation_id,
74            input.confidence,
75            input.reason,
76            now
77        ],
78    )
79    .context("insert memory edge")?;
80    Ok(conn.last_insert_rowid())
81}
82
83pub fn insert_replacement_edges(
84    conn: &Connection,
85    edge_type: MemoryEdgeType,
86    from_memory_ids: &[i64],
87    to_memory_id: i64,
88    context: MemoryEdgeWriteContext<'_>,
89) -> Result<usize> {
90    let state_key_id = context
91        .state_key_id
92        .or(memory_state_key_id(conn, to_memory_id)?);
93    let mut seen = std::collections::HashSet::with_capacity(from_memory_ids.len());
94    let mut inserted = 0usize;
95    for from_memory_id in from_memory_ids
96        .iter()
97        .copied()
98        .filter(|id| *id != to_memory_id && seen.insert(*id))
99    {
100        insert_memory_edge(
101            conn,
102            &MemoryEdgeInput {
103                edge_type,
104                from_memory_id: Some(from_memory_id),
105                to_memory_id: Some(to_memory_id),
106                state_key_id,
107                source_candidate_id: context.source_candidate_id,
108                evidence_event_ids: context.evidence_event_ids,
109                source_operation_id: context.source_operation_id,
110                confidence: context.confidence,
111                reason: context.reason,
112            },
113        )?;
114        inserted += 1;
115    }
116    Ok(inserted)
117}
118
119pub fn insert_supersedes_edges(
120    conn: &Connection,
121    from_memory_ids: &[i64],
122    to_memory_id: i64,
123    context: MemoryEdgeWriteContext<'_>,
124) -> Result<usize> {
125    insert_replacement_edges(
126        conn,
127        MemoryEdgeType::Supersedes,
128        from_memory_ids,
129        to_memory_id,
130        context,
131    )
132}
133
134pub fn insert_merged_into_edges(
135    conn: &Connection,
136    from_memory_ids: &[i64],
137    to_memory_id: i64,
138    context: MemoryEdgeWriteContext<'_>,
139) -> Result<usize> {
140    insert_replacement_edges(
141        conn,
142        MemoryEdgeType::MergedInto,
143        from_memory_ids,
144        to_memory_id,
145        context,
146    )
147}
148
149pub fn insert_conflicts_edges(
150    conn: &Connection,
151    from_memory_ids: &[i64],
152    to_memory_id: i64,
153    context: MemoryEdgeWriteContext<'_>,
154) -> Result<usize> {
155    insert_replacement_edges(
156        conn,
157        MemoryEdgeType::Conflicts,
158        from_memory_ids,
159        to_memory_id,
160        context,
161    )
162}
163
164pub fn insert_pairwise_conflict_edges(
165    conn: &Connection,
166    memory_ids: &[i64],
167    context: MemoryEdgeWriteContext<'_>,
168) -> Result<usize> {
169    let mut ids = memory_ids.to_vec();
170    ids.sort_unstable();
171    ids.dedup();
172
173    let mut inserted = 0usize;
174    for (idx, from_memory_id) in ids.iter().copied().enumerate() {
175        for to_memory_id in ids.iter().copied().skip(idx + 1) {
176            for state_key_id in
177                conflict_edge_state_keys(conn, from_memory_id, to_memory_id, context.state_key_id)?
178            {
179                insert_memory_edge(
180                    conn,
181                    &MemoryEdgeInput {
182                        edge_type: MemoryEdgeType::Conflicts,
183                        from_memory_id: Some(from_memory_id),
184                        to_memory_id: Some(to_memory_id),
185                        state_key_id,
186                        source_candidate_id: context.source_candidate_id,
187                        evidence_event_ids: context.evidence_event_ids,
188                        source_operation_id: context.source_operation_id,
189                        confidence: context.confidence,
190                        reason: context.reason,
191                    },
192                )?;
193                inserted += 1;
194            }
195        }
196    }
197    Ok(inserted)
198}
199
200fn conflict_edge_state_keys(
201    conn: &Connection,
202    from_memory_id: i64,
203    to_memory_id: i64,
204    explicit_state_key_id: Option<i64>,
205) -> Result<Vec<Option<i64>>> {
206    if explicit_state_key_id.is_some() {
207        return Ok(vec![explicit_state_key_id]);
208    }
209    let from_state_key_id = memory_state_key_id(conn, from_memory_id)?;
210    let to_state_key_id = memory_state_key_id(conn, to_memory_id)?;
211    if from_state_key_id == to_state_key_id {
212        return Ok(vec![from_state_key_id]);
213    }
214    let mut state_key_ids = Vec::new();
215    if from_state_key_id.is_some() {
216        state_key_ids.push(from_state_key_id);
217    }
218    if to_state_key_id.is_some() {
219        state_key_ids.push(to_state_key_id);
220    }
221    if state_key_ids.is_empty() {
222        state_key_ids.push(None);
223    }
224    Ok(state_key_ids)
225}
226
227fn memory_state_key_id(conn: &Connection, memory_id: i64) -> Result<Option<i64>> {
228    Ok(conn
229        .query_row(
230            "SELECT state_key_id FROM memories WHERE id = ?1",
231            [memory_id],
232            |row| row.get::<_, Option<i64>>(0),
233        )
234        .optional()
235        .with_context(|| format!("load state_key_id for memory edge target id={memory_id}"))?
236        .flatten())
237}
238
239#[derive(Debug, Clone, Serialize, PartialEq)]
240pub struct MemoryEdgeSummary {
241    pub incoming_count: usize,
242    pub outgoing_count: usize,
243    #[serde(skip_serializing_if = "Vec::is_empty")]
244    pub incoming: Vec<MemoryEdgeReference>,
245    #[serde(skip_serializing_if = "Vec::is_empty")]
246    pub outgoing: Vec<MemoryEdgeReference>,
247}
248
249impl MemoryEdgeSummary {
250    pub fn has_edges(&self) -> bool {
251        self.incoming_count > 0 || self.outgoing_count > 0
252    }
253}
254
255#[derive(Debug, Clone, Serialize, PartialEq)]
256pub struct MemoryEdgeReference {
257    pub id: i64,
258    pub edge_type: String,
259    pub from_memory_id: Option<i64>,
260    pub to_memory_id: Option<i64>,
261    #[serde(skip_serializing_if = "Option::is_none")]
262    pub state_key_id: Option<i64>,
263    #[serde(skip_serializing_if = "Option::is_none")]
264    pub source_candidate_id: Option<i64>,
265    #[serde(skip_serializing_if = "Vec::is_empty")]
266    pub evidence_event_ids: Vec<i64>,
267    #[serde(skip_serializing_if = "Option::is_none")]
268    pub source_operation_id: Option<i64>,
269    #[serde(skip_serializing_if = "Option::is_none")]
270    pub confidence: Option<f64>,
271    #[serde(skip_serializing_if = "Option::is_none")]
272    pub reason: Option<String>,
273    pub created_at_epoch: i64,
274}
275
276pub fn load_memory_edge_summary(conn: &Connection, memory_id: i64) -> Result<MemoryEdgeSummary> {
277    let incoming_count = count_edges(conn, "to_memory_id", memory_id)?;
278    let outgoing_count = count_edges(conn, "from_memory_id", memory_id)?;
279    Ok(MemoryEdgeSummary {
280        incoming_count,
281        outgoing_count,
282        incoming: load_edge_refs(conn, "to_memory_id", memory_id)?,
283        outgoing: load_edge_refs(conn, "from_memory_id", memory_id)?,
284    })
285}
286
287fn count_edges(conn: &Connection, column: &str, memory_id: i64) -> Result<usize> {
288    let sql = format!("SELECT COUNT(*) FROM memory_edges WHERE {column} = ?1");
289    let count: i64 = conn.query_row(&sql, [memory_id], |row| row.get(0))?;
290    Ok(count as usize)
291}
292
293fn load_edge_refs(
294    conn: &Connection,
295    column: &str,
296    memory_id: i64,
297) -> Result<Vec<MemoryEdgeReference>> {
298    let sql = format!(
299        "SELECT id, edge_type, from_memory_id, to_memory_id, state_key_id,
300                source_candidate_id, evidence_event_ids, source_operation_id,
301                confidence, reason, created_at_epoch
302         FROM memory_edges
303         WHERE {column} = ?1
304         ORDER BY created_at_epoch DESC, id DESC
305         LIMIT 25"
306    );
307    let mut stmt = conn.prepare(&sql)?;
308    let rows = stmt.query_map([memory_id], |row| {
309        let evidence_json: Option<String> = row.get(6)?;
310        let evidence_event_ids = match evidence_json {
311            Some(json) => serde_json::from_str::<Vec<i64>>(&json).map_err(|err| {
312                rusqlite::Error::FromSqlConversionFailure(
313                    6,
314                    rusqlite::types::Type::Text,
315                    Box::new(err),
316                )
317            })?,
318            None => Vec::new(),
319        };
320        Ok(MemoryEdgeReference {
321            id: row.get(0)?,
322            edge_type: row.get(1)?,
323            from_memory_id: row.get(2)?,
324            to_memory_id: row.get(3)?,
325            state_key_id: row.get(4)?,
326            source_candidate_id: row.get(5)?,
327            evidence_event_ids,
328            source_operation_id: row.get(7)?,
329            confidence: row.get(8)?,
330            reason: row.get(9)?,
331            created_at_epoch: row.get(10)?,
332        })
333    })?;
334    crate::db::query::collect_rows(rows).context("load memory edge references")
335}