759 lines
23 KiB
Python
759 lines
23 KiB
Python
# 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,
|
|
)
|