use std::collections::{HashMap, HashSet, VecDeque};
use std::sync::Arc;
use async_trait::async_trait;
use chrono::{DateTime, TimeZone, Utc};
use rusqlite::OptionalExtension;
use uuid::Uuid;
use khive_storage::error::StorageError;
use khive_storage::types::{
BatchWriteSummary, DeleteMode, DirectedNeighborHit, Direction, Edge, EdgeEndpointBaseCounts,
EdgeFilter, EdgeSeekPage, EdgeSortField, EdgeUpsertDisposition, EdgeUpsertRefusal,
EdgeUpsertRequest, EdgeUpsertResult, GraphPath, GuardedBatchOutcome, GuardedBatchRefusal,
GuardedEdgeBatchRefusal, GuardedEdgeBatchUpsertOutcome, GuardedEdgeUpsertOutcome,
GuardedWriteOutcome, MissingEndpoints, NeighborCursor, NeighborHit, NeighborQuery, Page,
PageRequest, PathNode, SeekCursor, SeekPage, SortDirection, SortOrder, SqlStatement, SqlValue,
TraversalExecutionBudget, TraversalOptions, TraversalRequest,
};
use khive_storage::GraphStore;
use khive_storage::LinkId;
use khive_storage::StorageCapability;
use khive_types::EdgeRelation;
use crate::error::SqliteError;
use crate::pool::ConnectionPool;
use crate::sql_bridge::bind_params;
use crate::writer_task::WriterTaskHandle;
fn map_err(e: rusqlite::Error, op: &'static str) -> StorageError {
StorageError::driver(StorageCapability::Graph, op, e)
}
fn map_sqlite_err(e: SqliteError, op: &'static str) -> StorageError {
StorageError::driver(StorageCapability::Graph, op, e)
}
fn resurrection_required_error(operation: &'static str, edge: &Edge) -> StorageError {
StorageError::Conflict {
capability: StorageCapability::Graph,
operation: operation.into(),
message: format!(
"edge {} is soft-deleted; explicit resurrection is required",
edge.id
),
}
}
const NAMESPACE_COUNT_CHUNK_SIZE: usize = 500;
const LATEST_ANNOTATING_NOTE_SQL: &str = r#"WITH incident AS MATERIALIZED (
SELECT source_id, deleted_at
FROM graph_edges INDEXED BY idx_graph_edges_ns_tgt_rel
WHERE namespace = ?1 AND target_id = ?2 AND relation = 'annotates'
LIMIT 3
)
SELECT result.id, result.created_at
FROM notes AS result
WHERE result.id = CASE WHEN (SELECT count(*) FROM incident) <= 2 THEN (
SELECT n.id
FROM incident AS e CROSS JOIN notes AS n
WHERE n.id = e.source_id AND e.deleted_at IS NULL
AND n.deleted_at IS NULL AND n.kind = ?3
AND EXISTS (SELECT 1 FROM json_each(CASE
WHEN json_type(n.properties, '$.tags') = 'array'
THEN json_extract(n.properties, '$.tags') ELSE '[]' END) AS tag
WHERE tag.type = 'text' AND tag.value = ?4 COLLATE BINARY)
ORDER BY n.created_at DESC, n.id ASC LIMIT 1
) ELSE (
SELECT n.id
FROM notes AS n INDEXED BY idx_notes_created
WHERE n.deleted_at IS NULL AND n.kind = ?3
AND EXISTS (SELECT 1 FROM json_each(CASE
WHEN json_type(n.properties, '$.tags') = 'array'
THEN json_extract(n.properties, '$.tags') ELSE '[]' END) AS tag
WHERE tag.type = 'text' AND tag.value = ?4 COLLATE BINARY)
AND EXISTS (SELECT 1 FROM graph_edges AS e INDEXED BY idx_graph_edges_unique_triple
WHERE e.namespace = ?1 AND e.source_id = n.id AND e.target_id = ?2
AND e.relation = 'annotates' AND e.deleted_at IS NULL)
ORDER BY n.created_at DESC, n.id ASC LIMIT 1
) END"#;
fn edge_conflict_clause(resurrect: bool) -> String {
let deleted_at = if resurrect {
"NULL"
} else {
"graph_edges.deleted_at"
};
let predicate = if resurrect {
""
} else {
" WHERE graph_edges.deleted_at IS NULL"
};
format!(
"weight = excluded.weight, \
updated_at = excluded.updated_at, \
deleted_at = {deleted_at}, \
metadata = excluded.metadata, \
target_backend = excluded.target_backend{predicate}"
)
}
fn endpoint_exists_clause(id_param: &str) -> String {
format!(
"EXISTS (SELECT 1 FROM entities WHERE id = {id_param} AND deleted_at IS NULL) \
OR EXISTS (SELECT 1 FROM notes WHERE id = {id_param} AND deleted_at IS NULL) \
OR EXISTS (SELECT 1 FROM events WHERE id = {id_param}) \
OR EXISTS (SELECT 1 FROM graph_edges WHERE id = {id_param} AND deleted_at IS NULL)"
)
}
pub fn edge_snapshot_assertion_statement(edge: &Edge, require_endpoints: bool) -> SqlStatement {
let mut sql = "SELECT 1 FROM graph_edges WHERE id=?1 AND namespace=?2 \
AND source_id=?3 AND target_id=?4 AND relation=?5 \
AND updated_at=?6 AND deleted_at IS NULL"
.to_string();
if require_endpoints {
sql.push_str(&format!(
" AND ({}) AND ({})",
endpoint_exists_clause("?3"),
endpoint_exists_clause("?4")
));
}
SqlStatement {
sql,
params: vec![
SqlValue::Text(Uuid::from(edge.id).to_string()),
SqlValue::Text(edge.namespace.clone()),
SqlValue::Text(edge.source_id.to_string()),
SqlValue::Text(edge.target_id.to_string()),
SqlValue::Text(edge.relation.to_string()),
SqlValue::Integer(edge.updated_at.timestamp_micros()),
],
label: Some("edge-snapshot-assertion".into()),
}
}
pub fn edge_upsert_statement(edge: &Edge) -> SqlStatement {
edge_upsert_statement_with_resurrection(edge, false)
}
pub fn edge_upsert_statement_with_resurrection(edge: &Edge, resurrect: bool) -> SqlStatement {
let (source_id, target_id) =
canonical_edge_endpoints(edge.relation, edge.source_id, edge.target_id);
let metadata_str = edge
.metadata
.as_ref()
.map(|v| serde_json::to_string(v).unwrap_or_default());
let conflict_clause = edge_conflict_clause(resurrect);
SqlStatement {
sql: format!(
"INSERT INTO graph_edges \
(namespace, id, source_id, target_id, relation, weight, \
created_at, updated_at, deleted_at, metadata, target_backend) \
VALUES (?1, ?2, ?3, ?4, ?5, ?6, ?7, ?8, ?9, ?10, ?11) \
ON CONFLICT(namespace, id) DO UPDATE SET \
source_id = excluded.source_id, \
target_id = excluded.target_id, \
relation = excluded.relation, \
{conflict_clause} \
ON CONFLICT(namespace, source_id, target_id, relation) DO UPDATE SET \
{conflict_clause}"
),
params: vec![
SqlValue::Text(edge.namespace.clone()),
SqlValue::Text(Uuid::from(edge.id).to_string()),
SqlValue::Text(source_id.to_string()),
SqlValue::Text(target_id.to_string()),
SqlValue::Text(edge.relation.to_string()),
SqlValue::Float(edge.weight),
SqlValue::Integer(edge.created_at.timestamp_micros()),
SqlValue::Integer(edge.updated_at.timestamp_micros()),
match edge.deleted_at {
Some(t) => SqlValue::Integer(t.timestamp_micros()),
None => SqlValue::Null,
},
match metadata_str {
Some(m) => SqlValue::Text(m),
None => SqlValue::Null,
},
match &edge.target_backend {
Some(b) => SqlValue::Text(b.clone()),
None => SqlValue::Null,
},
],
label: Some("edge-upsert".to_string()),
}
}
pub fn edge_insert_only_guarded_by_endpoints_statement(edge: &Edge) -> SqlStatement {
let mut statement = edge_upsert_statement(edge);
let src_exists = endpoint_exists_clause("?3");
let tgt_exists = endpoint_exists_clause("?4");
statement.sql = format!(
"INSERT INTO graph_edges \
(namespace, id, source_id, target_id, relation, weight, \
created_at, updated_at, deleted_at, metadata, target_backend) \
SELECT ?1, ?2, ?3, ?4, ?5, ?6, ?7, ?8, ?9, ?10, ?11 \
WHERE ({src_exists}) AND ({tgt_exists})"
);
statement.label = Some("edge-insert-only-where-endpoints-exist".to_string());
statement
}
pub fn edge_insert_if_absent_statement(edge: &Edge) -> SqlStatement {
let mut statement = edge_upsert_statement(edge);
statement.sql = "INSERT INTO graph_edges \
(namespace, id, source_id, target_id, relation, weight, \
created_at, updated_at, deleted_at, metadata, target_backend) \
VALUES (?1, ?2, ?3, ?4, ?5, ?6, ?7, ?8, ?9, ?10, ?11) \
ON CONFLICT DO NOTHING"
.to_string();
statement.label = Some("edge-insert-if-absent".to_string());
statement
}
pub fn edge_replace_if_unchanged_statement(
edge: &Edge,
expected_updated_at: DateTime<Utc>,
expected_deleted_at: Option<DateTime<Utc>>,
) -> SqlStatement {
let (source_id, target_id) =
canonical_edge_endpoints(edge.relation, edge.source_id, edge.target_id);
let metadata_str = edge
.metadata
.as_ref()
.map(|v| serde_json::to_string(v).unwrap_or_default());
SqlStatement {
sql: "UPDATE graph_edges SET \
namespace = ?1, source_id = ?2, target_id = ?3, relation = ?4, weight = ?5, \
updated_at = ?6, deleted_at = ?7, metadata = ?8, target_backend = ?9 \
WHERE id = ?10 AND updated_at = ?11 AND deleted_at IS ?12 \
AND ?6 > updated_at"
.to_string(),
params: vec![
SqlValue::Text(edge.namespace.clone()),
SqlValue::Text(source_id.to_string()),
SqlValue::Text(target_id.to_string()),
SqlValue::Text(edge.relation.to_string()),
SqlValue::Float(edge.weight),
SqlValue::Integer(edge.updated_at.timestamp_micros()),
match edge.deleted_at {
Some(t) => SqlValue::Integer(t.timestamp_micros()),
None => SqlValue::Null,
},
match metadata_str {
Some(m) => SqlValue::Text(m),
None => SqlValue::Null,
},
match &edge.target_backend {
Some(b) => SqlValue::Text(b.clone()),
None => SqlValue::Null,
},
SqlValue::Text(Uuid::from(edge.id).to_string()),
SqlValue::Integer(expected_updated_at.timestamp_micros()),
match expected_deleted_at {
Some(value) => SqlValue::Integer(value.timestamp_micros()),
None => SqlValue::Null,
},
],
label: Some("edge-replace-if-unchanged".to_string()),
}
}
#[allow(clippy::too_many_arguments)]
pub fn edge_insert_guarded_by_endpoints_statement(
namespace: &str,
edge_id: Uuid,
source_id: Uuid,
target_id: Uuid,
relation: EdgeRelation,
weight: f64,
now: i64,
metadata: Option<&str>,
) -> SqlStatement {
edge_insert_guarded_by_endpoints_with_resurrection_statement(
namespace, source_id, target_id, edge_id, relation, weight, now, metadata, false,
)
}
#[allow(clippy::too_many_arguments)]
pub fn edge_insert_guarded_by_endpoints_with_resurrection_statement(
namespace: &str,
source_id: Uuid,
target_id: Uuid,
edge_id: Uuid,
relation: EdgeRelation,
weight: f64,
now: i64,
metadata: Option<&str>,
resurrect: bool,
) -> SqlStatement {
let src_exists = endpoint_exists_clause("?3");
let tgt_exists = endpoint_exists_clause("?4");
let conflict_clause = edge_conflict_clause(resurrect);
SqlStatement {
sql: format!(
"INSERT INTO graph_edges \
(namespace, id, source_id, target_id, relation, weight, \
created_at, updated_at, metadata) \
SELECT ?1, ?2, ?3, ?4, ?5, ?6, ?7, ?7, ?8 \
WHERE ({src_exists}) AND ({tgt_exists}) \
ON CONFLICT(namespace, source_id, target_id, relation) DO UPDATE SET \
{conflict_clause}"
),
params: vec![
SqlValue::Text(namespace.to_string()),
SqlValue::Text(edge_id.to_string()),
SqlValue::Text(source_id.to_string()),
SqlValue::Text(target_id.to_string()),
SqlValue::Text(relation.as_str().to_string()),
SqlValue::Float(weight),
SqlValue::Integer(now),
match metadata {
Some(m) => SqlValue::Text(m.to_string()),
None => SqlValue::Null,
},
],
label: Some("atomic-link-insert-edge-where-exists".to_string()),
}
}
#[allow(clippy::too_many_arguments)]
pub fn edge_insert_new_guarded_by_endpoints_statement(
namespace: &str,
edge_id: Uuid,
source_id: Uuid,
target_id: Uuid,
relation: EdgeRelation,
weight: f64,
now: i64,
metadata: Option<&str>,
) -> SqlStatement {
let src_exists = endpoint_exists_clause("?3");
let tgt_exists = endpoint_exists_clause("?4");
SqlStatement {
sql: format!(
"INSERT INTO graph_edges \
(namespace, id, source_id, target_id, relation, weight, \
created_at, updated_at, metadata) \
SELECT ?1, ?2, ?3, ?4, ?5, ?6, ?7, ?7, ?8 \
WHERE ({src_exists}) AND ({tgt_exists}) \
ON CONFLICT DO NOTHING"
),
params: vec![
SqlValue::Text(namespace.to_string()),
SqlValue::Text(edge_id.to_string()),
SqlValue::Text(source_id.to_string()),
SqlValue::Text(target_id.to_string()),
SqlValue::Text(relation.to_string()),
SqlValue::Float(weight),
SqlValue::Integer(now),
metadata
.map(|value| SqlValue::Text(value.to_string()))
.unwrap_or(SqlValue::Null),
],
label: Some("edge-link-create-if-absent-and-endpoints-exist".to_string()),
}
}
pub fn edge_link_replace_if_unchanged_and_endpoints_exist_statement(
previous: &Edge,
weight: f64,
now: i64,
metadata: Option<&str>,
) -> SqlStatement {
let src_exists = endpoint_exists_clause("?6");
let tgt_exists = endpoint_exists_clause("?7");
SqlStatement {
sql: format!(
"UPDATE graph_edges SET \
weight = ?1, updated_at = ?2, deleted_at = NULL, \
metadata = ?3, target_backend = NULL \
WHERE namespace = ?4 AND id = ?5 \
AND source_id = ?6 AND target_id = ?7 AND relation = ?8 \
AND updated_at = ?9 AND deleted_at IS ?10 AND ?2 > updated_at \
AND ({src_exists}) AND ({tgt_exists})"
),
params: vec![
SqlValue::Float(weight),
SqlValue::Integer(now),
metadata
.map(|value| SqlValue::Text(value.to_string()))
.unwrap_or(SqlValue::Null),
SqlValue::Text(previous.namespace.clone()),
SqlValue::Text(Uuid::from(previous.id).to_string()),
SqlValue::Text(previous.source_id.to_string()),
SqlValue::Text(previous.target_id.to_string()),
SqlValue::Text(previous.relation.to_string()),
SqlValue::Integer(previous.updated_at.timestamp_micros()),
previous
.deleted_at
.map(|value| SqlValue::Integer(value.timestamp_micros()))
.unwrap_or(SqlValue::Null),
],
label: Some("edge-link-replace-if-unchanged-and-endpoints-exist".to_string()),
}
}
pub fn edge_soft_delete_statement(id: Uuid, now: i64) -> SqlStatement {
SqlStatement {
sql: "UPDATE graph_edges SET deleted_at = ?2, updated_at = ?2 \
WHERE id = ?1 AND deleted_at IS NULL"
.to_string(),
params: vec![SqlValue::Text(id.to_string()), SqlValue::Integer(now)],
label: Some("edge-delete-soft".to_string()),
}
}
pub fn edge_hard_delete_statement(id: Uuid) -> SqlStatement {
SqlStatement {
sql: "DELETE FROM graph_edges WHERE id = ?1".to_string(),
params: vec![SqlValue::Text(id.to_string())],
label: Some("edge-delete-hard".to_string()),
}
}
pub fn purge_incident_edges_statement(node_id: Uuid) -> SqlStatement {
SqlStatement {
sql: "DELETE FROM graph_edges WHERE source_id = ?1 OR target_id = ?1".to_string(),
params: vec![SqlValue::Text(node_id.to_string())],
label: Some("edge-purge-incident".to_string()),
}
}
pub const EDGE_SYMMETRIC_CONFLICT_PROBE_SQL: &str = "SELECT id FROM graph_edges \
WHERE namespace = ?1 AND source_id = ?2 AND target_id = ?3 \
AND relation = ?4 AND id != ?5";
pub const EDGE_SYMMETRIC_DELETE_NONCANONICAL_SQL: &str =
"DELETE FROM graph_edges WHERE namespace = ?1 AND id = ?2";
pub const EDGE_SYMMETRIC_DELETE_NONCANONICAL_GUARDED_SQL: &str =
"DELETE FROM graph_edges WHERE namespace = ?1 AND id = ?2 \
AND updated_at = ?3 AND deleted_at IS ?4";
pub const EDGE_SYMMETRIC_UPDATE_INPLACE_SQL: &str = "UPDATE graph_edges SET \
source_id = ?1, target_id = ?2, relation = ?3, \
weight = ?4, updated_at = ?5, metadata = ?6 \
WHERE namespace = ?7 AND id = ?8 \
AND updated_at = ?9 AND deleted_at IS ?10 \
AND ?5 > updated_at";
pub fn edge_symmetric_conflict_probe_statement(
namespace: &str,
canon_src: Uuid,
canon_tgt: Uuid,
relation: EdgeRelation,
exclude_id: Uuid,
) -> SqlStatement {
SqlStatement {
sql: EDGE_SYMMETRIC_CONFLICT_PROBE_SQL.to_string(),
params: vec![
SqlValue::Text(namespace.to_string()),
SqlValue::Text(canon_src.to_string()),
SqlValue::Text(canon_tgt.to_string()),
SqlValue::Text(relation.to_string()),
SqlValue::Text(exclude_id.to_string()),
],
label: Some("edge-symmetric-conflict-probe".to_string()),
}
}
pub fn edge_symmetric_delete_noncanonical_statement(namespace: &str, id: Uuid) -> SqlStatement {
SqlStatement {
sql: EDGE_SYMMETRIC_DELETE_NONCANONICAL_SQL.to_string(),
params: vec![
SqlValue::Text(namespace.to_string()),
SqlValue::Text(id.to_string()),
],
label: Some("edge-symmetric-delete-noncanonical".to_string()),
}
}
#[allow(clippy::too_many_arguments)]
pub fn edge_symmetric_update_inplace_statement(
namespace: &str,
id: Uuid,
canon_src: Uuid,
canon_tgt: Uuid,
relation: EdgeRelation,
weight: f64,
updated_at_micros: i64,
metadata: Option<&str>,
expected_updated_at_micros: i64,
expected_deleted_at_micros: Option<i64>,
) -> SqlStatement {
SqlStatement {
sql: EDGE_SYMMETRIC_UPDATE_INPLACE_SQL.to_string(),
params: vec![
SqlValue::Text(canon_src.to_string()),
SqlValue::Text(canon_tgt.to_string()),
SqlValue::Text(relation.to_string()),
SqlValue::Float(weight),
SqlValue::Integer(updated_at_micros),
match metadata {
Some(m) => SqlValue::Text(m.to_string()),
None => SqlValue::Null,
},
SqlValue::Text(namespace.to_string()),
SqlValue::Text(id.to_string()),
SqlValue::Integer(expected_updated_at_micros),
match expected_deleted_at_micros {
Some(value) => SqlValue::Integer(value),
None => SqlValue::Null,
},
],
label: Some("edge-symmetric-update-inplace".to_string()),
}
}
#[allow(clippy::too_many_arguments)]
pub fn edge_symmetric_delete_if_conflict_statement(
namespace: &str,
id: Uuid,
canon_src: Uuid,
canon_tgt: Uuid,
relation: EdgeRelation,
expected_updated_at_micros: i64,
expected_deleted_at_micros: Option<i64>,
) -> SqlStatement {
SqlStatement {
sql: "DELETE FROM graph_edges \
WHERE namespace = ?1 AND id = ?2 \
AND updated_at = ?6 AND deleted_at IS ?7 \
AND EXISTS ( \
SELECT 1 FROM graph_edges \
WHERE namespace = ?1 AND source_id = ?3 AND target_id = ?4 \
AND relation = ?5 AND id != ?2 \
)"
.to_string(),
params: vec![
SqlValue::Text(namespace.to_string()),
SqlValue::Text(id.to_string()),
SqlValue::Text(canon_src.to_string()),
SqlValue::Text(canon_tgt.to_string()),
SqlValue::Text(relation.to_string()),
SqlValue::Integer(expected_updated_at_micros),
match expected_deleted_at_micros {
Some(value) => SqlValue::Integer(value),
None => SqlValue::Null,
},
],
label: Some("edge-symmetric-delete-if-conflict".to_string()),
}
}
#[allow(clippy::too_many_arguments)]
pub fn edge_symmetric_absorb_or_update_inplace_statement(
namespace: &str,
id: Uuid,
canon_src: Uuid,
canon_tgt: Uuid,
relation: EdgeRelation,
weight: f64,
updated_at_micros: i64,
metadata: Option<&str>,
target_backend: Option<&str>,
expected_updated_at_micros: i64,
expected_deleted_at_micros: Option<i64>,
) -> SqlStatement {
SqlStatement {
sql: "UPDATE graph_edges SET \
source_id = CASE WHEN id = ?2 THEN ?3 ELSE source_id END, \
target_id = CASE WHEN id = ?2 THEN ?4 ELSE target_id END, \
relation = CASE WHEN id = ?2 THEN ?5 ELSE relation END, \
weight = CASE WHEN id = ?2 THEN ?6 ELSE weight END, \
updated_at = CASE WHEN id = ?2 THEN ?7 ELSE updated_at END, \
deleted_at = CASE WHEN id = ?2 THEN NULL ELSE deleted_at END, \
metadata = CASE WHEN id = ?2 THEN ?8 ELSE metadata END, \
target_backend = CASE WHEN id = ?2 THEN ?9 ELSE target_backend END \
WHERE namespace = ?1 \
AND ( \
(id = ?2 AND changes() = 0 AND updated_at = ?10 AND deleted_at IS ?11 \
AND ?7 > updated_at) \
OR (source_id = ?3 AND target_id = ?4 AND relation = ?5 \
AND id != ?2 AND changes() = 1) \
)"
.to_string(),
params: vec![
SqlValue::Text(namespace.to_string()),
SqlValue::Text(id.to_string()),
SqlValue::Text(canon_src.to_string()),
SqlValue::Text(canon_tgt.to_string()),
SqlValue::Text(relation.to_string()),
SqlValue::Float(weight),
SqlValue::Integer(updated_at_micros),
match metadata {
Some(m) => SqlValue::Text(m.to_string()),
None => SqlValue::Null,
},
match target_backend {
Some(b) => SqlValue::Text(b.to_string()),
None => SqlValue::Null,
},
SqlValue::Integer(expected_updated_at_micros),
match expected_deleted_at_micros {
Some(value) => SqlValue::Integer(value),
None => SqlValue::Null,
},
],
label: Some("edge-symmetric-absorb-or-update-inplace".to_string()),
}
}
pub struct SqlGraphStore {
pool: Arc<ConnectionPool>,
is_file_backed: bool,
namespace: String,
writer_task: Option<WriterTaskHandle>,
}
impl SqlGraphStore {
pub fn new_scoped(
pool: Arc<ConnectionPool>,
is_file_backed: bool,
namespace: impl Into<String>,
) -> Self {
let writer_task = pool.writer_task_handle().ok().flatten();
Self {
pool,
is_file_backed,
namespace: namespace.into(),
writer_task,
}
}
fn open_standalone_writer(&self) -> Result<rusqlite::Connection, StorageError> {
self.pool
.open_standalone_writer()
.map_err(|e| map_sqlite_err(e, "open_graph_writer"))
}
fn current_writer_task(
&self,
operation: &'static str,
) -> Result<Option<WriterTaskHandle>, StorageError> {
self.pool
.writer_task_for_write(self.writer_task.as_ref(), operation)
}
async fn with_writer<F, R>(&self, op: &'static str, f: F) -> Result<R, StorageError>
where
F: FnOnce(&rusqlite::Connection) -> Result<R, rusqlite::Error> + Send + 'static,
R: Send + 'static,
{
if let Some(writer_task) = self.current_writer_task(op)? {
return writer_task
.send_bounded(move |conn| f(conn).map_err(|e| map_err(e, op)))
.await;
}
self.pool
.record_direct_route(crate::timeout_sink::Site::DirectRouteGraphGeneralWrite);
if self.is_file_backed {
let conn = self.open_standalone_writer()?;
let db = crate::timeout_sink::db_label(&self.pool);
tokio::task::spawn_blocking(move || {
f(&conn).map_err(|e| {
crate::timeout_sink::maybe_emit_busy(
&db,
crate::timeout_sink::Site::StandaloneGraph,
&e,
);
map_err(e, op)
})
})
.await
.map_err(|e| StorageError::driver(StorageCapability::Graph, op, e))?
} else {
let pool = Arc::clone(&self.pool);
tokio::task::spawn_blocking(move || {
let guard = pool.try_writer().map_err(|e| map_sqlite_err(e, op))?;
f(guard.conn()).map_err(|e| map_err(e, op))
})
.await
.map_err(|e| StorageError::driver(StorageCapability::Graph, op, e))?
}
}
async fn observed_edge_write(
&self,
operation: &'static str,
request: EdgeUpsertRequest,
guard_endpoints: bool,
) -> Result<GuardedEdgeUpsertOutcome, StorageError> {
if let Some(writer_task) = self.current_writer_task(operation)? {
return writer_task
.send_bounded(move |conn| {
observed_edge_upsert(conn, &request, guard_endpoints)
.map_err(|error| map_err(error, operation))
})
.await;
}
let origin = self.pool.origin();
self.with_writer(operation, move |conn| {
conn.execute_batch("BEGIN IMMEDIATE")?;
let _tx_handle =
khive_storage::tx_registry::register_scoped(Some(operation.to_string()), origin);
let outcome = match observed_edge_upsert(conn, &request, guard_endpoints) {
Ok(outcome) => outcome,
Err(error) => {
let _ = conn.execute_batch("ROLLBACK");
return Err(error);
}
};
if let Err(error) = conn.execute_batch("COMMIT") {
let _ = conn.execute_batch("ROLLBACK");
return Err(error);
}
Ok(outcome)
})
.await
}
async fn observed_edge_batch_write(
&self,
operation: &'static str,
requests: Vec<EdgeUpsertRequest>,
guard_endpoints: bool,
) -> Result<GuardedEdgeBatchUpsertOutcome, StorageError> {
if let Some(writer_task) = self.current_writer_task(operation)? {
return writer_task
.send_bounded(move |conn| {
observed_edge_batch_upsert(conn, &requests, guard_endpoints)
.map_err(|error| map_err(error, operation))
})
.await;
}
let origin = self.pool.origin();
self.with_writer(operation, move |conn| {
conn.execute_batch("BEGIN IMMEDIATE")?;
let _tx_handle =
khive_storage::tx_registry::register_scoped(Some(operation.to_string()), origin);
let outcome = match observed_edge_batch_upsert(conn, &requests, guard_endpoints) {
Ok(outcome) => outcome,
Err(error) => {
let _ = conn.execute_batch("ROLLBACK");
return Err(error);
}
};
if let Err(error) = conn.execute_batch("COMMIT") {
let _ = conn.execute_batch("ROLLBACK");
return Err(error);
}
Ok(outcome)
})
.await
}
async fn with_reader<F, R>(&self, op: &'static str, f: F) -> Result<R, StorageError>
where
F: FnOnce(&rusqlite::Connection) -> Result<R, rusqlite::Error> + Send + 'static,
R: Send + 'static,
{
super::run_pooled_store_read(
Arc::clone(&self.pool),
StorageCapability::Graph,
op,
move |conn| f(conn).map_err(|error| map_err(error, op)),
)
.await
}
}
fn report_graph_usage(queries: &std::sync::atomic::AtomicU64, rows: &std::sync::atomic::AtomicU64) {
khive_storage::usage::count(
khive_storage::usage::UsageUnit::DbRoundTrips,
queries.load(std::sync::atomic::Ordering::Relaxed),
);
khive_storage::usage::count(
khive_storage::usage::UsageUnit::GraphHops,
rows.load(std::sync::atomic::Ordering::Relaxed),
);
}
fn read_edge(row: &rusqlite::Row<'_>) -> Result<Edge, rusqlite::Error> {
let namespace: String = row.get(0)?;
let id_str: String = row.get(1)?;
let source_str: String = row.get(2)?;
let target_str: String = row.get(3)?;
let relation_str: String = row.get(4)?;
let weight: f64 = row.get(5)?;
let created_micros: i64 = row.get(6)?;
let updated_micros: i64 = row.get(7)?;
let deleted_micros: Option<i64> = row.get(8)?;
let metadata_str: Option<String> = row.get(9)?;
let target_backend: Option<String> = row.get(10)?;
let id = parse_uuid(&id_str)?;
let source_id = parse_uuid(&source_str)?;
let target_id = parse_uuid(&target_str)?;
let created_at = micros_to_datetime(created_micros);
let relation = relation_str.parse::<EdgeRelation>().map_err(|e| {
rusqlite::Error::FromSqlConversionFailure(4, rusqlite::types::Type::Text, Box::new(e))
})?;
let metadata = match metadata_str {
Some(s) => {
let v = serde_json::from_str(&s).map_err(|e| {
rusqlite::Error::FromSqlConversionFailure(
9,
rusqlite::types::Type::Text,
Box::new(e),
)
})?;
Some(v)
}
None => None,
};
Ok(Edge {
id: id.into(),
namespace,
source_id,
target_id,
relation,
weight,
created_at,
updated_at: micros_to_datetime(updated_micros),
deleted_at: deleted_micros.map(micros_to_datetime),
metadata,
target_backend,
})
}
fn parse_uuid(s: &str) -> Result<Uuid, rusqlite::Error> {
Uuid::parse_str(s).map_err(|e| {
rusqlite::Error::FromSqlConversionFailure(0, rusqlite::types::Type::Text, Box::new(e))
})
}
fn neighbor_extra_clause(
query: &NeighborQuery,
start_param_idx: usize,
after: Option<&NeighborCursor>,
neighbor_kinds: Option<&[String]>,
) -> (String, String, Vec<Box<dyn rusqlite::types::ToSql>>) {
let mut conditions: Vec<String> = Vec::new();
let mut extra_params: Vec<Box<dyn rusqlite::types::ToSql>> = Vec::new();
let mut param_idx = start_param_idx;
if let Some(ref rels) = query.relations {
if !rels.is_empty() {
let placeholders: Vec<String> = rels
.iter()
.map(|r| {
extra_params.push(Box::new(r.to_string()));
let p = format!("?{}", param_idx);
param_idx += 1;
p
})
.collect();
conditions.push(format!("relation IN ({})", placeholders.join(",")));
}
}
if let Some(min_w) = query.min_weight {
extra_params.push(Box::new(min_w));
conditions.push(format!("weight >= ?{}", param_idx));
param_idx += 1;
}
if let Some(cursor) = after {
extra_params.push(Box::new(cursor.weight));
let weight_idx = param_idx;
param_idx += 1;
extra_params.push(Box::new(cursor.node_id.to_string()));
let node_idx = param_idx;
param_idx += 1;
extra_params.push(Box::new(cursor.edge_id.to_string()));
let edge_idx = param_idx;
param_idx += 1;
conditions.push(format!(
"(weight < ?{weight_idx} OR (weight = ?{weight_idx} AND node_id > ?{node_idx}) OR (weight = ?{weight_idx} AND node_id = ?{node_idx} AND edge_id > ?{edge_idx}))"
));
}
if let Some(kinds) = neighbor_kinds.filter(|kinds| !kinds.is_empty()) {
let placeholders: Vec<String> = kinds
.iter()
.map(|kind| {
extra_params.push(Box::new(kind.clone()));
let p = format!("?{param_idx}");
param_idx += 1;
p
})
.collect();
let entity_placeholders = placeholders.join(",");
let note_placeholders: Vec<String> = kinds
.iter()
.map(|kind| {
extra_params.push(Box::new(kind.clone()));
let p = format!("?{param_idx}");
param_idx += 1;
p
})
.collect();
conditions.push(format!(
"(EXISTS (SELECT 1 FROM entities AS neighbor_entities WHERE neighbor_entities.id = node_id AND neighbor_entities.namespace = ?1 AND neighbor_entities.deleted_at IS NULL AND neighbor_entities.kind IN ({entity_placeholders})) OR EXISTS (SELECT 1 FROM notes AS neighbor_notes WHERE neighbor_notes.id = node_id AND neighbor_notes.namespace = ?1 AND neighbor_notes.deleted_at IS NULL AND neighbor_notes.kind IN ({})))",
note_placeholders.join(",")
));
}
let where_extra = if conditions.is_empty() {
String::new()
} else {
format!(" WHERE {}", conditions.join(" AND "))
};
let limit_clause = if let Some(lim) = query.limit {
extra_params.push(Box::new(lim as i64));
format!(" LIMIT ?{}", param_idx)
} else {
String::new()
};
(where_extra, limit_clause, extra_params)
}
#[cfg(test)]
static NEIGHBOR_SELECT_COUNT: std::sync::atomic::AtomicUsize =
std::sync::atomic::AtomicUsize::new(0);
#[cfg(test)]
fn count_neighbor_select() {
NEIGHBOR_SELECT_COUNT.fetch_add(1, std::sync::atomic::Ordering::Relaxed);
}
#[cfg(not(test))]
fn count_neighbor_select() {}
#[cfg(test)]
pub(crate) fn reset_neighbor_select_count() {
NEIGHBOR_SELECT_COUNT.store(0, std::sync::atomic::Ordering::Relaxed);
}
#[cfg(test)]
pub(crate) fn neighbor_select_count() -> usize {
NEIGHBOR_SELECT_COUNT.load(std::sync::atomic::Ordering::Relaxed)
}
fn micros_to_datetime(micros: i64) -> DateTime<Utc> {
Utc.timestamp_micros(micros)
.single()
.unwrap_or_else(Utc::now)
}
fn edge_order_clause(sort: &[SortOrder<EdgeSortField>]) -> String {
if sort.is_empty() {
return " ORDER BY created_at DESC, id DESC".to_string();
}
let mut parts: Vec<String> = sort
.iter()
.map(|s| {
let dir = match s.direction {
SortDirection::Asc => "ASC",
SortDirection::Desc => "DESC",
};
format!("{} {}", edge_sort_col(&s.field), dir)
})
.collect();
let dir = match sort.last().map(|s| &s.direction) {
Some(SortDirection::Asc) => "ASC",
_ => "DESC",
};
parts.push(format!("id {dir}"));
format!(" ORDER BY {}", parts.join(", "))
}
fn endpoint_base_case(column: &str) -> String {
format!(
"CASE WHEN EXISTS (SELECT 1 FROM entities be WHERE be.id = graph_edges.{column}) \
THEN 'entity' \
WHEN EXISTS (SELECT 1 FROM notes bn WHERE bn.id = graph_edges.{column}) \
THEN 'note' ELSE 'none' END"
)
}
fn fold_endpoint_base_row(counts: &mut EdgeEndpointBaseCounts, source: &str, target: &str, n: u64) {
let slot = match (source, target) {
("entity", "entity") => &mut counts.entity_entity,
("entity", "note") => &mut counts.entity_note,
("note", "entity") => &mut counts.note_entity,
("note", "note") => &mut counts.note_note,
_ => &mut counts.unresolved,
};
*slot = slot.saturating_add(n);
}
const LIVE_ENDPOINTS_CONDITION: &str = "NOT EXISTS (SELECT 1 FROM entities le \
WHERE le.id = graph_edges.source_id AND le.deleted_at IS NOT NULL) \
AND NOT EXISTS (SELECT 1 FROM entities le \
WHERE le.id = graph_edges.target_id AND le.deleted_at IS NOT NULL) \
AND NOT EXISTS (SELECT 1 FROM notes ln \
WHERE ln.id = graph_edges.source_id AND ln.deleted_at IS NOT NULL) \
AND NOT EXISTS (SELECT 1 FROM notes ln \
WHERE ln.id = graph_edges.target_id AND ln.deleted_at IS NOT NULL)";
fn with_live_endpoints(where_clause: &str) -> String {
format!("{where_clause} AND {LIVE_ENDPOINTS_CONDITION}")
}
fn build_edge_filter_sql(
namespace: &str,
filter: &EdgeFilter,
) -> (String, Vec<Box<dyn rusqlite::types::ToSql>>) {
build_edge_filter_sql_for_namespaces(&[namespace.to_string()], filter)
}
fn build_edge_filter_sql_for_namespaces(
namespaces: &[String],
filter: &EdgeFilter,
) -> (String, Vec<Box<dyn rusqlite::types::ToSql>>) {
let params: Vec<Box<dyn rusqlite::types::ToSql>> = namespaces
.iter()
.map(|namespace| -> Box<dyn rusqlite::types::ToSql> { Box::new(namespace.clone()) })
.collect();
let namespace_condition = match namespaces.len() {
0 => "0".to_string(),
1 => "namespace = ?1".to_string(),
_ => {
let placeholders: Vec<String> =
(1..=namespaces.len()).map(|i| format!("?{i}")).collect();
format!("namespace IN ({})", placeholders.join(", "))
}
};
build_edge_filter_conditions(namespace_condition, params, filter)
}
fn build_edge_filter_sql_for_namespaces_json(
namespaces_json: &str,
filter: &EdgeFilter,
) -> (String, Vec<Box<dyn rusqlite::types::ToSql>>) {
let params: Vec<Box<dyn rusqlite::types::ToSql>> = vec![Box::new(namespaces_json.to_string())];
let namespace_condition = "namespace IN (SELECT value FROM json_each(?1))".to_string();
build_edge_filter_conditions(namespace_condition, params, filter)
}
fn build_edge_filter_conditions(
namespace_condition: String,
mut params: Vec<Box<dyn rusqlite::types::ToSql>>,
filter: &EdgeFilter,
) -> (String, Vec<Box<dyn rusqlite::types::ToSql>>) {
let mut conditions = vec![namespace_condition, "deleted_at IS NULL".to_string()];
if !filter.ids.is_empty() {
let placeholders: Vec<String> = filter
.ids
.iter()
.map(|id| {
params.push(Box::new(id.to_string()));
format!("?{}", params.len())
})
.collect();
conditions.push(format!("id IN ({})", placeholders.join(",")));
}
if !filter.source_ids.is_empty() {
let placeholders: Vec<String> = filter
.source_ids
.iter()
.map(|id| {
params.push(Box::new(id.to_string()));
format!("?{}", params.len())
})
.collect();
conditions.push(format!("source_id IN ({})", placeholders.join(",")));
}
if !filter.target_ids.is_empty() {
let placeholders: Vec<String> = filter
.target_ids
.iter()
.map(|id| {
params.push(Box::new(id.to_string()));
format!("?{}", params.len())
})
.collect();
conditions.push(format!("target_id IN ({})", placeholders.join(",")));
}
if !filter.relations.is_empty() {
let placeholders: Vec<String> = filter
.relations
.iter()
.map(|r| {
params.push(Box::new(r.to_string()));
format!("?{}", params.len())
})
.collect();
conditions.push(format!("relation IN ({})", placeholders.join(",")));
}
if let Some(min_w) = filter.min_weight {
params.push(Box::new(min_w));
conditions.push(format!("weight >= ?{}", params.len()));
}
if let Some(max_w) = filter.max_weight {
params.push(Box::new(max_w));
conditions.push(format!("weight <= ?{}", params.len()));
}
if let Some(ref time_range) = filter.created_at {
if let Some(start) = time_range.start {
params.push(Box::new(start.timestamp_micros()));
conditions.push(format!("created_at >= ?{}", params.len()));
}
if let Some(end) = time_range.end {
params.push(Box::new(end.timestamp_micros()));
conditions.push(format!("created_at < ?{}", params.len()));
}
}
let clause = format!(" WHERE {}", conditions.join(" AND "));
(clause, params)
}
fn edge_sort_col(field: &EdgeSortField) -> &'static str {
match field {
EdgeSortField::CreatedAt => "created_at",
EdgeSortField::Weight => "weight",
EdgeSortField::Relation => "relation",
}
}
fn canonical_edge_endpoints(
relation: EdgeRelation,
source_id: Uuid,
target_id: Uuid,
) -> (Uuid, Uuid) {
if relation.is_symmetric() && target_id < source_id {
(target_id, source_id)
} else {
(source_id, target_id)
}
}
fn edge_endpoints_exist(
conn: &rusqlite::Connection,
source_id: Uuid,
target_id: Uuid,
) -> Result<MissingEndpoints, rusqlite::Error> {
let src_exists = endpoint_exists_clause("?1");
let tgt_exists = endpoint_exists_clause("?2");
let sql = format!("SELECT ({src_exists}), ({tgt_exists})");
conn.query_row(
&sql,
rusqlite::params![source_id.to_string(), target_id.to_string()],
|row| {
let src_exists: bool = row.get(0)?;
let tgt_exists: bool = row.get(1)?;
Ok(MissingEndpoints {
source: !src_exists,
target: !tgt_exists,
})
},
)
}
fn edge_by_natural_key_including_deleted(
conn: &rusqlite::Connection,
edge: &Edge,
) -> Result<Option<Edge>, rusqlite::Error> {
let (source_id, target_id) =
canonical_edge_endpoints(edge.relation, edge.source_id, edge.target_id);
conn.query_row(
"SELECT namespace, id, source_id, target_id, relation, weight, \
created_at, updated_at, deleted_at, metadata, target_backend \
FROM graph_edges \
WHERE namespace = ?1 AND source_id = ?2 AND target_id = ?3 AND relation = ?4",
rusqlite::params![
&edge.namespace,
source_id.to_string(),
target_id.to_string(),
edge.relation.as_str(),
],
read_edge,
)
.optional()
}
fn observed_edge_upsert(
conn: &rusqlite::Connection,
request: &EdgeUpsertRequest,
guard_endpoints: bool,
) -> Result<GuardedEdgeUpsertOutcome, rusqlite::Error> {
let (source_id, target_id) = canonical_edge_endpoints(
request.edge.relation,
request.edge.source_id,
request.edge.target_id,
);
if guard_endpoints {
#[cfg(test)]
tests::insert_probe_seam::hook((source_id, target_id));
let missing = edge_endpoints_exist(conn, source_id, target_id)?;
if missing.any() {
return Ok(GuardedEdgeUpsertOutcome::Refused(
EdgeUpsertRefusal::MissingEndpoints(missing),
));
}
}
let previous = edge_by_natural_key_including_deleted(conn, &request.edge)?;
if let Some(edge) = previous.as_ref() {
if edge.deleted_at.is_some() && !request.resurrect {
return Ok(GuardedEdgeUpsertOutcome::Refused(
EdgeUpsertRefusal::ResurrectionRequired { edge: edge.clone() },
));
}
}
let statement = edge_upsert_statement_with_resurrection(&request.edge, request.resurrect);
let mut stmt = conn.prepare(&statement.sql)?;
bind_params(&mut stmt, &statement.params)?;
let affected = stmt.raw_execute()?;
if affected == 0 {
let edge = edge_by_natural_key_including_deleted(conn, &request.edge)?
.ok_or(rusqlite::Error::QueryReturnedNoRows)?;
return Ok(GuardedEdgeUpsertOutcome::Refused(
EdgeUpsertRefusal::ResurrectionRequired { edge },
));
}
let edge = edge_by_natural_key_including_deleted(conn, &request.edge)?
.ok_or(rusqlite::Error::QueryReturnedNoRows)?;
let disposition = match previous.as_ref().and_then(|edge| edge.deleted_at) {
None if previous.is_none() => EdgeUpsertDisposition::Created,
None => EdgeUpsertDisposition::Updated,
Some(_) => EdgeUpsertDisposition::Resurrected,
};
Ok(GuardedEdgeUpsertOutcome::Written(EdgeUpsertResult {
edge,
disposition,
previous,
}))
}
fn observed_edge_batch_upsert(
conn: &rusqlite::Connection,
requests: &[EdgeUpsertRequest],
guard_endpoints: bool,
) -> Result<GuardedEdgeBatchUpsertOutcome, rusqlite::Error> {
for (entry_index, request) in requests.iter().enumerate() {
let (source_id, target_id) = canonical_edge_endpoints(
request.edge.relation,
request.edge.source_id,
request.edge.target_id,
);
if guard_endpoints {
let missing = edge_endpoints_exist(conn, source_id, target_id)?;
if missing.any() {
return Ok(GuardedEdgeBatchUpsertOutcome {
rows: Vec::new(),
refusal: Some(GuardedEdgeBatchRefusal {
entry_index,
reason: EdgeUpsertRefusal::MissingEndpoints(missing),
}),
});
}
}
if let Some(edge) = edge_by_natural_key_including_deleted(conn, &request.edge)? {
if edge.deleted_at.is_some() && !request.resurrect {
return Ok(GuardedEdgeBatchUpsertOutcome {
rows: Vec::new(),
refusal: Some(GuardedEdgeBatchRefusal {
entry_index,
reason: EdgeUpsertRefusal::ResurrectionRequired { edge },
}),
});
}
}
}
let mut rows = Vec::with_capacity(requests.len());
for request in requests {
match observed_edge_upsert(conn, request, false)? {
GuardedEdgeUpsertOutcome::Written(row) => rows.push(row),
GuardedEdgeUpsertOutcome::Refused(_) => {
return Err(rusqlite::Error::ExecuteReturnedResults)
}
}
}
Ok(GuardedEdgeBatchUpsertOutcome {
rows,
refusal: None,
})
}
fn traversal_neighbor_sql(
direction: Direction,
relation_count: usize,
has_min_weight: bool,
) -> String {
let (node_column, endpoint_column, index) = match direction {
Direction::Out => ("target_id", "source_id", "idx_graph_edges_ns_src_rel"),
Direction::In => ("source_id", "target_id", "idx_graph_edges_ns_tgt_rel"),
Direction::Both => unreachable!("Direction::Both is split into indexed Out/In seeks"),
};
let mut sql = format!(
"SELECT {node_column}, id, weight \
FROM graph_edges INDEXED BY {index} \
WHERE namespace = ?1 AND {endpoint_column} = ?2 AND deleted_at IS NULL"
);
if relation_count > 0 {
let placeholders = (0..relation_count)
.map(|offset| format!("?{}", 4 + offset))
.collect::<Vec<_>>()
.join(",");
sql.push_str(&format!(" AND relation IN ({placeholders})"));
}
if has_min_weight {
sql.push_str(&format!(" AND weight >= ?{}", 4 + relation_count));
}
sql.push_str(" LIMIT ?3");
sql
}
fn traversal_timeout_error(budget: &TraversalExecutionBudget) -> StorageError {
StorageError::Timeout {
operation: format!(
"traverse ({}ms execution budget)",
budget.max_duration().as_millis()
)
.into(),
}
}
fn traversal_work_error(budget: &TraversalExecutionBudget) -> StorageError {
StorageError::InvalidInput {
capability: StorageCapability::Graph,
operation: "traverse".into(),
message: format!(
"traversal work budget exceeded after {} adjacency rows; \
narrow roots, depth, relations, or result limit",
budget.work_limit()
),
}
}
#[derive(Clone, Copy)]
struct TraversalFrontierNode {
node_id: Uuid,
depth: usize,
total_weight: f64,
}
#[allow(clippy::too_many_arguments)]
fn run_bounded_traversal(
conn: &rusqlite::Connection,
roots: Vec<Uuid>,
opts: TraversalOptions,
include_roots: bool,
namespace: String,
origin: khive_storage::tx_registry::TxOrigin,
budget: TraversalExecutionBudget,
counted_rows: &std::sync::atomic::AtomicU64,
counted_queries: &std::sync::atomic::AtomicU64,
) -> Result<Vec<GraphPath>, StorageError> {
let progress_timed_out = Arc::new(std::sync::atomic::AtomicBool::new(false));
let callback_timed_out = Arc::clone(&progress_timed_out);
let callback_budget = budget.clone();
#[cfg(test)]
let progress_seam_root = roots.first().copied();
conn.progress_handler(
1_000,
Some(move || {
if crate::read_cancellation::current_read_should_interrupt() {
return true;
}
#[cfg(test)]
if tests::traverse_progress_seam::hook(progress_seam_root) {
callback_timed_out.store(true, std::sync::atomic::Ordering::Relaxed);
return true;
}
let expired = callback_budget.is_expired();
if expired {
callback_timed_out.store(true, std::sync::atomic::Ordering::Relaxed);
}
expired
}),
)
.map_err(|e| map_err(e, "traverse_progress_handler"))?;
let result = (|| {
let result_limit = opts.effective_limit() as usize;
let relation_count = opts.relations.as_ref().map_or(0, Vec::len);
let directions = match opts.direction {
Direction::Out => vec![Direction::Out],
Direction::In => vec![Direction::In],
Direction::Both => vec![Direction::Out, Direction::In],
};
let statements = directions
.into_iter()
.map(|direction| {
traversal_neighbor_sql(direction, relation_count, opts.min_weight.is_some())
})
.collect::<Vec<_>>();
let map_sql_error = |error| {
if progress_timed_out.load(std::sync::atomic::Ordering::Relaxed) {
traversal_timeout_error(&budget)
} else {
map_err(error, "traverse")
}
};
let mut all_paths = Vec::with_capacity(roots.len());
for root_id in roots {
let mut seen = HashSet::new();
seen.insert(root_id);
let mut frontier = VecDeque::from([TraversalFrontierNode {
node_id: root_id,
depth: 0,
total_weight: 0.0,
}]);
let mut nodes = Vec::with_capacity(result_limit + usize::from(include_roots));
if include_roots {
nodes.push(PathNode {
node_id: root_id,
via_edge: None,
depth: 0,
name: None,
kind: None,
properties: None,
weight: 0.0,
});
}
let mut non_root_count = 0usize;
'root_walk: while non_root_count < result_limit {
let Some(current) = frontier.pop_front() else {
break;
};
if current.depth >= opts.max_depth {
continue;
}
if budget.is_expired() {
return Err(traversal_timeout_error(&budget));
}
for sql in &statements {
let row_cap = budget.remaining_work().saturating_add(1);
let mut params: Vec<Box<dyn rusqlite::types::ToSql>> = vec![
Box::new(namespace.clone()),
Box::new(current.node_id.to_string()),
Box::new(row_cap as i64),
];
if let Some(relations) = &opts.relations {
params.extend(relations.iter().map(|relation| {
Box::new(relation.to_string()) as Box<dyn rusqlite::types::ToSql>
}));
}
if let Some(min_weight) = opts.min_weight {
params.push(Box::new(min_weight));
}
let param_refs = params
.iter()
.map(|param| param.as_ref())
.collect::<Vec<&dyn rusqlite::types::ToSql>>();
let _snapshot = khive_storage::tx_registry::register_scoped(
Some("graph_traverse_read".to_string()),
origin.clone(),
);
let mut stmt = conn.prepare(sql).map_err(&map_sql_error)?;
counted_queries.fetch_add(1, std::sync::atomic::Ordering::Relaxed);
let mut rows = stmt.query(param_refs.as_slice()).map_err(&map_sql_error)?;
while let Some(row) = rows.next().map_err(&map_sql_error)? {
#[cfg(test)]
tests::traverse_snapshot_seam::hook(current.node_id);
counted_rows.fetch_add(1, std::sync::atomic::Ordering::Relaxed);
if budget.is_expired() {
return Err(traversal_timeout_error(&budget));
}
if !budget.try_consume_row() {
return Err(traversal_work_error(&budget));
}
let node_str: String = row.get(0).map_err(&map_sql_error)?;
let edge_str: String = row.get(1).map_err(&map_sql_error)?;
let edge_weight: f64 = row.get(2).map_err(&map_sql_error)?;
let node_id = parse_uuid(&node_str).map_err(&map_sql_error)?;
if !seen.insert(node_id) {
continue;
}
let via_edge = parse_uuid(&edge_str).map_err(&map_sql_error)?;
let depth = current.depth + 1;
let total_weight = current.total_weight + edge_weight;
nodes.push(PathNode {
node_id,
via_edge: Some(via_edge),
depth,
name: None,
kind: None,
properties: None,
weight: total_weight,
});
non_root_count += 1;
if depth < opts.max_depth {
frontier.push_back(TraversalFrontierNode {
node_id,
depth,
total_weight,
});
}
if non_root_count == result_limit {
break 'root_walk;
}
}
}
}
if !nodes.is_empty() {
let total_weight = nodes.iter().map(|node| node.weight).fold(0.0_f64, f64::max);
all_paths.push(GraphPath {
root_id,
nodes,
total_weight,
});
}
}
Ok(all_paths)
})();
conn.progress_handler(0, None::<fn() -> bool>)
.map_err(|e| map_err(e, "traverse_progress_handler_clear"))?;
result
}
impl SqlGraphStore {
async fn query_neighbors_page(
&self,
operation: &'static str,
node_id: Uuid,
query: NeighborQuery,
after: Option<NeighborCursor>,
neighbor_kinds: Option<Vec<String>>,
) -> Result<Vec<NeighborHit>, StorageError> {
count_neighbor_select();
let namespace = self.namespace.clone();
let node_str = node_id.to_string();
let counted_queries = Arc::new(std::sync::atomic::AtomicU64::new(0));
let counted_rows = Arc::new(std::sync::atomic::AtomicU64::new(0));
let closure_queries = Arc::clone(&counted_queries);
let closure_rows = Arc::clone(&counted_rows);
let result = self
.with_reader(operation, move |conn| {
let base_out = "SELECT target_id AS node_id, id AS edge_id, relation, weight \
FROM graph_edges \
WHERE namespace = ?1 AND source_id = ?2 AND deleted_at IS NULL";
let base_in = "SELECT source_id AS node_id, id AS edge_id, relation, weight \
FROM graph_edges \
WHERE namespace = ?1 AND target_id = ?2 AND deleted_at IS NULL";
let sql = match query.direction {
Direction::Out => base_out.to_string(),
Direction::In => base_in.to_string(),
Direction::Both => format!("{} UNION ALL {}", base_out, base_in),
};
let (where_extra, limit_clause, extra_params) =
neighbor_extra_clause(&query, 3, after.as_ref(), neighbor_kinds.as_deref());
let full_sql = format!(
"SELECT node_id, edge_id, relation, weight FROM ({}){} \
ORDER BY weight DESC, node_id ASC, edge_id ASC{}",
sql, where_extra, limit_clause
);
let mut stmt = conn.prepare(&full_sql)?;
let mut all_params: Vec<Box<dyn rusqlite::types::ToSql>> = Vec::new();
all_params.push(Box::new(namespace.clone()));
all_params.push(Box::new(node_str.clone()));
all_params.extend(extra_params);
let param_refs: Vec<&dyn rusqlite::types::ToSql> =
all_params.iter().map(|p| p.as_ref()).collect();
closure_queries.fetch_add(1, std::sync::atomic::Ordering::Relaxed);
let rows = stmt.query_map(param_refs.as_slice(), |row| {
let nid_str: String = row.get(0)?;
let eid_str: String = row.get(1)?;
let relation_str: String = row.get(2)?;
let weight: f64 = row.get(3)?;
Ok((nid_str, eid_str, relation_str, weight))
})?;
let mut hits = Vec::new();
for row in rows {
let (nid_str, eid_str, relation_str, weight) = row?;
closure_rows.fetch_add(1, std::sync::atomic::Ordering::Relaxed);
let relation = relation_str.parse::<EdgeRelation>().map_err(|e| {
rusqlite::Error::FromSqlConversionFailure(
2,
rusqlite::types::Type::Text,
Box::new(e),
)
})?;
hits.push(NeighborHit {
node_id: parse_uuid(&nid_str)?,
edge_id: parse_uuid(&eid_str)?,
relation,
weight,
name: None,
kind: None,
entity_type: None,
});
}
Ok(hits)
})
.await;
report_graph_usage(&counted_queries, &counted_rows);
result
}
}
#[async_trait]
impl GraphStore for SqlGraphStore {
async fn latest_annotating_note(
&self,
node_id: Uuid,
kind: &str,
tag: &str,
) -> Result<Option<(Uuid, i64)>, StorageError> {
let namespace = self.namespace.clone();
let node_id = node_id.to_string();
let kind = kind.to_owned();
let tag = tag.to_owned();
self.with_reader("latest_annotating_note", move |conn| {
conn.query_row(
LATEST_ANNOTATING_NOTE_SQL,
rusqlite::params![namespace, node_id, kind, tag],
|row| {
let id: String = row.get(0)?;
Ok((parse_uuid(&id)?, row.get(1)?))
},
)
.optional()
})
.await
}
async fn upsert_edge(&self, edge: Edge) -> Result<(), StorageError> {
self.upsert_edge_observed(EdgeUpsertRequest {
edge,
resurrect: false,
})
.await
.map(|_| ())
}
async fn upsert_edge_observed(
&self,
request: EdgeUpsertRequest,
) -> Result<EdgeUpsertResult, StorageError> {
match self
.observed_edge_write("upsert_edge_observed", request, false)
.await?
{
GuardedEdgeUpsertOutcome::Written(result) => Ok(result),
GuardedEdgeUpsertOutcome::Refused(EdgeUpsertRefusal::ResurrectionRequired { edge }) => {
Err(resurrection_required_error("upsert_edge_observed", &edge))
}
GuardedEdgeUpsertOutcome::Refused(EdgeUpsertRefusal::MissingEndpoints(_)) => {
Err(StorageError::Conflict {
capability: StorageCapability::Graph,
operation: "upsert_edge_observed".into(),
message: "unguarded edge upsert reported a missing-endpoint refusal".into(),
})
}
}
}
async fn insert_edge_if_absent(&self, edge: Edge) -> Result<bool, StorageError> {
let statement = edge_insert_if_absent_statement(&edge);
self.with_writer("insert_edge_if_absent", move |conn| {
let mut stmt = conn.prepare(&statement.sql)?;
bind_params(&mut stmt, &statement.params)?;
Ok(stmt.raw_execute()? > 0)
})
.await
}
async fn replace_edge_if_unchanged(
&self,
edge: Edge,
expected_updated_at: DateTime<Utc>,
expected_deleted_at: Option<DateTime<Utc>>,
) -> Result<bool, StorageError> {
let statement =
edge_replace_if_unchanged_statement(&edge, expected_updated_at, expected_deleted_at);
self.with_writer("replace_edge_if_unchanged", move |conn| {
let mut stmt = conn.prepare(&statement.sql)?;
bind_params(&mut stmt, &statement.params)?;
Ok(stmt.raw_execute()? > 0)
})
.await
}
async fn upsert_edges(&self, edges: Vec<Edge>) -> Result<BatchWriteSummary, StorageError> {
let attempted = edges.len() as u64;
let requests = edges
.into_iter()
.map(|edge| EdgeUpsertRequest {
edge,
resurrect: false,
})
.collect();
let outcome = self
.observed_edge_batch_write("upsert_edges", requests, false)
.await?;
if let Some(refusal) = outcome.refusal {
return match refusal.reason {
EdgeUpsertRefusal::ResurrectionRequired { edge } => {
Err(resurrection_required_error("upsert_edges", &edge))
}
EdgeUpsertRefusal::MissingEndpoints(_) => Err(StorageError::Conflict {
capability: StorageCapability::Graph,
operation: "upsert_edges".into(),
message: "unguarded edge batch reported a missing-endpoint refusal".into(),
}),
};
}
Ok(BatchWriteSummary {
attempted,
affected: outcome.rows.len() as u64,
..BatchWriteSummary::default()
})
}
async fn upsert_edge_guarded(&self, edge: Edge) -> Result<GuardedWriteOutcome, StorageError> {
match self
.upsert_edge_guarded_observed(EdgeUpsertRequest {
edge,
resurrect: false,
})
.await?
{
GuardedEdgeUpsertOutcome::Written(_) => Ok(GuardedWriteOutcome::Written),
GuardedEdgeUpsertOutcome::Refused(EdgeUpsertRefusal::MissingEndpoints(missing)) => {
Ok(GuardedWriteOutcome::Refused(missing))
}
GuardedEdgeUpsertOutcome::Refused(EdgeUpsertRefusal::ResurrectionRequired { edge }) => {
Err(resurrection_required_error("upsert_edge_guarded", &edge))
}
}
}
async fn upsert_edge_guarded_observed(
&self,
request: EdgeUpsertRequest,
) -> Result<GuardedEdgeUpsertOutcome, StorageError> {
self.observed_edge_write("upsert_edge_guarded_observed", request, true)
.await
}
async fn upsert_edges_guarded(
&self,
edges: Vec<Edge>,
) -> Result<GuardedBatchOutcome, StorageError> {
let attempted = edges.len() as u64;
let requests = edges
.iter()
.cloned()
.map(|edge| EdgeUpsertRequest {
edge,
resurrect: false,
})
.collect();
let outcome = self.upsert_edges_guarded_observed(requests).await?;
match outcome.refusal {
None => Ok(GuardedBatchOutcome {
summary: BatchWriteSummary {
attempted,
affected: outcome.rows.len() as u64,
..BatchWriteSummary::default()
},
refused: None,
}),
Some(refusal) => match refusal.reason {
EdgeUpsertRefusal::MissingEndpoints(missing) => {
let index = refusal.entry_index;
let edge = &edges[index];
let (source_id, target_id) =
canonical_edge_endpoints(edge.relation, edge.source_id, edge.target_id);
let message = format!(
"batch entry {index}: edge endpoint no longer exists at write time: source \
{source_id} or target {target_id}"
);
let mut summary = BatchWriteSummary {
attempted,
..BatchWriteSummary::default()
};
summary.first_error = message.clone();
let refusal = GuardedBatchRefusal {
entry_index: index,
missing,
};
for (failed_index, failed_edge) in edges.iter().enumerate() {
refusal.record_failure(&mut summary, failed_index, failed_edge, &message);
}
Ok(GuardedBatchOutcome {
summary,
refused: Some(refusal),
})
}
EdgeUpsertRefusal::ResurrectionRequired { edge } => {
Err(resurrection_required_error("upsert_edges_guarded", &edge))
}
},
}
}
async fn upsert_edges_guarded_observed(
&self,
requests: Vec<EdgeUpsertRequest>,
) -> Result<GuardedEdgeBatchUpsertOutcome, StorageError> {
self.observed_edge_batch_write("upsert_edges_guarded_observed", requests, true)
.await
}
async fn get_edge(&self, id: LinkId) -> Result<Option<Edge>, StorageError> {
let id_str = Uuid::from(id).to_string();
self.with_reader("get_edge", move |conn| {
let mut stmt = conn.prepare(
"SELECT namespace, id, source_id, target_id, relation, weight, \
created_at, updated_at, deleted_at, metadata, target_backend \
FROM graph_edges WHERE id = ?1 AND deleted_at IS NULL",
)?;
let mut rows = stmt.query(rusqlite::params![id_str])?;
match rows.next()? {
Some(row) => Ok(Some(read_edge(row)?)),
None => Ok(None),
}
})
.await
}
async fn get_edge_including_deleted(&self, id: LinkId) -> Result<Option<Edge>, StorageError> {
let id_str = Uuid::from(id).to_string();
self.with_reader("get_edge_including_deleted", move |conn| {
let mut stmt = conn.prepare(
"SELECT namespace, id, source_id, target_id, relation, weight, \
created_at, updated_at, deleted_at, metadata, target_backend \
FROM graph_edges WHERE id = ?1",
)?;
let mut rows = stmt.query(rusqlite::params![id_str])?;
match rows.next()? {
Some(row) => Ok(Some(read_edge(row)?)),
None => Ok(None),
}
})
.await
}
async fn edge_sequence(&self, id: Uuid) -> Result<Option<i64>, StorageError> {
let id = id.to_string();
self.with_reader("edge_sequence", move |conn| {
conn.query_row(
"SELECT seq FROM graph_edges_seq WHERE edge_id = ?1",
rusqlite::params![id],
|row| row.get(0),
)
.optional()
})
.await
}
async fn edge_sequences(&self, ids: &[Uuid]) -> Result<Vec<(Uuid, i64)>, StorageError> {
if ids.is_empty() {
return Ok(Vec::new());
}
let ids = ids.to_vec();
self.with_reader("edge_sequences", move |conn| {
const CHUNK: usize = 900;
let mut resolved = Vec::with_capacity(ids.len());
for chunk in ids.chunks(CHUNK) {
let placeholders = (1..=chunk.len())
.map(|index| format!("?{index}"))
.collect::<Vec<_>>()
.join(", ");
let sql = format!(
"SELECT edge_id, seq FROM graph_edges_seq WHERE edge_id IN ({placeholders})"
);
let strings = chunk.iter().map(Uuid::to_string).collect::<Vec<_>>();
let params = strings
.iter()
.map(|id| id as &dyn rusqlite::types::ToSql)
.collect::<Vec<_>>();
let mut stmt = conn.prepare(&sql)?;
let rows = stmt.query_map(params.as_slice(), |row| {
let id: String = row.get(0)?;
Ok((parse_uuid(&id)?, row.get::<_, i64>(1)?))
})?;
resolved.extend(rows.collect::<Result<Vec<_>, _>>()?);
}
Ok(resolved)
})
.await
}
async fn get_edge_by_natural_key_including_deleted(
&self,
namespace: &str,
source_id: Uuid,
target_id: Uuid,
relation: EdgeRelation,
) -> Result<Option<Edge>, StorageError> {
let namespace = namespace.to_string();
let source_str = source_id.to_string();
let target_str = target_id.to_string();
let relation_str = relation.to_string();
self.with_reader("get_edge_by_natural_key_including_deleted", move |conn| {
let mut stmt = conn.prepare(
"SELECT namespace, id, source_id, target_id, relation, weight, \
created_at, updated_at, deleted_at, metadata, target_backend \
FROM graph_edges \
WHERE namespace = ?1 AND source_id = ?2 AND target_id = ?3 AND relation = ?4",
)?;
let mut rows = stmt.query(rusqlite::params![
namespace,
source_str,
target_str,
relation_str
])?;
match rows.next()? {
Some(row) => Ok(Some(read_edge(row)?)),
None => Ok(None),
}
})
.await
}
async fn get_edges(&self, ids: &[LinkId]) -> Result<Vec<Edge>, StorageError> {
if ids.is_empty() {
return Ok(Vec::new());
}
const CHUNK: usize = 900;
let id_strs: Vec<String> = ids.iter().map(|id| Uuid::from(*id).to_string()).collect();
let mut result: Vec<Edge> = Vec::with_capacity(ids.len());
for chunk in id_strs.chunks(CHUNK) {
let chunk_owned: Vec<String> = chunk.to_vec();
let edges = self
.with_reader("get_edges", move |conn| {
let placeholders: Vec<String> =
(1..=chunk_owned.len()).map(|i| format!("?{}", i)).collect();
let sql = format!(
"SELECT namespace, id, source_id, target_id, relation, weight, \
created_at, updated_at, deleted_at, metadata, target_backend \
FROM graph_edges WHERE id IN ({}) AND deleted_at IS NULL",
placeholders.join(",")
);
let mut stmt = conn.prepare(&sql)?;
let params: Vec<&dyn rusqlite::types::ToSql> = chunk_owned
.iter()
.map(|s| s as &dyn rusqlite::types::ToSql)
.collect();
let rows = stmt.query_map(params.as_slice(), read_edge)?;
let mut edges = Vec::new();
for row in rows {
edges.push(row?);
}
Ok(edges)
})
.await?;
result.extend(edges);
}
Ok(result)
}
async fn batch_neighbors(
&self,
sources: &[Uuid],
query: NeighborQuery,
) -> Result<Vec<(Uuid, NeighborHit)>, StorageError> {
use khive_storage::types::Direction;
if sources.is_empty() {
return Ok(Vec::new());
}
let mut seen_sources = HashSet::with_capacity(sources.len());
let unique_sources: Vec<Uuid> = sources
.iter()
.copied()
.filter(|source| seen_sources.insert(*source))
.collect();
const CHUNK_SIZE: usize = 880;
let namespace = self.namespace.clone();
let mut result: Vec<(Uuid, NeighborHit)> = Vec::new();
let counted_queries = Arc::new(std::sync::atomic::AtomicU64::new(0));
let counted_rows = Arc::new(std::sync::atomic::AtomicU64::new(0));
for chunk in unique_sources.chunks(CHUNK_SIZE) {
let chunk_owned: Vec<Uuid> = chunk.to_vec();
let query_clone = query.clone();
let ns = namespace.clone();
let closure_queries = Arc::clone(&counted_queries);
let closure_rows = Arc::clone(&counted_rows);
let chunk_result = self
.with_reader("batch_neighbors", move |conn| {
let src_strs: Vec<String> = chunk_owned.iter().map(|u| u.to_string()).collect();
let sources_json = serde_json::to_string(&src_strs).map_err(|error| {
rusqlite::Error::ToSqlConversionFailure(Box::new(error))
})?;
let build_inner_sql =
|direction_out: bool,
q: &NeighborQuery|
-> (String, Vec<String>, Option<f64>) {
let (filter_col, node_col) = if direction_out {
("source_id", "target_id")
} else {
("target_id", "source_id")
};
let mut rel_params: Vec<String> = Vec::new();
let mut conditions: Vec<String> = Vec::new();
let mut param_idx = 3;
if let Some(ref rels) = q.relations {
if !rels.is_empty() {
let ps: Vec<String> = rels
.iter()
.map(|r| {
rel_params.push(r.to_string());
let p = format!("?{param_idx}");
param_idx += 1;
p
})
.collect();
conditions
.push(format!("edges.relation IN ({})", ps.join(",")));
}
}
let min_weight_val = if let Some(min_w) = q.min_weight {
conditions.push(format!("edges.weight >= ?{param_idx}"));
Some(min_w)
} else {
None
};
let where_extra = if conditions.is_empty() {
String::new()
} else {
format!(" AND {}", conditions.join(" AND "))
};
let sql = format!(
"SELECT requested.origin_id, edges.{node_col} AS node_id, \
edges.id AS edge_id, edges.relation, edges.weight \
FROM requested CROSS JOIN graph_edges AS edges \
ON edges.{filter_col} = requested.origin_id \
WHERE edges.namespace = ?1 \
AND edges.deleted_at IS NULL{where_extra}",
);
(sql, rel_params, min_weight_val)
};
let mut all_params: Vec<Box<dyn rusqlite::types::ToSql>> = Vec::new();
all_params.push(Box::new(ns.to_string()));
all_params.push(Box::new(sources_json));
let (combined_inner, rel_params, min_weight_val) = match query_clone.direction {
Direction::Out => build_inner_sql(true, &query_clone),
Direction::In => build_inner_sql(false, &query_clone),
Direction::Both => {
let (out_sql, rel_params, min_weight_val) =
build_inner_sql(true, &query_clone);
let (in_sql, _, _) = build_inner_sql(false, &query_clone);
(
format!("{out_sql} UNION ALL {in_sql}"),
rel_params,
min_weight_val,
)
}
};
for relation in rel_params {
all_params.push(Box::new(relation));
}
if let Some(min_weight) = min_weight_val {
all_params.push(Box::new(min_weight));
}
let limit_param_idx = all_params.len() + 1;
let full_sql = if let Some(lim) = query_clone.limit {
all_params.push(Box::new(lim as i64));
format!(
"WITH requested(origin_id) AS (\
SELECT value FROM json_each(?2)\
) SELECT origin_id, node_id, edge_id, relation, weight \
FROM (SELECT *, ROW_NUMBER() OVER (PARTITION BY origin_id \
ORDER BY weight DESC, node_id ASC) AS rn \
FROM ({combined_inner})) WHERE rn <= ?{limit_param_idx}",
)
} else {
format!(
"WITH requested(origin_id) AS (\
SELECT value FROM json_each(?2)\
) SELECT origin_id, node_id, edge_id, relation, weight \
FROM ({combined_inner})",
)
};
let param_refs: Vec<&dyn rusqlite::types::ToSql> =
all_params.iter().map(|p| p.as_ref()).collect();
let mut stmt = conn.prepare(&full_sql)?;
closure_queries.fetch_add(1, std::sync::atomic::Ordering::Relaxed);
let rows = stmt.query_map(param_refs.as_slice(), |row| {
let origin_str: String = row.get(0)?;
let nid_str: String = row.get(1)?;
let eid_str: String = row.get(2)?;
let relation_str: String = row.get(3)?;
let weight: f64 = row.get(4)?;
Ok((origin_str, nid_str, eid_str, relation_str, weight))
})?;
let mut pairs = Vec::new();
for row in rows {
let (origin_str, nid_str, eid_str, relation_str, weight) = row?;
closure_rows.fetch_add(1, std::sync::atomic::Ordering::Relaxed);
let origin = parse_uuid(&origin_str)?;
let node_id = parse_uuid(&nid_str)?;
let edge_id = parse_uuid(&eid_str)?;
let relation = relation_str.parse::<EdgeRelation>().map_err(|e| {
rusqlite::Error::FromSqlConversionFailure(
3,
rusqlite::types::Type::Text,
Box::new(e),
)
})?;
pairs.push((
origin,
NeighborHit {
node_id,
edge_id,
relation,
weight,
name: None,
kind: None,
entity_type: None,
},
));
}
Ok(pairs)
})
.await;
let pairs = match chunk_result {
Ok(pairs) => pairs,
Err(e) => {
report_graph_usage(&counted_queries, &counted_rows);
return Err(e);
}
};
result.extend(pairs);
}
report_graph_usage(&counted_queries, &counted_rows);
let requested: HashSet<Uuid> = unique_sources.iter().copied().collect();
let mut grouped: HashMap<Uuid, Vec<NeighborHit>> =
HashMap::with_capacity(unique_sources.len());
for (origin, hit) in result {
if !requested.contains(&origin) {
return Err(StorageError::Internal(format!(
"batch_neighbors returned unrequested origin {origin}"
)));
}
grouped.entry(origin).or_default().push(hit);
}
for hits in grouped.values_mut() {
hits.sort_by(|a, b| {
b.weight
.partial_cmp(&a.weight)
.unwrap_or(std::cmp::Ordering::Equal)
.then(a.node_id.cmp(&b.node_id))
.then(a.edge_id.cmp(&b.edge_id))
});
}
let mut ordered = Vec::new();
for &source in sources {
if let Some(hits) = grouped.get(&source) {
ordered.extend(hits.iter().cloned().map(|hit| (source, hit)));
}
}
Ok(ordered)
}
async fn delete_edge(&self, id: LinkId, mode: DeleteMode) -> Result<bool, StorageError> {
let id = Uuid::from(id);
let statement = match mode {
DeleteMode::Soft => {
edge_soft_delete_statement(id, chrono::Utc::now().timestamp_micros())
}
DeleteMode::Hard => edge_hard_delete_statement(id),
};
self.with_writer("delete_edge", move |conn| {
let mut stmt = conn.prepare(&statement.sql)?;
bind_params(&mut stmt, &statement.params)?;
Ok(stmt.raw_execute()? > 0)
})
.await
}
async fn query_edges(
&self,
filter: EdgeFilter,
sort: Vec<SortOrder<EdgeSortField>>,
page: PageRequest,
) -> Result<Page<Edge>, StorageError> {
let namespace = self.namespace.clone();
let limit_i64 = i64::from(page.limit);
let offset_i64 = i64::try_from(page.offset).map_err(|_| StorageError::InvalidInput {
capability: StorageCapability::Graph,
operation: "query_edges".into(),
message: format!(
"PageRequest: offset must be <= i64::MAX, got {}",
page.offset
),
})?;
self.with_reader("query_edges", move |conn| {
let (where_clause, mut all_params) = build_edge_filter_sql(&namespace, &filter);
let order_clause = edge_order_clause(&sort);
all_params.push(Box::new(limit_i64));
all_params.push(Box::new(offset_i64));
let limit_idx = all_params.len() - 1;
let offset_idx = all_params.len();
let data_sql = format!(
"SELECT namespace, id, source_id, target_id, relation, weight, \
created_at, updated_at, deleted_at, metadata, target_backend \
FROM graph_edges{}{} LIMIT ?{} OFFSET ?{}",
where_clause, order_clause, limit_idx, offset_idx,
);
let mut stmt = conn.prepare(&data_sql)?;
let param_refs: Vec<&dyn rusqlite::types::ToSql> =
all_params.iter().map(|p| p.as_ref()).collect();
let rows = stmt.query_map(param_refs.as_slice(), read_edge)?;
let mut items = Vec::new();
for row in rows {
items.push(row?);
}
Ok(Page { items, total: None })
})
.await
}
async fn count_edges(&self, filter: EdgeFilter) -> Result<u64, StorageError> {
let namespace = self.namespace.clone();
self.with_reader("count_edges", move |conn| {
let (where_clause, params) = build_edge_filter_sql(&namespace, &filter);
let sql = format!(
"SELECT COUNT(*) FROM graph_edges{}",
with_live_endpoints(&where_clause)
);
let mut stmt = conn.prepare(&sql)?;
let param_refs: Vec<&dyn rusqlite::types::ToSql> =
params.iter().map(|p| p.as_ref()).collect();
let count: i64 = stmt.query_row(param_refs.as_slice(), |row| row.get(0))?;
Ok(count as u64)
})
.await
}
async fn count_edges_in_namespaces(
&self,
namespaces: &[String],
filter: EdgeFilter,
) -> Result<u64, StorageError> {
let namespaces: Vec<String> = namespaces
.iter()
.cloned()
.collect::<HashSet<_>>()
.into_iter()
.collect();
self.with_reader("count_edges_in_namespaces", move |conn| {
let mut total = 0;
for chunk in namespaces.chunks(NAMESPACE_COUNT_CHUNK_SIZE) {
let (where_clause, params) = build_edge_filter_sql_for_namespaces(chunk, &filter);
let sql = format!(
"SELECT COUNT(*) FROM graph_edges{}",
with_live_endpoints(&where_clause)
);
let mut stmt = conn.prepare(&sql)?;
let param_refs: Vec<&dyn rusqlite::types::ToSql> =
params.iter().map(|p| p.as_ref()).collect();
let count: i64 = stmt.query_row(param_refs.as_slice(), |row| row.get(0))?;
total += count as u64;
}
Ok(total)
})
.await
}
async fn query_edges_in_namespaces(
&self,
namespaces: &[String],
filter: EdgeFilter,
sort: Vec<SortOrder<EdgeSortField>>,
page: PageRequest,
) -> Result<Page<Edge>, StorageError> {
let namespaces: Vec<String> = {
let mut seen = HashSet::new();
namespaces
.iter()
.filter(|ns| seen.insert((*ns).clone()))
.cloned()
.collect()
};
let limit_i64 = i64::from(page.limit);
let offset_i64 = i64::try_from(page.offset).map_err(|_| StorageError::InvalidInput {
capability: StorageCapability::Graph,
operation: "query_edges_in_namespaces".into(),
message: format!(
"PageRequest: offset must be <= i64::MAX, got {}",
page.offset
),
})?;
self.with_reader("query_edges_in_namespaces", move |conn| {
let namespaces_json = serde_json::to_string(&namespaces)
.map_err(|error| rusqlite::Error::ToSqlConversionFailure(Box::new(error)))?;
let (where_clause, mut all_params) =
build_edge_filter_sql_for_namespaces_json(&namespaces_json, &filter);
let order_clause = edge_order_clause(&sort);
all_params.push(Box::new(limit_i64));
all_params.push(Box::new(offset_i64));
let limit_idx = all_params.len() - 1;
let offset_idx = all_params.len();
let data_sql = format!(
"SELECT namespace, id, source_id, target_id, relation, weight, \
created_at, updated_at, deleted_at, metadata, target_backend \
FROM graph_edges{}{} LIMIT ?{} OFFSET ?{}",
where_clause, order_clause, limit_idx, offset_idx,
);
let mut stmt = conn.prepare(&data_sql)?;
let param_refs: Vec<&dyn rusqlite::types::ToSql> =
all_params.iter().map(|p| p.as_ref()).collect();
let rows = stmt.query_map(param_refs.as_slice(), read_edge)?;
let mut items = Vec::new();
for row in rows {
items.push(row?);
}
Ok(Page { items, total: None })
})
.await
}
async fn count_edges_by_relation(&self) -> Result<Vec<(EdgeRelation, u64)>, StorageError> {
let namespace = self.namespace.clone();
self.with_reader("count_edges_by_relation", move |conn| {
let sql = format!(
"SELECT relation, COUNT(*) FROM graph_edges \
WHERE namespace = ?1 AND deleted_at IS NULL AND {LIVE_ENDPOINTS_CONDITION} \
GROUP BY relation"
);
let mut stmt = conn.prepare(&sql)?;
let rows = stmt.query_map([&namespace], |row| {
let relation_str: String = row.get(0)?;
let count: i64 = row.get(1)?;
Ok((relation_str, count))
})?;
let mut out = Vec::new();
for row in rows {
let (relation_str, count) = row?;
let relation = relation_str.parse::<EdgeRelation>().map_err(|e| {
rusqlite::Error::FromSqlConversionFailure(
0,
rusqlite::types::Type::Text,
Box::new(e),
)
})?;
out.push((relation, count as u64));
}
Ok(out)
})
.await
}
async fn count_edges_by_relation_in_namespaces(
&self,
namespaces: &[String],
) -> Result<Vec<(EdgeRelation, u64)>, StorageError> {
let namespaces: Vec<String> = namespaces
.iter()
.cloned()
.collect::<HashSet<_>>()
.into_iter()
.collect();
self.with_reader("count_edges_by_relation_in_namespaces", move |conn| {
let mut totals = HashMap::new();
for chunk in namespaces.chunks(NAMESPACE_COUNT_CHUNK_SIZE) {
let (where_clause, params) =
build_edge_filter_sql_for_namespaces(chunk, &EdgeFilter::default());
let sql = format!(
"SELECT relation, COUNT(*) FROM graph_edges{} GROUP BY relation",
with_live_endpoints(&where_clause)
);
let mut stmt = conn.prepare(&sql)?;
let param_refs: Vec<&dyn rusqlite::types::ToSql> =
params.iter().map(|p| p.as_ref()).collect();
let rows = stmt.query_map(param_refs.as_slice(), |row| {
let relation_str: String = row.get(0)?;
let count: i64 = row.get(1)?;
Ok((relation_str, count))
})?;
for row in rows {
let (relation_str, count) = row?;
let relation = relation_str.parse::<EdgeRelation>().map_err(|e| {
rusqlite::Error::FromSqlConversionFailure(
0,
rusqlite::types::Type::Text,
Box::new(e),
)
})?;
*totals.entry(relation).or_insert(0) += count as u64;
}
}
Ok(totals.into_iter().collect())
})
.await
}
async fn count_edges_by_endpoint_base(&self) -> Result<EdgeEndpointBaseCounts, StorageError> {
let namespace = self.namespace.clone();
self.with_reader("count_edges_by_endpoint_base", move |conn| {
let source_case = endpoint_base_case("source_id");
let target_case = endpoint_base_case("target_id");
let sql = format!(
"SELECT {source_case} AS source_base, {target_case} AS target_base, COUNT(*) \
FROM graph_edges \
WHERE namespace = ?1 AND deleted_at IS NULL AND {LIVE_ENDPOINTS_CONDITION} \
GROUP BY source_base, target_base"
);
let mut stmt = conn.prepare(&sql)?;
let rows = stmt.query_map([&namespace], |row| {
let source: String = row.get(0)?;
let target: String = row.get(1)?;
let count: i64 = row.get(2)?;
Ok((source, target, count))
})?;
let mut counts = EdgeEndpointBaseCounts::default();
for row in rows {
let (source, target, count) = row?;
fold_endpoint_base_row(&mut counts, &source, &target, count as u64);
}
Ok(counts)
})
.await
}
async fn count_edges_by_endpoint_base_in_namespaces(
&self,
namespaces: &[String],
) -> Result<EdgeEndpointBaseCounts, StorageError> {
let namespaces: Vec<String> = namespaces
.iter()
.cloned()
.collect::<HashSet<_>>()
.into_iter()
.collect();
self.with_reader("count_edges_by_endpoint_base_in_namespaces", move |conn| {
let source_case = endpoint_base_case("source_id");
let target_case = endpoint_base_case("target_id");
let mut counts = EdgeEndpointBaseCounts::default();
for chunk in namespaces.chunks(NAMESPACE_COUNT_CHUNK_SIZE) {
let (where_clause, params) =
build_edge_filter_sql_for_namespaces(chunk, &EdgeFilter::default());
let sql = format!(
"SELECT {source_case} AS source_base, {target_case} AS target_base, COUNT(*) \
FROM graph_edges{} GROUP BY source_base, target_base",
with_live_endpoints(&where_clause)
);
let mut stmt = conn.prepare(&sql)?;
let param_refs: Vec<&dyn rusqlite::types::ToSql> =
params.iter().map(|p| p.as_ref()).collect();
let rows = stmt.query_map(param_refs.as_slice(), |row| {
let source: String = row.get(0)?;
let target: String = row.get(1)?;
let count: i64 = row.get(2)?;
Ok((source, target, count))
})?;
for row in rows {
let (source, target, count) = row?;
fold_endpoint_base_row(&mut counts, &source, &target, count as u64);
}
}
Ok(counts)
})
.await
}
async fn query_edges_after(
&self,
filter: EdgeFilter,
after: Option<Uuid>,
limit: u32,
) -> Result<EdgeSeekPage, StorageError> {
let namespace = self.namespace.clone();
let limit_usize = limit as usize;
let probe_limit_i64 = i64::from(limit) + 1;
self.with_reader("query_edges_after", move |conn| {
let (mut where_clause, mut params) = build_edge_filter_sql(&namespace, &filter);
if let Some(cursor) = after {
params.push(Box::new(cursor.to_string()));
where_clause.push_str(&format!(" AND id > ?{}", params.len()));
}
params.push(Box::new(probe_limit_i64));
let limit_idx = params.len();
let data_sql = format!(
"SELECT namespace, id, source_id, target_id, relation, weight, \
created_at, updated_at, deleted_at, metadata, target_backend \
FROM graph_edges{} ORDER BY id ASC LIMIT ?{}",
where_clause, limit_idx,
);
let mut stmt = conn.prepare(&data_sql)?;
let param_refs: Vec<&dyn rusqlite::types::ToSql> =
params.iter().map(|p| p.as_ref()).collect();
let rows = stmt.query_map(param_refs.as_slice(), read_edge)?;
let mut items = Vec::new();
for row in rows {
items.push(row?);
}
let has_more = items.len() > limit_usize;
if has_more {
items.truncate(limit_usize);
}
let next_after = if has_more {
items.last().map(|e| Uuid::from(e.id))
} else {
None
};
Ok(EdgeSeekPage { items, next_after })
})
.await
}
async fn query_edges_sequence_after(
&self,
filter: EdgeFilter,
after: Option<SeekCursor>,
limit: u32,
) -> Result<SeekPage<Edge>, StorageError> {
if limit == 0 {
return Ok(SeekPage::default());
}
let namespace = self.namespace.clone();
let limit_usize = limit as usize;
let probe_limit_i64 = i64::from(limit) + 1;
self.with_reader("query_edges_sequence_after", move |conn| {
let (mut where_clause, mut params) = build_edge_filter_sql(&namespace, &filter);
if let Some(cursor) = after {
params.push(Box::new(cursor.sequence));
where_clause.push_str(&format!(" AND graph_edges_seq.seq > ?{}", params.len()));
}
params.push(Box::new(probe_limit_i64));
let limit_idx = params.len();
let sql = format!(
"SELECT namespace, id, source_id, target_id, relation, weight, \
created_at, updated_at, deleted_at, metadata, target_backend, \
graph_edges_seq.seq \
FROM graph_edges_seq CROSS JOIN graph_edges \
ON graph_edges.id = graph_edges_seq.edge_id{where_clause} \
ORDER BY graph_edges_seq.seq ASC LIMIT ?{limit_idx}"
);
let mut stmt = conn.prepare(&sql)?;
let param_refs: Vec<&dyn rusqlite::types::ToSql> =
params.iter().map(|param| param.as_ref()).collect();
let rows = stmt.query_map(param_refs.as_slice(), |row| {
Ok((read_edge(row)?, row.get::<_, i64>(11)?))
})?;
let mut entries = rows.collect::<Result<Vec<_>, _>>()?;
let has_more = entries.len() > limit_usize;
if has_more {
entries.truncate(limit_usize);
}
let next_after = if has_more {
entries.last().map(|(edge, sequence)| SeekCursor {
sequence: *sequence,
id: Uuid::from(edge.id),
})
} else {
None
};
let items = entries.into_iter().map(|(edge, _)| edge).collect();
Ok(SeekPage { items, next_after })
})
.await
}
async fn neighbors(
&self,
node_id: Uuid,
query: NeighborQuery,
) -> Result<Vec<NeighborHit>, StorageError> {
self.query_neighbors_page("neighbors", node_id, query, None, None)
.await
}
async fn neighbors_page(
&self,
node_id: Uuid,
query: NeighborQuery,
after: Option<NeighborCursor>,
neighbor_kinds: Option<Vec<String>>,
) -> Result<Vec<NeighborHit>, StorageError> {
self.query_neighbors_page("neighbors_page", node_id, query, after, neighbor_kinds)
.await
}
async fn neighbors_both_directions(
&self,
node_id: Uuid,
query: NeighborQuery,
) -> Result<Vec<DirectedNeighborHit>, StorageError> {
count_neighbor_select();
let namespace = self.namespace.clone();
let node_str = node_id.to_string();
let counted_queries = Arc::new(std::sync::atomic::AtomicU64::new(0));
let counted_rows = Arc::new(std::sync::atomic::AtomicU64::new(0));
let closure_queries = Arc::clone(&counted_queries);
let closure_rows = Arc::clone(&counted_rows);
let result = self
.with_reader("neighbors_both_directions", move |conn| {
let base_out = "SELECT target_id AS node_id, id AS edge_id, relation, weight, \
'out' AS dir \
FROM graph_edges \
WHERE namespace = ?1 AND source_id = ?2 AND deleted_at IS NULL";
let base_in = "SELECT source_id AS node_id, id AS edge_id, relation, weight, \
'in' AS dir \
FROM graph_edges \
WHERE namespace = ?1 AND target_id = ?2 AND deleted_at IS NULL";
let sql = format!("{} UNION ALL {}", base_out, base_in);
let (where_extra, limit_clause, extra_params) =
neighbor_extra_clause(&query, 3, None, None);
let full_sql = format!(
"SELECT node_id, edge_id, relation, weight, dir FROM ({}){} \
ORDER BY weight DESC, node_id ASC, \
CASE dir WHEN 'out' THEN 0 ELSE 1 END ASC, edge_id ASC{}",
sql, where_extra, limit_clause
);
let mut stmt = conn.prepare(&full_sql)?;
let mut all_params: Vec<Box<dyn rusqlite::types::ToSql>> = Vec::new();
all_params.push(Box::new(namespace.clone()));
all_params.push(Box::new(node_str.clone()));
all_params.extend(extra_params);
let param_refs: Vec<&dyn rusqlite::types::ToSql> =
all_params.iter().map(|p| p.as_ref()).collect();
closure_queries.fetch_add(1, std::sync::atomic::Ordering::Relaxed);
let rows = stmt.query_map(param_refs.as_slice(), |row| {
let nid_str: String = row.get(0)?;
let eid_str: String = row.get(1)?;
let relation_str: String = row.get(2)?;
let weight: f64 = row.get(3)?;
let dir_str: String = row.get(4)?;
Ok((nid_str, eid_str, relation_str, weight, dir_str))
})?;
let mut hits = Vec::new();
for row in rows {
let (nid_str, eid_str, relation_str, weight, dir_str) = row?;
closure_rows.fetch_add(1, std::sync::atomic::Ordering::Relaxed);
let relation = relation_str.parse::<EdgeRelation>().map_err(|e| {
rusqlite::Error::FromSqlConversionFailure(
2,
rusqlite::types::Type::Text,
Box::new(e),
)
})?;
let direction = if dir_str == "out" {
Direction::Out
} else {
Direction::In
};
hits.push(DirectedNeighborHit {
hit: NeighborHit {
node_id: parse_uuid(&nid_str)?,
edge_id: parse_uuid(&eid_str)?,
relation,
weight,
name: None,
kind: None,
entity_type: None,
},
direction,
});
}
Ok(hits)
})
.await;
report_graph_usage(&counted_queries, &counted_rows);
result
}
async fn traverse(&self, request: TraversalRequest) -> Result<Vec<GraphPath>, StorageError> {
request
.validate()
.map_err(|message| StorageError::InvalidInput {
capability: StorageCapability::Graph,
operation: "traverse".into(),
message,
})?;
if request.roots.is_empty() {
return Ok(Vec::new());
}
let mut distinct_roots = HashSet::with_capacity(request.roots.len());
let roots = request
.roots
.iter()
.copied()
.filter(|root| distinct_roots.insert(*root))
.collect::<Vec<_>>();
let opts = request.options;
let include_roots = request.include_roots;
let namespace = self.namespace.clone();
let origin = self.pool.origin();
let budget = request.execution_budget;
let counted_rows = Arc::new(std::sync::atomic::AtomicU64::new(0));
let counted_queries = Arc::new(std::sync::atomic::AtomicU64::new(0));
let closure_rows = Arc::clone(&counted_rows);
let closure_queries = Arc::clone(&counted_queries);
let result = self
.with_reader("traverse", move |conn| {
Ok(run_bounded_traversal(
conn,
roots,
opts,
include_roots,
namespace,
origin,
budget,
closure_rows.as_ref(),
closure_queries.as_ref(),
))
})
.await
.and_then(|inner| inner);
khive_storage::usage::count(
khive_storage::usage::UsageUnit::DbRoundTrips,
counted_queries.load(std::sync::atomic::Ordering::Relaxed),
);
khive_storage::usage::count(
khive_storage::usage::UsageUnit::GraphHops,
counted_rows.load(std::sync::atomic::Ordering::Relaxed),
);
result
}
async fn purge_incident_edges(&self, node_id: Uuid) -> Result<u64, StorageError> {
let statement = purge_incident_edges_statement(node_id);
self.with_writer("purge_incident_edges", move |conn| {
let mut stmt = conn.prepare(&statement.sql)?;
bind_params(&mut stmt, &statement.params)?;
Ok(stmt.raw_execute()? as u64)
})
.await
}
}
const GRAPH_DDL: &str = include_str!("../../sql/graph-ddl.sql");
pub(crate) fn ensure_graph_schema(conn: &rusqlite::Connection) -> Result<(), rusqlite::Error> {
conn.execute_batch(GRAPH_DDL)
}
#[cfg(test)]
#[path = "graph_annotation_tests.rs"]
mod annotation_tests;
#[cfg(test)]
#[path = "graph_tests.rs"]
mod tests;