from __future__ import annotations from typing import Any from neo4j import GraphDatabase from sleep_db.config import Settings from sleep_db.rules import load_rules_directory from sleep_db.schema import ALLOWED_LABELS, ALLOWED_REL_TYPES def primary_rule_content_node_id(rule: dict[str, Any]) -> str | None: nodes = rule.get("nodes", []) for node in nodes: if node.get("label") == "ConditionGroup": return node.get("id") if nodes: return nodes[0].get("id") return None class KnowledgeGraphSync: def __init__(self, settings: Settings): self.driver = GraphDatabase.driver( settings.neo4j_uri, auth=(settings.neo4j_user, settings.neo4j_password), ) def close(self) -> None: self.driver.close() def create_constraints(self) -> None: queries = [ "CREATE CONSTRAINT metric_id IF NOT EXISTS FOR (n:Metric) REQUIRE n.id IS UNIQUE", "CREATE CONSTRAINT condition_id IF NOT EXISTS FOR (n:ThresholdCondition) REQUIRE n.id IS UNIQUE", "CREATE CONSTRAINT condition_group_id IF NOT EXISTS FOR (n:ConditionGroup) REQUIRE n.id IS UNIQUE", "CREATE CONSTRAINT state_id IF NOT EXISTS FOR (n:State) REQUIRE n.id IS UNIQUE", "CREATE CONSTRAINT effect_id IF NOT EXISTS FOR (n:Effect) REQUIRE n.id IS UNIQUE", "CREATE CONSTRAINT advice_id IF NOT EXISTS FOR (n:Advice) REQUIRE n.id IS UNIQUE", "CREATE CONSTRAINT activity_id IF NOT EXISTS FOR (n:Activity) REQUIRE n.id IS UNIQUE", "CREATE CONSTRAINT environment_id IF NOT EXISTS FOR (n:Environment) REQUIRE n.id IS UNIQUE", "CREATE CONSTRAINT nutrition_id IF NOT EXISTS FOR (n:Nutrition) REQUIRE n.id IS UNIQUE", "CREATE CONSTRAINT doc_doc_id IF NOT EXISTS FOR (n:Doc) REQUIRE n.doc_id IS UNIQUE", "CREATE CONSTRAINT rule_version_id IF NOT EXISTS FOR (n:RuleVersion) REQUIRE n.rule_id IS UNIQUE", ] with self.driver.session() as session: for query in queries: session.run(query) def sync_rules_directory(self, rules_dir) -> dict[str, list[str]]: self.create_constraints() current = {rule["id"]: (rule, rule_hash) for rule, rule_hash in load_rules_directory(rules_dir)} with self.driver.session() as session: rows = session.run( "MATCH (rv:RuleVersion) RETURN rv.rule_id AS id, rv.yaml_hash AS hash" ) existing = {row["id"]: row["hash"] for row in rows} report = {"added": [], "modified": [], "deleted": [], "unchanged": []} for rule_id, (rule, rule_hash) in current.items(): if rule_id not in existing: self.replace_rule(rule, rule_hash) report["added"].append(rule_id) elif existing[rule_id] != rule_hash: self.replace_rule(rule, rule_hash) report["modified"].append(rule_id) else: self.ensure_rule_version_edge(rule) report["unchanged"].append(rule_id) for deleted_rule_id in sorted(set(existing) - set(current)): self.delete_rule(deleted_rule_id) report["deleted"].append(deleted_rule_id) return report def replace_rule(self, rule: dict[str, Any], rule_hash: str) -> None: self.delete_rule(rule["id"]) self.upsert_rule(rule, rule_hash) def upsert_rule(self, rule: dict[str, Any], rule_hash: str) -> None: with self.driver.session() as session: for doc_id in rule.get("source_doc_ids", []): session.execute_write(self._merge_doc, doc_id) for node in rule.get("nodes", []): props = dict(node.get("properties", {})) props.update( { "name": node.get("name"), "review_status": rule.get("review_status"), "confidence": rule.get("confidence"), "source_doc_ids": rule.get("source_doc_ids", []), } ) props = {key: value for key, value in props.items() if value is not None} session.execute_write( self._merge_node, node["label"], node["id"], props, rule["id"], ) for edge in rule.get("edges", []): props = dict(edge.get("properties", {})) props.update( { "review_status": rule.get("review_status"), "confidence": props.get("confidence", rule.get("confidence")), "source_doc_ids": rule.get("source_doc_ids", []), } ) session.execute_write( self._merge_edge, edge["from"], edge["to"], edge["type"], props, rule["id"], ) for node in rule.get("nodes", []): for doc_id in rule.get("source_doc_ids", []): session.execute_write(self._merge_refers_to, node["id"], doc_id, rule["id"]) session.run( """ MERGE (rv:RuleVersion {rule_id: $rule_id}) SET rv.version = $version, rv.yaml_hash = $yaml_hash, rv.review_status = $review_status, rv.updated_at = datetime() """, rule_id=rule["id"], version=rule.get("version", 1), yaml_hash=rule_hash, review_status=rule.get("review_status"), ) self._ensure_rule_version_edge(session, rule) def ensure_rule_version_edge(self, rule: dict[str, Any]) -> None: with self.driver.session() as session: self._ensure_rule_version_edge(session, rule) def _ensure_rule_version_edge(self, session: Any, rule: dict[str, Any]) -> None: content_node_id = primary_rule_content_node_id(rule) if not content_node_id: return session.execute_write( self._merge_rule_version_defines, rule["id"], content_node_id, { "review_status": rule.get("review_status"), "confidence": rule.get("confidence"), "source_doc_ids": rule.get("source_doc_ids", []), }, ) def delete_rule(self, rule_id: str) -> None: with self.driver.session() as session: session.run( """ MATCH ()-[r]->() WHERE $rule_id IN coalesce(r.source_rule_ids, []) AND size(r.source_rule_ids) = 1 DELETE r """, rule_id=rule_id, ) session.run( """ MATCH ()-[r]->() WHERE $rule_id IN coalesce(r.source_rule_ids, []) SET r.source_rule_ids = [id IN r.source_rule_ids WHERE id <> $rule_id] """, rule_id=rule_id, ) session.run( """ MATCH (n) WHERE $rule_id IN coalesce(n.source_rule_ids, []) SET n.source_rule_ids = [id IN n.source_rule_ids WHERE id <> $rule_id] """, rule_id=rule_id, ) session.run( """ MATCH (n) WHERE n.source_rule_ids = [] AND NOT (n)--() DELETE n """, rule_id=rule_id, ) session.run( "MATCH (rv:RuleVersion {rule_id: $rule_id}) DELETE rv", rule_id=rule_id, ) @staticmethod def _merge_doc(tx, doc_id: str) -> None: tx.run("MERGE (:Doc {doc_id: $doc_id})", doc_id=doc_id) @staticmethod def _merge_node(tx, label: str, node_id: str, props: dict[str, Any], rule_id: str) -> None: if label not in ALLOWED_LABELS: raise ValueError(f"Unsupported label: {label}") query = f""" MERGE (n:{label} {{id: $node_id}}) SET n += $props SET n.source_rule_ids = CASE WHEN n.source_rule_ids IS NULL THEN [$rule_id] WHEN NOT $rule_id IN n.source_rule_ids THEN n.source_rule_ids + $rule_id ELSE n.source_rule_ids END """ tx.run(query, node_id=node_id, props=props, rule_id=rule_id) @staticmethod def _merge_edge( tx, from_id: str, to_id: str, rel_type: str, props: dict[str, Any], rule_id: str, ) -> None: if rel_type not in ALLOWED_REL_TYPES: raise ValueError(f"Unsupported relationship type: {rel_type}") query = f""" MATCH (a {{id: $from_id}}) MATCH (b {{id: $to_id}}) MERGE (a)-[r:{rel_type}]->(b) SET r += $props SET r.source_rule_ids = CASE WHEN r.source_rule_ids IS NULL THEN [$rule_id] WHEN NOT $rule_id IN r.source_rule_ids THEN r.source_rule_ids + $rule_id ELSE r.source_rule_ids END """ tx.run(query, from_id=from_id, to_id=to_id, props=props, rule_id=rule_id) @staticmethod def _merge_refers_to(tx, node_id: str, doc_id: str, rule_id: str) -> None: tx.run( """ MATCH (n {id: $node_id}) MATCH (d:Doc {doc_id: $doc_id}) MERGE (n)-[r:REFERS_TO]->(d) SET r.source_rule_ids = CASE WHEN r.source_rule_ids IS NULL THEN [$rule_id] WHEN NOT $rule_id IN r.source_rule_ids THEN r.source_rule_ids + $rule_id ELSE r.source_rule_ids END """, node_id=node_id, doc_id=doc_id, rule_id=rule_id, ) @staticmethod def _merge_rule_version_defines( tx, rule_id: str, node_id: str, props: dict[str, Any], ) -> None: tx.run( """ MATCH (rv:RuleVersion {rule_id: $rule_id}) MATCH (n {id: $node_id}) MERGE (rv)-[r:DEFINES]->(n) SET r += $props SET r.source_rule_ids = CASE WHEN r.source_rule_ids IS NULL THEN [$rule_id] WHEN NOT $rule_id IN r.source_rule_ids THEN r.source_rule_ids + $rule_id ELSE r.source_rule_ids END """, rule_id=rule_id, node_id=node_id, props=props, )