diff --git a/src/neo4j_graphrag/experimental/components/kg_writer.py b/src/neo4j_graphrag/experimental/components/kg_writer.py index c9c73c2f5..9df159ce5 100644 --- a/src/neo4j_graphrag/experimental/components/kg_writer.py +++ b/src/neo4j_graphrag/experimental/components/kg_writer.py @@ -215,9 +215,18 @@ def __init__( self.is_version_5_24_or_above = is_version_5_24_or_above(version_tuple) def _db_setup(self) -> None: - self.driver.execute_query(""" - CREATE INDEX __entity__tmp_internal_id IF NOT EXISTS FOR (n:__KGBuilder__) ON (n.__tmp_internal_id) - """) + self.driver.execute_query( + "DROP INDEX __entity__tmp_internal_id IF EXISTS", + database_=self.neo4j_database, + ) + self.driver.execute_query( + """ + CREATE CONSTRAINT __entity__tmp_internal_id_unique IF NOT EXISTS + FOR (n:__KGBuilder__) + REQUIRE n.__tmp_internal_id IS UNIQUE + """, + database_=self.neo4j_database, + ) @staticmethod def _nodes_to_rows( diff --git a/tests/unit/experimental/components/test_kg_writer.py b/tests/unit/experimental/components/test_kg_writer.py index 1f13bdf4e..75318b6c4 100644 --- a/tests/unit/experimental/components/test_kg_writer.py +++ b/tests/unit/experimental/components/test_kg_writer.py @@ -102,6 +102,34 @@ def test_get_unique_properties_for_node_type_deprecation_warning() -> None: ] +@mock.patch( + "neo4j_graphrag.experimental.components.kg_writer.get_version", + return_value=((5, 22, 0), False, False), +) +def test_neo4j_writer_db_setup_uses_unique_constraint( + _: Mock, driver: MagicMock +) -> None: + neo4j_writer = Neo4jWriter(driver=driver, neo4j_database="my_db") + driver.execute_query.reset_mock() + + neo4j_writer._db_setup() + + assert driver.execute_query.call_count == 2 + drop_call = driver.execute_query.call_args_list[0] + assert drop_call.args == ("DROP INDEX __entity__tmp_internal_id IF EXISTS",) + assert drop_call.kwargs == {"database_": "my_db"} + + constraint_call = driver.execute_query.call_args_list[1] + constraint_query = constraint_call.args[0] + assert ( + "CREATE CONSTRAINT __entity__tmp_internal_id_unique IF NOT EXISTS" + in constraint_query + ) + assert "FOR (n:__KGBuilder__)" in constraint_query + assert "REQUIRE n.__tmp_internal_id IS UNIQUE" in constraint_query + assert constraint_call.kwargs == {"database_": "my_db"} + + # --- FilenameCollisionHandler tests ---