"""Guarded graph-write plans. No caller can submit Cypher. The writer accepts only normalized node/edge records and executes fixed parameterized templates when explicitly enabled. """ from __future__ import annotations import hashlib from dataclasses import dataclass from typing import Any, Protocol class GraphTransaction(Protocol): async def run(self, query: str, parameters: dict[str, Any]) -> Any: ... @dataclass(frozen=True, slots=True) class GraphWritePlan: scope: str nodes: tuple[dict[str, Any], ...] edges: tuple[dict[str, Any], ...] NODE_QUERY = """ MERGE (n:DeerFlowKnowledge {scope: $scope, entity_key: $entity_key}) SET n.display_name = $display_name, n.entity_type = $entity_type, n.binding_id = $binding_id, n.snapshot_id = $snapshot_id, n.job_id = $job_id, n.confidence = $confidence, n.updated_at = datetime() """.strip() EDGE_QUERY = """ MATCH (a:DeerFlowKnowledge {scope: $scope, entity_key: $source_key}) MATCH (b:DeerFlowKnowledge {scope: $scope, entity_key: $target_key}) MERGE (a)-[r:DEERFLOW_RELATION {edge_key: $edge_key}]->(b) SET r.predicate = $predicate, r.binding_id = $binding_id, r.snapshot_id = $snapshot_id, r.job_id = $job_id, r.confidence = $confidence, r.evidence = $evidence, r.updated_at = datetime() """.strip() def _key(*parts: str) -> str: return hashlib.sha256("\x1f".join(parts).encode()).hexdigest() def plan_from_ir( ir: dict[str, Any], *, target_type: str, target_id: str, binding_id: str, snapshot_id: str, job_id: str, ) -> GraphWritePlan: """Build a stable write plan; callers cannot supply executable Cypher.""" scope = f"{target_type}:{target_id}" local_keys: dict[str, str] = {} nodes: list[dict[str, Any]] = [] for entity in ir.get("entities") or []: entity_type = str(entity.get("type") or "concept") normalized_name = str(entity.get("normalized_name") or entity.get("name") or "").strip().casefold() if not normalized_name: continue entity_key = _key(entity_type, normalized_name) local_keys[str(entity.get("local_id") or normalized_name)] = entity_key nodes.append( { "entity_key": entity_key, "display_name": str(entity.get("name") or normalized_name), "entity_type": entity_type, "binding_id": binding_id, "snapshot_id": snapshot_id, "job_id": job_id, "confidence": float(entity.get("confidence") or 1.0), } ) edges: list[dict[str, Any]] = [] for relation in ir.get("relations") or []: source_key = local_keys.get(str(relation.get("source_local_id") or "")) target_key = local_keys.get(str(relation.get("target_local_id") or "")) evidence = list(relation.get("evidence") or []) confidence = float(relation.get("confidence") or 0) predicate = str(relation.get("predicate") or "related_to") if not source_key or not target_key or confidence < 0.75 or not evidence: continue edges.append( { "edge_key": _key(source_key, predicate, target_key, binding_id, snapshot_id), "source_key": source_key, "target_key": target_key, "predicate": predicate, "binding_id": binding_id, "snapshot_id": snapshot_id, "job_id": job_id, "confidence": confidence, "evidence": evidence, } ) return GraphWritePlan(scope=scope, nodes=tuple(nodes), edges=tuple(edges)) def parameterized_statements(plan: GraphWritePlan) -> list[dict[str, Any]]: """Return Neo4j transactional HTTP statements using only fixed templates.""" statements = [{"statement": NODE_QUERY, "parameters": {"scope": plan.scope, **node}} for node in plan.nodes] statements.extend({"statement": EDGE_QUERY, "parameters": {"scope": plan.scope, **edge}} for edge in plan.edges) return statements async def execute_plan(tx: GraphTransaction, plan: GraphWritePlan, *, enabled: bool) -> dict[str, int]: if not enabled: return {"nodes": 0, "edges": 0} for node in plan.nodes: await tx.run(NODE_QUERY, {"scope": plan.scope, **node}) for edge in plan.edges: await tx.run(EDGE_QUERY, {"scope": plan.scope, **edge}) return {"nodes": len(plan.nodes), "edges": len(plan.edges)}