dp-zp-agent/scripts/build_graphrag.py
2026-09-29 18:42:17 +02:00

464 lines
8.5 KiB
Python

from __future__ import annotations
import argparse
import json
from pathlib import Path
from typing import Any
from scripts.graph_db import (
create_neo4j_driver,
get_neo4j_settings,
)
from scripts.graphrag import (
GRAPH_SCOPE,
build_graph_payload,
)
PROJECT_ROOT = Path(
__file__
).resolve().parents[1]
DOCUMENTS_FILE = (
PROJECT_ROOT
/ "data"
/ "documents.json"
)
CHUNKS_FILE = (
PROJECT_ROOT
/ "data"
/ "chunks.json"
)
CONSTRAINTS = (
(
"zpwiki_document_path_unique",
"Document",
"path",
),
(
"zpwiki_person_id_unique",
"Person",
"id",
),
(
"zpwiki_author_id_unique",
"Author",
"id",
),
(
"zpwiki_category_name_unique",
"Category",
"name",
),
(
"zpwiki_topic_id_unique",
"Topic",
"id",
),
(
"zpwiki_work_id_unique",
"Work",
"id",
),
)
NODE_QUERIES = {
"people": """
UNWIND $rows AS row
MERGE (n:Person {id: row.id})
SET n += row
""",
"authors": """
UNWIND $rows AS row
MERGE (n:Author {id: row.id})
SET n += row
""",
"documents": """
UNWIND $rows AS row
MERGE (n:Document {path: row.path})
SET n += row
""",
"categories": """
UNWIND $rows AS row
MERGE (n:Category {name: row.name})
SET n += row
""",
"topics": """
UNWIND $rows AS row
MERGE (n:Topic {id: row.id})
SET n += row
""",
"works": """
UNWIND $rows AS row
MERGE (n:Work {id: row.id})
SET n += row
""",
}
RELATIONSHIP_QUERIES = {
"person_documents": """
UNWIND $rows AS row
MATCH (a:Person {id: row.person_id})
MATCH (b:Document {
path: row.document_path
})
MERGE (a)-[:HAS_DOCUMENT]->(b)
""",
"author_documents": """
UNWIND $rows AS row
MATCH (a:Author {id: row.author_id})
MATCH (b:Document {
path: row.document_path
})
MERGE (a)-[:AUTHORED]->(b)
""",
"document_categories": """
UNWIND $rows AS row
MATCH (a:Document {
path: row.document_path
})
MATCH (b:Category {
name: row.category_name
})
MERGE (a)-[:IN_CATEGORY]->(b)
""",
"document_topics": """
UNWIND $rows AS row
MATCH (a:Document {
path: row.document_path
})
MATCH (b:Topic {
id: row.topic_id
})
MERGE (a)-[:HAS_TOPIC]->(b)
""",
"document_describes_topics": """
UNWIND $rows AS row
MATCH (a:Document {
path: row.document_path
})
MATCH (b:Topic {
id: row.topic_id
})
MERGE (a)-[:DESCRIBES]->(b)
""",
"person_works": """
UNWIND $rows AS row
MATCH (a:Person {
id: row.person_id
})
MATCH (b:Work {
id: row.work_id
})
MERGE (a)-[:HAS_WORK]->(b)
""",
"work_documents": """
UNWIND $rows AS row
MATCH (a:Work {
id: row.work_id
})
MATCH (b:Document {
path: row.document_path
})
MERGE (a)-[:EVIDENCED_BY]->(b)
""",
}
def load_json(
path: Path,
) -> list[dict[str, Any]]:
if not path.exists():
raise FileNotFoundError(
f"Missing file: {path}"
)
data = json.loads(
path.read_text(
encoding="utf-8"
)
)
if not isinstance(
data,
list,
):
raise ValueError(
f"Expected list in {path}"
)
return data
def create_constraints(
session,
) -> None:
for (
name,
label,
property_name,
) in CONSTRAINTS:
session.run(
(
f"CREATE CONSTRAINT "
f"{name} IF NOT EXISTS "
f"FOR (n:{label}) "
f"REQUIRE n.{property_name} "
f"IS UNIQUE"
)
).consume()
def clear_managed_graph(
session,
) -> None:
session.run(
"""
MATCH (n)
WHERE n.graph_scope = $scope
DETACH DELETE n
""",
scope=GRAPH_SCOPE,
).consume()
def write_rows(
session,
query: str,
rows: list[dict[str, Any]],
) -> None:
if not rows:
return
session.run(
query,
rows=rows,
).consume()
def graph_counts(
session,
) -> tuple[
list[tuple[str, int]],
list[tuple[str, int]],
]:
node_rows = session.run(
"""
MATCH (n)
WHERE n.graph_scope = $scope
UNWIND labels(n) AS label
RETURN label, count(*) AS count
ORDER BY label
""",
scope=GRAPH_SCOPE,
)
nodes = [
(
str(record["label"]),
int(record["count"]),
)
for record in node_rows
]
relationship_rows = session.run(
"""
MATCH (a)-[r]->(b)
WHERE
a.graph_scope = $scope
AND b.graph_scope = $scope
RETURN
type(r) AS relationship,
count(*) AS count
ORDER BY relationship
""",
scope=GRAPH_SCOPE,
)
relationships = [
(
str(
record[
"relationship"
]
),
int(record["count"]),
)
for record in relationship_rows
]
return nodes, relationships
def work_quality_counts(
session,
) -> tuple[int, int]:
record = session.run(
"""
MATCH (w:Work)
WHERE w.graph_scope = $scope
RETURN
count(w) AS total,
count(w.title) AS with_title
""",
scope=GRAPH_SCOPE,
).single()
if record is None:
return 0, 0
return (
int(record["total"]),
int(record["with_title"]),
)
def print_payload_summary(
payload: dict[
str,
list[dict[str, Any]],
],
) -> None:
print()
print("Prepared GraphRAG payload")
print("=" * 60)
for key in (
"people",
"authors",
"documents",
"categories",
"topics",
"works",
):
print(
f"{key:<24}"
f"{len(payload[key]):>6}"
)
print("=" * 60)
def print_database_summary(
session,
) -> None:
nodes, relationships = (
graph_counts(
session
)
)
print()
print("Neo4j GraphRAG knowledge graph")
print("=" * 60)
print("Nodes")
for label, count in nodes:
print(
f" {label:<22}"
f"{count:>6}"
)
print()
print("Relationships")
for relationship, count in relationships:
print(
f" {relationship:<22}"
f"{count:>6}"
)
total, with_title = (
work_quality_counts(
session
)
)
print()
print(
"Works with extracted title: "
f"{with_title}/{total}"
)
print("=" * 60)
def main() -> None:
parser = argparse.ArgumentParser()
parser.add_argument(
"--replace",
action=argparse.BooleanOptionalAction,
default=True,
)
args = parser.parse_args()
documents = load_json(
DOCUMENTS_FILE
)
chunks = load_json(
CHUNKS_FILE
)
payload = build_graph_payload(
documents,
chunks,
)
print_payload_summary(
payload
)
settings = get_neo4j_settings()
with create_neo4j_driver(
settings
) as driver:
driver.verify_connectivity()
with driver.session(
database=settings.database
) as session:
create_constraints(
session
)
if args.replace:
clear_managed_graph(
session
)
for key, query in (
NODE_QUERIES.items()
):
write_rows(
session,
query,
payload[key],
)
for key, query in (
RELATIONSHIP_QUERIES.items()
):
write_rows(
session,
query,
payload[key],
)
print_database_summary(
session
)
if __name__ == "__main__":
main()