116 lines
4.4 KiB
Python
116 lines
4.4 KiB
Python
"""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)}
|