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