#!/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)