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}