use std::collections::{HashMap, HashSet};
use sqlx::postgres::{PgPoolOptions, PgRow};
use sqlx::{PgPool, Postgres, QueryBuilder, Row};
use super::schema::GRAPH_SCHEMA_VERSION;
use super::types::{EdgeDirection, GraphEdge, GraphNode, escape_like, join_csv, split_csv};
use crate::error::{KernelError, Result};
const NODE_COLUMNS: &str = "id, node_type, title, tags, projects, agents, \
created, updated, body, importance, access_count, accessed_at";
fn pg_err(e: sqlx::Error) -> KernelError {
KernelError::Store(format!("sqlx: {e:?}"))
}
fn row_to_node(row: &PgRow) -> GraphNode {
let tags: String = row.get("tags");
let projects: String = row.get("projects");
let agents: String = row.get("agents");
GraphNode {
id: row.get("id"),
node_type: row.get("node_type"),
title: row.get("title"),
tags: split_csv(&tags),
projects: split_csv(&projects),
agents: split_csv(&agents),
created: row.get("created"),
updated: row.get("updated"),
body: row.get("body"),
importance: row.get("importance"),
access_count: row.get("access_count"),
accessed_at: row.get("accessed_at"),
}
}
fn row_to_edge(row: &PgRow) -> GraphEdge {
GraphEdge {
id: row.get("id"),
source: row.get("source"),
target: row.get("target"),
relation: row.get("relation"),
weight: row.get("weight"),
ts: row.get("ts"),
}
}
fn search_patterns(query: &str) -> Vec<String> {
query
.split_whitespace()
.map(|t| format!("%{}%", escape_like(t)))
.collect()
}
fn is_identifier_safe(prefix: &str) -> bool {
let mut chars = prefix.chars();
match chars.next() {
None => true,
Some(first) if first == '_' || first.is_ascii_alphabetic() => {
chars.all(|c| c.is_ascii_alphanumeric() || c == '_')
}
Some(_) => false,
}
}
pub struct SqlxPgGraph {
pool: PgPool,
table_prefix: String,
}
impl SqlxPgGraph {
pub async fn from_pool(pool: PgPool) -> Result<Self> {
Self::from_pool_with_prefix(pool, "").await
}
pub async fn from_pool_with_prefix(pool: PgPool, prefix: &str) -> Result<Self> {
if !is_identifier_safe(prefix) {
return Err(KernelError::Store(format!(
"invalid table prefix {prefix:?}: only ASCII letters, digits, and underscore are allowed (and the first character must not be a digit)"
)));
}
let graph = Self {
pool,
table_prefix: prefix.to_string(),
};
graph.init_schema().await?;
graph.migrate().await?;
Ok(graph)
}
pub async fn connect(url: &str) -> Result<Self> {
Self::connect_with_prefix(url, "").await
}
pub async fn connect_with_prefix(url: &str, prefix: &str) -> Result<Self> {
let pool = PgPoolOptions::new()
.max_connections(8)
.connect(url)
.await
.map_err(pg_err)?;
Self::from_pool_with_prefix(pool, prefix).await
}
pub fn pool(&self) -> &PgPool {
&self.pool
}
fn nodes_tbl(&self) -> String {
format!("{}nodes", self.table_prefix)
}
fn edges_tbl(&self) -> String {
format!("{}edges", self.table_prefix)
}
fn meta_tbl(&self) -> String {
format!("{}_meta", self.table_prefix)
}
async fn init_schema(&self) -> Result<()> {
let p = &self.table_prefix;
let nodes = self.nodes_tbl();
let edges = self.edges_tbl();
let meta = self.meta_tbl();
let idx = |name: &str| format!("{p}{name}");
let stmts: Vec<String> = vec![
format!(
"CREATE TABLE IF NOT EXISTS {nodes} (
id TEXT PRIMARY KEY,
node_type TEXT NOT NULL,
title TEXT NOT NULL,
tags TEXT NOT NULL DEFAULT '',
projects TEXT NOT NULL DEFAULT '',
agents TEXT NOT NULL DEFAULT '',
created TEXT NOT NULL,
updated TEXT NOT NULL,
body TEXT NOT NULL DEFAULT '',
importance DOUBLE PRECISION NOT NULL DEFAULT 0.5,
access_count BIGINT NOT NULL DEFAULT 0,
accessed_at TEXT NOT NULL DEFAULT ''
)"
),
format!(
"CREATE TABLE IF NOT EXISTS {edges} (
id TEXT PRIMARY KEY,
source TEXT NOT NULL,
target TEXT NOT NULL,
relation TEXT NOT NULL DEFAULT 'related',
weight DOUBLE PRECISION NOT NULL DEFAULT 1.0,
ts TEXT NOT NULL
)"
),
format!(
"CREATE INDEX IF NOT EXISTS {} ON {edges}(source)",
idx("idx_edges_source")
),
format!(
"CREATE INDEX IF NOT EXISTS {} ON {edges}(target)",
idx("idx_edges_target")
),
format!(
"CREATE UNIQUE INDEX IF NOT EXISTS {} ON {edges}(source, target, relation)",
idx("idx_edges_src_tgt_rel")
),
format!(
"CREATE INDEX IF NOT EXISTS {} ON {edges}(source, relation)",
idx("idx_edges_src_rel")
),
format!(
"CREATE INDEX IF NOT EXISTS {} ON {edges}(target, relation)",
idx("idx_edges_tgt_rel")
),
format!(
"CREATE INDEX IF NOT EXISTS {} ON {nodes}(node_type)",
idx("idx_nodes_type")
),
format!(
"CREATE INDEX IF NOT EXISTS {} ON {nodes}(updated DESC)",
idx("idx_nodes_updated")
),
format!(
"CREATE INDEX IF NOT EXISTS {} ON {nodes}(importance DESC)",
idx("idx_nodes_importance")
),
format!(
"CREATE INDEX IF NOT EXISTS {} ON {nodes}(accessed_at DESC)",
idx("idx_nodes_accessed")
),
format!(
"CREATE INDEX IF NOT EXISTS {} ON {nodes}(created)",
idx("idx_nodes_created")
),
format!(
"CREATE TABLE IF NOT EXISTS {meta} (key TEXT PRIMARY KEY, value TEXT NOT NULL)"
),
format!(
"INSERT INTO {meta} (key, value) VALUES ('graph_schema_version', '{}')
ON CONFLICT (key) DO NOTHING",
GRAPH_SCHEMA_VERSION
),
];
for ddl in &stmts {
sqlx::query(ddl).execute(&self.pool).await.map_err(pg_err)?;
}
Ok(())
}
pub async fn current_version(&self) -> Result<u32> {
let meta = self.meta_tbl();
let row = sqlx::query(&format!(
"SELECT value FROM {meta} WHERE key = 'graph_schema_version'"
))
.fetch_optional(&self.pool)
.await
.map_err(pg_err)?;
match row {
Some(r) => {
let s: String = r.get("value");
Ok(s.parse().unwrap_or(0))
}
None => Ok(0),
}
}
pub async fn migrate(&self) -> Result<u32> {
let current = self.current_version().await?;
if current >= GRAPH_SCHEMA_VERSION {
return Ok(current);
}
let p = &self.table_prefix;
let nodes = self.nodes_tbl();
let edges = self.edges_tbl();
let meta = self.meta_tbl();
let idx = |name: &str| format!("{p}{name}");
let mut tx = self.pool.begin().await.map_err(pg_err)?;
let mut v = current;
if v < 2 {
sqlx::query(&format!(
"CREATE INDEX IF NOT EXISTS {} ON {nodes}(created)",
idx("idx_nodes_created")
))
.execute(&mut *tx)
.await
.map_err(pg_err)?;
v = 2;
}
if v < 3 {
sqlx::query(&format!(
"CREATE INDEX IF NOT EXISTS {} ON {edges}(source, relation)",
idx("idx_edges_src_rel")
))
.execute(&mut *tx)
.await
.map_err(pg_err)?;
sqlx::query(&format!(
"CREATE INDEX IF NOT EXISTS {} ON {edges}(target, relation)",
idx("idx_edges_tgt_rel")
))
.execute(&mut *tx)
.await
.map_err(pg_err)?;
v = 3;
}
sqlx::query(&format!(
"UPDATE {meta} SET value = $1 WHERE key = 'graph_schema_version'"
))
.bind(v.to_string())
.execute(&mut *tx)
.await
.map_err(pg_err)?;
tx.commit().await.map_err(pg_err)?;
Ok(v)
}
pub async fn upsert_node(&self, node: &GraphNode) -> Result<()> {
let tags = join_csv(&node.tags);
let projects = join_csv(&node.projects);
let agents = join_csv(&node.agents);
let nodes = self.nodes_tbl();
sqlx::query(&format!(
"INSERT INTO {nodes} (id, node_type, title, tags, projects, agents, created, updated, body, importance, access_count, accessed_at)
VALUES ($1,$2,$3,$4,$5,$6,$7,$8,$9,$10,$11,$12)
ON CONFLICT (id) DO UPDATE SET
node_type=EXCLUDED.node_type, title=EXCLUDED.title, tags=EXCLUDED.tags,
projects=EXCLUDED.projects, agents=EXCLUDED.agents, created=EXCLUDED.created,
updated=EXCLUDED.updated, body=EXCLUDED.body, importance=EXCLUDED.importance,
access_count=EXCLUDED.access_count, accessed_at=EXCLUDED.accessed_at"
))
.bind(node.id.as_str())
.bind(node.node_type.as_str())
.bind(node.title.as_str())
.bind(tags.as_str())
.bind(projects.as_str())
.bind(agents.as_str())
.bind(node.created.as_str())
.bind(node.updated.as_str())
.bind(node.body.as_str())
.bind(node.importance)
.bind(node.access_count)
.bind(node.accessed_at.as_str())
.execute(&self.pool)
.await
.map_err(pg_err)?;
Ok(())
}
pub async fn read_node(&self, id: &str) -> Result<Option<GraphNode>> {
let nodes = self.nodes_tbl();
let row = sqlx::query(&format!("SELECT {NODE_COLUMNS} FROM {nodes} WHERE id = $1"))
.bind(id)
.fetch_optional(&self.pool)
.await
.map_err(pg_err)?;
Ok(row.as_ref().map(row_to_node))
}
pub async fn delete_node(&self, id: &str) -> Result<bool> {
let nodes = self.nodes_tbl();
let res = sqlx::query(&format!("DELETE FROM {nodes} WHERE id = $1"))
.bind(id)
.execute(&self.pool)
.await
.map_err(pg_err)?;
Ok(res.rows_affected() > 0)
}
pub async fn append_edge(&self, edge: &GraphEdge) -> Result<()> {
let edges = self.edges_tbl();
sqlx::query(&format!(
"INSERT INTO {edges} (id, source, target, relation, weight, ts)
VALUES ($1,$2,$3,$4,$5,$6) ON CONFLICT DO NOTHING"
))
.bind(edge.id.as_str())
.bind(edge.source.as_str())
.bind(edge.target.as_str())
.bind(edge.relation.as_str())
.bind(edge.weight)
.bind(edge.ts.as_str())
.execute(&self.pool)
.await
.map_err(pg_err)?;
Ok(())
}
pub async fn edges_for_node(&self, node_id: &str) -> Result<Vec<GraphEdge>> {
let edges = self.edges_tbl();
let rows = sqlx::query(&format!(
"SELECT id, source, target, relation, weight, ts FROM {edges} \
WHERE source = $1 OR target = $1"
))
.bind(node_id)
.fetch_all(&self.pool)
.await
.map_err(pg_err)?;
Ok(rows.iter().map(row_to_edge).collect())
}
pub async fn delete_edge(&self, id: &str) -> Result<bool> {
let edges = self.edges_tbl();
let res = sqlx::query(&format!("DELETE FROM {edges} WHERE id = $1"))
.bind(id)
.execute(&self.pool)
.await
.map_err(pg_err)?;
Ok(res.rows_affected() > 0)
}
pub async fn append_edges(&self, edges: &[GraphEdge]) -> Result<()> {
if edges.is_empty() {
return Ok(());
}
let edges_tbl = self.edges_tbl();
let sql = format!(
"INSERT INTO {edges_tbl} (id, source, target, relation, weight, ts)
VALUES ($1,$2,$3,$4,$5,$6) ON CONFLICT DO NOTHING"
);
const CHUNK: usize = 5000;
for chunk in edges.chunks(CHUNK) {
let mut tx = self.pool.begin().await.map_err(pg_err)?;
for e in chunk {
sqlx::query(&sql)
.bind(e.id.as_str())
.bind(e.source.as_str())
.bind(e.target.as_str())
.bind(e.relation.as_str())
.bind(e.weight)
.bind(e.ts.as_str())
.execute(&mut *tx)
.await
.map_err(pg_err)?;
}
tx.commit().await.map_err(pg_err)?;
}
Ok(())
}
pub async fn append_edges_in_tx(
&self,
tx: &mut sqlx::PgConnection,
edges: &[GraphEdge],
) -> Result<()> {
if edges.is_empty() {
return Ok(());
}
let edges_tbl = self.edges_tbl();
let sql = format!(
"INSERT INTO {edges_tbl} (id, source, target, relation, weight, ts)
VALUES ($1,$2,$3,$4,$5,$6) ON CONFLICT DO NOTHING"
);
for e in edges {
sqlx::query(&sql)
.bind(e.id.as_str())
.bind(e.source.as_str())
.bind(e.target.as_str())
.bind(e.relation.as_str())
.bind(e.weight)
.bind(e.ts.as_str())
.execute(&mut *tx)
.await
.map_err(pg_err)?;
}
Ok(())
}
pub async fn edges_for_node_dir(
&self,
node_id: &str,
dir: EdgeDirection,
relation: Option<&str>,
) -> Result<Vec<GraphEdge>> {
let edges = self.edges_tbl();
let dir_clause = match dir {
EdgeDirection::Out => "source = $1",
EdgeDirection::In => "target = $1",
EdgeDirection::Both => "(source = $1 OR target = $1)",
};
let rows = if let Some(r) = relation {
let sql = format!(
"SELECT id, source, target, relation, weight, ts FROM {edges} \
WHERE {dir_clause} AND relation = $2 ORDER BY weight DESC"
);
sqlx::query(&sql)
.bind(node_id)
.bind(r)
.fetch_all(&self.pool)
.await
.map_err(pg_err)?
} else {
let sql = format!(
"SELECT id, source, target, relation, weight, ts FROM {edges} \
WHERE {dir_clause} ORDER BY weight DESC"
);
sqlx::query(&sql)
.bind(node_id)
.fetch_all(&self.pool)
.await
.map_err(pg_err)?
};
Ok(rows.iter().map(row_to_edge).collect())
}
pub async fn neighbors_weighted(
&self,
seed_ids: &[String],
dir: EdgeDirection,
relation: Option<&str>,
) -> Result<Vec<(String, f64)>> {
if seed_ids.is_empty() {
return Ok(vec![]);
}
let edges = self.edges_tbl();
const MAX_SEEDS: usize = 100;
let seed_ids = if seed_ids.len() > MAX_SEEDS {
&seed_ids[..MAX_SEEDS]
} else {
seed_ids
};
let seed_arr: Vec<String> = seed_ids.to_vec();
let seed_set: HashSet<&str> = seed_ids.iter().map(String::as_str).collect();
let halves: &[&str] = match dir {
EdgeDirection::Out => &["source"],
EdgeDirection::In => &["target"],
EdgeDirection::Both => &["source", "target"],
};
let mut weights: HashMap<String, f64> = HashMap::new();
for &follow in halves {
let select_col = if follow == "source" {
"target"
} else {
"source"
};
let rel_clause = relation.map(|_| " AND relation = $2").unwrap_or("");
let sql = format!(
"SELECT {select_col} AS nb, SUM(weight) AS w FROM {edges} \
WHERE {follow} = ANY($1){rel_clause} GROUP BY {select_col}"
);
let rows = if let Some(r) = relation {
sqlx::query(&sql)
.bind(&seed_arr)
.bind(r)
.fetch_all(&self.pool)
.await
.map_err(pg_err)?
} else {
sqlx::query(&sql)
.bind(&seed_arr)
.fetch_all(&self.pool)
.await
.map_err(pg_err)?
};
for row in &rows {
let nb: String = row.get("nb");
let w: f64 = row.get("w");
if !seed_set.contains(nb.as_str()) {
*weights.entry(nb).or_default() += w;
}
}
}
let mut result: Vec<(String, f64)> = weights.into_iter().collect();
result.sort_by(|a, b| b.1.partial_cmp(&a.1).unwrap_or(std::cmp::Ordering::Equal));
Ok(result)
}
pub async fn remove_edges_for_node(&self, node_id: &str) -> Result<()> {
let edges = self.edges_tbl();
sqlx::query(&format!(
"DELETE FROM {edges} WHERE source = $1 OR target = $1"
))
.bind(node_id)
.execute(&self.pool)
.await
.map_err(pg_err)?;
Ok(())
}
pub async fn remove_edges_for_node_in_tx(
&self,
tx: &mut sqlx::PgConnection,
node_id: &str,
) -> Result<()> {
let edges = self.edges_tbl();
sqlx::query(&format!(
"DELETE FROM {edges} WHERE source = $1 OR target = $1"
))
.bind(node_id)
.execute(tx)
.await
.map_err(pg_err)?;
Ok(())
}
pub async fn search_nodes(&self, query: &str, limit: usize) -> Result<Vec<GraphNode>> {
let terms = search_patterns(query);
if terms.is_empty() {
return Ok(vec![]);
}
let nodes = self.nodes_tbl();
let mut qb = QueryBuilder::<Postgres>::new("SELECT ");
qb.push(NODE_COLUMNS);
qb.push(" FROM ");
qb.push(nodes.as_str());
qb.push(" WHERE ");
for (i, term) in terms.iter().enumerate() {
if i > 0 {
qb.push(" AND ");
}
qb.push("(title || ' ' || body || ' ' || tags) ILIKE ");
qb.push_bind(term.clone());
qb.push(" ESCAPE '\\'");
}
qb.push(" ORDER BY importance DESC, updated DESC LIMIT ");
qb.push_bind(limit as i64);
let rows = qb.build().fetch_all(&self.pool).await.map_err(pg_err)?;
Ok(rows.iter().map(row_to_node).collect())
}
pub async fn related_nodes(&self, start_id: &str, depth: usize) -> Result<Vec<String>> {
let edges = self.edges_tbl();
let sql = format!(
"WITH RECURSIVE bfs(node_id, lvl) AS (
SELECT nb.node_id, 1 FROM (
SELECT target AS node_id FROM {edges} WHERE source = $1
UNION
SELECT source AS node_id FROM {edges} WHERE target = $1
) nb
UNION
SELECT CASE WHEN e.source = bfs.node_id THEN e.target ELSE e.source END,
bfs.lvl + 1
FROM bfs
JOIN {edges} e ON e.source = bfs.node_id OR e.target = bfs.node_id
WHERE bfs.lvl < $2
)
SELECT DISTINCT node_id FROM bfs WHERE node_id <> $1 LIMIT 500"
);
let rows = sqlx::query(&sql)
.bind(start_id)
.bind(depth as i32)
.fetch_all(&self.pool)
.await
.map_err(pg_err)?;
Ok(rows
.iter()
.map(|r| {
let id: String = r.get("node_id");
id
})
.collect())
}
}
#[cfg(test)]
mod tests {
use super::*;
fn pg_url() -> Option<String> {
std::env::var("LLMKERNEL_PG_URL").ok()
}
fn sample_node(id: &str) -> GraphNode {
GraphNode {
id: id.to_string(),
node_type: "concept".to_string(),
title: format!("Node {id}"),
body: "sqlx pg test body".to_string(),
tags: vec!["sqlx".to_string()],
projects: vec![],
agents: vec![],
created: "2026-01-01T00:00:00Z".to_string(),
updated: "2026-01-01T00:00:00Z".to_string(),
importance: 0.5,
access_count: 0,
accessed_at: String::new(),
}
}
fn sample_edge(id: &str, src: &str, tgt: &str, rel: &str, w: f64) -> GraphEdge {
GraphEdge {
id: id.to_string(),
source: src.to_string(),
target: tgt.to_string(),
relation: rel.to_string(),
weight: w,
ts: "2026-01-01T00:00:00Z".to_string(),
}
}
async fn cleanup(pool: &PgPool, prefix: &str) {
let _ = sqlx::query(&format!("DROP TABLE IF EXISTS {}nodes", prefix))
.execute(pool)
.await;
let _ = sqlx::query(&format!("DROP TABLE IF EXISTS {}edges", prefix))
.execute(pool)
.await;
let _ = sqlx::query(&format!("DROP TABLE IF EXISTS {}_meta", prefix))
.execute(pool)
.await;
}
#[test]
fn search_patterns_escapes_and_wraps() {
assert!(search_patterns("").is_empty());
assert_eq!(search_patterns("rust"), vec!["%rust%".to_string()]);
assert_eq!(search_patterns("100%"), vec!["%100\\%%".to_string()]);
assert_eq!(search_patterns("a_b"), vec!["%a\\_b%".to_string()]);
assert_eq!(
search_patterns("rust db"),
vec!["%rust%".to_string(), "%db%".to_string()]
);
}
#[tokio::test]
async fn basic_crud_roundtrip() {
let Some(url) = pg_url() else {
eprintln!("skip: LLMKERNEL_PG_URL unset");
return;
};
let prefix = "lk_sg1_";
let g = SqlxPgGraph::connect_with_prefix(&url, prefix)
.await
.expect("connect");
assert!(g.read_node("n1").await.unwrap().is_none());
g.upsert_node(&sample_node("n1")).await.unwrap();
let loaded = g.read_node("n1").await.unwrap().unwrap();
assert_eq!(loaded.title, "Node n1");
assert_eq!(loaded.tags, vec!["sqlx".to_string()]);
let mut updated = sample_node("n1");
updated.title = "Updated".into();
g.upsert_node(&updated).await.unwrap();
assert_eq!(g.read_node("n1").await.unwrap().unwrap().title, "Updated");
g.upsert_node(&sample_node("n2")).await.unwrap();
g.append_edge(&sample_edge("e1", "n1", "n2", "related", 1.0))
.await
.unwrap();
assert_eq!(g.edges_for_node("n1").await.unwrap().len(), 1);
assert!(g.delete_edge("e1").await.unwrap());
assert!(!g.delete_edge("e1").await.unwrap());
assert_eq!(g.edges_for_node("n1").await.unwrap().len(), 0);
assert!(g.delete_node("n1").await.unwrap());
assert!(!g.delete_node("n1").await.unwrap());
assert!(g.read_node("n1").await.unwrap().is_none());
cleanup(g.pool(), prefix).await;
}
#[tokio::test]
async fn batch_edges_dedup() {
let Some(url) = pg_url() else {
eprintln!("skip: LLMKERNEL_PG_URL unset");
return;
};
let prefix = "lk_sg2_";
let g = SqlxPgGraph::connect_with_prefix(&url, prefix)
.await
.expect("connect");
for n in &["a", "b", "c"] {
g.upsert_node(&sample_node(n)).await.unwrap();
}
let edges = vec![
sample_edge("e1", "a", "b", "cites", 1.0),
sample_edge("e2", "a", "c", "cites", 0.8),
sample_edge("e1dup", "a", "b", "cites", 2.0),
];
g.append_edges(&edges).await.unwrap();
let out = g
.edges_for_node_dir("a", EdgeDirection::Out, Some("cites"))
.await
.unwrap();
assert_eq!(out.len(), 2, "duplicate (src,tgt,rel) edge ignored");
g.append_edges(&[]).await.unwrap();
cleanup(g.pool(), prefix).await;
}
#[tokio::test]
async fn edges_dir_and_neighbors() {
let Some(url) = pg_url() else {
eprintln!("skip: LLMKERNEL_PG_URL unset");
return;
};
let prefix = "lk_sg3_";
let g = SqlxPgGraph::connect_with_prefix(&url, prefix)
.await
.expect("connect");
for n in &["seed", "t1", "t2", "t3"] {
g.upsert_node(&sample_node(n)).await.unwrap();
}
g.append_edges(&[
sample_edge("e1", "seed", "t1", "cites", 1.0),
sample_edge("e2", "seed", "t2", "cites", 0.5),
sample_edge("e3", "seed", "t3", "cites", 0.3),
])
.await
.unwrap();
g.append_edge(&sample_edge("e4", "t1", "seed", "cites", 2.0))
.await
.unwrap();
g.append_edge(&sample_edge("e5", "seed", "t1", "related", 1.0))
.await
.unwrap();
let out_cites = g
.edges_for_node_dir("seed", EdgeDirection::Out, Some("cites"))
.await
.unwrap();
assert_eq!(out_cites.len(), 3);
let in_cites = g
.edges_for_node_dir("seed", EdgeDirection::In, Some("cites"))
.await
.unwrap();
assert_eq!(in_cites.len(), 1);
assert_eq!(in_cites[0].source, "t1");
let neighbors = g
.neighbors_weighted(&["seed".to_string()], EdgeDirection::Out, Some("cites"))
.await
.unwrap();
assert_eq!(neighbors.len(), 3);
assert_eq!(neighbors[0].0, "t1");
assert!((neighbors[0].1 - 1.0).abs() < 1e-9);
assert_eq!(neighbors[2].0, "t3");
cleanup(g.pool(), prefix).await;
}
#[tokio::test]
async fn remove_edges_for_node_drops_only_touching() {
let Some(url) = pg_url() else {
eprintln!("skip: LLMKERNEL_PG_URL unset");
return;
};
let prefix = "lk_sg4_";
let g = SqlxPgGraph::connect_with_prefix(&url, prefix)
.await
.expect("connect");
for n in &["x", "y", "z"] {
g.upsert_node(&sample_node(n)).await.unwrap();
}
g.append_edges(&[
sample_edge("e1", "x", "y", "cites", 1.0),
sample_edge("e2", "z", "x", "cites", 1.0),
sample_edge("e3", "y", "z", "cites", 1.0),
])
.await
.unwrap();
assert_eq!(g.edges_for_node("x").await.unwrap().len(), 2);
g.remove_edges_for_node("x").await.unwrap();
assert_eq!(g.edges_for_node("x").await.unwrap().len(), 0);
assert_eq!(g.edges_for_node("y").await.unwrap().len(), 1);
cleanup(g.pool(), prefix).await;
}
#[tokio::test]
async fn tx_atomicity_commit_and_rollback() {
let Some(url) = pg_url() else {
eprintln!("skip: LLMKERNEL_PG_URL unset");
return;
};
let prefix = "lk_sg5_";
let g = SqlxPgGraph::connect_with_prefix(&url, prefix)
.await
.expect("connect");
for n in &["p", "q"] {
g.upsert_node(&sample_node(n)).await.unwrap();
}
let mut tx = g.pool().begin().await.expect("begin tx");
g.append_edges_in_tx(&mut tx, &[sample_edge("e1", "p", "q", "cites", 1.0)])
.await
.unwrap();
tx.commit().await.expect("commit");
assert_eq!(g.edges_for_node("p").await.unwrap().len(), 1);
let mut tx2 = g.pool().begin().await.expect("begin tx2");
g.append_edges_in_tx(&mut tx2, &[sample_edge("e2", "p", "q", "cites", 1.0)])
.await
.unwrap();
tx2.rollback().await.expect("rollback");
assert_eq!(
g.edges_for_node("p").await.unwrap().len(),
1,
"rolled-back edge not visible"
);
let mut tx3 = g.pool().begin().await.expect("begin tx3");
g.remove_edges_for_node_in_tx(&mut tx3, "p").await.unwrap();
tx3.rollback().await.expect("rollback tx3");
assert_eq!(g.edges_for_node("p").await.unwrap().len(), 1);
let mut tx4 = g.pool().begin().await.expect("begin tx4");
g.remove_edges_for_node_in_tx(&mut tx4, "p").await.unwrap();
tx4.commit().await.expect("commit tx4");
assert_eq!(g.edges_for_node("p").await.unwrap().len(), 0);
cleanup(g.pool(), prefix).await;
}
#[tokio::test]
async fn search_related_and_version() {
let Some(url) = pg_url() else {
eprintln!("skip: LLMKERNEL_PG_URL unset");
return;
};
let prefix = "lk_sg6_";
let g = SqlxPgGraph::connect_with_prefix(&url, prefix)
.await
.expect("connect");
let mut rust = sample_node("rust");
rust.title = "Rust ownership".into();
rust.body = "borrow checker".into();
g.upsert_node(&rust).await.unwrap();
g.upsert_node(&sample_node("py")).await.unwrap();
let hits = g.search_nodes("rust", 10).await.unwrap();
assert_eq!(hits.len(), 1);
assert_eq!(hits[0].id, "rust");
assert!(g.search_nodes("", 10).await.unwrap().is_empty());
g.append_edge(&sample_edge("e1", "rust", "py", "related", 1.0))
.await
.unwrap();
let related = g.related_nodes("rust", 2).await.unwrap();
assert!(related.contains(&"py".to_string()));
assert_eq!(g.current_version().await.unwrap(), GRAPH_SCHEMA_VERSION);
assert_eq!(g.migrate().await.unwrap(), GRAPH_SCHEMA_VERSION);
cleanup(g.pool(), prefix).await;
}
#[tokio::test]
async fn invalid_prefix_rejected() {
let Some(url) = pg_url() else {
eprintln!("skip: LLMKERNEL_PG_URL unset");
return;
};
assert!(
SqlxPgGraph::connect_with_prefix(&url, "lk; drop")
.await
.is_err()
);
assert!(SqlxPgGraph::connect_with_prefix(&url, "1lk").await.is_err());
}
}