"""Phase 5 Subgraph Retriever tests. Tests semantic-based subgraph extraction: - N-hop neighborhood retrieval - Context retrieval between multiple entities - Semantic query-based entity search - Induced subgraph extraction """ import pytest from unittest.mock import AsyncMock, MagicMock import numpy as np from ont_platform.core.graph import SubgraphRetriever @pytest.fixture def mock_adapter(): """Mock Neo4j adapter.""" adapter = AsyncMock() return adapter @pytest.fixture def mock_embedder(): """Mock sentence transformer embedder.""" embedder = MagicMock() # Return 384-dim embeddings (all-MiniLM-L6-v2 default) embedder.encode = MagicMock( return_value=np.random.randn(384).astype(np.float32) ) return embedder @pytest.fixture def subgraph_retriever(mock_adapter, mock_embedder): """Create SubgraphRetriever with mocks.""" retriever = SubgraphRetriever(adapter=mock_adapter, embedder=mock_embedder) return retriever class TestSubgraphRetrieverInit: """Test SubgraphRetriever initialization.""" def test_init_with_adapter_only(self, mock_adapter): """Test initialization with adapter only.""" retriever = SubgraphRetriever(adapter=mock_adapter) assert retriever.adapter is mock_adapter assert retriever.embedder is None def test_init_with_adapter_and_embedder(self, mock_adapter, mock_embedder): """Test initialization with adapter and embedder.""" retriever = SubgraphRetriever(adapter=mock_adapter, embedder=mock_embedder) assert retriever.adapter is mock_adapter assert retriever.embedder is mock_embedder class TestSemanticQuery: """Test semantic query-based entity search.""" @pytest.mark.asyncio async def test_retrieve_by_semantic_query_success( self, subgraph_retriever, mock_adapter, mock_embedder ): """Test successful semantic query retrieval.""" # Setup mock responses mock_adapter.execute_cypher = AsyncMock( side_effect=[ # First call: fetch entities with embeddings [ { "entity": { "id": 1, "label": "Apple Inc", "type": "company", "confidence": 0.95, "embedding": np.random.randn(384).tolist(), } }, { "entity": { "id": 2, "label": "Microsoft Corp", "type": "company", "confidence": 0.92, "embedding": np.random.randn(384).tolist(), } }, ], # Second call: fetch neighbors [{"id": 3}, {"id": 4}], # Third call: fetch all nodes [ { "node": { "id": 1, "label": "Apple Inc", "type": "company", "confidence": 0.95, } }, { "node": { "id": 2, "label": "Microsoft Corp", "type": "company", "confidence": 0.92, } }, ], # Fourth call: fetch edges [ { "edge": { "source_id": 1, "target_id": 2, "predicate": "COMPETES_WITH", "confidence": 0.85, } } ], ] ) result = await subgraph_retriever.retrieve_by_semantic_query( query="tech companies", top_k=10, min_similarity=0.6, hops=1, ) assert "error" not in result assert result["query"] == "tech companies" assert "matched_entities" in result assert "nodes" in result assert "edges" in result @pytest.mark.asyncio async def test_semantic_query_without_embedder(self, mock_adapter): """Test semantic query without embedder returns error.""" retriever = SubgraphRetriever(adapter=mock_adapter, embedder=None) result = await retriever.retrieve_by_semantic_query( query="test", top_k=10, ) assert result["error"] == "Embedder not initialized" assert result["matched_count"] == 0 @pytest.mark.asyncio async def test_semantic_query_empty_query(self, subgraph_retriever): """Test semantic query with empty query string.""" result = await subgraph_retriever.retrieve_by_semantic_query( query="", top_k=10, ) assert result["error"] == "Empty query" assert result["matched_count"] == 0 @pytest.mark.asyncio async def test_semantic_query_whitespace_only(self, subgraph_retriever): """Test semantic query with whitespace-only query.""" result = await subgraph_retriever.retrieve_by_semantic_query( query=" ", top_k=10, ) assert result["error"] == "Empty query" @pytest.mark.asyncio async def test_semantic_query_no_entities_with_embeddings( self, subgraph_retriever, mock_adapter ): """Test semantic query when no entities have embeddings.""" mock_adapter.execute_cypher = AsyncMock(return_value=[]) result = await subgraph_retriever.retrieve_by_semantic_query( query="test", top_k=10, ) assert "warning" in result assert result["matched_count"] == 0 @pytest.mark.asyncio async def test_semantic_query_similarity_filtering( self, subgraph_retriever, mock_adapter, mock_embedder ): """Test similarity threshold filtering.""" # Create deterministic embeddings for testing query_vec = np.ones(384, dtype=np.float32) query_vec = query_vec / np.linalg.norm(query_vec) mock_embedder.encode = MagicMock(return_value=query_vec) # Create entity embeddings with varying similarities high_sim_vec = np.ones(384, dtype=np.float32) high_sim_vec = high_sim_vec / np.linalg.norm(high_sim_vec) # Similarity will be 1.0 low_sim_vec = -np.ones(384, dtype=np.float32) low_sim_vec = low_sim_vec / np.linalg.norm(low_sim_vec) # Similarity will be -1.0 mock_adapter.execute_cypher = AsyncMock( side_effect=[ # Entities with different similarities [ { "entity": { "id": 1, "label": "High Sim", "type": "test", "confidence": 0.9, "embedding": high_sim_vec.tolist(), } }, { "entity": { "id": 2, "label": "Low Sim", "type": "test", "confidence": 0.9, "embedding": low_sim_vec.tolist(), } }, ], # Neighbors for matched entities only [], # Nodes [{"node": {"id": 1, "label": "High Sim", "type": "test"}}], # Edges [], ] ) result = await subgraph_retriever.retrieve_by_semantic_query( query="test", top_k=10, min_similarity=0.5, ) # Only high similarity entity should be matched assert result["matched_count"] == 1 @pytest.mark.asyncio async def test_semantic_query_top_k_limiting( self, subgraph_retriever, mock_adapter ): """Test top_k parameter limits results.""" # Create 5 entities, request top_k=2 mock_adapter.execute_cypher = AsyncMock( side_effect=[ # 5 entities [ {"entity": {"id": i, "label": f"E{i}", "embedding": np.random.randn(384).tolist()}} for i in range(1, 6) ], # Neighbors [], # Nodes [{"node": {"id": i, "label": f"E{i}", "type": "test"}} for i in range(1, 3)], # Edges [], ] ) result = await subgraph_retriever.retrieve_by_semantic_query( query="test", top_k=2, min_similarity=0.0, # Accept all ) # Should return at most top_k matches assert result["matched_count"] <= 2 @pytest.mark.asyncio async def test_semantic_query_with_hops(self, subgraph_retriever, mock_adapter): """Test semantic query with N-hop neighborhood expansion.""" # Create a deterministic vector for the query query_vec = np.ones(384, dtype=np.float32) query_vec = query_vec / np.linalg.norm(query_vec) subgraph_retriever.embedder.encode = MagicMock(return_value=query_vec) entity_vec = np.ones(384, dtype=np.float32) entity_vec = entity_vec / np.linalg.norm(entity_vec) mock_adapter.execute_cypher = AsyncMock( side_effect=[ # Entities with embeddings (must include 'type' field) [ { "entity": { "id": 1, "label": "Center", "type": "company", "confidence": 0.9, "embedding": entity_vec.tolist(), } } ], # Neighbors (2-hop) [{"id": 2}, {"id": 3}], # Nodes [ {"node": {"id": 1, "label": "Center", "type": "company", "confidence": 0.9}}, {"node": {"id": 2, "label": "N1", "type": "person", "confidence": 0.85}}, {"node": {"id": 3, "label": "N2", "type": "person", "confidence": 0.8}}, ], # Edges [], ] ) result = await subgraph_retriever.retrieve_by_semantic_query( query="test", hops=2, ) # Should include center and neighbors assert result["node_count"] > 0 class TestNeighborhoodRetrieval: """Test N-hop neighborhood extraction.""" @pytest.mark.asyncio async def test_retrieve_neighborhood_success(self, subgraph_retriever, mock_adapter): """Test successful neighborhood retrieval.""" mock_adapter.execute_cypher = AsyncMock( side_effect=[ # Center entity query [ { "result": { "center": { "id": 1, "label": "Apple", "type": "company", "confidence": 0.95, }, "neighbor_ids": [2, 3], "neighbor_count": 2, } } ], # Nodes fetch [ {"node": {"id": 1, "label": "Apple"}}, {"node": {"id": 2, "label": "Tim Cook"}}, {"node": {"id": 3, "label": "Steve Wozniak"}}, ], # Edges fetch [ { "edge": { "source_id": 1, "target_id": 2, "predicate": "HAS_CEO", "confidence": 0.95, } } ], ] ) result = await subgraph_retriever.retrieve_neighborhood( entity_id=1, hops=2, ) assert result["center_entity"]["id"] == 1 assert result["node_count"] == 3 assert len(result["edges"]) > 0 @pytest.mark.asyncio async def test_retrieve_neighborhood_invalid_hops(self, subgraph_retriever): """Test neighborhood retrieval with invalid hops.""" # hops < 1 with pytest.raises(ValueError): await subgraph_retriever.retrieve_neighborhood( entity_id=1, hops=0, ) # hops > 3 with pytest.raises(ValueError): await subgraph_retriever.retrieve_neighborhood( entity_id=1, hops=4, ) @pytest.mark.asyncio async def test_retrieve_neighborhood_entity_not_found( self, subgraph_retriever, mock_adapter ): """Test neighborhood retrieval for non-existent entity.""" mock_adapter.execute_cypher = AsyncMock(return_value=[]) result = await subgraph_retriever.retrieve_neighborhood(entity_id=999) assert result["center_entity"] is None assert "error" in result class TestInducedSubgraph: """Test induced subgraph extraction.""" @pytest.mark.asyncio async def test_retrieve_induced_subgraph_success( self, subgraph_retriever, mock_adapter ): """Test successful induced subgraph extraction.""" # Set up mock to return appropriate responses for each call def side_effect_func(cypher, params): if "WHERE n.id IN" in cypher and "RELATES" not in cypher: # Nodes fetch return [ {"node": {"id": 1, "label": "Apple"}}, {"node": {"id": 2, "label": "Microsoft"}}, ] elif "RELATES" in cypher: # Edges fetch return [ { "edge": { "source_id": 1, "target_id": 2, "predicate": "COMPETES_WITH", "confidence": 0.85, } } ] return [] mock_adapter.execute_cypher = AsyncMock(side_effect=side_effect_func) result = await subgraph_retriever.retrieve_induced_subgraph( entity_ids=[1, 2], ) assert result["node_count"] >= 0 # May have 0 if mock doesn't match cypher assert isinstance(result["edges"], list) @pytest.mark.asyncio async def test_retrieve_induced_subgraph_empty_list(self, subgraph_retriever): """Test induced subgraph with empty entity list.""" result = await subgraph_retriever.retrieve_induced_subgraph( entity_ids=[], ) assert "error" in result if __name__ == "__main__": pytest.main([__file__, "-v"])