1use 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#[derive(Debug, thiserror::Error)]
13pub enum StoreError {
14 #[error("sqlite error: {0}")]
16 Sqlite(#[from] rusqlite::Error),
17 #[error("json error: {0}")]
19 Json(#[from] serde_json::Error),
20 #[error("unknown node key: {0}")]
22 UnknownNode(String),
23 #[error("invalid edge: {0}")]
25 InvalidEdge(String),
26 #[error("corrupt store: {0}")]
28 Corrupt(String),
29}
30
31const 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
35const 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
40pub struct Store {
42 conn: Connection,
43}
44
45impl Store {
46 pub fn open(path: &Path) -> Result<Self, StoreError> {
52 let conn = Connection::open(path)?;
53 Self::from_conn(conn)
54 }
55
56 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 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 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 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 pub fn upsert_node(&self, node: &Node) -> Result<(), StoreError> {
112 upsert_node(&self.conn, node)
113 }
114
115 pub fn insert_edge(&self, edge: &Edge) -> Result<(), StoreError> {
122 insert_edge(&self.conn, edge)
123 }
124
125 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 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 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 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 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 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 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 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
233fn 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 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 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}