Skip to main content

rto_graph/
store.rs

1//! SQLite-backed graph store.
2
3use std::path::Path;
4
5use rusqlite::{Connection, OptionalExtension, params};
6
7use crate::migrations;
8use crate::model::{Direction, Edge, EdgeKind, FactSet, Node, NodeKind, Span};
9use crate::provenance::Provenance;
10
11/// Errors raised by the store.
12#[derive(Debug, thiserror::Error)]
13pub enum StoreError {
14    /// Underlying `SQLite` failure.
15    #[error("sqlite error: {0}")]
16    Sqlite(#[from] rusqlite::Error),
17    /// A node's `meta` could not be (de)serialized as JSON.
18    #[error("json error: {0}")]
19    Json(#[from] serde_json::Error),
20    /// An edge referenced a node key that does not exist in the store.
21    #[error("unknown node key: {0}")]
22    UnknownNode(String),
23    /// An edge violated the provenance/confidence invariant.
24    #[error("invalid edge: {0}")]
25    InvalidEdge(String),
26    /// A stored value could not be interpreted (database corruption).
27    #[error("corrupt store: {0}")]
28    Corrupt(String),
29}
30
31/// Qualified node columns for `SELECT`s that alias the `nodes` table as `n`.
32const NODE_COLS: &str =
33    "n.key, n.kind, n.name, n.path, n.lang, n.blob_hash, n.span_start, n.span_end, n.meta";
34
35/// `SELECT` prefix that yields an [`Edge`] row (endpoints resolved back to keys).
36const EDGE_SELECT: &str = "SELECT ns.key AS src, nd.key AS dst, e.kind, e.provenance, \
37     e.confidence, e.src_ref \
38     FROM edges e JOIN nodes ns ON ns.id = e.src JOIN nodes nd ON nd.id = e.dst";
39
40/// A Roteiro graph store backed by a single `SQLite` database.
41pub struct Store {
42    conn: Connection,
43}
44
45impl Store {
46    /// Open (creating if absent) a store at `path` and apply pending migrations.
47    ///
48    /// # Errors
49    /// Returns [`StoreError::Sqlite`] if the database cannot be opened or a
50    /// migration fails.
51    pub fn open(path: &Path) -> Result<Self, StoreError> {
52        let conn = Connection::open(path)?;
53        Self::from_conn(conn)
54    }
55
56    /// Open an in-memory store (tests, previews).
57    ///
58    /// # Errors
59    /// Returns [`StoreError::Sqlite`] if a migration fails.
60    pub fn open_in_memory() -> Result<Self, StoreError> {
61        let conn = Connection::open_in_memory()?;
62        Self::from_conn(conn)
63    }
64
65    fn from_conn(mut conn: Connection) -> Result<Self, StoreError> {
66        conn.execute_batch("PRAGMA foreign_keys = ON;")?;
67        migrations::apply(&mut conn)?;
68        Ok(Self { conn })
69    }
70
71    /// The schema version this store has been migrated to.
72    ///
73    /// # Errors
74    /// Returns [`StoreError::Sqlite`] on query failure.
75    pub fn schema_version(&self) -> Result<u32, StoreError> {
76        let v: i64 = self.conn.query_row(
77            "SELECT COALESCE(MAX(version), 0) FROM schema_migrations",
78            [],
79            |r| r.get(0),
80        )?;
81        Ok(u32::try_from(v).unwrap_or(0))
82    }
83
84    /// Number of nodes currently in the store.
85    ///
86    /// # Errors
87    /// Returns [`StoreError::Sqlite`] on query failure.
88    pub fn node_count(&self) -> Result<u64, StoreError> {
89        let n: i64 = self
90            .conn
91            .query_row("SELECT COUNT(*) FROM nodes", [], |r| r.get(0))?;
92        Ok(u64::try_from(n).unwrap_or(0))
93    }
94
95    /// Number of edges currently in the store.
96    ///
97    /// # Errors
98    /// Returns [`StoreError::Sqlite`] on query failure.
99    pub fn edge_count(&self) -> Result<u64, StoreError> {
100        let n: i64 = self
101            .conn
102            .query_row("SELECT COUNT(*) FROM edges", [], |r| r.get(0))?;
103        Ok(u64::try_from(n).unwrap_or(0))
104    }
105
106    /// Insert or update a node, keyed by its natural [`Node::key`].
107    ///
108    /// # Errors
109    /// Returns [`StoreError::Json`] if `meta` cannot be serialized, or
110    /// [`StoreError::Sqlite`] on write failure.
111    pub fn upsert_node(&self, node: &Node) -> Result<(), StoreError> {
112        upsert_node(&self.conn, node)
113    }
114
115    /// Insert an edge. Both endpoints must already resolve to nodes.
116    ///
117    /// # Errors
118    /// Returns [`StoreError::InvalidEdge`] if the provenance/confidence
119    /// invariant is violated, [`StoreError::UnknownNode`] if an endpoint key is
120    /// absent, or [`StoreError::Sqlite`] on write failure.
121    pub fn insert_edge(&self, edge: &Edge) -> Result<(), StoreError> {
122        insert_edge(&self.conn, edge)
123    }
124
125    /// Apply a fact set atomically: all nodes are upserted, then all edges are
126    /// inserted, in a single transaction. On any error nothing is committed.
127    ///
128    /// # Errors
129    /// Returns the first error encountered (see [`Store::upsert_node`] and
130    /// [`Store::insert_edge`]); the transaction is rolled back.
131    pub fn apply_factset(&mut self, facts: &FactSet) -> Result<(), StoreError> {
132        let tx = self.conn.transaction()?;
133        for node in &facts.nodes {
134            upsert_node(&tx, node)?;
135        }
136        for edge in &facts.edges {
137            insert_edge(&tx, edge)?;
138        }
139        tx.commit()?;
140        Ok(())
141    }
142
143    /// Fetch a node by its natural key.
144    ///
145    /// # Errors
146    /// Returns [`StoreError::Sqlite`], [`StoreError::Json`], or
147    /// [`StoreError::Corrupt`] if a stored value cannot be decoded.
148    pub fn get_node(&self, key: &str) -> Result<Option<Node>, StoreError> {
149        let sql = format!("SELECT {NODE_COLS} FROM nodes n WHERE n.key = ?1");
150        let mut stmt = self.conn.prepare(&sql)?;
151        let mut rows = stmt.query([key])?;
152        match rows.next()? {
153            Some(row) => Ok(Some(row_to_node(row)?)),
154            None => Ok(None),
155        }
156    }
157
158    /// All nodes of a given kind.
159    ///
160    /// # Errors
161    /// Returns [`StoreError::Sqlite`], [`StoreError::Json`], or
162    /// [`StoreError::Corrupt`] on decode failure.
163    pub fn nodes_by_kind(&self, kind: &NodeKind) -> Result<Vec<Node>, StoreError> {
164        let sql = format!("SELECT {NODE_COLS} FROM nodes n WHERE n.kind = ?1 ORDER BY n.key");
165        let mut stmt = self.conn.prepare(&sql)?;
166        let mut rows = stmt.query([kind.as_str()])?;
167        collect_nodes(&mut rows)
168    }
169
170    /// Edges whose source is the node with the given key.
171    ///
172    /// # Errors
173    /// Returns [`StoreError::Sqlite`] or [`StoreError::Corrupt`] on failure.
174    pub fn edges_from(&self, key: &str) -> Result<Vec<Edge>, StoreError> {
175        let sql = format!("{EDGE_SELECT} WHERE ns.key = ?1 ORDER BY e.id");
176        let mut stmt = self.conn.prepare(&sql)?;
177        let mut rows = stmt.query([key])?;
178        collect_edges(&mut rows)
179    }
180
181    /// Edges whose destination is the node with the given key.
182    ///
183    /// # Errors
184    /// Returns [`StoreError::Sqlite`] or [`StoreError::Corrupt`] on failure.
185    pub fn edges_to(&self, key: &str) -> Result<Vec<Edge>, StoreError> {
186        let sql = format!("{EDGE_SELECT} WHERE nd.key = ?1 ORDER BY e.id");
187        let mut stmt = self.conn.prepare(&sql)?;
188        let mut rows = stmt.query([key])?;
189        collect_edges(&mut rows)
190    }
191
192    /// All edges with the given provenance.
193    ///
194    /// # Errors
195    /// Returns [`StoreError::Sqlite`] or [`StoreError::Corrupt`] on failure.
196    pub fn edges_by_provenance(&self, provenance: Provenance) -> Result<Vec<Edge>, StoreError> {
197        let sql = format!("{EDGE_SELECT} WHERE e.provenance = ?1 ORDER BY e.id");
198        let mut stmt = self.conn.prepare(&sql)?;
199        let mut rows = stmt.query([provenance.as_str()])?;
200        collect_edges(&mut rows)
201    }
202
203    /// Neighbouring nodes reachable from `key` in the given direction. Returns
204    /// an empty vector if the node does not exist.
205    ///
206    /// # Errors
207    /// Returns [`StoreError::Sqlite`], [`StoreError::Json`], or
208    /// [`StoreError::Corrupt`] on failure.
209    pub fn neighbors(&self, key: &str, dir: Direction) -> Result<Vec<Node>, StoreError> {
210        let out = format!(
211            "SELECT {NODE_COLS} FROM nodes n JOIN edges e ON n.id = e.dst \
212             JOIN nodes s ON s.id = e.src WHERE s.key = ?1"
213        );
214        let inc = format!(
215            "SELECT {NODE_COLS} FROM nodes n JOIN edges e ON n.id = e.src \
216             JOIN nodes d ON d.id = e.dst WHERE d.key = ?1"
217        );
218        // Order by output column 1 (the node key) so results are deterministic
219        // across SQLite versions/plans. Positional ordering avoids both the
220        // ambiguity of a bare `key` (present in every joined table) and the fact
221        // that a table-qualified name cannot be used after the `Both` UNION.
222        let sql = match dir {
223            Direction::Outgoing => format!("{out} ORDER BY 1"),
224            Direction::Incoming => format!("{inc} ORDER BY 1"),
225            Direction::Both => format!("{out} UNION {inc} ORDER BY 1"),
226        };
227        let mut stmt = self.conn.prepare(&sql)?;
228        let mut rows = stmt.query([key])?;
229        collect_nodes(&mut rows)
230    }
231}
232
233// --- Free helpers operating on a `Connection` (a `Transaction` derefs to one) ---
234
235fn node_row_id(conn: &Connection, key: &str) -> rusqlite::Result<Option<i64>> {
236    conn.query_row("SELECT id FROM nodes WHERE key = ?1", [key], |r| r.get(0))
237        .optional()
238}
239
240fn upsert_node(conn: &Connection, node: &Node) -> Result<(), StoreError> {
241    let meta = serde_json::to_string(&node.meta)?;
242    let (span_start, span_end) = match node.span {
243        Some(s) => (Some(i64::from(s.start)), Some(i64::from(s.end))),
244        None => (None, None),
245    };
246    conn.execute(
247        "INSERT INTO nodes (key, kind, name, path, lang, blob_hash, span_start, span_end, meta)
248         VALUES (?1, ?2, ?3, ?4, ?5, ?6, ?7, ?8, ?9)
249         ON CONFLICT(key) DO UPDATE SET
250             kind = excluded.kind, name = excluded.name, path = excluded.path,
251             lang = excluded.lang, blob_hash = excluded.blob_hash,
252             span_start = excluded.span_start, span_end = excluded.span_end,
253             meta = excluded.meta",
254        params![
255            node.key,
256            node.kind.as_str(),
257            node.name,
258            node.path,
259            node.lang,
260            node.blob_hash,
261            span_start,
262            span_end,
263            meta,
264        ],
265    )?;
266    Ok(())
267}
268
269fn insert_edge(conn: &Connection, edge: &Edge) -> Result<(), StoreError> {
270    if !edge.is_valid() {
271        return Err(StoreError::InvalidEdge(format!(
272            "confidence must be present iff provenance is inferred (src={}, dst={})",
273            edge.src, edge.dst
274        )));
275    }
276    let src_id =
277        node_row_id(conn, &edge.src)?.ok_or_else(|| StoreError::UnknownNode(edge.src.clone()))?;
278    let dst_id =
279        node_row_id(conn, &edge.dst)?.ok_or_else(|| StoreError::UnknownNode(edge.dst.clone()))?;
280    conn.execute(
281        "INSERT INTO edges (src, dst, kind, provenance, confidence, src_ref)
282         VALUES (?1, ?2, ?3, ?4, ?5, ?6)",
283        params![
284            src_id,
285            dst_id,
286            edge.kind.as_str(),
287            edge.provenance.as_str(),
288            edge.confidence,
289            edge.src_ref,
290        ],
291    )?;
292    Ok(())
293}
294
295fn collect_nodes(rows: &mut rusqlite::Rows) -> Result<Vec<Node>, StoreError> {
296    let mut out = Vec::new();
297    while let Some(row) = rows.next()? {
298        out.push(row_to_node(row)?);
299    }
300    Ok(out)
301}
302
303fn collect_edges(rows: &mut rusqlite::Rows) -> Result<Vec<Edge>, StoreError> {
304    let mut out = Vec::new();
305    while let Some(row) = rows.next()? {
306        out.push(row_to_edge(row)?);
307    }
308    Ok(out)
309}
310
311fn row_to_node(row: &rusqlite::Row) -> Result<Node, StoreError> {
312    let kind: String = row.get("kind")?;
313    let span_start: Option<i64> = row.get("span_start")?;
314    let span_end: Option<i64> = row.get("span_end")?;
315    let span = match (span_start, span_end) {
316        (Some(s), Some(e)) => Some(Span::new(to_u32(s)?, to_u32(e)?)),
317        _ => None,
318    };
319    let meta: String = row.get("meta")?;
320    Ok(Node {
321        key: row.get("key")?,
322        kind: NodeKind::from_token(&kind),
323        name: row.get("name")?,
324        path: row.get("path")?,
325        lang: row.get("lang")?,
326        blob_hash: row.get("blob_hash")?,
327        span,
328        meta: serde_json::from_str(&meta)?,
329    })
330}
331
332fn row_to_edge(row: &rusqlite::Row) -> Result<Edge, StoreError> {
333    let kind: String = row.get("kind")?;
334    let provenance: String = row.get("provenance")?;
335    let provenance = Provenance::from_token(&provenance)
336        .ok_or_else(|| StoreError::Corrupt(format!("unknown provenance: {provenance}")))?;
337    Ok(Edge {
338        src: row.get("src")?,
339        dst: row.get("dst")?,
340        kind: EdgeKind::from_token(&kind),
341        provenance,
342        confidence: row.get("confidence")?,
343        src_ref: row.get("src_ref")?,
344    })
345}
346
347fn to_u32(v: i64) -> Result<u32, StoreError> {
348    u32::try_from(v).map_err(|_| StoreError::Corrupt(format!("span offset out of range: {v}")))
349}
350
351#[cfg(test)]
352mod tests {
353    use super::Store;
354    use crate::model::{Direction, Edge, EdgeKind, FactSet, Node, NodeKind, Span};
355    use crate::provenance::Provenance;
356
357    fn sample_node(key: &str) -> Node {
358        Node {
359            key: key.to_owned(),
360            kind: NodeKind::Fn,
361            name: "sample".to_owned(),
362            path: Some("src/lib.rs".to_owned()),
363            lang: Some("rust".to_owned()),
364            blob_hash: Some("deadbeef".to_owned()),
365            span: Some(Span::new(10, 42)),
366            meta: serde_json::json!({"vis": "pub"}),
367        }
368    }
369
370    #[test]
371    fn open_in_memory_applies_schema() {
372        let store = Store::open_in_memory().expect("open");
373        assert_eq!(store.node_count().expect("count"), 0);
374        assert_eq!(store.schema_version().expect("version"), 1);
375    }
376
377    #[test]
378    fn upsert_and_get_round_trips_all_fields() {
379        let store = Store::open_in_memory().expect("open");
380        let node = sample_node("sym:rust:src/lib.rs#sample");
381        store.upsert_node(&node).expect("upsert");
382        let got = store.get_node(&node.key).expect("get").expect("present");
383        assert_eq!(got, node);
384    }
385
386    #[test]
387    fn upsert_updates_in_place() {
388        let store = Store::open_in_memory().expect("open");
389        let mut node = sample_node("k");
390        store.upsert_node(&node).expect("insert");
391        node.name = "renamed".to_owned();
392        node.kind = NodeKind::Struct;
393        store.upsert_node(&node).expect("update");
394        assert_eq!(store.node_count().expect("count"), 1);
395        let got = store.get_node("k").expect("get").expect("present");
396        assert_eq!(got.name, "renamed");
397        assert_eq!(got.kind, NodeKind::Struct);
398    }
399
400    #[test]
401    fn edge_with_unknown_endpoint_is_rejected() {
402        let store = Store::open_in_memory().expect("open");
403        store
404            .upsert_node(&Node::new("a", NodeKind::Fn, "a"))
405            .expect("a");
406        let edge = Edge::derived("a", "missing", EdgeKind::Calls);
407        let err = store.insert_edge(&edge).expect_err("should reject");
408        assert!(matches!(err, super::StoreError::UnknownNode(k) if k == "missing"));
409    }
410
411    #[test]
412    fn inferred_edge_requires_confidence() {
413        let store = Store::open_in_memory().expect("open");
414        store
415            .upsert_node(&Node::new("a", NodeKind::Fn, "a"))
416            .expect("a");
417        store
418            .upsert_node(&Node::new("b", NodeKind::Fn, "b"))
419            .expect("b");
420        // Hand-build an inferred edge with no confidence to violate the invariant.
421        let bad = Edge {
422            src: "a".to_owned(),
423            dst: "b".to_owned(),
424            kind: EdgeKind::References,
425            provenance: Provenance::Inferred,
426            confidence: None,
427            src_ref: None,
428        };
429        assert!(matches!(
430            store.insert_edge(&bad).expect_err("reject"),
431            super::StoreError::InvalidEdge(_)
432        ));
433    }
434
435    #[test]
436    fn apply_factset_is_atomic() {
437        let mut store = Store::open_in_memory().expect("open");
438        // Second edge references a missing node, so the whole set must roll back.
439        let facts = FactSet::new()
440            .with_node(Node::new("a", NodeKind::Fn, "a"))
441            .with_node(Node::new("b", NodeKind::Fn, "b"))
442            .with_edge(Edge::derived("a", "b", EdgeKind::Calls))
443            .with_edge(Edge::derived("a", "ghost", EdgeKind::Calls));
444        assert!(store.apply_factset(&facts).is_err());
445        assert_eq!(store.node_count().expect("count"), 0, "rolled back");
446        assert_eq!(store.edge_count().expect("count"), 0, "rolled back");
447    }
448
449    #[test]
450    fn neighbors_and_provenance_queries() {
451        let mut store = Store::open_in_memory().expect("open");
452        let facts = FactSet::new()
453            .with_node(Node::new("a", NodeKind::Fn, "a"))
454            .with_node(Node::new("b", NodeKind::Fn, "b"))
455            .with_node(Node::new("c", NodeKind::Fn, "c"))
456            .with_edge(Edge::derived("a", "b", EdgeKind::Calls))
457            .with_edge(Edge::inferred("a", "c", EdgeKind::References, 0.5));
458        store.apply_factset(&facts).expect("apply");
459
460        let out = store.neighbors("a", Direction::Outgoing).expect("out");
461        let mut keys: Vec<_> = out.iter().map(|n| n.key.clone()).collect();
462        keys.sort();
463        assert_eq!(keys, ["b", "c"]);
464
465        assert!(
466            store
467                .neighbors("b", Direction::Outgoing)
468                .expect("b out")
469                .is_empty()
470        );
471        assert_eq!(
472            store
473                .neighbors("b", Direction::Incoming)
474                .expect("b in")
475                .len(),
476            1
477        );
478
479        let inferred = store
480            .edges_by_provenance(Provenance::Inferred)
481            .expect("inf");
482        assert_eq!(inferred.len(), 1);
483        assert_eq!(inferred[0].confidence, Some(0.5));
484    }
485
486    #[test]
487    fn neighbors_of_absent_node_is_empty() {
488        let store = Store::open_in_memory().expect("open");
489        assert!(
490            store
491                .neighbors("nope", Direction::Both)
492                .expect("q")
493                .is_empty()
494        );
495    }
496
497    #[test]
498    fn get_missing_node_is_none() {
499        let store = Store::open_in_memory().expect("open");
500        assert!(store.get_node("absent").expect("get").is_none());
501    }
502
503    #[test]
504    fn nodes_by_kind_and_edges_to() {
505        let mut store = Store::open_in_memory().expect("open");
506        let facts = FactSet::new()
507            .with_node(Node::new("f1", NodeKind::Fn, "f1"))
508            .with_node(Node::new("f2", NodeKind::Fn, "f2"))
509            .with_node(Node::new("s1", NodeKind::Struct, "s1"))
510            .with_edge(Edge::derived("f1", "s1", EdgeKind::References))
511            .with_edge(Edge::derived("f2", "s1", EdgeKind::References));
512        store.apply_factset(&facts).expect("apply");
513
514        let fns = store.nodes_by_kind(&NodeKind::Fn).expect("fns");
515        assert_eq!(
516            fns.iter().map(|n| n.key.as_str()).collect::<Vec<_>>(),
517            ["f1", "f2"]
518        );
519        assert!(
520            store
521                .nodes_by_kind(&NodeKind::Enum)
522                .expect("enums")
523                .is_empty()
524        );
525
526        let into_s1 = store.edges_to("s1").expect("edges_to");
527        assert_eq!(into_s1.len(), 2);
528        assert!(into_s1.iter().all(|e| e.dst == "s1"));
529    }
530
531    #[test]
532    fn open_persists_across_reopen() {
533        let path =
534            std::env::temp_dir().join(format!("roteiro-open-test-{}.db", std::process::id()));
535        std::fs::remove_file(&path).ok();
536        {
537            let store = Store::open(&path).expect("open");
538            store
539                .upsert_node(&sample_node("persisted"))
540                .expect("upsert");
541        }
542        {
543            let store = Store::open(&path).expect("reopen");
544            assert_eq!(store.node_count().expect("count"), 1);
545            assert_eq!(store.schema_version().expect("version"), 1);
546            assert!(store.get_node("persisted").expect("get").is_some());
547        }
548        std::fs::remove_file(&path).expect("cleanup");
549    }
550}