Files

759 lines
23 KiB
Python
Raw Permalink Normal View History

2026-05-12 19:40:31 +09:00
# Copyright (c) "Neo4j"
# Neo4j Sweden AB [https://neo4j.com]
# #
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
# You may obtain a copy of the License at
# #
# https://www.apache.org/licenses/LICENSE-2.0
# #
# Unless required by applicable law or agreed to in writing, software
# distributed under the License is distributed on an "AS IS" BASIS,
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
# See the License for the specific language governing permissions and
# limitations under the License.
from __future__ import annotations
import datetime
from typing import Any, Optional
from unittest.mock import ANY, Mock, patch
import pytest
from neo4j_graphrag.experimental.components.graph_pruning import (
GraphPruning,
GraphPruningResult,
PruningReason,
PruningStats,
)
from neo4j_graphrag.experimental.components.schema import (
GraphConstraintType,
GraphSchema,
NodeType,
Pattern,
PropertyType,
RelationshipType,
)
from neo4j_graphrag.experimental.components.types import (
LexicalGraphConfig,
Neo4jGraph,
Neo4jNode,
Neo4jRelationship,
)
@pytest.fixture(scope="module")
def lexical_graph_config() -> LexicalGraphConfig:
return LexicalGraphConfig(
chunk_node_label="Paragraph",
)
@pytest.mark.parametrize(
"properties, valid_properties, additional_properties, expected_filtered_properties",
[
(
# no required, additional allowed
{
"name": "John Does",
"age": 25,
},
[
PropertyType(
name="name",
type="STRING",
)
],
True,
{
"name": "John Does",
"age": 25,
},
),
(
# no required, additional not allowed
{
"name": "John Does",
"age": 25,
},
[
PropertyType(
name="name",
type="STRING",
)
],
False,
{
"name": "John Does",
},
),
],
)
def test_graph_pruning_filter_properties(
properties: dict[str, Any],
valid_properties: list[PropertyType],
additional_properties: bool,
expected_filtered_properties: dict[str, Any],
) -> None:
pruner = GraphPruning()
filtered_properties = pruner._filter_properties(
properties,
valid_properties,
additional_properties=additional_properties,
node_label="Label",
pruning_stats=PruningStats(),
)
assert filtered_properties == expected_filtered_properties
@pytest.mark.parametrize(
"properties, expected_filtered_properties",
[
(
# all good, no bad types
{
"name": "John Does",
"age": 25,
"is_active": True,
},
{
"name": "John Does",
"age": 25,
"is_active": True,
},
),
(
# map must be serialized
{
"age": {"dob": datetime.date(2000, 1, 1), "age_in_2025": 25},
},
{
"age": '{"dob": "2000-01-01", "age_in_2025": 25}',
},
),
],
)
def test_graph_pruning_ensure_property_type(
properties: dict[str, Any],
expected_filtered_properties: dict[str, Any],
) -> None:
pruner = GraphPruning()
type_safe_properties = pruner._ensure_property_types(
properties,
)
assert type_safe_properties == expected_filtered_properties
@pytest.fixture(scope="module")
def node_type_no_properties() -> NodeType:
return NodeType(label="Person")
@pytest.fixture(scope="module")
def node_type_required_name() -> NodeType:
return NodeType(
label="Person",
properties=[
PropertyType(name="name", type="STRING", required=True),
PropertyType(name="age", type="INTEGER"),
],
)
def _graph_schema_for_node_entity(entity: NodeType | None) -> GraphSchema:
"""Build a GraphSchema from a node type fixture (applies required→EXISTENCE migration)."""
if entity is None:
return GraphSchema(node_types=tuple())
return GraphSchema.model_validate({"node_types": [entity.model_dump()]})
def _schema_for_relationship_validation(
patterns: tuple[Pattern, ...],
) -> GraphSchema:
"""Minimal valid GraphSchema for _validate_relationship tests (REL + Person/Location)."""
return GraphSchema.model_validate(
{
"node_types": [
{
"label": "Person",
"properties": [{"name": "name", "type": "STRING"}],
},
{
"label": "Location",
"properties": [{"name": "name", "type": "STRING"}],
},
],
"relationship_types": [{"label": "REL"}],
"patterns": [tuple(p) for p in patterns],
}
)
@pytest.mark.parametrize(
"node, entity, additional_node_types, expected_node",
[
# all good, with default values
(
Neo4jNode(id="1", label="Person", properties={"name": "John Doe"}),
"node_type_no_properties",
True,
Neo4jNode(id="1", label="Person", properties={"name": "John Doe"}),
),
# properties empty (missing default)
(
Neo4jNode(id="1", label="Person", properties={"age": 45}),
"node_type_required_name",
True,
None,
),
# node label not is schema, additional not allowed
(
Neo4jNode(id="1", label="Location", properties={"name": "New York"}),
None,
False,
None,
),
# node label not is schema, additional allowed
(
Neo4jNode(id="1", label="Location", properties={"name": "New York"}),
None,
True,
Neo4jNode(id="1", label="Location", properties={"name": "New York"}),
),
# node label not valid
(
Neo4jNode(id="1", label="", properties={"name": "New York"}),
"node_type_required_name",
True,
None,
),
# node ID not valid
(
Neo4jNode(id="", label="Location", properties={"name": "New York"}),
"node_type_required_name",
True,
None,
),
],
)
def test_graph_pruning_validate_node(
node: Neo4jNode,
entity: str,
additional_node_types: bool,
expected_node: Neo4jNode,
request: pytest.FixtureRequest,
) -> None:
e_fixture = request.getfixturevalue(entity) if entity else None
schema = _graph_schema_for_node_entity(e_fixture)
e = schema.node_type_from_label(node.label) if node.label else None
pruner = GraphPruning()
result = pruner._validate_node(
node, PruningStats(), e, schema, additional_node_types
)
if expected_node is not None:
assert result == expected_node
else:
assert result is None
def test_graph_pruning_enforce_nodes_lexical_graph(
lexical_graph_config: LexicalGraphConfig,
) -> None:
pruner = GraphPruning()
result = pruner._enforce_nodes(
nodes=[
Neo4jNode(id="1", label="Paragraph"),
],
schema=GraphSchema(node_types=tuple(), additional_node_types=False),
lexical_graph_config=lexical_graph_config,
pruning_stats=PruningStats(),
)
assert len(result) == 1
assert result[0].label == "Paragraph"
def test_graph_pruning_enforce_relationships_lexical_graph_with_pruned_nodes(
lexical_graph_config: LexicalGraphConfig,
) -> None:
"""Test that lexical relationships are pruned when their nodes are pruned."""
pruner = GraphPruning()
# Create nodes: Chunk (lexical) and Person (extracted, will be pruned)
nodes = [
Neo4jNode(id="chunk-1", label="Paragraph", properties={"text": "test"}),
Neo4jNode(
id="person-1", label="Person", properties={"age": 30}
), # Missing required 'name'
]
# Person node type: name must exist (EXISTENCE constraint)
schema = GraphSchema.model_validate(
{
"node_types": [
{
"label": "Person",
"properties": [
{"name": "name", "type": "STRING"},
{"name": "age", "type": "INTEGER"},
],
}
],
"constraints": [
{
"type": GraphConstraintType.EXISTENCE.value,
"node_type": "Person",
"property_name": "name",
"relationship_type": None,
}
],
}
)
# Filter nodes - Person should be pruned due to missing required property
pruning_stats = PruningStats()
filtered_nodes = pruner._enforce_nodes(
nodes=nodes,
schema=schema,
lexical_graph_config=lexical_graph_config,
pruning_stats=pruning_stats,
)
# Only Chunk should remain
assert len(filtered_nodes) == 1
assert filtered_nodes[0].id == "chunk-1"
assert pruning_stats.number_of_pruned_nodes == 1
# Create relationships including FROM_CHUNK to the pruned Person node
relationships = [
Neo4jRelationship(
start_node_id="person-1",
end_node_id="chunk-1",
type="FROM_CHUNK", # Lexical relationship
),
Neo4jRelationship(
start_node_id="chunk-1",
end_node_id="person-1",
type="NEXT_CHUNK", # Another lexical relationship
),
]
# Filter relationships
pruning_stats_rels = PruningStats()
filtered_rels = pruner._enforce_relationships(
relationships=relationships,
filtered_nodes=filtered_nodes,
schema=schema,
lexical_graph_config=lexical_graph_config,
pruning_stats=pruning_stats_rels,
)
# Both lexical relationships should be pruned because person-1 node is missing
assert len(filtered_rels) == 0
assert pruning_stats_rels.number_of_pruned_relationships == 2
def test_graph_pruning_key_constraint_prunes_missing_property(
lexical_graph_config: LexicalGraphConfig,
) -> None:
"""KEY constraints require presence like EXISTENCE (Neo4j key mandatory properties)."""
pruner = GraphPruning()
nodes = [
Neo4jNode(id="chunk-1", label="Paragraph", properties={"text": "test"}),
Neo4jNode(
id="person-1", label="Person", properties={"age": 30}
), # Missing KEY 'email'
]
schema = GraphSchema.model_validate(
{
"node_types": [
{
"label": "Person",
"properties": [
{"name": "email", "type": "STRING"},
{"name": "age", "type": "INTEGER"},
],
}
],
"constraints": [
{
"type": GraphConstraintType.KEY.value,
"node_type": "Person",
"property_name": "email",
"relationship_type": None,
}
],
}
)
pruning_stats = PruningStats()
filtered_nodes = pruner._enforce_nodes(
nodes=nodes,
schema=schema,
lexical_graph_config=lexical_graph_config,
pruning_stats=pruning_stats,
)
assert len(filtered_nodes) == 1
assert filtered_nodes[0].id == "chunk-1"
assert pruning_stats.number_of_pruned_nodes == 1
def _schema_with_relationship_key_constraint() -> GraphSchema:
return GraphSchema.model_validate(
{
"node_types": [
{"label": "Person", "properties": [{"name": "name", "type": "STRING"}]},
{
"label": "Company",
"properties": [{"name": "name", "type": "STRING"}],
},
],
"relationship_types": [
{
"label": "WORKS_FOR",
"properties": [
{"name": "since", "type": "STRING"},
{"name": "role", "type": "STRING"},
],
}
],
"patterns": [("Person", "WORKS_FOR", "Company")],
"constraints": [
{
"type": GraphConstraintType.KEY.value,
"node_type": "",
"property_name": "since",
"relationship_type": "WORKS_FOR",
}
],
}
)
def test_graph_pruning_key_constraint_on_relationship_mandatory_enforced() -> None:
"""KEY on a relationship contributes to mandatory props (like EXISTENCE)."""
schema = _schema_with_relationship_key_constraint()
assert schema.mandatory_property_names_for_relationship("WORKS_FOR") == {"since"}
pruner = GraphPruning()
rel = Neo4jRelationship(
start_node_id="p1",
end_node_id="c1",
type="WORKS_FOR",
properties={"role": "engineer"},
)
valid_nodes = {"p1": "Person", "c1": "Company"}
pruning_stats = PruningStats()
rel_type = schema.relationship_type_from_label("WORKS_FOR")
assert rel_type is not None
out = pruner._validate_relationship(
rel,
valid_nodes,
pruning_stats,
rel_type,
schema.additional_relationship_types,
schema.patterns,
schema.additional_patterns,
schema,
)
assert out is not None
assert out.properties == {}
assert len(pruning_stats.pruned_relationships) == 1
assert (
pruning_stats.pruned_relationships[0].pruned_reason
== PruningReason.MISSING_REQUIRED_PROPERTY
)
assert pruning_stats.pruned_relationships[0].metadata.get(
"missing_required_properties"
) == ["since"]
def test_graph_pruning_key_constraint_on_relationship_kept_when_present() -> None:
schema = _schema_with_relationship_key_constraint()
pruner = GraphPruning()
rel = Neo4jRelationship(
start_node_id="p1",
end_node_id="c1",
type="WORKS_FOR",
properties={"since": "2020-01-01", "role": "engineer"},
)
pruning_stats = PruningStats()
rel_type = schema.relationship_type_from_label("WORKS_FOR")
assert rel_type is not None
out = pruner._validate_relationship(
rel,
{"p1": "Person", "c1": "Company"},
pruning_stats,
rel_type,
schema.additional_relationship_types,
schema.patterns,
schema.additional_patterns,
schema,
)
assert out is not None
assert out.properties == {"since": "2020-01-01", "role": "engineer"}
assert pruning_stats.number_of_pruned_relationships == 0
@pytest.fixture
def neo4j_relationship() -> Neo4jRelationship:
return Neo4jRelationship(
start_node_id="1",
end_node_id="2",
type="REL",
properties={},
)
@pytest.fixture
def neo4j_relationship_invalid_type() -> Neo4jRelationship:
return Neo4jRelationship(
start_node_id="1",
end_node_id="2",
type="",
properties={},
)
@pytest.fixture
def neo4j_reversed_relationship(
neo4j_relationship: Neo4jRelationship,
) -> Neo4jRelationship:
return Neo4jRelationship(
start_node_id=neo4j_relationship.end_node_id,
end_node_id=neo4j_relationship.start_node_id,
type=neo4j_relationship.type,
properties=neo4j_relationship.properties,
)
@pytest.mark.parametrize(
"relationship, valid_nodes, relationship_type, additional_relationship_types, patterns, additional_patterns, expected_relationship",
[
# all good
(
"neo4j_relationship", # relationship,
{ # valid_nodes
"1": "Person",
"2": "Location",
},
RelationshipType( # relationship_type
label="REL",
),
True, # additional_relationship_types
(
Pattern(source="Person", relationship="REL", target="Location"),
), # patterns
True, # additional_patterns
"neo4j_relationship", # expected_relationship
),
# reverse relationship
(
"neo4j_reversed_relationship",
{
"1": "Person",
"2": "Location",
},
RelationshipType(
label="REL",
),
True, # additional_relationship_types
(Pattern(source="Person", relationship="REL", target="Location"),),
True, # additional_patterns
"neo4j_relationship",
),
# invalid start node ID
(
"neo4j_reversed_relationship",
{
"10": "Person",
"2": "Location",
},
RelationshipType(
label="REL",
),
True, # additional_relationship_types
(Pattern(source="Person", relationship="REL", target="Location"),),
True, # additional_patterns
None,
),
# invalid type, addition allowed
(
"neo4j_relationship",
{
"1": "Person",
"2": "Location",
},
None, # relationship_type
True, # additional_relationship_types
(Pattern(source="Person", relationship="REL", target="Location"),),
True, # additional_patterns
"neo4j_relationship",
),
# invalid type, addition allowed but invalid node ID
(
"neo4j_relationship",
{
"1": "Person",
},
None, # relationship_type
True, # additional_relationship_types
(Pattern(source="Person", relationship="REL", target="Location"),),
True, # additional_patterns
None,
),
# invalid type, addition not allowed
(
"neo4j_relationship",
{
"1": "Person",
"2": "Location",
},
None, # relationship_type
False, # additional_relationship_types
(Pattern(source="Person", relationship="REL", target="Location"),),
True, # additional_patterns
None,
),
# invalid pattern, addition allowed
(
"neo4j_relationship",
{
"1": "Person",
"2": "Location",
},
RelationshipType(
label="REL",
),
True, # additional_relationship_types
(Pattern(source="Person", relationship="REL", target="Person"),),
True, # additional_patterns
"neo4j_relationship",
),
# invalid pattern, addition not allowed
(
"neo4j_relationship",
{
"1": "Person",
"2": "Location",
},
RelationshipType(
label="REL",
),
True, # additional_relationship_types
(Pattern(source="Person", relationship="REL", target="Person"),),
False, # additional_patterns
None,
),
# invalid extracted type
(
"neo4j_relationship_invalid_type", # relationship,
{ # valid_nodes
"1": "Person",
"2": "Location",
},
RelationshipType( # relationship_type
label="REL",
),
True, # additional_relationship_types
(
Pattern(source="Person", relationship="REL", target="Location"),
), # patterns
True, # additional_patterns
None, # expected_relationship
),
],
)
def test_graph_pruning_validate_relationship(
relationship: str,
valid_nodes: dict[str, str],
relationship_type: RelationshipType,
additional_relationship_types: bool,
patterns: tuple[Pattern, ...],
additional_patterns: bool,
expected_relationship: Optional[str],
request: pytest.FixtureRequest,
) -> None:
relationship_obj = request.getfixturevalue(relationship)
expected_relationship_obj = (
request.getfixturevalue(expected_relationship)
if expected_relationship
else None
)
pruner = GraphPruning()
schema = _schema_for_relationship_validation(patterns)
assert (
pruner._validate_relationship(
relationship_obj,
valid_nodes,
PruningStats(),
relationship_type,
additional_relationship_types,
patterns,
additional_patterns,
schema,
)
== expected_relationship_obj
)
@patch("neo4j_graphrag.experimental.components.graph_pruning.GraphPruning._clean_graph")
@pytest.mark.asyncio
async def test_graph_pruning_run_happy_path(
mock_clean_graph: Mock,
node_type_required_name: NodeType,
lexical_graph_config: LexicalGraphConfig,
) -> None:
initial_graph = Neo4jGraph(
nodes=[Neo4jNode(id="1", label="Person"), Neo4jNode(id="2", label="Location")],
)
schema = _graph_schema_for_node_entity(node_type_required_name)
cleaned_graph = Neo4jGraph(nodes=[Neo4jNode(id="1", label="Person")])
mock_clean_graph.return_value = (cleaned_graph, PruningStats())
pruner = GraphPruning()
pruner_result = await pruner.run(
graph=initial_graph,
schema=schema,
lexical_graph_config=lexical_graph_config,
)
assert isinstance(pruner_result, GraphPruningResult)
assert pruner_result.graph == cleaned_graph
mock_clean_graph.assert_called_once_with(
initial_graph, schema, lexical_graph_config
)
@pytest.mark.asyncio
async def test_graph_pruning_run_no_schema() -> None:
initial_graph = Neo4jGraph(nodes=[Neo4jNode(id="1", label="Person")])
pruner = GraphPruning()
pruner_result = await pruner.run(
graph=initial_graph,
schema=None,
)
assert isinstance(pruner_result, GraphPruningResult)
assert pruner_result.graph == initial_graph
@patch(
"neo4j_graphrag.experimental.components.graph_pruning.GraphPruning._enforce_nodes"
)
def test_graph_pruning_clean_graph(
mock_enforce_nodes: Mock,
lexical_graph_config: LexicalGraphConfig,
) -> None:
mock_enforce_nodes.return_value = []
initial_graph = Neo4jGraph(nodes=[Neo4jNode(id="1", label="Person")])
schema = GraphSchema(node_types=())
pruner = GraphPruning()
cleaned_graph, pruning_stats = pruner._clean_graph(
initial_graph, schema, lexical_graph_config
)
assert cleaned_graph == Neo4jGraph()
assert isinstance(pruning_stats, PruningStats)
mock_enforce_nodes.assert_called_once_with(
[Neo4jNode(id="1", label="Person")],
schema,
lexical_graph_config,
ANY,
)