use rustc_hash::FxHashMap;
use std::collections::{HashSet, VecDeque};
use std::num::NonZeroUsize;
use std::path::Path;
use std::sync::atomic::{AtomicI64, AtomicUsize, Ordering};
use std::time::Duration;
use parking_lot::{Mutex, MutexGuard};
use rusqlite::{Connection, OpenFlags, params};
use crate::errors::{MCSError, Result};
use crate::mutation::{
MutationContext, MutationRequest, MutationResult, MutationService, ObservationUpdate,
};
use crate::storage::{Durability, SqliteTuning};
use crate::types::{
Degree, Entity, EntityDescription, EntityInput, Observation, ObservationInput, Relation,
};
const OBSERVATION_JSON: &str = "json_object('body',o.body,'createdAtUs',o.created_us,'occurredAtUs',o.occurred_us,'originEntityName',o.origin_entity_name)";
const MAX_TRAVERSAL_ENTITIES: usize = 500_000;
const MAX_TRAVERSAL_RELS: usize = 2_000_000;
fn sqlite_err(e: rusqlite::Error) -> MCSError {
MCSError::IoError(std::io::Error::other(e))
}
const fn is_not_found(e: &rusqlite::Error) -> bool {
matches!(e, rusqlite::Error::QueryReturnedNoRows)
}
#[inline(always)]
pub fn name_hash(name: &str) -> i64 {
let mut h: u64 = 0xcbf29ce484222325;
for b in name.bytes() {
h ^= u64::from(b);
h = h.wrapping_mul(0x100000001b3);
}
h as i64
}
fn entity_name_lookup(conn: &Connection, name: &str) -> Result<Option<i64>> {
let h = name_hash(name);
let mut stmt = conn
.prepare_cached("SELECT id FROM entity WHERE name_hash = ?1 AND name = ?2 AND flags = 0")
.map_err(sqlite_err)?;
match stmt.query_row(params![h, name], |row| row.get::<_, i64>(0)) {
Ok(id) => Ok(Some(id)),
Err(e) if is_not_found(&e) => Ok(None),
Err(e) => Err(sqlite_err(e)),
}
}
fn lookup_type_id(conn: &Connection, type_name: &str, kind: i64) -> Option<i64> {
conn.prepare_cached("SELECT id FROM type_dict WHERE kind = ?1 AND name = ?2")
.ok()?
.query_row(params![kind, type_name], |row| row.get::<_, i64>(0))
.ok()
}
fn read_graph_stat(conn: &Connection, key: &str) -> Result<i64> {
conn.query_row(
"SELECT value FROM graph_stat WHERE key = ?1",
params![key],
|row| row.get(0),
)
.map_err(sqlite_err)
}
fn select_all_types(conn: &Connection, kind: i64) -> Result<Vec<(String, usize)>> {
let mut stmt = conn
.prepare_cached(
"SELECT name, count FROM type_dict WHERE kind = ?1 AND count > 0 ORDER BY count DESC",
)
.map_err(sqlite_err)?;
let rows = stmt
.query_map(params![kind], |row| {
Ok((row.get::<_, String>(0)?, row.get::<_, i64>(1)? as usize))
})
.map_err(sqlite_err)?
.filter_map(|r| r.ok())
.collect();
Ok(rows)
}
fn int_csv(ids: &[i64]) -> String {
use std::fmt::Write as _;
let mut s = String::with_capacity(ids.len() * 8);
for (i, id) in ids.iter().enumerate() {
if i > 0 {
s.push(',');
}
let _ = write!(s, "{id}");
}
s
}
fn rel_values_literal(rels: &HashSet<(i64, i64, i64)>) -> String {
use std::fmt::Write as _;
let mut s = String::with_capacity(rels.len() * 16);
for (i, (f, t, tp)) in rels.iter().enumerate() {
if i > 0 {
s.push(',');
}
let _ = write!(s, "({f},{t},{tp})");
}
s
}
fn batch_entities_by_ids(conn: &Connection, ids: &[i64]) -> FxHashMap<i64, Entity> {
let mut map = FxHashMap::default();
if ids.is_empty() {
return map;
}
let sql = format!(
"SELECT e.id, e.name, t.name,
COALESCE((SELECT json_group_array({OBSERVATION_JSON} ORDER BY o.idx, o.id)
FROM observation o WHERE o.entity_id = e.id), '[]')
FROM entity e JOIN type_dict t ON t.id = e.type_id
WHERE e.id IN ({}) AND e.flags = 0",
int_csv(ids)
);
if let Ok(mut stmt) = conn.prepare(&sql)
&& let Ok(rows) = stmt.query_map([], |row| {
Ok((
row.get::<_, i64>(0)?,
row.get::<_, String>(1)?,
row.get::<_, String>(2)?,
row.get::<_, String>(3)?,
))
})
{
for (id, name, etype, obs_json) in rows.flatten() {
let observations: Vec<Observation> =
serde_json::from_str(&obs_json).unwrap_or_default();
map.insert(
id,
Entity {
name,
entity_type: etype,
observations,
},
);
}
}
map
}
fn batch_entity_lite_by_ids(
conn: &Connection,
ids: &[i64],
) -> FxHashMap<i64, (String, String, i64)> {
let mut map = FxHashMap::default();
if ids.is_empty() {
return map;
}
let sql = format!(
"SELECT e.id, e.name, t.name, e.obs_count
FROM entity e JOIN type_dict t ON t.id = e.type_id
WHERE e.id IN ({}) AND e.flags = 0",
int_csv(ids)
);
if let Ok(mut stmt) = conn.prepare(&sql)
&& let Ok(rows) = stmt.query_map([], |row| {
Ok((
row.get::<_, i64>(0)?,
row.get::<_, String>(1)?,
row.get::<_, String>(2)?,
row.get::<_, i64>(3)?,
))
})
{
for (id, name, etype, oc) in rows.flatten() {
map.insert(id, (name, etype, oc));
}
}
map
}
fn fts_candidate_ids(conn: &Connection, query: &str, cap: usize) -> Vec<i64> {
let mut ids: Vec<i64> = Vec::new();
let mut seen: HashSet<i64> = HashSet::new();
let cap_i64 = cap as i64;
if let Ok(mut stmt) =
conn.prepare("SELECT rowid FROM name_fts WHERE name_fts MATCH ?1 ORDER BY rank LIMIT ?2")
&& let Ok(rows) = stmt.query_map(params![query, cap_i64], |row| row.get::<_, i64>(0))
{
for id in rows.flatten() {
if seen.insert(id) {
ids.push(id);
}
}
}
if let Ok(mut stmt) = conn.prepare(
"SELECT entity_id FROM obs_fts JOIN observation ON obs_fts.rowid = observation.id
WHERE obs_fts MATCH ?1
GROUP BY entity_id
LIMIT ?2",
) && let Ok(rows) = stmt.query_map(params![query, cap_i64], |row| row.get::<_, i64>(0))
{
for id in rows.flatten() {
if seen.insert(id) {
ids.push(id);
}
}
}
ids
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum Direction {
Outgoing,
Incoming,
Both,
}
impl Direction {
pub fn parse(s: Option<&str>) -> Self {
match s {
Some("OUTGOING") => Direction::Outgoing,
Some("INCOMING") => Direction::Incoming,
_ => Direction::Both,
}
}
}
pub fn push_json_str(buf: &mut String, raw: &str) {
buf.push('"');
let mut start = 0;
let bytes = raw.as_bytes();
for (i, &b) in bytes.iter().enumerate() {
let esc: u8 = match b {
b'"' => b'"',
b'\\' => b'\\',
b'\n' => b'n',
b'\r' => b'r',
b'\t' => b't',
0x08 => b'b',
0x0C => b'f',
0x00..=0x07 | 0x0B | 0x0E..=0x1F => continue, _ => continue,
};
buf.push_str(&raw[start..i]);
buf.push('\\');
buf.push(esc as char);
start = i + 1;
}
for (i, &b) in bytes.iter().enumerate().skip(start) {
if b <= 0x07 || b == 0x0B || (0x0E..=0x1F).contains(&b) {
buf.push_str(&raw[start..i]);
write_escape_unicode(buf, b);
start = i + 1;
}
}
buf.push_str(&raw[start..]);
buf.push('"');
}
#[inline(never)]
fn write_escape_unicode(buf: &mut String, b: u8) {
use std::fmt::Write;
write!(buf, "\\u{:04x}", b).unwrap();
}
pub(crate) struct TxGuard<'a> {
conn: &'a Connection,
done: bool,
}
impl<'a> TxGuard<'a> {
pub(crate) fn begin(conn: &'a Connection) -> Result<Self> {
conn.execute_batch("BEGIN IMMEDIATE").map_err(sqlite_err)?;
Ok(Self { conn, done: false })
}
pub(crate) fn commit(mut self) -> Result<()> {
self.conn.execute_batch("COMMIT").map_err(sqlite_err)?;
self.done = true;
Ok(())
}
}
impl Drop for TxGuard<'_> {
fn drop(&mut self) {
if !self.done {
let _ = self.conn.execute_batch("ROLLBACK");
}
}
}
struct ReaderPool {
conns: Vec<Mutex<Connection>>,
next: AtomicUsize,
}
impl ReaderPool {
fn get(&self) -> MutexGuard<'_, Connection> {
for c in &self.conns {
if let Some(g) = c.try_lock() {
return g;
}
}
let i = self.next.fetch_add(1, Ordering::Relaxed) % self.conns.len();
self.conns[i].lock()
}
}
pub struct GraphHandle {
pub(crate) writer: Mutex<Connection>,
readers: ReaderPool,
seq_entity: AtomicI64,
seq_obs: AtomicI64,
}
fn open_reader(path: &Path, tuning: &SqliteTuning) -> Result<Connection> {
let conn = Connection::open_with_flags(
path,
OpenFlags::SQLITE_OPEN_READ_WRITE
| OpenFlags::SQLITE_OPEN_NO_MUTEX
| OpenFlags::SQLITE_OPEN_URI,
)
.map_err(sqlite_err)?;
conn.busy_timeout(Duration::from_millis(tuning.busy_timeout_ms))
.map_err(sqlite_err)?;
conn.execute_batch(&format!(
"PRAGMA query_only = ON;
PRAGMA cache_size = -{};
PRAGMA temp_store = MEMORY;
PRAGMA mmap_size = {};",
tuning.cache_size_kb, tuning.mmap_size
))
.map_err(sqlite_err)?;
Ok(conn)
}
impl GraphHandle {
pub fn new(
path: &Path,
durability: Durability,
tuning: SqliteTuning,
_lru_cache_size: NonZeroUsize,
read_pool_size: usize,
) -> Result<Self> {
let conn = Connection::open(path).map_err(sqlite_err)?;
conn.busy_timeout(Duration::from_millis(tuning.busy_timeout_ms))
.map_err(sqlite_err)?;
conn.execute_batch(&format!(
"PRAGMA page_size = {};
PRAGMA auto_vacuum = INCREMENTAL;",
tuning.page_size
))
.map_err(sqlite_err)?;
conn.execute_batch(&format!(
"PRAGMA journal_mode = WAL;
PRAGMA foreign_keys = OFF;
PRAGMA cache_size = -{};
PRAGMA temp_store = MEMORY;
PRAGMA busy_timeout = {};
PRAGMA synchronous = NORMAL;
PRAGMA journal_size_limit = {};",
tuning.cache_size_kb, tuning.busy_timeout_ms, tuning.journal_size_limit
))
.map_err(sqlite_err)?;
crate::schema::initialize_database(&conn)?;
conn.execute_batch(&format!("PRAGMA mmap_size = {};", tuning.mmap_size))
.map_err(sqlite_err)?;
let sync_pragma = match durability {
Durability::Sync => "PRAGMA synchronous = FULL",
Durability::Async => "PRAGMA synchronous = NORMAL",
};
conn.execute_batch(sync_pragma).map_err(sqlite_err)?;
conn.execute_batch("PRAGMA analysis_limit = 400;")
.map_err(sqlite_err)?;
conn.execute_batch("PRAGMA optimize;").map_err(sqlite_err)?;
let seq_entity = read_graph_stat(&conn, "entity_seq").unwrap_or(0);
let seq_obs = read_graph_stat(&conn, "obs_seq").unwrap_or(0);
let pool_size = read_pool_size.max(1);
let mut conns = Vec::with_capacity(pool_size);
for _ in 0..pool_size {
conns.push(Mutex::new(open_reader(path, &tuning)?));
}
let readers = ReaderPool {
conns,
next: AtomicUsize::new(0),
};
Ok(Self {
writer: Mutex::new(conn),
readers,
seq_entity: AtomicI64::new(seq_entity),
seq_obs: AtomicI64::new(seq_obs),
})
}
pub(crate) fn next_entity_id(&self) -> i64 {
self.seq_entity.fetch_add(1, Ordering::Relaxed) + 1
}
pub(crate) fn refresh_seqs(&self, conn: &Connection) -> Result<()> {
self.seq_entity
.fetch_max(read_graph_stat(conn, "entity_seq")?, Ordering::Relaxed);
self.seq_obs
.fetch_max(read_graph_stat(conn, "obs_seq")?, Ordering::Relaxed);
Ok(())
}
pub(crate) fn next_obs_id(&self) -> i64 {
self.seq_obs.fetch_add(1, Ordering::Relaxed) + 1
}
fn get_entity_id(&self, conn: &Connection, name: &str) -> Result<Option<(i64, i64, i64, i64)>> {
use rusqlite::OptionalExtension;
conn.query_row(
"SELECT id, type_id, out_deg, in_deg FROM entity WHERE name_hash = ?1 AND name = ?2 AND flags = 0",
params![name_hash(name), name],
|row| Ok((row.get(0)?, row.get(1)?, row.get(2)?, row.get(3)?)),
).optional().map_err(sqlite_err)
}
pub(crate) fn sync_seqs(&self, conn: &Connection) -> Result<()> {
let seq_e = self.seq_entity.load(Ordering::Relaxed);
let seq_o = self.seq_obs.load(Ordering::Relaxed);
conn.execute(
"UPDATE graph_stat SET value = CASE key WHEN 'entity_seq' THEN ?1 WHEN 'obs_seq' THEN ?2 ELSE value END
WHERE key IN ('entity_seq', 'obs_seq')",
params![seq_e, seq_o],
)
.map_err(sqlite_err)?;
Ok(())
}
pub fn get_entity(&self, name: &str) -> Result<Option<Entity>> {
let conn = self.readers.get();
let tx = conn.unchecked_transaction().map_err(sqlite_err)?;
let entity = crate::mutation::read_entity(&tx, name)?.map(|snapshot| snapshot.entity());
tx.commit().map_err(sqlite_err)?;
Ok(entity)
}
fn mutate(&self, request: MutationRequest) -> Result<MutationResult> {
MutationService::new(self)
.apply_with_result(request, MutationContext::local())
.map(|(_, result)| result)
}
pub fn create_entities(&self, entities: &[EntityInput]) -> Result<Vec<Entity>> {
match self.mutate(MutationRequest::CreateEntities {
entities: entities.to_vec(),
})? {
MutationResult::Entities(result) => Ok(result),
_ => unreachable!("create_entities always returns entities"),
}
}
pub fn upsert_entities(&self, entities: &[EntityInput]) -> Result<Vec<Entity>> {
match self.mutate(MutationRequest::UpsertEntities {
entities: entities.to_vec(),
})? {
MutationResult::Entities(result) => Ok(result),
_ => unreachable!("upsert_entities always returns entities"),
}
}
pub fn delete_entities(&self, names: &[String]) -> Result<()> {
self.mutate(MutationRequest::DeleteEntities {
names: names.to_vec(),
})
.map(|_| ())
}
pub fn create_relations(&self, relations: &[Relation]) -> Result<Vec<Relation>> {
match self.mutate(MutationRequest::CreateRelations {
relations: relations.to_vec(),
})? {
MutationResult::Relations(result) => Ok(result),
_ => unreachable!("create_relations always returns relations"),
}
}
pub fn delete_relations(&self, relations: &[Relation]) -> Result<()> {
self.mutate(MutationRequest::DeleteRelations {
relations: relations.to_vec(),
})
.map(|_| ())
}
pub fn add_observations(
&self,
entity_name: &str,
contents: &[ObservationInput],
) -> Result<Vec<Observation>> {
match self.mutate(MutationRequest::AddObservations {
observations: vec![ObservationUpdate {
entity_name: entity_name.into(),
contents: contents.to_vec(),
}],
})? {
MutationResult::Observations(mut result) => Ok(result.remove(0).added_observations),
_ => unreachable!("add_observations always returns observations"),
}
}
pub fn delete_observations(
&self,
entity_name: &str,
observations: &[ObservationInput],
) -> Result<()> {
self.mutate(MutationRequest::DeleteObservations {
observations: vec![ObservationUpdate {
entity_name: entity_name.into(),
contents: observations.to_vec(),
}],
})
.map(|_| ())
}
pub fn merge_entities(&self, source: &str, target: &str) -> Result<Entity> {
match self.mutate(MutationRequest::MergeEntities {
source: source.into(),
target: target.into(),
})? {
MutationResult::Entity(result) => Ok(result),
_ => unreachable!("merge_entities always returns an entity"),
}
}
pub fn rename_entity(&self, old_name: &str, new_name: &str) -> Result<Entity> {
match self.mutate(MutationRequest::RenameEntity {
old_name: old_name.into(),
new_name: new_name.into(),
})? {
MutationResult::Entity(result) => Ok(result),
_ => unreachable!("rename_entity always returns an entity"),
}
}
pub fn code_purge_file(&self, rel_path: &str) -> Result<usize> {
match self.mutate(MutationRequest::PurgeDefinedEntities {
name: rel_path.into(),
})? {
MutationResult::Count(count) => Ok(count),
_ => unreachable!("purge always returns a count"),
}
}
pub fn search_nodes_filtered(
&self,
query: &str,
filter_type: Option<&str>,
offset: usize,
limit: usize,
) -> Vec<Entity> {
if query.is_empty() {
return Vec::new();
}
let conn = self.readers.get();
let cap = offset.saturating_add(limit);
let candidates = fts_candidate_ids(&conn, query, cap);
let mut by_id = batch_entities_by_ids(&conn, &candidates);
let mut results = Vec::new();
let mut count: usize = 0;
for eid in candidates {
let Some(entity) = by_id.remove(&eid) else {
continue;
};
if let Some(ft) = filter_type
&& !ft.is_empty()
&& entity.entity_type != ft
{
continue;
}
if count < offset {
count += 1;
continue;
}
if results.len() >= limit {
break;
}
results.push(entity);
count += 1;
}
results
}
pub fn search_nodes_lite_json(
&self,
query: &str,
filter_type: Option<&str>,
offset: usize,
limit: usize,
) -> (String, usize, bool) {
use std::fmt::Write as _;
if query.is_empty() {
return ("[]".to_string(), 0, false);
}
let conn = self.readers.get();
let cap = offset.saturating_add(limit).saturating_add(1);
let candidates = fts_candidate_ids(&conn, query, cap);
let by_id = batch_entity_lite_by_ids(&conn, &candidates);
let ft = filter_type.filter(|s| !s.is_empty());
let mut arr = String::from("[");
let mut count: usize = 0; let mut returned: usize = 0;
let mut has_more = false;
for eid in candidates {
let Some((name, etype, oc)) = by_id.get(&eid) else {
continue;
};
if let Some(f) = ft
&& etype != f
{
continue;
}
if count < offset {
count += 1;
continue;
}
if returned >= limit {
has_more = true;
break;
}
if returned > 0 {
arr.push(',');
}
arr.push_str("{\"name\":");
push_json_str(&mut arr, name);
arr.push_str(",\"entityType\":");
push_json_str(&mut arr, etype);
let _ = write!(arr, ",\"obsCount\":{oc}}}");
returned += 1;
count += 1;
}
arr.push(']');
(arr, returned, has_more)
}
pub fn read_graph_filtered(
&self,
filter_type: Option<&str>,
offset: usize,
limit: usize,
) -> Result<String> {
self.read_graph_page(filter_type, offset, limit, true)
.map(|(json, _)| json)
}
pub fn read_graph_filtered_lite(
&self,
filter_type: Option<&str>,
offset: usize,
limit: usize,
) -> Result<(String, usize)> {
self.read_graph_page(filter_type, offset, limit, false)
}
fn read_graph_page(
&self,
filter_type: Option<&str>,
offset: usize,
limit: usize,
include_obs: bool,
) -> Result<(String, usize)> {
let conn = self.readers.get();
let limit_sql: i64 = if limit == usize::MAX {
-1
} else {
limit.min(i64::MAX as usize) as i64
};
let offset_sql: i64 = offset as i64;
let filter = filter_type.filter(|ft| !ft.is_empty());
let ids: Vec<i64> = if let Some(ft) = filter {
let mut stmt = conn
.prepare_cached(
"SELECT e.id FROM entity e
WHERE e.type_id = (SELECT id FROM type_dict WHERE kind = 0 AND name = ?1)
AND e.flags = 0
ORDER BY e.id LIMIT ?2 OFFSET ?3",
)
.map_err(sqlite_err)?;
stmt.query_map(params![ft, limit_sql, offset_sql], |r| r.get::<_, i64>(0))
.map_err(sqlite_err)?
.filter_map(|r| r.ok())
.collect()
} else {
let mut stmt = conn
.prepare_cached(
"SELECT e.id FROM entity e WHERE e.flags = 0
ORDER BY e.id LIMIT ?1 OFFSET ?2",
)
.map_err(sqlite_err)?;
stmt.query_map(params![limit_sql, offset_sql], |r| r.get::<_, i64>(0))
.map_err(sqlite_err)?
.filter_map(|r| r.ok())
.collect()
};
if ids.is_empty() {
return Ok((r#"{"entities":[],"relations":[]}"#.to_string(), 0));
}
let idlist = int_csv(&ids);
let returned = ids.len();
let obs_field = if include_obs {
format!("'observations', COALESCE((SELECT json_group_array({OBSERVATION_JSON} ORDER BY o.idx, o.id)
FROM observation o WHERE o.entity_id = e.id), json('[]'))")
} else {
"'obsCount', e.obs_count".to_owned()
};
let entities_json: String = {
let sql = format!(
"SELECT COALESCE(json_group_array(json_object(
'name', e.name,
'entityType', t.name,
{obs_field}
) ORDER BY e.id), json('[]'))
FROM entity e
JOIN type_dict t ON t.id = e.type_id
WHERE e.id IN ({idlist}) AND e.flags = 0"
);
conn.query_row(&sql, [], |row| row.get::<_, String>(0))
.map_err(sqlite_err)?
};
let relations_json: String = {
let sql = format!(
"SELECT COALESCE(json_group_array(json_object(
'from', e1.name,
'to', e2.name,
'relationType', t.name
)), json('[]'))
FROM relation r
JOIN entity e1 ON e1.id = r.from_id
JOIN entity e2 ON e2.id = r.to_id
JOIN type_dict t ON t.id = r.type_id
WHERE r.from_id IN ({idlist}) AND r.to_id IN ({idlist})
AND e1.flags = 0 AND e2.flags = 0"
);
conn.query_row(&sql, [], |row| row.get::<_, String>(0))
.map_err(sqlite_err)?
};
let mut out = String::with_capacity(32 + entities_json.len() + relations_json.len());
out.push_str("{\"entities\":");
out.push_str(&entities_json);
out.push_str(",\"relations\":");
out.push_str(&relations_json);
out.push('}');
Ok((out, returned))
}
pub fn open_nodes(&self, names: &[String]) -> String {
let conn = self.readers.get();
let mut entity_ids: Vec<i64> = Vec::new();
for name in names {
let h = name_hash(name);
if let Ok(Some(id)) = conn
.query_row(
"SELECT id FROM entity WHERE name_hash = ?1 AND name = ?2 AND flags = 0",
params![h, name],
|row| row.get::<_, i64>(0),
)
.map(Some)
.or_else(|e| {
if is_not_found(&e) {
Ok(None)
} else {
Err(sqlite_err(e))
}
})
{
entity_ids.push(id);
}
}
if entity_ids.is_empty() {
return r#"{"entities":[],"relations":[]}"#.to_string();
}
let placeholders: Vec<String> = entity_ids.iter().map(|_| "?".to_string()).collect();
let ids_str = placeholders.join(",");
let entities_json: String = {
let sql = format!(
"SELECT COALESCE(json_group_array(json_object(
'name', e.name,
'entityType', t.name,
'observations', COALESCE((
SELECT json_group_array({OBSERVATION_JSON} ORDER BY o.idx, o.id)
FROM observation o WHERE o.entity_id = e.id
), json('[]'))
) ORDER BY e.id), json('[]'))
FROM entity e
JOIN type_dict t ON t.id = e.type_id
WHERE e.id IN ({ids_str}) AND e.flags = 0"
);
conn.query_row(&sql, rusqlite::params_from_iter(&entity_ids), |row| {
row.get::<_, String>(0)
})
.unwrap_or_else(|_| "[]".to_string())
};
let relations_json: String = {
let sql = format!(
"SELECT COALESCE(json_group_array(json_object(
'from', e1.name,
'to', e2.name,
'relationType', t.name
)), json('[]'))
FROM relation r
JOIN entity e1 ON e1.id = r.from_id
JOIN entity e2 ON e2.id = r.to_id
JOIN type_dict t ON t.id = r.type_id
WHERE (r.from_id IN ({ids_str}) OR r.to_id IN ({ids_str}))
AND e1.flags = 0 AND e2.flags = 0"
);
let all_params: Vec<&dyn rusqlite::types::ToSql> = entity_ids
.iter()
.map(|id| id as &dyn rusqlite::types::ToSql)
.chain(
entity_ids
.iter()
.map(|id| id as &dyn rusqlite::types::ToSql),
)
.collect();
let mut stmt = conn.prepare(&sql).unwrap();
stmt.query_row(all_params.as_slice(), |row| row.get::<_, String>(0))
.unwrap_or_else(|_| "[]".to_string())
};
let mut out = String::with_capacity(32 + entities_json.len() + relations_json.len());
out.push_str("{\"entities\":");
out.push_str(&entities_json);
out.push_str(",\"relations\":");
out.push_str(&relations_json);
out.push('}');
out
}
pub fn entities_exist(&self, names: &[String]) -> Result<Vec<bool>> {
let conn = self.readers.get();
let mut results = Vec::with_capacity(names.len());
for name in names {
let h = name_hash(name);
let exists: bool = conn
.query_row(
"SELECT 1 FROM entity WHERE name_hash = ?1 AND name = ?2 AND flags = 0",
params![h, name],
|_| Ok(()),
)
.is_ok();
results.push(exists);
}
Ok(results)
}
pub fn degree(&self, name: &str, direction: Direction) -> Result<usize> {
let conn = self.readers.get();
let (_, _, out_d, in_d) = match self.get_entity_id(&conn, name)? {
Some(v) => v,
None => {
return Err(MCSError::InvalidParams(format!(
"Entity '{name}' not found"
)));
}
};
Ok(match direction {
Direction::Outgoing => out_d as usize,
Direction::Incoming => in_d as usize,
Direction::Both => (out_d + in_d) as usize,
})
}
pub fn get_entity_count(&self) -> Result<usize> {
let conn = self.readers.get();
read_graph_stat(&conn, "entities")
.map(|v| v as usize)
.map_err(|_| MCSError::MemoryError("Failed to read entity count".into()))
}
pub fn get_relation_count(&self) -> Result<usize> {
let conn = self.readers.get();
read_graph_stat(&conn, "relations")
.map(|v| v as usize)
.map_err(|_| MCSError::MemoryError("Failed to read relation count".into()))
}
pub fn search_relations(
&self,
from: Option<&str>,
to: Option<&str>,
rtype: Option<&str>,
limit: Option<usize>,
) -> Vec<Relation> {
let conn = self.readers.get();
let mut results = Vec::new();
let from_id = from
.filter(|f| !f.is_empty())
.map(|f| entity_name_lookup(&conn, f).ok().flatten().unwrap_or(-1));
let to_id = to
.filter(|t| !t.is_empty())
.map(|t| entity_name_lookup(&conn, t).ok().flatten().unwrap_or(-1));
let type_id = rtype
.filter(|rt| !rt.is_empty())
.map(|rt| lookup_type_id(&conn, rt, 1).unwrap_or(-1));
match (from_id, to_id, type_id) {
(Some(fid), Some(tid), Some(tpid)) => {
if let Ok(mut stmt) = conn.prepare_cached(
"SELECT e1.name, e2.name, t.name
FROM relation r
JOIN entity e1 ON e1.id = r.from_id
JOIN entity e2 ON e2.id = r.to_id
JOIN type_dict t ON t.id = r.type_id
WHERE r.from_id = ?1 AND r.to_id = ?2 AND r.type_id = ?3
AND e1.flags = 0 AND e2.flags = 0
ORDER BY r.from_id, r.to_id",
) && let Ok(rows) = stmt.query_map(params![fid, tid, tpid], |row| {
Ok(Relation {
from: row.get(0)?,
to: row.get(1)?,
relation_type: row.get(2)?,
})
}) {
for row in rows.flatten() {
results.push(row);
}
}
}
(Some(fid), Some(tid), None) => {
if let Ok(mut stmt) = conn.prepare_cached(
"SELECT e1.name, e2.name, t.name
FROM relation r
JOIN entity e1 ON e1.id = r.from_id
JOIN entity e2 ON e2.id = r.to_id
JOIN type_dict t ON t.id = r.type_id
WHERE r.from_id = ?1 AND r.to_id = ?2
AND e1.flags = 0 AND e2.flags = 0
ORDER BY r.from_id, r.to_id",
) && let Ok(rows) = stmt.query_map(params![fid, tid], |row| {
Ok(Relation {
from: row.get(0)?,
to: row.get(1)?,
relation_type: row.get(2)?,
})
}) {
for row in rows.flatten() {
results.push(row);
}
}
}
(Some(fid), None, Some(tpid)) => {
if let Ok(mut stmt) = conn.prepare_cached(
"SELECT e1.name, e2.name, t.name
FROM relation r
JOIN entity e1 ON e1.id = r.from_id
JOIN entity e2 ON e2.id = r.to_id
JOIN type_dict t ON t.id = r.type_id
WHERE r.from_id = ?1 AND r.type_id = ?2
AND e1.flags = 0 AND e2.flags = 0
ORDER BY r.from_id, r.to_id",
) && let Ok(rows) = stmt.query_map(params![fid, tpid], |row| {
Ok(Relation {
from: row.get(0)?,
to: row.get(1)?,
relation_type: row.get(2)?,
})
}) {
for row in rows.flatten() {
results.push(row);
}
}
}
(None, Some(tid), Some(tpid)) => {
if let Ok(mut stmt) = conn.prepare_cached(
"SELECT e1.name, e2.name, t.name
FROM relation r
JOIN entity e1 ON e1.id = r.from_id
JOIN entity e2 ON e2.id = r.to_id
JOIN type_dict t ON t.id = r.type_id
WHERE r.to_id = ?1 AND r.type_id = ?2
AND e1.flags = 0 AND e2.flags = 0
ORDER BY r.from_id, r.to_id",
) && let Ok(rows) = stmt.query_map(params![tid, tpid], |row| {
Ok(Relation {
from: row.get(0)?,
to: row.get(1)?,
relation_type: row.get(2)?,
})
}) {
for row in rows.flatten() {
results.push(row);
}
}
}
(Some(fid), None, None) => {
if let Ok(mut stmt) = conn.prepare_cached(
"SELECT e1.name, e2.name, t.name
FROM relation r
JOIN entity e1 ON e1.id = r.from_id
JOIN entity e2 ON e2.id = r.to_id
JOIN type_dict t ON t.id = r.type_id
WHERE r.from_id = ?1
AND e1.flags = 0 AND e2.flags = 0
ORDER BY r.from_id, r.to_id",
) && let Ok(rows) = stmt.query_map(params![fid], |row| {
Ok(Relation {
from: row.get(0)?,
to: row.get(1)?,
relation_type: row.get(2)?,
})
}) {
for row in rows.flatten() {
results.push(row);
}
}
}
(None, Some(tid), None) => {
if let Ok(mut stmt) = conn.prepare_cached(
"SELECT e1.name, e2.name, t.name
FROM relation r
JOIN entity e1 ON e1.id = r.from_id
JOIN entity e2 ON e2.id = r.to_id
JOIN type_dict t ON t.id = r.type_id
WHERE r.to_id = ?1
AND e1.flags = 0 AND e2.flags = 0
ORDER BY r.from_id, r.to_id",
) && let Ok(rows) = stmt.query_map(params![tid], |row| {
Ok(Relation {
from: row.get(0)?,
to: row.get(1)?,
relation_type: row.get(2)?,
})
}) {
for row in rows.flatten() {
results.push(row);
}
}
}
(None, None, Some(tpid)) => {
if let Ok(mut stmt) = conn.prepare_cached(
"SELECT e1.name, e2.name, t.name
FROM relation r
JOIN entity e1 ON e1.id = r.from_id
JOIN entity e2 ON e2.id = r.to_id
JOIN type_dict t ON t.id = r.type_id
WHERE r.type_id = ?1
AND e1.flags = 0 AND e2.flags = 0
ORDER BY r.from_id, r.to_id",
) && let Ok(rows) = stmt.query_map(params![tpid], |row| {
Ok(Relation {
from: row.get(0)?,
to: row.get(1)?,
relation_type: row.get(2)?,
})
}) {
for row in rows.flatten() {
results.push(row);
}
}
}
(None, None, None) => {
if let Ok(mut stmt) = conn.prepare_cached(
"SELECT e1.name, e2.name, t.name
FROM relation r
JOIN entity e1 ON e1.id = r.from_id
JOIN entity e2 ON e2.id = r.to_id
JOIN type_dict t ON t.id = r.type_id
WHERE e1.flags = 0 AND e2.flags = 0
ORDER BY r.from_id, r.to_id",
) && let Ok(rows) = stmt.query_map([], |row| {
Ok(Relation {
from: row.get(0)?,
to: row.get(1)?,
relation_type: row.get(2)?,
})
}) {
for row in rows.flatten() {
results.push(row);
}
}
}
}
if let Some(lim) = limit {
results.truncate(lim);
}
results
}
pub fn find_path(&self, from: &str, to: &str) -> Result<Option<Vec<String>>> {
let conn = self.readers.get();
let (from_id, _, _, _) = match self.get_entity_id(&conn, from)? {
Some(v) => v,
None => {
return Err(MCSError::InvalidParams(format!(
"Source entity '{from}' not found"
)));
}
};
let (to_id, _, _, _) = match self.get_entity_id(&conn, to)? {
Some(v) => v,
None => {
return Err(MCSError::InvalidParams(format!(
"Target entity '{to}' not found"
)));
}
};
if from_id == to_id {
return Ok(Some(vec![from.to_string()]));
}
let mut visited = HashSet::new();
let mut parent: FxHashMap<i64, i64> = FxHashMap::default();
let mut queue = VecDeque::new();
visited.insert(from_id);
queue.push_back(from_id);
while let Some(cur) = queue.pop_front() {
if cur == to_id {
break;
}
if let Ok(mut stmt) =
conn.prepare_cached("SELECT to_id FROM relation WHERE from_id = ?1")
&& let Ok(rows) = stmt.query_map(params![cur], |row| row.get::<_, i64>(0))
{
for row in rows.flatten() {
if visited.insert(row) {
parent.insert(row, cur);
queue.push_back(row);
}
}
}
if let Ok(mut stmt) =
conn.prepare_cached("SELECT from_id FROM relation WHERE to_id = ?1")
&& let Ok(rows) = stmt.query_map(params![cur], |row| row.get::<_, i64>(0))
{
for row in rows.flatten() {
if visited.insert(row) {
parent.insert(row, cur);
queue.push_back(row);
}
}
}
}
if !parent.contains_key(&to_id) && to_id != from_id {
return Ok(None);
}
let mut path = Vec::new();
let mut cur = to_id;
path.push(cur);
while let Some(&p) = parent.get(&cur) {
path.push(p);
cur = p;
if cur == from_id {
break;
}
}
path.reverse();
let placeholders: Vec<String> = path.iter().map(|_| "?".to_string()).collect();
let sql = format!(
"SELECT id, name FROM entity WHERE id IN ({})",
placeholders.join(",")
);
let name_map: FxHashMap<i64, String> = if let Ok(mut stmt) = conn.prepare(&sql)
&& let Ok(rows) = stmt.query_map(rusqlite::params_from_iter(&path), |row| {
Ok((row.get::<_, i64>(0)?, row.get::<_, String>(1)?))
}) {
rows.flatten().collect()
} else {
FxHashMap::default()
};
let name_path: Vec<String> = path
.iter()
.filter_map(|id| name_map.get(id).cloned())
.collect();
Ok(Some(name_path))
}
pub fn compact(&self) -> Result<()> {
self.mutate(MutationRequest::Compact).map(|_| ())
}
pub fn neighbors(
&self,
name: &str,
direction: Direction,
rtype: Option<&str>,
depth: u32,
) -> Result<String> {
self._traverse(name, direction, rtype, depth, true)
}
pub fn extract_subgraph(&self, names: &[String], depth: u32) -> Result<String> {
if names.is_empty() {
return Ok(r#"{"entities":[],"relations":[]}"#.to_string());
}
let conn = self.readers.get();
let mut all_entity_ids: HashSet<i64> = HashSet::new();
let mut frontier: HashSet<i64> = HashSet::new();
let mut all_rel_pairs: HashSet<(i64, i64, i64)> = HashSet::new();
for name in names {
let h = name_hash(name);
if let Ok(Some(id)) = conn
.query_row(
"SELECT id FROM entity WHERE name_hash = ?1 AND name = ?2 AND flags = 0",
params![h, name],
|row| row.get::<_, i64>(0),
)
.map(Some)
.or_else(|e| {
if is_not_found(&e) {
Ok(None)
} else {
Err(sqlite_err(e))
}
})
{
all_entity_ids.insert(id);
frontier.insert(id);
}
}
let mut current_depth = 0u32;
while current_depth < depth && !frontier.is_empty() {
let mut next_frontier: HashSet<i64> = HashSet::new();
const CHUNK: usize = 500;
let frontier_ids: Vec<i64> = frontier.iter().copied().collect();
for chunk in frontier_ids.chunks(CHUNK) {
let placeholders: Vec<String> = chunk.iter().map(|_| "?".to_string()).collect();
let in_clause = placeholders.join(",");
if let Ok(mut stmt) = conn.prepare(&format!(
"SELECT from_id, to_id, type_id FROM relation WHERE from_id IN ({in_clause})",
)) && let Ok(rows) = stmt.query_map(rusqlite::params_from_iter(chunk), |row| {
Ok((
row.get::<_, i64>(0)?,
row.get::<_, i64>(1)?,
row.get::<_, i64>(2)?,
))
}) {
for row in rows.flatten() {
let (from_id, to_id, type_id) = row;
all_rel_pairs.insert((from_id, to_id, type_id));
if all_entity_ids.insert(to_id) {
next_frontier.insert(to_id);
}
}
}
if let Ok(mut stmt) = conn.prepare(&format!(
"SELECT from_id, to_id, type_id FROM relation WHERE to_id IN ({in_clause})",
)) && let Ok(rows) = stmt.query_map(rusqlite::params_from_iter(chunk), |row| {
Ok((
row.get::<_, i64>(0)?,
row.get::<_, i64>(1)?,
row.get::<_, i64>(2)?,
))
}) {
for row in rows.flatten() {
let (from_id, to_id, type_id) = row;
all_rel_pairs.insert((from_id, to_id, type_id));
if all_entity_ids.insert(from_id) {
next_frontier.insert(from_id);
}
}
}
}
if all_entity_ids.len() > MAX_TRAVERSAL_ENTITIES
|| all_rel_pairs.len() > MAX_TRAVERSAL_RELS
{
break;
}
frontier = next_frontier;
current_depth += 1;
}
let entities_json: String = if all_entity_ids.is_empty() {
"[]".to_string()
} else {
let ids: Vec<i64> = all_entity_ids.iter().copied().collect();
let sql = format!(
"SELECT COALESCE(json_group_array(json_object(
'name', e.name,
'entityType', t.name,
'observations', COALESCE((
SELECT json_group_array({OBSERVATION_JSON} ORDER BY o.idx, o.id)
FROM observation o WHERE o.entity_id = e.id
), json('[]'))
) ORDER BY e.id), json('[]'))
FROM entity e
JOIN type_dict t ON t.id = e.type_id
WHERE e.id IN ({}) AND e.flags = 0",
int_csv(&ids)
);
conn.query_row(&sql, [], |row| row.get::<_, String>(0))
.map_err(sqlite_err)?
};
let relations_json: String = if all_rel_pairs.is_empty() {
"[]".to_string()
} else {
let sql = format!(
"WITH r(from_id, to_id, type_id) AS (VALUES {})
SELECT COALESCE(json_group_array(json_object(
'from', e1.name,
'to', e2.name,
'relationType', t.name
)), json('[]'))
FROM r
JOIN entity e1 ON e1.id = r.from_id
JOIN entity e2 ON e2.id = r.to_id
JOIN type_dict t ON t.id = r.type_id
WHERE e1.flags = 0 AND e2.flags = 0",
rel_values_literal(&all_rel_pairs)
);
conn.query_row(&sql, [], |row| row.get::<_, String>(0))
.map_err(sqlite_err)?
};
let mut out = String::with_capacity(32 + entities_json.len() + relations_json.len());
out.push_str("{\"entities\":");
out.push_str(&entities_json);
out.push_str(",\"relations\":");
out.push_str(&relations_json);
out.push('}');
Ok(out)
}
pub fn describe_entity(&self, name: &str) -> Result<EntityDescription> {
let conn = self.readers.get();
let tx = conn.unchecked_transaction().map_err(sqlite_err)?;
let entity = crate::mutation::read_entity(&tx, name)?
.ok_or_else(|| MCSError::InvalidParams(format!("Entity '{name}' not found")))?;
let relations = crate::mutation::relations_for(&tx, name)?;
let mut neighbors: Vec<String> = relations
.iter()
.map(|relation| {
if relation.from == name {
relation.to.clone()
} else {
relation.from.clone()
}
})
.collect();
neighbors.sort();
neighbors.dedup();
let incoming = relations
.iter()
.filter(|relation| relation.to == name)
.count() as i64;
let outgoing = relations
.iter()
.filter(|relation| relation.from == name)
.count() as i64;
tx.commit().map_err(sqlite_err)?;
Ok(EntityDescription {
name: entity.name,
entity_type: entity.entity_type,
observations: entity.observations,
relations,
neighbors,
degree: Degree { incoming, outgoing },
})
}
pub fn entity_type_counts(&self) -> Vec<(String, usize)> {
let conn = self.readers.get();
select_all_types(&conn, 0).unwrap_or_default()
}
pub fn ui_meta(&self) -> (Vec<(String, usize)>, usize, usize) {
let conn = self.readers.get();
let types = select_all_types(&conn, 0).unwrap_or_default();
let entities = read_graph_stat(&conn, "entities").unwrap_or(0).max(0) as usize;
let relations = read_graph_stat(&conn, "relations").unwrap_or(0).max(0) as usize;
(types, entities, relations)
}
pub fn relation_type_counts(&self) -> Vec<(String, usize)> {
let conn = self.readers.get();
select_all_types(&conn, 1).unwrap_or_default()
}
pub fn batch_get_entities(&self, names: &[String]) -> Vec<Option<Entity>> {
names
.iter()
.map(|n| self.get_entity(n).unwrap_or(None))
.collect()
}
pub fn find_all_paths(
&self,
from: &str,
to: &str,
max_depth: usize,
max_paths: usize,
) -> Result<Vec<Vec<String>>> {
let conn = self.readers.get();
let (from_id, _, _, _) = match self.get_entity_id(&conn, from)? {
Some(v) => v,
None => {
return Err(MCSError::InvalidParams(format!(
"Source entity '{from}' not found"
)));
}
};
let (to_id, _, _, _) = match self.get_entity_id(&conn, to)? {
Some(v) => v,
None => {
return Err(MCSError::InvalidParams(format!(
"Target entity '{to}' not found"
)));
}
};
if from_id == to_id {
return Ok(vec![vec![from.to_string()]]);
}
let mut all_paths: Vec<Vec<i64>> = Vec::new();
let mut queue: VecDeque<(i64, Vec<i64>)> = VecDeque::new();
queue.push_back((from_id, vec![from_id]));
const MAX_QUEUE_SIZE: usize = 10_000_000;
while let Some((cur, path)) = queue.pop_front() {
if all_paths.len() >= max_paths {
break;
}
if path.len() > max_depth {
continue;
}
if let Ok(mut stmt) =
conn.prepare_cached("SELECT to_id FROM relation WHERE from_id = ?1")
&& let Ok(rows) = stmt.query_map(params![cur], |row| row.get::<_, i64>(0))
{
for next_id in rows.flatten() {
if next_id == to_id {
let mut full_path = path.clone();
full_path.push(next_id);
all_paths.push(full_path);
if all_paths.len() >= max_paths {
break;
}
} else if !path.contains(&next_id) && path.len() < max_depth {
if queue.len() >= MAX_QUEUE_SIZE {
return Err(MCSError::InvalidParams(
"Path exploration queue exceeded limit (too many paths on highly connected graph)".to_string()
));
}
let mut new_path = path.clone();
new_path.push(next_id);
queue.push_back((next_id, new_path));
}
}
}
if let Ok(mut stmt) =
conn.prepare_cached("SELECT from_id FROM relation WHERE to_id = ?1")
&& let Ok(rows) = stmt.query_map(params![cur], |row| row.get::<_, i64>(0))
{
for next_id in rows.flatten() {
if next_id == to_id {
let mut full_path = path.clone();
full_path.push(next_id);
all_paths.push(full_path);
if all_paths.len() >= max_paths {
break;
}
} else if !path.contains(&next_id) && path.len() < max_depth {
if queue.len() >= MAX_QUEUE_SIZE {
return Err(MCSError::InvalidParams(
"Path exploration queue exceeded limit (too many paths on highly connected graph)".to_string()
));
}
let mut new_path = path.clone();
new_path.push(next_id);
queue.push_back((next_id, new_path));
}
}
}
}
let all_ids: HashSet<i64> = all_paths.iter().flat_map(|p| p.iter()).copied().collect();
let id_list: Vec<i64> = all_ids.into_iter().collect();
let name_map: FxHashMap<i64, String> = if id_list.is_empty() {
FxHashMap::default()
} else {
let placeholders: Vec<String> = id_list.iter().map(|_| "?".to_string()).collect();
let sql = format!(
"SELECT id, name FROM entity WHERE id IN ({})",
placeholders.join(",")
);
if let Ok(mut stmt) = conn.prepare(&sql)
&& let Ok(rows) = stmt.query_map(rusqlite::params_from_iter(&id_list), |row| {
Ok((row.get::<_, i64>(0)?, row.get::<_, String>(1)?))
})
{
rows.flatten().collect()
} else {
FxHashMap::default()
}
};
let mut named_paths: Vec<Vec<String>> = Vec::with_capacity(all_paths.len());
for path_ids in all_paths {
let named: Vec<String> = path_ids
.iter()
.filter_map(|id| name_map.get(id).cloned())
.collect();
named_paths.push(named);
}
Ok(named_paths)
}
pub fn export(&self, _format: &str, max_rows: i64) -> Result<String> {
let conn = self.readers.get();
conn.query_row(
&format!(
"SELECT json_object(
'entities', COALESCE((
SELECT json_group_array(json_object(
'name', e.name,
'entityType', t.name,
'observations', COALESCE((
SELECT json_group_array({OBSERVATION_JSON} ORDER BY o.idx, o.id)
FROM observation o WHERE o.entity_id = e.id
), json('[]'))
) ORDER BY e.id)
FROM (
SELECT id, name, type_id FROM entity
WHERE flags = 0 ORDER BY id LIMIT ?1
) e
JOIN type_dict t ON t.id = e.type_id
), json('[]')),
'relations', COALESCE((
SELECT json_group_array(json_object(
'from', e1.name,
'to', e2.name,
'relationType', t.name
))
FROM (
SELECT from_id, to_id, type_id FROM relation LIMIT ?1
) r
JOIN entity e1 ON e1.id = r.from_id
JOIN entity e2 ON e2.id = r.to_id
JOIN type_dict t ON t.id = r.type_id
WHERE e1.flags = 0 AND e2.flags = 0
), json('[]'))
)"
),
params![max_rows],
|row| row.get::<_, String>(0),
)
.map_err(sqlite_err)
}
pub fn wipe(&self) -> Result<()> {
self.mutate(MutationRequest::Wipe).map(|_| ())
}
pub fn run_maintenance(&self) -> Result<()> {
let conn = self.writer.lock();
conn.execute_batch("PRAGMA wal_checkpoint(TRUNCATE);")
.map_err(sqlite_err)?;
conn.execute_batch("PRAGMA optimize(0x10000);")
.map_err(sqlite_err)?;
let tx = TxGuard::begin(&conn)?;
conn.execute_batch(
"INSERT INTO name_fts(name_fts) VALUES('optimize');
INSERT INTO obs_fts(obs_fts) VALUES('optimize');",
)
.map_err(sqlite_err)?;
tx.commit()?;
Ok(())
}
pub fn checkpoint_passive(&self) -> Result<()> {
let conn = self.writer.lock();
conn.execute_batch("PRAGMA wal_checkpoint(PASSIVE);")
.map_err(sqlite_err)?;
Ok(())
}
fn _traverse(
&self,
name: &str,
direction: Direction,
rtype: Option<&str>,
depth: u32,
_include_relations: bool,
) -> Result<String> {
let conn = self.readers.get();
let (start_id, _, _, _) = match self.get_entity_id(&conn, name)? {
Some(v) => v,
None => {
return Err(MCSError::InvalidParams(format!(
"Entity '{name}' not found"
)));
}
};
let mut all_ids: HashSet<i64> = HashSet::new();
let mut all_rels: HashSet<(i64, i64, i64)> = HashSet::new();
let mut frontier: HashSet<i64> = HashSet::new();
all_ids.insert(start_id);
frontier.insert(start_id);
let type_filter: Option<i64> = rtype
.filter(|rt| !rt.is_empty())
.map(|rt| lookup_type_id(&conn, rt, 1).unwrap_or(-1));
let mut q_out_t = conn.prepare_cached(
"SELECT to_id, type_id FROM relation WHERE from_id = ?1 AND type_id = ?2",
);
let mut q_out =
conn.prepare_cached("SELECT to_id, type_id FROM relation WHERE from_id = ?1");
let mut q_in_t = conn.prepare_cached(
"SELECT from_id, type_id FROM relation WHERE to_id = ?1 AND type_id = ?2",
);
let mut q_in =
conn.prepare_cached("SELECT from_id, type_id FROM relation WHERE to_id = ?1");
let mut cur_depth = 0u32;
while cur_depth < depth && !frontier.is_empty() {
let mut next_frontier: HashSet<i64> = HashSet::new();
for &fid in &frontier {
if direction == Direction::Outgoing || direction == Direction::Both {
if let Some(tid) = type_filter {
if let Ok(ref mut stmt) = q_out_t
&& let Ok(rows) = stmt.query_map(params![fid, tid], |row| {
Ok((row.get::<_, i64>(0)?, row.get::<_, i64>(1)?))
})
{
for row in rows.flatten() {
let (to_id, t_id) = row;
all_rels.insert((fid, to_id, t_id));
if all_ids.insert(to_id) {
next_frontier.insert(to_id);
}
}
}
} else if let Ok(ref mut stmt) = q_out
&& let Ok(rows) = stmt.query_map(params![fid], |row| {
Ok((row.get::<_, i64>(0)?, row.get::<_, i64>(1)?))
})
{
for row in rows.flatten() {
let (to_id, t_id) = row;
all_rels.insert((fid, to_id, t_id));
if all_ids.insert(to_id) {
next_frontier.insert(to_id);
}
}
}
}
if direction == Direction::Incoming || direction == Direction::Both {
if let Some(tid) = type_filter {
if let Ok(ref mut stmt) = q_in_t
&& let Ok(rows) = stmt.query_map(params![fid, tid], |row| {
Ok((row.get::<_, i64>(0)?, row.get::<_, i64>(1)?))
})
{
for row in rows.flatten() {
let (from_id, t_id) = row;
all_rels.insert((from_id, fid, t_id));
if all_ids.insert(from_id) {
next_frontier.insert(from_id);
}
}
}
} else if let Ok(ref mut stmt) = q_in
&& let Ok(rows) = stmt.query_map(params![fid], |row| {
Ok((row.get::<_, i64>(0)?, row.get::<_, i64>(1)?))
})
{
for row in rows.flatten() {
let (from_id, t_id) = row;
all_rels.insert((from_id, fid, t_id));
if all_ids.insert(from_id) {
next_frontier.insert(from_id);
}
}
}
}
}
if all_ids.len() > MAX_TRAVERSAL_ENTITIES || all_rels.len() > MAX_TRAVERSAL_RELS {
break;
}
frontier = next_frontier;
cur_depth += 1;
}
let entities_json: String = if all_ids.is_empty() {
"[]".to_string()
} else {
let ids: Vec<i64> = all_ids.iter().copied().collect();
let sql = format!(
"SELECT COALESCE(json_group_array(json_object(
'name', e.name,
'entityType', t.name,
'observations', COALESCE((
SELECT json_group_array({OBSERVATION_JSON} ORDER BY o.idx, o.id)
FROM observation o WHERE o.entity_id = e.id
), json('[]'))
) ORDER BY e.id), json('[]'))
FROM entity e
JOIN type_dict t ON t.id = e.type_id
WHERE e.id IN ({}) AND e.flags = 0",
int_csv(&ids)
);
conn.query_row(&sql, [], |row| row.get::<_, String>(0))
.map_err(sqlite_err)?
};
let relations_json: String = if all_rels.is_empty() {
"[]".to_string()
} else {
let sql = format!(
"WITH r(from_id, to_id, type_id) AS (VALUES {})
SELECT COALESCE(json_group_array(json_object(
'from', e1.name,
'to', e2.name,
'relationType', t.name
)), json('[]'))
FROM r
JOIN entity e1 ON e1.id = r.from_id
JOIN entity e2 ON e2.id = r.to_id
JOIN type_dict t ON t.id = r.type_id
WHERE e1.flags = 0 AND e2.flags = 0",
rel_values_literal(&all_rels)
);
conn.query_row(&sql, [], |row| row.get::<_, String>(0))
.map_err(sqlite_err)?
};
let mut out = String::with_capacity(32 + entities_json.len() + relations_json.len());
out.push_str("{\"entities\":");
out.push_str(&entities_json);
out.push_str(",\"relations\":");
out.push_str(&relations_json);
out.push('}');
Ok(out)
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::types::EntityInput as Entity;
use serde_json::Value;
use std::ops::Deref;
use std::path::PathBuf;
struct TestKg(GraphHandle, PathBuf);
impl Deref for TestKg {
type Target = GraphHandle;
fn deref(&self) -> &GraphHandle {
&self.0
}
}
impl Drop for TestKg {
fn drop(&mut self) {
cleanup_db(&self.1);
}
}
fn cleanup_db(path: &std::path::Path) {
let _ = std::fs::remove_file(path);
let _ = std::fs::remove_file(path.with_extension("db-wal"));
let _ = std::fs::remove_file(path.with_extension("db-shm"));
}
fn new_kg() -> TestKg {
use std::sync::atomic::AtomicU64;
use std::sync::atomic::Ordering;
static COUNTER: AtomicU64 = AtomicU64::new(0);
let n = COUNTER.fetch_add(1, Ordering::SeqCst);
let dir = std::env::temp_dir();
let path = dir.join(format!("kg_test_{}_{}.db", std::process::id(), n));
cleanup_db(&path);
let kg = GraphHandle::new(
&path,
Durability::Async,
SqliteTuning::default(),
NonZeroUsize::new(10000).unwrap(),
4,
)
.expect("create KG");
TestKg(kg, path)
}
#[test]
fn test_create_and_get_entity() {
let kg = new_kg();
let entities = vec![Entity {
name: "test".into(),
entity_type: "person".into(),
observations: vec!["obs1".into(), "obs2".into()],
}];
let created = kg.create_entities(&entities).unwrap();
assert_eq!(created.len(), 1);
let got = kg.get_entity("test").unwrap().unwrap();
assert_eq!(got.name, "test");
assert_eq!(got.entity_type, "person");
assert_eq!(
got.observations
.iter()
.map(|o| o.body.as_str())
.collect::<Vec<_>>(),
vec!["obs1", "obs2"]
);
}
#[test]
fn test_get_nonexistent() {
let kg = new_kg();
assert!(kg.get_entity("nonexistent").unwrap().is_none());
}
#[test]
fn test_delete_entity() {
let kg = new_kg();
kg.create_entities(&[Entity {
name: "del".into(),
entity_type: "t".into(),
observations: vec![],
}])
.unwrap();
assert!(kg.get_entity("del").unwrap().is_some());
kg.delete_entities(&["del".to_string()]).unwrap();
assert!(kg.get_entity("del").unwrap().is_none());
}
#[test]
fn test_add_and_delete_observations() {
let kg = new_kg();
kg.create_entities(&[Entity {
name: "obs_test".into(),
entity_type: "t".into(),
observations: vec!["a".into()],
}])
.unwrap();
let added = kg
.add_observations("obs_test", &["b".into(), "c".into()])
.unwrap();
assert_eq!(added.len(), 2);
let ent = kg.get_entity("obs_test").unwrap().unwrap();
assert!(ent.observations.iter().any(|o| o.body == "b"));
assert!(ent.observations.iter().any(|o| o.body == "c"));
kg.delete_observations("obs_test", &["b".into()]).unwrap();
let ent = kg.get_entity("obs_test").unwrap().unwrap();
assert!(!ent.observations.iter().any(|o| o.body == "b"));
assert!(ent.observations.iter().any(|o| o.body == "c"));
assert!(ent.observations.iter().any(|o| o.body == "a"));
}
#[test]
fn test_create_relations() {
let kg = new_kg();
kg.create_entities(&[
Entity {
name: "A".into(),
entity_type: "node".into(),
observations: vec![],
},
Entity {
name: "B".into(),
entity_type: "node".into(),
observations: vec![],
},
])
.unwrap();
let rels = kg
.create_relations(&[Relation {
from: "A".into(),
to: "B".into(),
relation_type: "edge".into(),
}])
.unwrap();
assert_eq!(rels.len(), 1);
assert_eq!(kg.get_entity_count().unwrap(), 2);
assert_eq!(kg.get_relation_count().unwrap(), 1);
}
#[test]
fn test_search_nodes() {
let kg = new_kg();
kg.create_entities(&[Entity {
name: "Einstein".into(),
entity_type: "scientist".into(),
observations: vec!["physics".into(), "relativity".into()],
}])
.unwrap();
let results = kg.search_nodes_filtered("physics", None, 0, 10);
assert_eq!(results.len(), 1);
assert_eq!(results[0].name, "Einstein");
let results = kg.search_nodes_filtered("physics", Some("scientist"), 0, 10);
assert_eq!(results.len(), 1);
let results = kg.search_nodes_filtered("physics", Some("nonexistent"), 0, 10);
assert_eq!(results.len(), 0);
}
#[test]
fn test_find_path() {
let kg = new_kg();
kg.create_entities(&[
Entity {
name: "A".into(),
entity_type: "n".into(),
observations: vec![],
},
Entity {
name: "B".into(),
entity_type: "n".into(),
observations: vec![],
},
Entity {
name: "C".into(),
entity_type: "n".into(),
observations: vec![],
},
])
.unwrap();
kg.create_relations(&[
Relation {
from: "A".into(),
to: "B".into(),
relation_type: "e".into(),
},
Relation {
from: "B".into(),
to: "C".into(),
relation_type: "e".into(),
},
])
.unwrap();
let path = kg.find_path("A", "C").unwrap().unwrap();
assert_eq!(path, vec!["A", "B", "C"]);
}
#[test]
fn test_degree() {
let kg = new_kg();
kg.create_entities(&[
Entity {
name: "A".into(),
entity_type: "n".into(),
observations: vec![],
},
Entity {
name: "B".into(),
entity_type: "n".into(),
observations: vec![],
},
Entity {
name: "C".into(),
entity_type: "n".into(),
observations: vec![],
},
])
.unwrap();
kg.create_relations(&[
Relation {
from: "A".into(),
to: "B".into(),
relation_type: "e".into(),
},
Relation {
from: "A".into(),
to: "C".into(),
relation_type: "e".into(),
},
])
.unwrap();
assert_eq!(kg.degree("A", Direction::Outgoing).unwrap(), 2);
assert_eq!(kg.degree("A", Direction::Incoming).unwrap(), 0);
assert_eq!(kg.degree("B", Direction::Incoming).unwrap(), 1);
}
#[test]
fn test_neighbors() {
let kg = new_kg();
kg.create_entities(&[
Entity {
name: "A".into(),
entity_type: "n".into(),
observations: vec![],
},
Entity {
name: "B".into(),
entity_type: "n".into(),
observations: vec![],
},
])
.unwrap();
kg.create_relations(&[Relation {
from: "A".into(),
to: "B".into(),
relation_type: "e".into(),
}])
.unwrap();
let result = kg.neighbors("A", Direction::Outgoing, None, 1).unwrap();
let v: Value = serde_json::from_str(&result).unwrap();
assert_eq!(v["entities"].as_array().unwrap().len(), 2);
assert_eq!(v["relations"].as_array().unwrap().len(), 1);
}
#[test]
fn test_open_nodes() {
let kg = new_kg();
kg.create_entities(&[
Entity {
name: "X".into(),
entity_type: "n".into(),
observations: vec!["obs_x".into()],
},
Entity {
name: "Y".into(),
entity_type: "n".into(),
observations: vec!["obs_y".into()],
},
])
.unwrap();
kg.create_relations(&[Relation {
from: "X".into(),
to: "Y".into(),
relation_type: "e".into(),
}])
.unwrap();
let result = kg.open_nodes(&["X".into()]);
let v: Value = serde_json::from_str(&result).unwrap();
assert_eq!(v["entities"].as_array().unwrap().len(), 1);
assert_eq!(v["relations"].as_array().unwrap().len(), 1);
}
#[test]
fn test_entities_exist() {
let kg = new_kg();
kg.create_entities(&[Entity {
name: "exists".into(),
entity_type: "t".into(),
observations: vec![],
}])
.unwrap();
let res = kg
.entities_exist(&["exists".into(), "missing".into()])
.unwrap();
assert_eq!(res, vec![true, false]);
}
#[test]
fn test_describe_entity() {
let kg = new_kg();
kg.create_entities(&[
Entity {
name: "A".into(),
entity_type: "t".into(),
observations: vec!["o".into()],
},
Entity {
name: "B".into(),
entity_type: "t".into(),
observations: vec![],
},
Entity {
name: "C".into(),
entity_type: "t".into(),
observations: vec![],
},
])
.unwrap();
kg.create_relations(&[
Relation {
from: "B".into(),
to: "A".into(),
relation_type: "inbound".into(),
},
Relation {
from: "A".into(),
to: "B".into(),
relation_type: "outbound".into(),
},
Relation {
from: "A".into(),
to: "C".into(),
relation_type: "other".into(),
},
Relation {
from: "A".into(),
to: "A".into(),
relation_type: "self".into(),
},
])
.unwrap();
let entity = kg.describe_entity("A").unwrap();
assert_eq!(entity.name, "A");
assert_eq!(entity.entity_type, "t");
assert_eq!(
entity
.observations
.iter()
.map(|o| o.body.as_str())
.collect::<Vec<_>>(),
["o"]
);
assert_eq!(entity.relations.len(), 4);
assert_eq!(
entity.relations,
vec![
Relation {
from: "A".into(),
to: "A".into(),
relation_type: "self".into(),
},
Relation {
from: "A".into(),
to: "B".into(),
relation_type: "outbound".into(),
},
Relation {
from: "A".into(),
to: "C".into(),
relation_type: "other".into(),
},
Relation {
from: "B".into(),
to: "A".into(),
relation_type: "inbound".into(),
},
]
);
assert_eq!(entity.neighbors, ["A", "B", "C"]);
assert_eq!(entity.degree.incoming, 2);
assert_eq!(entity.degree.outgoing, 3);
kg.writer
.lock()
.execute(
"UPDATE entity SET out_deg = 99, in_deg = 88 WHERE name = 'A'",
[],
)
.unwrap();
let entity = kg.describe_entity("A").unwrap();
assert_eq!(entity.degree.incoming, 2);
assert_eq!(entity.degree.outgoing, 3);
assert!(kg.describe_entity("missing").is_err());
}
#[test]
fn test_entity_type_counts() {
let kg = new_kg();
kg.create_entities(&[
Entity {
name: "a".into(),
entity_type: "person".into(),
observations: vec![],
},
Entity {
name: "b".into(),
entity_type: "person".into(),
observations: vec![],
},
Entity {
name: "c".into(),
entity_type: "place".into(),
observations: vec![],
},
])
.unwrap();
let counts = kg.entity_type_counts();
let map: FxHashMap<_, _> = counts.into_iter().collect();
assert_eq!(map.get("person"), Some(&2));
assert_eq!(map.get("place"), Some(&1));
}
#[test]
fn test_relation_type_counts() {
let kg = new_kg();
kg.create_entities(&[
Entity {
name: "a".into(),
entity_type: "n".into(),
observations: vec![],
},
Entity {
name: "b".into(),
entity_type: "n".into(),
observations: vec![],
},
Entity {
name: "c".into(),
entity_type: "n".into(),
observations: vec![],
},
])
.unwrap();
kg.create_relations(&[
Relation {
from: "a".into(),
to: "b".into(),
relation_type: "knows".into(),
},
Relation {
from: "a".into(),
to: "c".into(),
relation_type: "knows".into(),
},
])
.unwrap();
let counts = kg.relation_type_counts();
let map: FxHashMap<_, _> = counts.into_iter().collect();
assert_eq!(map.get("knows"), Some(&2));
}
#[test]
fn test_upsert_entities() {
let kg = new_kg();
kg.create_entities(&[Entity {
name: "A".into(),
entity_type: "OldType".into(),
observations: vec!["old".into()],
}])
.unwrap();
kg.create_relations(&[Relation {
from: "A".into(),
to: "A".into(),
relation_type: "self".into(),
}])
.unwrap();
kg.upsert_entities(&[Entity {
name: "A".into(),
entity_type: "NewType".into(),
observations: vec!["old".into(), "new".into()],
}])
.unwrap();
assert_eq!(kg.get_entity_count().unwrap(), 1);
let ent = kg.get_entity("A").unwrap().unwrap();
assert_eq!(ent.entity_type, "NewType");
assert_eq!(
ent.observations
.iter()
.map(|o| o.body.as_str())
.collect::<Vec<_>>(),
["old", "new"]
);
let type_counts: FxHashMap<_, _> = kg.entity_type_counts().into_iter().collect();
assert_eq!(type_counts.get("OldType"), None);
assert_eq!(type_counts.get("NewType"), Some(&1));
assert_eq!(
kg.search_relations(Some("A"), Some("A"), Some("self"), None),
[Relation {
from: "A".into(),
to: "A".into(),
relation_type: "self".into(),
}]
);
}
#[test]
fn test_merge_entities() {
let kg = new_kg();
kg.create_entities(&[
Entity {
name: "source".into(),
entity_type: "t".into(),
observations: vec!["src_obs".into()],
},
Entity {
name: "target".into(),
entity_type: "t".into(),
observations: vec!["tgt_obs".into()],
},
])
.unwrap();
kg.create_relations(&[Relation {
from: "source".into(),
to: "target".into(),
relation_type: "e".into(),
}])
.unwrap();
let merged = kg.merge_entities("source", "target").unwrap();
assert_eq!(merged.name, "target");
assert!(kg.get_entity("source").unwrap().is_none());
}
#[test]
fn test_find_all_paths() {
let kg = new_kg();
kg.create_entities(&[
Entity {
name: "A".into(),
entity_type: "n".into(),
observations: vec![],
},
Entity {
name: "B".into(),
entity_type: "n".into(),
observations: vec![],
},
Entity {
name: "C".into(),
entity_type: "n".into(),
observations: vec![],
},
])
.unwrap();
kg.create_relations(&[
Relation {
from: "A".into(),
to: "B".into(),
relation_type: "e".into(),
},
Relation {
from: "B".into(),
to: "C".into(),
relation_type: "e".into(),
},
Relation {
from: "A".into(),
to: "C".into(),
relation_type: "e".into(),
},
])
.unwrap();
let paths = kg.find_all_paths("A", "C", 5, 10).unwrap();
assert!(paths.len() >= 2);
}
#[test]
fn test_batch_get_entities() {
let kg = new_kg();
kg.create_entities(&[
Entity {
name: "a".into(),
entity_type: "t".into(),
observations: vec![],
},
Entity {
name: "b".into(),
entity_type: "t".into(),
observations: vec![],
},
])
.unwrap();
let results = kg.batch_get_entities(&["a".into(), "missing".into(), "b".into()]);
assert_eq!(results.len(), 3);
assert!(results[0].is_some());
assert!(results[1].is_none());
assert!(results[2].is_some());
}
#[test]
fn test_export_graph() {
let kg = new_kg();
kg.create_entities(&[Entity {
name: "exp".into(),
entity_type: "t".into(),
observations: vec!["o".into()],
}])
.unwrap();
let exported = kg.export("json", i64::MAX).unwrap();
assert!(exported.contains("exp"));
assert!(exported.contains("o"));
}
#[test]
fn test_graph_stats() {
let kg = new_kg();
assert_eq!(kg.get_entity_count().unwrap(), 0);
assert_eq!(kg.get_relation_count().unwrap(), 0);
kg.create_entities(&[Entity {
name: "s".into(),
entity_type: "t".into(),
observations: vec![],
}])
.unwrap();
assert_eq!(kg.get_entity_count().unwrap(), 1);
}
#[test]
fn test_read_graph_filtered() {
let kg = new_kg();
kg.create_entities(&[
Entity {
name: "p1".into(),
entity_type: "person".into(),
observations: vec![],
},
Entity {
name: "p2".into(),
entity_type: "place".into(),
observations: vec![],
},
])
.unwrap();
let out = kg.read_graph_filtered(Some("person"), 0, 10).unwrap();
let v: Value = serde_json::from_str(&out).unwrap();
assert_eq!(v["entities"].as_array().unwrap().len(), 1);
assert_eq!(v["entities"][0]["name"], "p1");
}
#[test]
fn test_wipe() {
let kg = new_kg();
kg.create_entities(&[Entity {
name: "w".into(),
entity_type: "t".into(),
observations: vec!["o".into()],
}])
.unwrap();
assert_eq!(kg.get_entity_count().unwrap(), 1);
kg.wipe().unwrap();
assert_eq!(kg.get_entity_count().unwrap(), 0);
}
#[test]
fn test_push_json_str() {
let mut buf = String::new();
push_json_str(&mut buf, "hello");
assert_eq!(buf, "\"hello\"");
let mut buf = String::new();
push_json_str(&mut buf, "he\"llo");
assert_eq!(buf, "\"he\\\"llo\"");
}
#[test]
fn test_create_entities_empty_input() {
let kg = new_kg();
let created = kg.create_entities(&[]).unwrap();
assert!(created.is_empty());
}
#[test]
fn test_create_entities_skip_empty_name() {
let kg = new_kg();
let created = kg
.create_entities(&[Entity {
name: "".into(),
entity_type: "t".into(),
observations: vec![],
}])
.unwrap();
assert!(created.is_empty());
assert_eq!(kg.get_entity_count().unwrap(), 0);
}
#[test]
fn test_create_entities_duplicate_names() {
let kg = new_kg();
let e = Entity {
name: "dup".into(),
entity_type: "t".into(),
observations: vec!["obs".into()],
};
let first = kg.create_entities(std::slice::from_ref(&e)).unwrap();
assert_eq!(first.len(), 1);
let second = kg.create_entities(&[e]).unwrap();
assert!(second.is_empty());
assert_eq!(kg.get_entity_count().unwrap(), 1);
}
#[test]
fn test_create_entities_partial_duplicates() {
let kg = new_kg();
let created = kg
.create_entities(&[
Entity {
name: "a".into(),
entity_type: "t".into(),
observations: vec![],
},
Entity {
name: "b".into(),
entity_type: "t".into(),
observations: vec![],
},
])
.unwrap();
assert_eq!(created.len(), 2);
let second = kg
.create_entities(&[
Entity {
name: "b".into(),
entity_type: "t".into(),
observations: vec![],
},
Entity {
name: "c".into(),
entity_type: "t".into(),
observations: vec![],
},
])
.unwrap();
assert_eq!(second.len(), 1); assert_eq!(second[0].name, "c");
assert_eq!(kg.get_entity_count().unwrap(), 3);
}
#[test]
fn test_create_entities_mixed_empty_and_valid() {
let kg = new_kg();
let created = kg
.create_entities(&[
Entity {
name: "".into(),
entity_type: "t".into(),
observations: vec![],
},
Entity {
name: "valid".into(),
entity_type: "t".into(),
observations: vec![],
},
Entity {
name: "".into(),
entity_type: "t".into(),
observations: vec![],
},
])
.unwrap();
assert_eq!(created.len(), 1);
assert_eq!(created[0].name, "valid");
assert_eq!(kg.get_entity_count().unwrap(), 1);
}
#[test]
fn test_create_entities_same_name_in_batch() {
let kg = new_kg();
let created = kg
.create_entities(&[
Entity {
name: "dup_in_batch".into(),
entity_type: "t".into(),
observations: vec![],
},
Entity {
name: "dup_in_batch".into(),
entity_type: "t".into(),
observations: vec![],
},
])
.unwrap();
assert_eq!(created.len(), 1);
assert_eq!(kg.get_entity_count().unwrap(), 1);
}
#[test]
fn test_create_relations_empty_input() {
let kg = new_kg();
let rels = kg.create_relations(&[]).unwrap();
assert!(rels.is_empty());
}
#[test]
fn test_create_relations_nonexistent_from() {
let kg = new_kg();
kg.create_entities(&[Entity {
name: "B".into(),
entity_type: "t".into(),
observations: vec![],
}])
.unwrap();
let rels = kg
.create_relations(&[Relation {
from: "A".into(),
to: "B".into(),
relation_type: "e".into(),
}])
.unwrap();
assert!(rels.is_empty());
assert_eq!(kg.get_relation_count().unwrap(), 0);
}
#[test]
fn test_create_relations_nonexistent_to() {
let kg = new_kg();
kg.create_entities(&[Entity {
name: "A".into(),
entity_type: "t".into(),
observations: vec![],
}])
.unwrap();
let rels = kg
.create_relations(&[Relation {
from: "A".into(),
to: "B".into(),
relation_type: "e".into(),
}])
.unwrap();
assert!(rels.is_empty());
assert_eq!(kg.get_relation_count().unwrap(), 0);
}
#[test]
fn test_create_relations_both_nonexistent() {
let kg = new_kg();
let rels = kg
.create_relations(&[Relation {
from: "A".into(),
to: "B".into(),
relation_type: "e".into(),
}])
.unwrap();
assert!(rels.is_empty());
}
#[test]
fn test_create_relations_self_loop() {
let kg = new_kg();
kg.create_entities(&[Entity {
name: "self".into(),
entity_type: "t".into(),
observations: vec![],
}])
.unwrap();
let rels = kg
.create_relations(&[Relation {
from: "self".into(),
to: "self".into(),
relation_type: "loop".into(),
}])
.unwrap();
assert_eq!(rels.len(), 1);
assert_eq!(kg.get_relation_count().unwrap(), 1);
assert_eq!(kg.degree("self", Direction::Outgoing).unwrap(), 1);
assert_eq!(kg.degree("self", Direction::Incoming).unwrap(), 1);
}
#[test]
fn test_create_relations_duplicate() {
let kg = new_kg();
kg.create_entities(&[
Entity {
name: "A".into(),
entity_type: "t".into(),
observations: vec![],
},
Entity {
name: "B".into(),
entity_type: "t".into(),
observations: vec![],
},
])
.unwrap();
let r = Relation {
from: "A".into(),
to: "B".into(),
relation_type: "e".into(),
};
let first = kg.create_relations(std::slice::from_ref(&r)).unwrap();
assert_eq!(first.len(), 1);
let second = kg.create_relations(&[r]).unwrap();
assert!(second.is_empty());
assert_eq!(kg.get_relation_count().unwrap(), 1);
}
#[test]
fn test_create_relations_new_type_auto_created() {
let kg = new_kg();
kg.create_entities(&[
Entity {
name: "A".into(),
entity_type: "t".into(),
observations: vec![],
},
Entity {
name: "B".into(),
entity_type: "t".into(),
observations: vec![],
},
])
.unwrap();
let rels = kg
.create_relations(&[Relation {
from: "A".into(),
to: "B".into(),
relation_type: "brand_new_type".into(),
}])
.unwrap();
assert_eq!(rels.len(), 1);
let counts = kg.relation_type_counts();
let map: FxHashMap<_, _> = counts.into_iter().collect();
assert_eq!(map.get("brand_new_type"), Some(&1));
}
#[test]
fn test_create_relations_degree_updates() {
let kg = new_kg();
kg.create_entities(&[
Entity {
name: "A".into(),
entity_type: "t".into(),
observations: vec![],
},
Entity {
name: "B".into(),
entity_type: "t".into(),
observations: vec![],
},
Entity {
name: "C".into(),
entity_type: "t".into(),
observations: vec![],
},
])
.unwrap();
kg.create_relations(&[
Relation {
from: "A".into(),
to: "B".into(),
relation_type: "e".into(),
},
Relation {
from: "A".into(),
to: "C".into(),
relation_type: "e".into(),
},
])
.unwrap();
assert_eq!(kg.degree("A", Direction::Outgoing).unwrap(), 2);
assert_eq!(kg.degree("A", Direction::Incoming).unwrap(), 0);
assert_eq!(kg.degree("B", Direction::Incoming).unwrap(), 1);
assert_eq!(kg.degree("C", Direction::Incoming).unwrap(), 1);
assert_eq!(kg.degree("A", Direction::Both).unwrap(), 2);
}
#[test]
fn test_create_relations_delete_and_recreate() {
let kg = new_kg();
kg.create_entities(&[
Entity {
name: "A".into(),
entity_type: "t".into(),
observations: vec![],
},
Entity {
name: "B".into(),
entity_type: "t".into(),
observations: vec![],
},
])
.unwrap();
let r = Relation {
from: "A".into(),
to: "B".into(),
relation_type: "e".into(),
};
kg.create_relations(std::slice::from_ref(&r)).unwrap();
assert_eq!(kg.get_relation_count().unwrap(), 1);
kg.delete_relations(std::slice::from_ref(&r)).unwrap();
assert_eq!(kg.get_relation_count().unwrap(), 0);
let re = kg.create_relations(&[r]).unwrap();
assert_eq!(re.len(), 1);
assert_eq!(kg.get_relation_count().unwrap(), 1);
}
#[test]
fn test_create_entities_then_relations_then_delete_entity_with_relations() {
let kg = new_kg();
kg.create_entities(&[
Entity {
name: "A".into(),
entity_type: "t".into(),
observations: vec![],
},
Entity {
name: "B".into(),
entity_type: "t".into(),
observations: vec![],
},
])
.unwrap();
kg.create_relations(&[Relation {
from: "A".into(),
to: "B".into(),
relation_type: "e".into(),
}])
.unwrap();
assert_eq!(kg.get_relation_count().unwrap(), 1);
kg.delete_entities(&["A".into()]).unwrap();
assert!(kg.get_entity("A").unwrap().is_none());
assert_eq!(kg.get_relation_count().unwrap(), 0);
}
#[test]
fn test_graph_stats_after_entity_with_observations() {
let kg = new_kg();
kg.create_entities(&[Entity {
name: "stat".into(),
entity_type: "t".into(),
observations: vec!["o1".into(), "o2".into(), "o3".into()],
}])
.unwrap();
let ecount = kg.get_entity_count().unwrap();
assert_eq!(ecount, 1);
kg.delete_entities(&["stat".into()]).unwrap();
assert_eq!(kg.get_entity_count().unwrap(), 0);
}
fn new_kg_with_pool(read_pool_size: usize) -> TestKg {
use std::sync::atomic::AtomicU64;
static COUNTER: AtomicU64 = AtomicU64::new(1_000_000);
let n = COUNTER.fetch_add(1, Ordering::SeqCst);
let path = std::env::temp_dir().join(format!("kg_pool_{}_{}.db", std::process::id(), n));
cleanup_db(&path);
let kg = GraphHandle::new(
&path,
Durability::Async,
SqliteTuning::default(),
NonZeroUsize::new(10_000).unwrap(),
read_pool_size,
)
.expect("create KG");
TestKg(kg, path)
}
fn seed_line(kg: &GraphHandle, n: usize) {
let entities: Vec<Entity> = (0..n)
.map(|i| Entity {
name: format!("n{i}"),
entity_type: "node".into(),
observations: vec![format!("obs of n{i}").into()],
})
.collect();
kg.create_entities(&entities).unwrap();
let rels: Vec<Relation> = (0..n.saturating_sub(1))
.map(|i| Relation {
from: format!("n{i}"),
to: format!("n{}", i + 1),
relation_type: "edge".into(),
})
.collect();
if !rels.is_empty() {
kg.create_relations(&rels).unwrap();
}
}
fn count_relations(graph_json: &str) -> usize {
let v: Value = serde_json::from_str(graph_json).unwrap();
v["relations"].as_array().unwrap().len()
}
fn count_entities(graph_json: &str) -> usize {
let v: Value = serde_json::from_str(graph_json).unwrap();
v["entities"].as_array().unwrap().len()
}
#[test]
fn test_pool_size_one_still_works() {
let kg = new_kg_with_pool(1);
seed_line(&kg, 5);
assert_eq!(kg.get_entity_count().unwrap(), 5);
assert!(kg.get_entity("n2").unwrap().is_some());
let g = kg.read_graph_filtered(None, 0, usize::MAX).unwrap();
assert_eq!(count_entities(&g), 5);
}
#[test]
fn test_reads_see_committed_writes() {
let kg = new_kg_with_pool(4);
kg.create_entities(&[Entity {
name: "fresh".into(),
entity_type: "t".into(),
observations: vec!["v".into()],
}])
.unwrap();
let got = kg.get_entity("fresh").unwrap().unwrap();
assert_eq!(
got.observations
.iter()
.map(|o| o.body.as_str())
.collect::<Vec<_>>(),
vec!["v"]
);
}
#[test]
fn test_concurrent_readers_consistent() {
let kg = new_kg_with_pool(4);
seed_line(&kg, 50);
std::thread::scope(|s| {
for _ in 0..8 {
s.spawn(|| {
for _ in 0..200 {
let _ = kg.get_entity("n10");
let _ = kg.search_nodes_filtered("obs", None, 0, 10);
let _ = kg.read_graph_filtered(None, 0, 100);
let _ = kg.get_entity_count();
let _ = kg.neighbors("n10", Direction::Both, None, 2);
}
});
}
s.spawn(|| {
for i in 100..160 {
kg.create_entities(&[Entity {
name: format!("w{i}"),
entity_type: "node".into(),
observations: vec![format!("w obs {i}").into()],
}])
.unwrap();
}
});
});
assert_eq!(kg.get_entity_count().unwrap(), 110);
assert!(kg.get_entity("w159").unwrap().is_some());
}
#[test]
fn test_reader_pool_rejects_writes_internally() {
let kg = new_kg_with_pool(1);
seed_line(&kg, 3);
std::thread::scope(|s| {
for _ in 0..4 {
s.spawn(|| {
for _ in 0..100 {
let _ = kg.read_graph_filtered(None, 0, 10);
}
});
}
});
assert_eq!(kg.get_entity_count().unwrap(), 3);
}
#[test]
fn test_read_graph_relations_scoped_to_page() {
let kg = new_kg_with_pool(2);
seed_line(&kg, 4);
let full = kg.read_graph_filtered(None, 0, usize::MAX).unwrap();
assert_eq!(count_entities(&full), 4);
assert_eq!(count_relations(&full), 3);
let page1 = kg.read_graph_filtered(None, 0, 1).unwrap();
assert_eq!(count_entities(&page1), 1);
assert_eq!(count_relations(&page1), 0);
let page2 = kg.read_graph_filtered(None, 0, 2).unwrap();
assert_eq!(count_entities(&page2), 2);
assert_eq!(count_relations(&page2), 1);
}
#[test]
fn test_read_graph_pagination_offset() {
let kg = new_kg_with_pool(2);
seed_line(&kg, 5);
let g = kg.read_graph_filtered(None, 2, 2).unwrap();
assert_eq!(count_entities(&g), 2);
assert!(!g.contains("\"n0\""));
assert!(!g.contains("\"n1\""));
assert!(g.contains("\"n2\""));
assert!(g.contains("\"n3\""));
}
#[test]
fn test_read_graph_empty() {
let kg = new_kg_with_pool(2);
let g = kg.read_graph_filtered(None, 0, usize::MAX).unwrap();
assert_eq!(g, r#"{"entities":[],"relations":[]}"#);
}
#[test]
fn test_read_graph_filtered_by_type() {
let kg = new_kg_with_pool(2);
kg.create_entities(&[
Entity {
name: "p1".into(),
entity_type: "person".into(),
observations: vec![],
},
Entity {
name: "q1".into(),
entity_type: "place".into(),
observations: vec![],
},
Entity {
name: "p2".into(),
entity_type: "person".into(),
observations: vec![],
},
])
.unwrap();
let g = kg
.read_graph_filtered(Some("person"), 0, usize::MAX)
.unwrap();
assert_eq!(count_entities(&g), 2);
assert!(g.contains("\"p1\""));
assert!(g.contains("\"p2\""));
assert!(!g.contains("\"q1\""));
}
#[test]
fn test_export_respects_max_rows() {
let kg = new_kg_with_pool(2);
seed_line(&kg, 5);
let full = kg.export("json", i64::MAX).unwrap();
assert_eq!(count_entities(&full), 5);
assert_eq!(count_relations(&full), 4);
let capped = kg.export("json", 2).unwrap();
assert_eq!(count_entities(&capped), 2);
assert_eq!(count_relations(&capped), 2);
}
#[test]
fn test_export_negative_max_rows_is_unbounded() {
let kg = new_kg_with_pool(2);
seed_line(&kg, 3);
let out = kg.export("json", -1).unwrap();
assert_eq!(count_entities(&out), 3);
}
#[test]
fn test_many_small_write_batches_stay_consistent() {
let kg = new_kg_with_pool(2);
for i in 0..100 {
kg.create_entities(&[Entity {
name: format!("e{i}"),
entity_type: "t".into(),
observations: vec![format!("o{i}").into()],
}])
.unwrap();
}
assert_eq!(kg.get_entity_count().unwrap(), 100);
let hits = kg.search_nodes_filtered("e57", None, 0, 10);
assert!(hits.iter().any(|e| e.name == "e57"));
}
#[test]
fn test_wipe_clears_name_and_obs_fts() {
let kg = new_kg_with_pool(2);
kg.create_entities(&[Entity {
name: "Einstein".into(),
entity_type: "scientist".into(),
observations: vec!["physics".into()],
}])
.unwrap();
assert_eq!(kg.search_nodes_filtered("Einstein", None, 0, 10).len(), 1);
assert_eq!(kg.search_nodes_filtered("physics", None, 0, 10).len(), 1);
kg.wipe().unwrap();
assert_eq!(kg.get_entity_count().unwrap(), 0);
assert!(kg.search_nodes_filtered("Einstein", None, 0, 10).is_empty());
assert!(kg.search_nodes_filtered("physics", None, 0, 10).is_empty());
}
#[test]
fn test_wipe_then_recreate_search_works() {
let kg = new_kg_with_pool(2);
kg.create_entities(&[Entity {
name: "Einstein".into(),
entity_type: "scientist".into(),
observations: vec!["physics".into()],
}])
.unwrap();
kg.wipe().unwrap();
kg.create_entities(&[Entity {
name: "Einstein".into(),
entity_type: "scientist".into(),
observations: vec!["physics".into(), "relativity".into()],
}])
.unwrap();
let by_name = kg.search_nodes_filtered("Einstein", None, 0, 10);
assert_eq!(by_name.len(), 1, "exactly one Einstein after recreate");
let by_obs = kg.search_nodes_filtered("relativity", None, 0, 10);
assert_eq!(by_obs.len(), 1);
assert_eq!(kg.get_entity_count().unwrap(), 1);
}
#[test]
fn test_search_relations_missing_type_returns_empty() {
let kg = new_kg_with_pool(2);
seed_line(&kg, 3); let r = kg.search_relations(None, None, Some("does_not_exist"), None);
assert!(r.is_empty());
let types = kg.relation_type_counts();
assert!(types.iter().all(|(t, _)| t != "does_not_exist"));
}
#[test]
fn test_search_relations_missing_from_returns_empty() {
let kg = new_kg_with_pool(2);
seed_line(&kg, 3);
let r = kg.search_relations(Some("ghost"), None, None, None);
assert!(r.is_empty(), "missing 'from' must not match every relation");
}
#[test]
fn test_search_relations_existing_filters_still_work() {
let kg = new_kg_with_pool(2);
seed_line(&kg, 3);
let r = kg.search_relations(Some("n0"), None, Some("edge"), None);
assert_eq!(r.len(), 1);
assert_eq!(r[0].from, "n0");
assert_eq!(r[0].to, "n1");
}
#[test]
fn test_neighbors_missing_type_returns_only_start() {
let kg = new_kg_with_pool(2);
seed_line(&kg, 3);
let json = kg
.neighbors("n0", Direction::Both, Some("nonexistent"), 2)
.unwrap();
assert_eq!(count_entities(&json), 1);
assert_eq!(count_relations(&json), 0);
}
#[test]
fn test_neighbors_existing_type_filters() {
let kg = new_kg_with_pool(2);
kg.create_entities(&[
Entity {
name: "a".into(),
entity_type: "n".into(),
observations: vec![],
},
Entity {
name: "b".into(),
entity_type: "n".into(),
observations: vec![],
},
Entity {
name: "c".into(),
entity_type: "n".into(),
observations: vec![],
},
])
.unwrap();
kg.create_relations(&[
Relation {
from: "a".into(),
to: "b".into(),
relation_type: "knows".into(),
},
Relation {
from: "a".into(),
to: "c".into(),
relation_type: "likes".into(),
},
])
.unwrap();
let json = kg
.neighbors("a", Direction::Outgoing, Some("knows"), 1)
.unwrap();
assert!(json.contains("\"b\""));
assert!(!json.contains("\"c\""));
assert_eq!(count_relations(&json), 1);
}
#[test]
fn test_sqlite_tuning_applied_to_fresh_db() {
use std::sync::atomic::AtomicU64;
static COUNTER: AtomicU64 = AtomicU64::new(2_000_000);
let n = COUNTER.fetch_add(1, Ordering::SeqCst);
let path = std::env::temp_dir().join(format!("kg_tuning_{}_{}.db", std::process::id(), n));
cleanup_db(&path);
let tuning = SqliteTuning {
page_size: 8192,
..SqliteTuning::default()
};
let kg = TestKg(
GraphHandle::new(
&path,
Durability::Async,
tuning,
NonZeroUsize::new(64).unwrap(),
2,
)
.expect("create KG"),
path.clone(),
);
kg.create_entities(&[Entity {
name: "a".into(),
entity_type: "n".into(),
observations: vec!["o".into()],
}])
.unwrap();
let probe = Connection::open(&path).unwrap();
let page_size: i64 = probe
.query_row("PRAGMA page_size", [], |r| r.get(0))
.unwrap();
assert_eq!(page_size, 8192);
let auto_vacuum: i64 = probe
.query_row("PRAGMA auto_vacuum", [], |r| r.get(0))
.unwrap();
assert_eq!(auto_vacuum, 2, "expected INCREMENTAL auto_vacuum");
let journal: String = probe
.query_row("PRAGMA journal_mode", [], |r| r.get(0))
.unwrap();
assert_eq!(journal.to_lowercase(), "wal");
}
#[test]
fn test_checkpoint_passive_is_noop_safe() {
let kg = new_kg();
kg.checkpoint_passive().unwrap();
kg.create_entities(&[Entity {
name: "a".into(),
entity_type: "n".into(),
observations: vec!["o".into()],
}])
.unwrap();
kg.checkpoint_passive().unwrap();
kg.checkpoint_passive().unwrap();
assert!(kg.get_entity("a").unwrap().is_some());
}
}