use std::collections::{BTreeMap, BTreeSet};
use kimetsu_core::KimetsuResult;
use rusqlite::Connection;
use crate::consolidate::parse_tags;
#[derive(Debug, Clone, PartialEq, Eq, PartialOrd, Ord)]
pub struct EdgeProposal {
pub src_id: String,
pub dst_id: String,
pub edge_type: String,
}
pub const RELATES_TO: &str = "relates_to";
pub const DEFAULT_MAX_FAN_OUT: usize = 8;
const MIN_KEYWORD_LEN: usize = 5;
const STOPWORDS: &[&str] = &[
"about", "above", "after", "again", "against", "always", "because", "before", "being", "below",
"between", "could", "default", "during", "every", "first", "found", "their", "there", "these",
"thing", "things", "those", "through", "under", "until", "using", "value", "where", "which",
"while", "would", "should", "while",
];
#[derive(Debug, Clone, Copy, PartialEq, Eq, PartialOrd, Ord)]
pub enum EntitySource {
Tag,
Term,
}
impl EntitySource {
pub fn as_str(self) -> &'static str {
match self {
EntitySource::Tag => "tag",
EntitySource::Term => "term",
}
}
}
pub fn extract_entities_with_source(text: &str) -> Vec<(String, EntitySource)> {
let mut sources: BTreeMap<String, EntitySource> = BTreeMap::new();
for tag in tag_entities(text) {
sources.insert(tag, EntitySource::Tag);
}
for term in term_entities(text) {
sources.entry(term).or_insert(EntitySource::Term);
}
sources.into_iter().collect()
}
fn tag_entities(text: &str) -> BTreeSet<String> {
let mut set = BTreeSet::new();
for t in parse_tags(text) {
for word in t.split_whitespace() {
let w = word.trim();
if w.len() >= 3 {
set.insert(w.to_string());
}
}
}
set
}
fn term_entities(text: &str) -> BTreeSet<String> {
let mut set = BTreeSet::new();
for raw in text.split(|c: char| !c.is_alphanumeric()) {
if raw.is_empty() {
continue;
}
let is_proper = raw.chars().next().is_some_and(|c| c.is_uppercase())
&& raw.chars().skip(1).any(|c| c.is_lowercase());
let lower = raw.to_ascii_lowercase();
let proper_kept = is_proper && lower.len() >= 3;
let informative = lower.len() >= MIN_KEYWORD_LEN
&& !STOPWORDS.contains(&lower.as_str())
&& lower.chars().any(|c| c.is_alphabetic());
if proper_kept || informative {
set.insert(lower);
}
}
set
}
pub fn extract_entities(text: &str) -> Vec<String> {
let mut set = tag_entities(text);
set.extend(term_entities(text));
set.into_iter().collect()
}
fn load_active_memories(conn: &Connection) -> KimetsuResult<Vec<(String, String)>> {
let mut stmt = conn.prepare(
"SELECT memory_id, text
FROM memories
WHERE invalidated_at IS NULL AND superseded_by IS NULL
ORDER BY memory_id",
)?;
let rows = stmt
.query_map([], |r| Ok((r.get::<_, String>(0)?, r.get::<_, String>(1)?)))?
.collect::<Result<Vec<_>, _>>()?;
Ok(rows)
}
pub fn build_relates_to_edges(
conn: &Connection,
max_fan_out: usize,
) -> KimetsuResult<Vec<EdgeProposal>> {
let cap = if max_fan_out == 0 {
DEFAULT_MAX_FAN_OUT
} else {
max_fan_out
};
let memories = load_active_memories(conn)?;
let mut by_entity: BTreeMap<String, Vec<String>> = BTreeMap::new();
for (id, text) in &memories {
for entity in extract_entities(text) {
by_entity.entry(entity).or_default().push(id.clone());
}
}
let mut pairs: BTreeSet<(String, String)> = BTreeSet::new();
for ids in by_entity.values() {
if ids.len() < 2 || ids.len() > cap.max(2) * 4 {
continue;
}
for i in 0..ids.len() {
for j in (i + 1)..ids.len() {
let (a, b) = if ids[i] < ids[j] {
(ids[i].clone(), ids[j].clone())
} else if ids[i] > ids[j] {
(ids[j].clone(), ids[i].clone())
} else {
continue; };
pairs.insert((a, b));
}
}
}
let mut fan_out: BTreeMap<String, usize> = BTreeMap::new();
let mut proposals: Vec<EdgeProposal> = Vec::new();
for (a, b) in pairs {
let ca = fan_out.entry(a.clone()).or_insert(0);
if *ca >= cap {
continue;
}
*ca += 1;
proposals.push(EdgeProposal {
src_id: a,
dst_id: b,
edge_type: RELATES_TO.to_string(),
});
}
Ok(proposals)
}
pub const INCREMENTAL_MIN_SHARED_ENTITIES: usize = 2;
pub fn project_entities(conn: &Connection, memory_id: &str, text: &str) -> KimetsuResult<usize> {
conn.execute(
"DELETE FROM memory_entities WHERE memory_id = ?1",
rusqlite::params![memory_id],
)?;
let entities = extract_entities_with_source(text);
let mut stmt = conn.prepare(
"INSERT OR REPLACE INTO memory_entities (memory_id, entity, source) VALUES (?1, ?2, ?3)",
)?;
for (entity, source) in &entities {
stmt.execute(rusqlite::params![memory_id, entity, source.as_str()])?;
}
Ok(entities.len())
}
pub fn forget_entities(conn: &Connection, memory_id: &str) -> KimetsuResult<()> {
conn.execute(
"DELETE FROM memory_entities WHERE memory_id = ?1",
rusqlite::params![memory_id],
)?;
Ok(())
}
pub fn incremental_edges_for_memory(
conn: &Connection,
memory_id: &str,
max_fan_out: usize,
) -> KimetsuResult<Vec<EdgeProposal>> {
let cap = if max_fan_out == 0 {
DEFAULT_MAX_FAN_OUT
} else {
max_fan_out
};
let mut stmt = conn.prepare(
"SELECT other.memory_id, COUNT(*) AS shared
FROM memory_entities AS mine
JOIN memory_entities AS other
ON other.entity = mine.entity AND other.memory_id != mine.memory_id
JOIN memories AS m
ON m.memory_id = other.memory_id
WHERE mine.memory_id = ?1
AND m.invalidated_at IS NULL
AND m.superseded_by IS NULL
GROUP BY other.memory_id
HAVING shared >= ?2
ORDER BY shared DESC, other.memory_id ASC
LIMIT ?3",
)?;
let neighbours = stmt
.query_map(
rusqlite::params![
memory_id,
INCREMENTAL_MIN_SHARED_ENTITIES as i64,
cap as i64
],
|row| row.get::<_, String>(0),
)?
.collect::<Result<Vec<_>, _>>()?;
Ok(neighbours
.into_iter()
.map(|other| {
let (src_id, dst_id) = if memory_id < other.as_str() {
(memory_id.to_string(), other)
} else {
(other, memory_id.to_string())
};
EdgeProposal {
src_id,
dst_id,
edge_type: RELATES_TO.to_string(),
}
})
.collect())
}
pub fn reproject_all_entities(conn: &Connection) -> KimetsuResult<usize> {
let memories = load_active_memories(conn)?;
let mut total = 0usize;
for (id, text) in &memories {
total += project_entities(conn, id, text)?;
}
Ok(total)
}
#[cfg(test)]
mod tests {
use super::*;
use crate::projector::add_memory_edges;
use crate::schema;
use rusqlite::params;
fn make_conn() -> Connection {
let conn = Connection::open_in_memory().expect("open_in_memory");
schema::initialize(&conn).expect("schema::initialize");
conn
}
fn insert_active_memory(conn: &Connection, id: &str, text: &str) {
conn.execute(
"INSERT INTO memories
(memory_id, scope, kind, text, normalized_text, confidence, provenance_snapshot_json, created_at)
VALUES (?1, 'global_user', 'fact', ?2, ?2, 0.85, '{}', '2024-01-01T00:00:00Z')",
params![id, text],
)
.expect("insert memory");
}
#[test]
fn extract_entities_picks_tags_and_salient_terms() {
let ents = extract_entities("[tags: rust mutex] Holding a Mutex across an await deadlocks");
assert!(ents.contains(&"rust".to_string()));
assert!(ents.contains(&"mutex".to_string()));
assert!(ents.contains(&"deadlocks".to_string()));
assert!(!ents.contains(&"a".to_string()));
assert!(!ents.contains(&"an".to_string()));
}
#[test]
fn extract_entities_is_sorted_and_deduped() {
let ents = extract_entities("Docker docker DOCKER mount mount");
let mut sorted = ents.clone();
sorted.sort();
assert_eq!(ents, sorted, "entities must be returned sorted");
let set: BTreeSet<&String> = ents.iter().collect();
assert_eq!(set.len(), ents.len(), "no duplicates");
}
#[test]
fn build_edges_links_shared_entity_and_skips_unrelated() {
let conn = make_conn();
insert_active_memory(
&conn,
"a",
"[tags: deadlock] holding a mutex guard deadlock risk",
);
insert_active_memory(
&conn,
"b",
"the async runtime can deadlock under contention",
);
insert_active_memory(
&conn,
"c",
"the website landing page uses a teal gradient hero",
);
let edges = build_relates_to_edges(&conn, 0).expect("build");
assert_eq!(edges.len(), 1, "got {edges:?}");
assert_eq!(edges[0].src_id, "a");
assert_eq!(edges[0].dst_id, "b");
assert_eq!(edges[0].edge_type, RELATES_TO);
}
#[test]
fn build_edges_persist_roundtrip() {
let conn = make_conn();
insert_active_memory(&conn, "a", "windows docker named pipe mount rule");
insert_active_memory(&conn, "b", "docker mount breaks under a tcp host");
let edges = build_relates_to_edges(&conn, 0).expect("build");
assert!(!edges.is_empty());
let tuples: Vec<(String, String, String)> = edges
.iter()
.map(|e| (e.src_id.clone(), e.dst_id.clone(), e.edge_type.clone()))
.collect();
let written = add_memory_edges(&conn, &tuples).expect("persist");
assert_eq!(written, edges.len());
let n: i64 = conn
.query_row(
"SELECT COUNT(*) FROM memory_edges WHERE edge_type='relates_to'",
[],
|r| r.get(0),
)
.unwrap();
assert_eq!(n as usize, edges.len());
}
#[test]
fn entity_source_prefers_the_author_supplied_tag() {
let pairs = extract_entities_with_source("[tags: mutex] Holding a mutex across an await");
let mutex = pairs
.iter()
.find(|(e, _)| e == "mutex")
.expect("mutex must be extracted");
assert_eq!(
mutex.1,
EntitySource::Tag,
"an entity that is both tagged and mentioned is a tag: the author said it out loud"
);
let holding = pairs.iter().find(|(e, _)| e == "holding");
assert_eq!(holding.map(|(_, s)| *s), Some(EntitySource::Term));
}
#[test]
fn project_entities_replaces_rather_than_accumulates() {
let conn = make_conn();
insert_active_memory(&conn, "a", "[tags: sqlite] vacuum reclaims dead pages");
project_entities(&conn, "a", "[tags: sqlite] vacuum reclaims dead pages").expect("project");
let first: i64 = conn
.query_row(
"SELECT COUNT(*) FROM memory_entities WHERE memory_id='a'",
[],
|r| r.get(0),
)
.unwrap();
assert!(first > 0, "entities must land");
project_entities(&conn, "a", "[tags: sqlite]").expect("reproject");
let entities: Vec<String> = conn
.prepare("SELECT entity FROM memory_entities WHERE memory_id='a' ORDER BY entity")
.unwrap()
.query_map([], |r| r.get(0))
.unwrap()
.collect::<Result<_, _>>()
.unwrap();
assert_eq!(entities, vec!["sqlite".to_string()]);
}
#[test]
fn incremental_edges_link_a_new_memory_to_its_neighbours() {
let conn = make_conn();
for (id, text) in [
(
"a",
"[tags: sqlite wal] WAL mode needs a checkpoint before backup",
),
(
"b",
"[tags: sqlite wal] Opening a WAL database read-only skips recovery",
),
("c", "[tags: rust] Prefer thiserror for library error types"),
] {
insert_active_memory(&conn, id, text);
project_entities(&conn, id, text).expect("project");
}
let edges = incremental_edges_for_memory(&conn, "b", 0).expect("incremental");
let linked: Vec<&str> = edges
.iter()
.map(|e| {
if e.src_id == "b" {
e.dst_id.as_str()
} else {
e.src_id.as_str()
}
})
.collect();
assert_eq!(
linked,
vec!["a"],
"two shared entities (sqlite, wal) links a-b; one topic in common does not link c"
);
let edge = &edges[0];
assert!(edge.src_id < edge.dst_id, "edges are stored src < dst");
assert_eq!(edge.edge_type, RELATES_TO);
}
#[test]
fn incremental_edges_need_more_than_one_shared_entity() {
let conn = make_conn();
for (id, text) in [
("a", "[tags: sqlite] checkpoint before backup"),
("b", "[tags: sqlite] different subject entirely"),
] {
insert_active_memory(&conn, id, text);
project_entities(&conn, id, text).expect("project");
}
let shared: i64 = conn
.query_row(
"SELECT COUNT(*) FROM memory_entities x JOIN memory_entities y
ON x.entity = y.entity AND x.memory_id='a' AND y.memory_id='b'",
[],
|r| r.get(0),
)
.unwrap();
assert_eq!(shared, 1, "fixture must share exactly one entity");
let edges = incremental_edges_for_memory(&conn, "b", 0).expect("incremental");
assert!(
edges.is_empty(),
"one shared entity is not enough: {edges:?}"
);
}
#[test]
fn incremental_edges_skip_inactive_neighbours() {
let conn = make_conn();
for (id, text) in [
("a", "[tags: sqlite wal] WAL mode needs a checkpoint"),
("b", "[tags: sqlite wal] WAL recovery on read-only open"),
] {
insert_active_memory(&conn, id, text);
project_entities(&conn, id, text).expect("project");
}
conn.execute(
"UPDATE memories SET invalidated_at = '2024-06-01T00:00:00Z' WHERE memory_id = 'a'",
[],
)
.unwrap();
let edges = incremental_edges_for_memory(&conn, "b", 0).expect("incremental");
assert!(
edges.is_empty(),
"an invalidated memory is not a reachable destination: {edges:?}"
);
}
#[test]
fn incremental_edges_respect_the_fan_out_cap() {
let conn = make_conn();
for i in 0..10 {
let id = format!("m{i}");
let text = "[tags: sqlite wal] shared topic";
insert_active_memory(&conn, &id, text);
project_entities(&conn, &id, text).expect("project");
}
let edges = incremental_edges_for_memory(&conn, "m0", 3).expect("incremental");
assert_eq!(edges.len(), 3, "fan-out cap must bound the write path");
}
#[test]
fn reproject_all_entities_backfills_an_existing_corpus() {
let conn = make_conn();
insert_active_memory(&conn, "a", "[tags: sqlite wal] checkpoint before backup");
insert_active_memory(&conn, "b", "[tags: sqlite wal] recovery on read-only open");
let before: i64 = conn
.query_row("SELECT COUNT(*) FROM memory_entities", [], |r| r.get(0))
.unwrap();
assert_eq!(before, 0);
reproject_all_entities(&conn).expect("backfill");
let edges = incremental_edges_for_memory(&conn, "b", 0).expect("incremental");
assert_eq!(
edges.len(),
1,
"backfilled index must be usable immediately"
);
}
#[test]
fn recording_memories_links_them_without_a_manual_graph_build() {
use kimetsu_core::memory::{MemoryKind, MemoryScope};
let ts = std::time::SystemTime::now()
.duration_since(std::time::UNIX_EPOCH)
.map(|d| d.as_nanos())
.unwrap_or(0);
let dir = std::env::temp_dir().join(format!("kimetsu-graph-ingest-{ts}"));
std::fs::create_dir_all(&dir).expect("create tmp");
kimetsu_core::paths::git_init_boundary(&dir);
crate::user_brain::with_user_brain_disabled(|| {
crate::project::init_project(&dir, true).expect("init");
crate::project::add_memory(
&dir,
MemoryScope::Project,
MemoryKind::Convention,
"[tags: sqlite wal] Checkpoint the WAL before copying brain.db",
)
.expect("add first");
crate::project::add_memory(
&dir,
MemoryScope::Project,
MemoryKind::FailurePattern,
"[tags: sqlite wal] Opening a WAL database read-only skips recovery",
)
.expect("add second");
let paths = kimetsu_core::paths::ProjectPaths::discover(&dir).expect("paths");
let conn = Connection::open(&paths.brain_db).expect("open brain");
let relates: i64 = conn
.query_row(
"SELECT COUNT(*) FROM memory_edges WHERE edge_type = 'relates_to'",
[],
|r| r.get(0),
)
.expect("count edges");
assert!(
relates >= 1,
"the write path must link related memories; got {relates} relates_to edges"
);
let entities: i64 = conn
.query_row("SELECT COUNT(*) FROM memory_entities", [], |r| r.get(0))
.expect("count entities");
assert!(entities > 0, "entities must be projected on write");
});
std::fs::remove_dir_all(dir).ok();
}
#[test]
fn build_edges_excludes_superseded() {
let conn = make_conn();
insert_active_memory(&conn, "a", "shared topic alpha beta gamma");
insert_active_memory(&conn, "b", "shared topic alpha beta gamma too");
conn.execute(
"UPDATE memories SET superseded_by = 'a' WHERE memory_id = 'b'",
[],
)
.unwrap();
let edges = build_relates_to_edges(&conn, 0).expect("build");
assert!(edges.is_empty(), "superseded memory must not be linked");
}
}