Files
AI/test_phase5_subgraph_retriever.py

225 lines
7.0 KiB
Python
Raw Normal View History

#!/usr/bin/env python3
"""Phase 5.1 Subgraph Retriever tests."""
import asyncio
import sys
from pathlib import Path
from unittest.mock import AsyncMock, MagicMock
sys.path.insert(0, str(Path(__file__).parent / "ontology_platform"))
from ont_platform.core.graph.subgraph_retriever import SubgraphRetriever
class MockAdapter:
"""Mock Neo4j adapter for testing."""
async def execute_cypher(self, cypher: str, params=None):
"""Mock Cypher execution."""
params = params or {}
# Simulate test data based on query
if "neighbor_count" in cypher:
# retrieve_neighborhood query
return [
{
"result": {
"center": {
"id": params.get("entity_id"),
"label": "Apple Inc.",
"type": "Company",
"confidence": 0.95,
},
"neighbor_ids": [2, 3, 4],
"neighbor_count": 3,
}
}
]
if "SELECT node" in cypher or "WHERE n.id IN" in cypher:
# Node fetching query
node_ids = params.get("ids", [])
return [
{"node": {"id": nid, "label": f"Entity_{nid}", "type": "Concept", "confidence": 0.9}}
for nid in node_ids
]
if "source:Entity" in cypher and "RELATES" in cypher:
# Edge fetching query
return [
{
"edge": {
"source_id": 1,
"target_id": 2,
"predicate": "related_to",
"confidence": 0.85,
}
},
{
"edge": {
"source_id": 1,
"target_id": 3,
"predicate": "mentions",
"confidence": 0.88,
}
},
]
return []
async def test_retrieve_neighborhood():
"""Test N-hop neighborhood extraction."""
print("\n[TEST 1] Retrieve Neighborhood")
adapter = MockAdapter()
retriever = SubgraphRetriever(adapter)
result = await retriever.retrieve_neighborhood(entity_id=1, hops=2)
assert result["center_entity"] is not None, "Center entity should be found"
assert result["center_entity"]["id"] == 1, "Center entity ID should match"
assert result["node_count"] > 0, "Should have nodes in neighborhood"
assert result["hop_count"] == 2, "Hop count should be preserved"
print(f" [OK] Center entity: {result['center_entity']['label']}")
print(f" [OK] Neighbor count: {result['node_count']}")
print(f" [OK] Edge count: {result['edge_count']}")
print(" [PASS]")
async def test_retrieve_context():
"""Test context retrieval for multiple entities."""
print("\n[TEST 2] Retrieve Context (Multi-Entity)")
adapter = MockAdapter()
retriever = SubgraphRetriever(adapter)
result = await retriever.retrieve_context(entity_ids=[1, 2, 3])
assert "seed_entities" in result, "Should have seed entities"
assert "common_neighbors" in result, "Should have common neighbors"
assert result["total_nodes"] >= 0, "Should have node count"
print(f" [OK] Seed entities: {len(result['seed_entities'])}")
print(f" [OK] Common neighbors: {len(result['common_neighbors'])}")
print(f" [OK] Total nodes: {result['total_nodes']}")
print(" [PASS]")
async def test_retrieve_induced_subgraph():
"""Test induced subgraph extraction."""
print("\n[TEST 3] Retrieve Induced Subgraph")
adapter = MockAdapter()
retriever = SubgraphRetriever(adapter)
result = await retriever.retrieve_induced_subgraph(entity_ids=[1, 2, 3, 4])
assert "nodes" in result, "Should have nodes"
assert "edges" in result, "Should have edges"
assert "node_count" in result, "Should have node count"
print(f" [OK] Induced nodes: {result['node_count']}")
print(f" [OK] Induced edges: {result['edge_count']}")
print(" [PASS]")
async def test_validate_hop_limit():
"""Test hop limit validation."""
print("\n[TEST 4] Hop Limit Validation")
adapter = MockAdapter()
retriever = SubgraphRetriever(adapter)
try:
await retriever.retrieve_neighborhood(entity_id=1, hops=5)
assert False, "Should raise ValueError for hops > 3"
except ValueError as e:
assert "hops must be between 1 and 3" in str(e)
print(" [OK] Correctly rejects hops > 3")
try:
await retriever.retrieve_neighborhood(entity_id=1, hops=0)
assert False, "Should raise ValueError for hops < 1"
except ValueError as e:
assert "hops must be between 1 and 3" in str(e)
print(" [OK] Correctly rejects hops < 1")
print(" [PASS]")
async def test_retrieve_context_validation():
"""Test context retrieval validation."""
print("\n[TEST 5] Context Retrieval Validation")
adapter = MockAdapter()
retriever = SubgraphRetriever(adapter)
result = await retriever.retrieve_context(entity_ids=[])
assert "error" in result, "Should error with empty entity_ids"
print(" [OK] Rejects empty entity_ids")
result = await retriever.retrieve_context(entity_ids=[1])
assert "error" in result, "Should error with single entity"
print(" [OK] Rejects single entity")
print(" [PASS]")
async def test_induced_subgraph_validation():
"""Test induced subgraph validation."""
print("\n[TEST 6] Induced Subgraph Validation")
adapter = MockAdapter()
retriever = SubgraphRetriever(adapter)
result = await retriever.retrieve_induced_subgraph(entity_ids=[])
assert "error" in result, "Should error with empty entity_ids"
print(" [OK] Rejects empty entity_ids")
print(" [PASS]")
async def main():
"""Run all tests."""
print("=" * 70)
print("Phase 5.1 Subgraph Retriever Tests")
print("=" * 70)
try:
await test_retrieve_neighborhood()
await test_retrieve_context()
await test_retrieve_induced_subgraph()
await test_validate_hop_limit()
await test_retrieve_context_validation()
await test_induced_subgraph_validation()
print("\n" + "=" * 70)
print("All tests passed!")
print("=" * 70)
print("\nPhase 5.1 Subgraph Retriever capabilities:")
print(" [OK] N-hop neighborhood extraction")
print(" [OK] Multi-entity context retrieval")
print(" [OK] Induced subgraph extraction")
print(" [OK] Input validation")
return True
except AssertionError as e:
print(f"\nTest failed: {e}")
import traceback
traceback.print_exc()
return False
except Exception as e:
print(f"\nUnexpected error: {e}")
import traceback
traceback.print_exc()
return False
if __name__ == "__main__":
success = asyncio.run(main())
sys.exit(0 if success else 1)