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()