459 lines
15 KiB
Python
459 lines
15 KiB
Python
|
|
"""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"])
|