참고소스 수정본
This commit is contained in:
@@ -0,0 +1,758 @@
|
||||
# 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,
|
||||
)
|
||||
Reference in New Issue
Block a user