Skip to main content

remem/memory/
facts.rs

1use anyhow::{bail, Result};
2use rusqlite::{params, Connection, OptionalExtension};
3
4#[derive(Debug, Clone, Copy, PartialEq, Eq)]
5pub enum FactPredicate {
6    FixedBy,
7    VerifiedBy,
8    Supersedes,
9    BlockedBy,
10    UsesFile,
11    UsesCommand,
12    AffectsProject,
13}
14
15impl FactPredicate {
16    pub fn db_value(self) -> &'static str {
17        match self {
18            Self::FixedBy => "fixed_by",
19            Self::VerifiedBy => "verified_by",
20            Self::Supersedes => "supersedes",
21            Self::BlockedBy => "blocked_by",
22            Self::UsesFile => "uses_file",
23            Self::UsesCommand => "uses_command",
24            Self::AffectsProject => "affects_project",
25        }
26    }
27
28    /// Parse an externally supplied predicate label (e.g. from the
29    /// extraction LLM). Same closed vocabulary as the DB encoding.
30    pub fn parse_public(raw: &str) -> Option<Self> {
31        Self::parse_db(raw.trim().to_ascii_lowercase().as_str())
32    }
33
34    fn parse_db(raw: &str) -> Option<Self> {
35        match raw {
36            "fixed_by" => Some(Self::FixedBy),
37            "verified_by" => Some(Self::VerifiedBy),
38            "supersedes" => Some(Self::Supersedes),
39            "blocked_by" => Some(Self::BlockedBy),
40            "uses_file" => Some(Self::UsesFile),
41            "uses_command" => Some(Self::UsesCommand),
42            "affects_project" => Some(Self::AffectsProject),
43            _ => None,
44        }
45    }
46}
47
48#[derive(Debug, Clone)]
49pub struct TemporalFactInput<'a> {
50    pub project: &'a str,
51    pub subject: &'a str,
52    pub predicate: FactPredicate,
53    pub object: &'a str,
54    pub valid_from_epoch: Option<i64>,
55    pub valid_to_epoch: Option<i64>,
56    pub learned_at_epoch: Option<i64>,
57    pub source_memory_id: Option<i64>,
58    pub source_observation_id: Option<i64>,
59    pub source_event_ids: &'a [i64],
60    pub confidence: f64,
61    pub supersedes_fact_id: Option<i64>,
62}
63
64#[derive(Debug, Clone, PartialEq)]
65pub struct TemporalFact {
66    pub id: i64,
67    pub project: String,
68    pub subject: String,
69    pub predicate: FactPredicate,
70    pub object: String,
71    pub valid_from_epoch: Option<i64>,
72    pub valid_to_epoch: Option<i64>,
73    pub learned_at_epoch: i64,
74    pub source_memory_id: Option<i64>,
75    pub source_observation_id: Option<i64>,
76    pub source_event_ids: Vec<i64>,
77    pub confidence: f64,
78    pub supersedes_fact_id: Option<i64>,
79    pub status: String,
80}
81
82pub(crate) fn invalidated_at_epoch_available(conn: &Connection) -> Result<bool> {
83    let exists: i64 = conn.query_row(
84        "SELECT EXISTS (
85             SELECT 1 FROM pragma_table_info('memory_facts')
86             WHERE name = 'invalidated_at_epoch'
87         )",
88        [],
89        |row| row.get(0),
90    )?;
91    Ok(exists != 0)
92}
93
94pub(crate) fn current_fact_filter_sql(alias: &str, has_invalidated_at_epoch: bool) -> String {
95    let alias = alias.trim();
96    let prefix = if alias.is_empty() {
97        String::new()
98    } else {
99        format!("{alias}.")
100    };
101    if has_invalidated_at_epoch {
102        format!("{prefix}status = 'active' AND {prefix}invalidated_at_epoch IS NULL")
103    } else {
104        format!("{prefix}status = 'active'")
105    }
106}
107
108pub(crate) fn as_of_validity_filter_sql(
109    alias: &str,
110    epoch_param_idx: usize,
111    has_invalidated_at_epoch: bool,
112) -> String {
113    let alias = alias.trim();
114    let prefix = if alias.is_empty() {
115        String::new()
116    } else {
117        format!("{alias}.")
118    };
119    if has_invalidated_at_epoch {
120        let outer_id = if alias.is_empty() {
121            "memory_facts.id".to_string()
122        } else {
123            format!("{alias}.id")
124        };
125        format!(
126            "({prefix}valid_to_epoch IS NULL OR {prefix}valid_to_epoch > ?{epoch_param_idx} \
127              OR ({prefix}invalidated_at_epoch IS NOT NULL \
128                  AND {prefix}invalidated_at_epoch > ?{epoch_param_idx} \
129                  AND NOT EXISTS (
130                      SELECT 1 FROM memory_facts AS replacement
131                      WHERE replacement.supersedes_fact_id = {outer_id}
132                        AND replacement.learned_at_epoch <= ?{epoch_param_idx}
133                  )))"
134        )
135    } else {
136        format!("({prefix}valid_to_epoch IS NULL OR {prefix}valid_to_epoch > ?{epoch_param_idx})")
137    }
138}
139
140pub fn insert_temporal_fact(conn: &mut Connection, input: &TemporalFactInput<'_>) -> Result<i64> {
141    validate_input(input)?;
142    let now = chrono::Utc::now().timestamp();
143
144    let tx = conn.transaction()?;
145    let id = insert_temporal_fact_in_current_tx(&tx, input, now)?;
146    tx.commit()?;
147    Ok(id)
148}
149
150pub(crate) fn insert_temporal_fact_in_current_tx(
151    conn: &Connection,
152    input: &TemporalFactInput<'_>,
153    now: i64,
154) -> Result<i64> {
155    validate_input(input)?;
156    let learned_at = input.learned_at_epoch.unwrap_or(now);
157    let superseded_at = input.valid_from_epoch.unwrap_or(learned_at);
158    let source_event_ids = serde_json::to_string(input.source_event_ids)?;
159
160    if let Some(old_id) = input.supersedes_fact_id {
161        let old_fact: Option<(String, Option<i64>)> = conn
162            .query_row(
163                "SELECT project, valid_from_epoch FROM memory_facts WHERE id = ?1",
164                [old_id],
165                |row| Ok((row.get(0)?, row.get(1)?)),
166            )
167            .optional()?;
168        match old_fact {
169            Some((project, old_valid_from)) if project == input.project => {
170                if let Some(old_from) = old_valid_from {
171                    if superseded_at < old_from {
172                        bail!(
173                            "cannot supersede fact {old_id}: cutoff {superseded_at} is before existing valid_from_epoch {}",
174                            old_from
175                        );
176                    }
177                }
178            }
179            Some((project, _)) => bail!(
180                "cannot supersede fact {old_id} from project '{project}' with project '{}'",
181                input.project
182            ),
183            None => bail!("cannot supersede missing memory fact {old_id}"),
184        }
185        conn.execute(
186            "UPDATE memory_facts
187             SET status = 'stale',
188                 valid_to_epoch = CASE
189                     WHEN valid_to_epoch IS NULL OR valid_to_epoch > ?1 THEN ?1
190                     ELSE valid_to_epoch
191                 END,
192                 invalidated_at_epoch = COALESCE(invalidated_at_epoch, ?2),
193                 updated_at_epoch = ?2
194             WHERE id = ?3",
195            params![superseded_at, now, old_id],
196        )?;
197    }
198
199    conn.execute(
200        "INSERT INTO memory_facts
201         (project, subject, predicate, object, valid_from_epoch, valid_to_epoch,
202          learned_at_epoch, source_memory_id, source_observation_id, source_event_ids,
203          confidence, supersedes_fact_id, status, created_at_epoch, updated_at_epoch)
204         VALUES (?1, ?2, ?3, ?4, ?5, ?6, ?7, ?8, ?9, ?10, ?11, ?12, 'active', ?13, ?13)",
205        params![
206            input.project,
207            input.subject,
208            input.predicate.db_value(),
209            input.object,
210            input.valid_from_epoch,
211            input.valid_to_epoch,
212            learned_at,
213            input.source_memory_id,
214            input.source_observation_id,
215            source_event_ids,
216            input.confidence,
217            input.supersedes_fact_id,
218            now
219        ],
220    )?;
221    let id = conn.last_insert_rowid();
222    Ok(id)
223}
224
225/// The currently valid fact for (project, subject, predicate), as
226/// `(id, object)`. Used by extraction writes to decide between no-op
227/// (same object) and supersede (contradicting object).
228pub(crate) fn find_active_fact(
229    conn: &Connection,
230    project: &str,
231    subject: &str,
232    predicate: FactPredicate,
233) -> Result<Option<(i64, String)>> {
234    let has_invalidated = invalidated_at_epoch_available(conn)?;
235    let current_filter = current_fact_filter_sql("f", has_invalidated);
236    let now = chrono::Utc::now().timestamp();
237    let row = conn
238        .query_row(
239            &format!(
240                "SELECT f.id, f.object FROM memory_facts f
241                 WHERE f.project = ?1 AND f.subject = ?2 AND f.predicate = ?3
242                   AND (f.valid_from_epoch IS NULL OR f.valid_from_epoch <= ?4)
243                   AND (f.valid_to_epoch IS NULL OR f.valid_to_epoch > ?4)
244                   AND {current_filter}
245                 ORDER BY f.id DESC
246                 LIMIT 1"
247            ),
248            params![project, subject, predicate.db_value(), now],
249            |row| Ok((row.get(0)?, row.get(1)?)),
250        )
251        .optional()?;
252    Ok(row)
253}
254
255pub fn list_current_facts(
256    conn: &Connection,
257    project: &str,
258    subject: Option<&str>,
259    predicate: Option<FactPredicate>,
260) -> Result<Vec<TemporalFact>> {
261    let now = chrono::Utc::now().timestamp();
262    query_facts(conn, project, subject, predicate, Some(now), true)
263}
264
265pub fn list_facts_as_of(
266    conn: &Connection,
267    project: &str,
268    as_of_epoch: i64,
269    subject: Option<&str>,
270    predicate: Option<FactPredicate>,
271) -> Result<Vec<TemporalFact>> {
272    query_facts(conn, project, subject, predicate, Some(as_of_epoch), false)
273}
274
275fn query_facts(
276    conn: &Connection,
277    project: &str,
278    subject: Option<&str>,
279    predicate: Option<FactPredicate>,
280    as_of_epoch: Option<i64>,
281    active_only: bool,
282) -> Result<Vec<TemporalFact>> {
283    let has_invalidated_at_epoch = invalidated_at_epoch_available(conn)?;
284    let mut conditions = vec!["project = ?1".to_string()];
285    let mut params: Vec<Box<dyn rusqlite::types::ToSql>> = vec![Box::new(project.to_string())];
286    let mut idx = 2;
287    if let Some(subject) = subject {
288        conditions.push(format!("subject = ?{idx}"));
289        params.push(Box::new(subject.to_string()));
290        idx += 1;
291    }
292    if let Some(predicate) = predicate {
293        conditions.push(format!("predicate = ?{idx}"));
294        params.push(Box::new(predicate.db_value().to_string()));
295        idx += 1;
296    }
297    if let Some(as_of_epoch) = as_of_epoch {
298        conditions.push(format!(
299            "(valid_from_epoch IS NULL OR valid_from_epoch <= ?{idx})"
300        ));
301        conditions.push(as_of_validity_filter_sql("", idx, has_invalidated_at_epoch));
302        conditions.push(format!("learned_at_epoch <= ?{idx}"));
303        if has_invalidated_at_epoch {
304            conditions.push(format!(
305                "(invalidated_at_epoch IS NULL OR invalidated_at_epoch > ?{idx})"
306            ));
307        }
308        params.push(Box::new(as_of_epoch));
309    }
310    if active_only {
311        conditions.push(current_fact_filter_sql("", has_invalidated_at_epoch));
312    }
313
314    let sql = format!(
315        "SELECT id, project, subject, predicate, object, valid_from_epoch,
316                valid_to_epoch, learned_at_epoch, source_memory_id,
317                source_observation_id, source_event_ids, confidence,
318                supersedes_fact_id, status
319         FROM memory_facts
320         WHERE {}
321         ORDER BY learned_at_epoch DESC, id DESC",
322        conditions.join(" AND ")
323    );
324    let mut stmt = conn.prepare(&sql)?;
325    let refs = crate::db::to_sql_refs(&params);
326    let rows = stmt.query_map(refs.as_slice(), map_fact_row)?;
327    crate::db::query::collect_rows(rows)
328}
329
330fn validate_input(input: &TemporalFactInput<'_>) -> Result<()> {
331    if input.project.trim().is_empty() {
332        bail!("memory fact project is required");
333    }
334    if input.subject.trim().is_empty() {
335        bail!("memory fact subject is required");
336    }
337    if input.object.trim().is_empty() {
338        bail!("memory fact object is required");
339    }
340    if !(0.0..=1.0).contains(&input.confidence) {
341        bail!("memory fact confidence out of range");
342    }
343    if let (Some(valid_from), Some(valid_to)) = (input.valid_from_epoch, input.valid_to_epoch) {
344        if valid_to < valid_from {
345            bail!("memory fact valid_to_epoch cannot be before valid_from_epoch");
346        }
347    }
348    Ok(())
349}
350
351fn map_fact_row(row: &rusqlite::Row<'_>) -> rusqlite::Result<TemporalFact> {
352    let predicate_raw: String = row.get(3)?;
353    let source_event_json: String = row.get(10)?;
354    let source_event_ids = serde_json::from_str(&source_event_json).map_err(|err| {
355        rusqlite::Error::FromSqlConversionFailure(10, rusqlite::types::Type::Text, Box::new(err))
356    })?;
357    let predicate = FactPredicate::parse_db(&predicate_raw).ok_or_else(|| {
358        rusqlite::Error::FromSqlConversionFailure(
359            3,
360            rusqlite::types::Type::Text,
361            Box::new(std::io::Error::new(
362                std::io::ErrorKind::InvalidData,
363                format!("unknown memory fact predicate: {predicate_raw}"),
364            )),
365        )
366    })?;
367    Ok(TemporalFact {
368        id: row.get(0)?,
369        project: row.get(1)?,
370        subject: row.get(2)?,
371        predicate,
372        object: row.get(4)?,
373        valid_from_epoch: row.get(5)?,
374        valid_to_epoch: row.get(6)?,
375        learned_at_epoch: row.get(7)?,
376        source_memory_id: row.get(8)?,
377        source_observation_id: row.get(9)?,
378        source_event_ids,
379        confidence: row.get(11)?,
380        supersedes_fact_id: row.get(12)?,
381        status: row.get(13)?,
382    })
383}
384
385#[cfg(test)]
386mod tests;