feat(library-desk): implement HybridRAG query system
HybridRAG Service: - Combine vector (Qdrant), graph (Neo4j), and web (SearXNG) search - Reciprocal Rank Fusion (RRF) for result merging - LLM re-ranking with mistral-nemo - Graph enrichment with related dossiers - Query enhancement with keyword/synonym extraction - Search result persistence for offline processing Router: - POST /query/hybrid endpoint - Configurable search limits per source - Enable/disable individual sources - Timing breakdown for performance monitoring Models: - HybridRAGRequest, HybridRAGResponse - HybridRAGResult with source tracking - KeywordExtraction for query analysis - TimingBreakdown for performance metrics Tests: - End-to-end HybridRAG query tests - RRF fusion algorithm validation - Multi-source result merging
This commit is contained in:
@@ -0,0 +1,87 @@
|
||||
"""
|
||||
HybridRAG models for multi-source search with RRF fusion.
|
||||
|
||||
Combines vector search (Qdrant), knowledge graph (Neo4j), and web search (SearXNG)
|
||||
with Reciprocal Rank Fusion and LLM re-ranking.
|
||||
"""
|
||||
|
||||
from pydantic import BaseModel, Field
|
||||
from typing import List, Optional, Dict, Any
|
||||
|
||||
|
||||
class HybridRAGConfig(BaseModel):
|
||||
"""Configuration for HybridRAG query."""
|
||||
vector_limit: int = Field(default=10, ge=1, le=50, description="Max vector results")
|
||||
graph_limit: int = Field(default=10, ge=1, le=50, description="Max graph results")
|
||||
web_limit: int = Field(default=5, ge=1, le=20, description="Max web results")
|
||||
enable_vector: bool = Field(default=True, description="Enable vector search")
|
||||
enable_graph: bool = Field(default=True, description="Enable graph search")
|
||||
enable_web: bool = Field(default=True, description="Enable web search")
|
||||
enable_reranking: bool = Field(default=True, description="Enable LLM re-ranking")
|
||||
enable_enrichment: bool = Field(default=True, description="Enable graph enrichment")
|
||||
final_result_count: int = Field(default=10, ge=1, le=50, description="Final results to return")
|
||||
rrf_k: int = Field(default=60, ge=1, le=100, description="RRF constant")
|
||||
|
||||
|
||||
class RelatedDossier(BaseModel):
|
||||
"""Related document metadata from graph enrichment."""
|
||||
page_id: int
|
||||
title: str
|
||||
path: str
|
||||
tag: str
|
||||
shared_entities: int
|
||||
|
||||
|
||||
class HybridRAGResult(BaseModel):
|
||||
"""Single result from HybridRAG query."""
|
||||
source_type: str = Field(..., description="Source: 'vector', 'graph', 'web'")
|
||||
title: str
|
||||
content: str
|
||||
url: Optional[str] = Field(None, description="URL for web results")
|
||||
page_id: Optional[int] = Field(None, description="Page ID for wiki results")
|
||||
page_path: Optional[str] = Field(None, description="Wiki page path")
|
||||
rrf_score: float = Field(..., description="Reciprocal Rank Fusion score")
|
||||
final_rank: int = Field(..., description="Final rank after re-ranking")
|
||||
sources: List[str] = Field(..., description="Which sources included this result")
|
||||
related_dossiers: List[RelatedDossier] = Field(default=[], description="Related documents via shared entities")
|
||||
metadata: Dict[str, Any] = Field(default={}, description="Additional metadata")
|
||||
|
||||
|
||||
class TimingBreakdown(BaseModel):
|
||||
"""Performance timing breakdown for each phase."""
|
||||
query_enhancement_ms: float = Field(..., description="Phase 0: Keyword/synonym extraction")
|
||||
vector_ms: float = Field(..., description="Phase 1: Vector search")
|
||||
graph_ms: float = Field(..., description="Phase 1: Graph search")
|
||||
web_ms: float = Field(..., description="Phase 1: Web search")
|
||||
fusion_ms: float = Field(..., description="Phase 2: RRF fusion")
|
||||
enrichment_ms: float = Field(..., description="Phase 3: Graph enrichment")
|
||||
reranking_ms: float = Field(..., description="Phase 4: LLM re-ranking")
|
||||
persistence_ms: float = Field(..., description="Phase 6: Search persistence")
|
||||
total_ms: float = Field(..., description="Total end-to-end time")
|
||||
|
||||
|
||||
class KeywordExtraction(BaseModel):
|
||||
"""Extracted keywords and synonyms from query enhancement."""
|
||||
core_keywords: List[str] = Field(default=[], description="Primary keywords")
|
||||
entities: List[str] = Field(default=[], description="Named entities")
|
||||
synonyms: Dict[str, List[str]] = Field(default={}, description="Synonyms map")
|
||||
expansions: Dict[str, List[str]] = Field(default={}, description="Abbreviation expansions")
|
||||
|
||||
|
||||
class HybridRAGResponse(BaseModel):
|
||||
"""Response from HybridRAG query."""
|
||||
query: str = Field(..., description="Original search query")
|
||||
keywords: KeywordExtraction = Field(..., description="Extracted keywords/synonyms")
|
||||
results: List[HybridRAGResult] = Field(..., description="Ranked search results")
|
||||
context: str = Field(..., description="Formatted context for LLM consumption")
|
||||
source_counts: Dict[str, int] = Field(..., description="Result counts by source")
|
||||
total_results: int = Field(..., description="Total number of results")
|
||||
timing: TimingBreakdown = Field(..., description="Performance breakdown")
|
||||
config_used: HybridRAGConfig = Field(..., description="Configuration used")
|
||||
search_id: Optional[str] = Field(None, description="Search ID for Librarian tracking")
|
||||
|
||||
|
||||
class HybridRAGRequest(BaseModel):
|
||||
"""Request for HybridRAG query."""
|
||||
query: str = Field(..., min_length=1, max_length=500, description="Search query")
|
||||
config: Optional[HybridRAGConfig] = Field(None, description="Custom configuration")
|
||||
@@ -0,0 +1,117 @@
|
||||
"""
|
||||
HybridRAG router for multi-source search API.
|
||||
|
||||
Provides endpoint for combining vector, graph, and web search
|
||||
with RRF fusion and LLM re-ranking.
|
||||
"""
|
||||
|
||||
from fastapi import APIRouter, HTTPException, Depends, Query
|
||||
import logging
|
||||
|
||||
from src.models.hybrid_rag import HybridRAGRequest, HybridRAGResponse
|
||||
from src.services.hybrid_rag_service import HybridRAGService
|
||||
from src.services.vector_service import VectorService
|
||||
from src.services.graph_service import GraphService
|
||||
from src.clients.searxng_client import SearXNGClient
|
||||
from src.clients.ollama_client import OllamaClient
|
||||
from src.core.dependencies import (
|
||||
Neo4jDep, WikiJSDep, QdrantDep, OllamaDep,
|
||||
SearXNGDep, verify_api_key, get_settings
|
||||
)
|
||||
from src.config import Settings
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
router = APIRouter(prefix="/query", tags=["HybridRAG"])
|
||||
|
||||
|
||||
# Dependency to get HybridRAG service
|
||||
def get_hybrid_rag_service(
|
||||
neo4j_client: Neo4jDep,
|
||||
wiki_client: WikiJSDep,
|
||||
qdrant_client: QdrantDep,
|
||||
ollama_client: OllamaDep,
|
||||
searxng_client: SearXNGDep,
|
||||
settings: Settings = Depends(get_settings)
|
||||
) -> HybridRAGService:
|
||||
"""Get HybridRAG service instance with all dependencies."""
|
||||
from src.services.vector_service import VectorService
|
||||
from src.services.graph_service import GraphService
|
||||
|
||||
# Create component services
|
||||
vector_service = VectorService(qdrant_client, wiki_client, ollama_client)
|
||||
graph_service = GraphService(neo4j_client, wiki_client)
|
||||
|
||||
# Create HybridRAG service
|
||||
return HybridRAGService(
|
||||
vector_service=vector_service,
|
||||
graph_service=graph_service,
|
||||
searxng_client=searxng_client,
|
||||
ollama_client=ollama_client,
|
||||
settings=settings
|
||||
)
|
||||
|
||||
|
||||
@router.post("/hybrid", response_model=HybridRAGResponse)
|
||||
async def hybrid_search(
|
||||
request: HybridRAGRequest,
|
||||
user: str = Query(default="jpmschweitzer", description="User identifier for multi-tenancy"),
|
||||
hybrid_rag_service: HybridRAGService = Depends(get_hybrid_rag_service),
|
||||
api_key: str = Depends(verify_api_key)
|
||||
):
|
||||
"""
|
||||
Execute HybridRAG query combining vector, graph, and web search.
|
||||
|
||||
**6-Phase Pipeline:**
|
||||
1. **Query Enhancement**: Extract keywords/synonyms with LLM
|
||||
2. **Parallel Retrieval**: Search vector (Qdrant), graph (Neo4j), web (SearXNG)
|
||||
3. **RRF Fusion**: Merge results with Reciprocal Rank Fusion
|
||||
4. **Enrichment**: Add related documents via shared entities
|
||||
5. **LLM Re-ranking**: Re-rank with mistral-nemo for relevance
|
||||
6. **Context Formatting**: Format for LLM consumption
|
||||
7. **Persistence**: Store for Librarian knowledge consolidation
|
||||
|
||||
**Example Request:**
|
||||
```json
|
||||
{
|
||||
"query": "How does Docker orchestration work with Kubernetes?",
|
||||
"user": "jpmschweitzer",
|
||||
"config": {
|
||||
"vector_limit": 10,
|
||||
"graph_limit": 10,
|
||||
"web_limit": 5,
|
||||
"enable_reranking": true,
|
||||
"final_result_count": 10
|
||||
}
|
||||
}
|
||||
```
|
||||
|
||||
**Returns:**
|
||||
- Ranked results from all sources
|
||||
- Extracted keywords/synonyms
|
||||
- Related dossiers (via graph)
|
||||
- Formatted context for LLM
|
||||
- Performance timing breakdown
|
||||
- Search ID for Librarian tracking
|
||||
"""
|
||||
try:
|
||||
logger.info(f"HybridRAG request: '{request.query}' for user '{user}'")
|
||||
|
||||
response = await hybrid_rag_service.search(
|
||||
query=request.query,
|
||||
user=user,
|
||||
config=request.config
|
||||
)
|
||||
|
||||
logger.info(
|
||||
f"HybridRAG completed: {response.total_results} results in {response.timing.total_ms:.0f}ms"
|
||||
)
|
||||
|
||||
return response
|
||||
|
||||
except ValueError as e:
|
||||
logger.error(f"Invalid request: {e}")
|
||||
raise HTTPException(status_code=400, detail=str(e))
|
||||
except Exception as e:
|
||||
logger.error(f"HybridRAG search failed: {e}", exc_info=True)
|
||||
raise HTTPException(status_code=500, detail="Search failed")
|
||||
@@ -0,0 +1,739 @@
|
||||
"""
|
||||
HybridRAG service combining vector, graph, and web search.
|
||||
|
||||
6-Phase Pipeline:
|
||||
0. Query Enhancement - Extract keywords/synonyms with LLM
|
||||
1. Parallel Retrieval - Vector + Graph + Web search
|
||||
2. RRF Fusion - Merge results with Reciprocal Rank Fusion
|
||||
3. Enrichment - Add related dossiers via graph
|
||||
4. LLM Re-ranking - Re-rank with mistral-nemo
|
||||
5. Context Formatting - Format for LLM consumption
|
||||
6. Persistence - Store for Librarian processing
|
||||
"""
|
||||
|
||||
import asyncio
|
||||
import time
|
||||
import json
|
||||
import uuid
|
||||
from typing import List, Dict, Any, Optional
|
||||
import logging
|
||||
|
||||
from src.services.vector_service import VectorService
|
||||
from src.services.graph_service import GraphService
|
||||
from src.clients.searxng_client import SearXNGClient
|
||||
from src.clients.ollama_client import OllamaClient
|
||||
from src.config import Settings
|
||||
from src.models.hybrid_rag import (
|
||||
HybridRAGConfig, HybridRAGRequest, HybridRAGResponse,
|
||||
HybridRAGResult, TimingBreakdown, KeywordExtraction,
|
||||
RelatedDossier
|
||||
)
|
||||
from src.core.multi_tenancy import get_neo4j_user_base_label, get_neo4j_user_label
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
class HybridRAGService:
|
||||
"""
|
||||
Service for HybridRAG multi-source search with fusion and re-ranking.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
vector_service: VectorService,
|
||||
graph_service: GraphService,
|
||||
searxng_client: SearXNGClient,
|
||||
ollama_client: OllamaClient,
|
||||
settings: Settings
|
||||
):
|
||||
"""
|
||||
Initialize HybridRAG service.
|
||||
|
||||
Args:
|
||||
vector_service: Service for Qdrant vector search
|
||||
graph_service: Service for Neo4j graph search
|
||||
searxng_client: Client for web search
|
||||
ollama_client: Client for LLM (keyword extraction, re-ranking)
|
||||
settings: Application settings
|
||||
"""
|
||||
self.vector = vector_service
|
||||
self.graph = graph_service
|
||||
self.searxng = searxng_client
|
||||
self.ollama = ollama_client
|
||||
self.settings = settings
|
||||
self.reranker_model = settings.reranker_model
|
||||
|
||||
async def search(
|
||||
self,
|
||||
query: str,
|
||||
user: str,
|
||||
config: Optional[HybridRAGConfig] = None
|
||||
) -> HybridRAGResponse:
|
||||
"""
|
||||
Execute HybridRAG search across all sources.
|
||||
|
||||
Args:
|
||||
query: Search query
|
||||
user: User identifier
|
||||
config: Optional configuration override
|
||||
|
||||
Returns:
|
||||
Complete search response with ranked results and timing
|
||||
"""
|
||||
start_time = time.time()
|
||||
timing = {}
|
||||
|
||||
# Use default config if not provided
|
||||
if not config:
|
||||
config = HybridRAGConfig()
|
||||
|
||||
logger.info(f"HybridRAG search: '{query}' for user '{user}'")
|
||||
|
||||
# Phase 0: Query Enhancement
|
||||
phase0_start = time.time()
|
||||
keywords_data = await self._extract_keywords_and_synonyms(query)
|
||||
timing["query_enhancement_ms"] = (time.time() - phase0_start) * 1000
|
||||
|
||||
# Phase 1: Parallel Retrieval
|
||||
phase1_start = time.time()
|
||||
raw_results = await self._retrieve_parallel(query, user, config, keywords_data)
|
||||
timing["vector_ms"] = raw_results.get("timing", {}).get("vector_ms", 0)
|
||||
timing["graph_ms"] = raw_results.get("timing", {}).get("graph_ms", 0)
|
||||
timing["web_ms"] = raw_results.get("timing", {}).get("web_ms", 0)
|
||||
|
||||
# Phase 2: RRF Fusion
|
||||
phase2_start = time.time()
|
||||
fused_results = self._reciprocal_rank_fusion(
|
||||
results_by_source={
|
||||
"vector": raw_results.get("vector", []),
|
||||
"graph": raw_results.get("graph", []),
|
||||
"web": raw_results.get("web", [])
|
||||
},
|
||||
k=config.rrf_k
|
||||
)
|
||||
timing["fusion_ms"] = (time.time() - phase2_start) * 1000
|
||||
|
||||
# Phase 3: Enrichment
|
||||
phase3_start = time.time()
|
||||
if config.enable_enrichment:
|
||||
enriched_results = await self._enrich_with_related_dossiers(fused_results, user)
|
||||
else:
|
||||
enriched_results = fused_results
|
||||
timing["enrichment_ms"] = (time.time() - phase3_start) * 1000
|
||||
|
||||
# Phase 4: LLM Re-ranking
|
||||
phase4_start = time.time()
|
||||
if config.enable_reranking and len(enriched_results) > 1:
|
||||
reranked_results = await self._rerank_with_llm(enriched_results[:20], query)
|
||||
else:
|
||||
reranked_results = enriched_results
|
||||
timing["reranking_ms"] = (time.time() - phase4_start) * 1000
|
||||
|
||||
# Limit to final result count
|
||||
final_results = reranked_results[:config.final_result_count]
|
||||
|
||||
# Update final ranks
|
||||
for i, result in enumerate(final_results, start=1):
|
||||
result["final_rank"] = i
|
||||
|
||||
# Convert to HybridRAGResult models
|
||||
result_models = self._convert_to_result_models(final_results)
|
||||
|
||||
# Phase 5: Context Formatting
|
||||
context = self._format_context_for_llm(result_models)
|
||||
|
||||
# Calculate source counts
|
||||
source_counts = {}
|
||||
for result in result_models:
|
||||
for source in result.sources:
|
||||
source_counts[source] = source_counts.get(source, 0) + 1
|
||||
|
||||
timing["total_ms"] = (time.time() - start_time) * 1000
|
||||
|
||||
# Phase 6: Persistence (async, non-blocking)
|
||||
phase6_start = time.time()
|
||||
search_id = await self._persist_search_for_librarian(
|
||||
query=query,
|
||||
user=user,
|
||||
keywords_data=keywords_data,
|
||||
raw_results=raw_results,
|
||||
final_results=final_results,
|
||||
timing=timing
|
||||
)
|
||||
timing["persistence_ms"] = (time.time() - phase6_start) * 1000
|
||||
|
||||
# Build response
|
||||
return HybridRAGResponse(
|
||||
query=query,
|
||||
keywords=KeywordExtraction(**keywords_data),
|
||||
results=result_models,
|
||||
context=context,
|
||||
source_counts=source_counts,
|
||||
total_results=len(result_models),
|
||||
timing=TimingBreakdown(**timing),
|
||||
config_used=config,
|
||||
search_id=search_id
|
||||
)
|
||||
|
||||
async def _extract_keywords_and_synonyms(self, query: str) -> Dict[str, Any]:
|
||||
"""
|
||||
Phase 0: Extract keywords, entities, and synonyms using LLM.
|
||||
|
||||
Args:
|
||||
query: Search query
|
||||
|
||||
Returns:
|
||||
Dictionary with keywords, entities, synonyms, expansions
|
||||
"""
|
||||
prompt = f"""Extract search terms from this query. For each important word, provide synonyms and expansions.
|
||||
|
||||
Query: "{query}"
|
||||
|
||||
Return ONLY valid JSON:
|
||||
{{
|
||||
"core_keywords": ["key", "words", "from", "query"],
|
||||
"synonyms": {{
|
||||
"word": ["alternative", "terms"]
|
||||
}}
|
||||
}}
|
||||
|
||||
Example for "Docker container hosting":
|
||||
{{
|
||||
"core_keywords": ["docker", "container", "hosting"],
|
||||
"synonyms": {{
|
||||
"docker": ["containerization", "container runtime"],
|
||||
"hosting": ["server", "infrastructure"]
|
||||
}}
|
||||
}}
|
||||
|
||||
JSON:"""
|
||||
|
||||
try:
|
||||
response = await self.ollama.generate_text(
|
||||
prompt=prompt,
|
||||
model=self.reranker_model
|
||||
)
|
||||
|
||||
# Parse JSON response (handle potential extra text)
|
||||
response_clean = response.strip()
|
||||
# Try to extract JSON if wrapped in text
|
||||
if '{' in response_clean:
|
||||
json_start = response_clean.find('{')
|
||||
json_end = response_clean.rfind('}') + 1
|
||||
response_clean = response_clean[json_start:json_end]
|
||||
|
||||
keywords_data = json.loads(response_clean)
|
||||
|
||||
# Ensure all required fields exist
|
||||
result = {
|
||||
"core_keywords": keywords_data.get("core_keywords", []),
|
||||
"entities": keywords_data.get("entities", []),
|
||||
"synonyms": keywords_data.get("synonyms", {}),
|
||||
"expansions": keywords_data.get("expansions", {})
|
||||
}
|
||||
|
||||
logger.info(f"Extracted keywords: {result['core_keywords'][:5]}, synonyms: {len(result['synonyms'])} terms")
|
||||
return result
|
||||
|
||||
except json.JSONDecodeError as e:
|
||||
logger.warning(f"Failed to parse LLM keyword extraction: {e}, using fallback")
|
||||
# Fallback to simple extraction
|
||||
words = query.split()
|
||||
return {
|
||||
"core_keywords": words,
|
||||
"entities": [],
|
||||
"synonyms": {},
|
||||
"expansions": {}
|
||||
}
|
||||
except Exception as e:
|
||||
logger.error(f"Keyword extraction failed: {e}", exc_info=True)
|
||||
return {
|
||||
"core_keywords": query.split(),
|
||||
"entities": [],
|
||||
"synonyms": {},
|
||||
"expansions": {}
|
||||
}
|
||||
|
||||
async def _retrieve_parallel(
|
||||
self,
|
||||
query: str,
|
||||
user: str,
|
||||
config: HybridRAGConfig,
|
||||
keywords_data: Dict[str, Any]
|
||||
) -> Dict[str, List]:
|
||||
"""
|
||||
Phase 1: Retrieve results from all sources in parallel.
|
||||
|
||||
Args:
|
||||
query: Search query
|
||||
user: User identifier
|
||||
config: Search configuration
|
||||
keywords_data: Extracted keywords/synonyms
|
||||
|
||||
Returns:
|
||||
Dictionary with results from each source and timing
|
||||
"""
|
||||
tasks = {}
|
||||
timing = {}
|
||||
|
||||
# Vector search
|
||||
if config.enable_vector:
|
||||
async def vector_search():
|
||||
start = time.time()
|
||||
try:
|
||||
response = await self.vector.search(
|
||||
query=query,
|
||||
user=user,
|
||||
limit=config.vector_limit
|
||||
)
|
||||
results = [
|
||||
{
|
||||
"page_id": r.page_id,
|
||||
"title": r.page_title,
|
||||
"content": r.content,
|
||||
"path": r.page_path,
|
||||
"score": r.score,
|
||||
"source": "vector"
|
||||
}
|
||||
for r in response.results
|
||||
]
|
||||
return results, (time.time() - start) * 1000
|
||||
except Exception as e:
|
||||
logger.error(f"Vector search failed: {e}", exc_info=True)
|
||||
return [], (time.time() - start) * 1000
|
||||
|
||||
tasks["vector"] = vector_search()
|
||||
|
||||
# Graph search
|
||||
if config.enable_graph:
|
||||
async def graph_search():
|
||||
start = time.time()
|
||||
try:
|
||||
results = await self.graph.search_documents(
|
||||
query=query,
|
||||
user=user,
|
||||
limit=config.graph_limit,
|
||||
keywords_data=keywords_data
|
||||
)
|
||||
formatted = [
|
||||
{
|
||||
"page_id": r["page_id"],
|
||||
"title": r["title"],
|
||||
"content": "", # Graph doesn't return content
|
||||
"path": r["path"],
|
||||
"entity_matches": r.get("entity_matches", 0),
|
||||
"matched_entities": r.get("matched_entities", []),
|
||||
"source": "graph"
|
||||
}
|
||||
for r in results
|
||||
]
|
||||
return formatted, (time.time() - start) * 1000
|
||||
except Exception as e:
|
||||
logger.error(f"Graph search failed: {e}", exc_info=True)
|
||||
return [], (time.time() - start) * 1000
|
||||
|
||||
tasks["graph"] = graph_search()
|
||||
|
||||
# Web search
|
||||
if config.enable_web:
|
||||
async def web_search():
|
||||
start = time.time()
|
||||
try:
|
||||
results = await self.searxng.search_general(
|
||||
query=query,
|
||||
limit=config.web_limit
|
||||
)
|
||||
formatted = [
|
||||
{
|
||||
"url": r.get("url"),
|
||||
"title": r.get("title", ""),
|
||||
"content": r.get("content", ""),
|
||||
"engine": r.get("engine", ""),
|
||||
"source": "web"
|
||||
}
|
||||
for r in results
|
||||
]
|
||||
return formatted, (time.time() - start) * 1000
|
||||
except Exception as e:
|
||||
logger.error(f"Web search failed: {e}", exc_info=True)
|
||||
return [], (time.time() - start) * 1000
|
||||
|
||||
tasks["web"] = web_search()
|
||||
|
||||
# Execute all searches in parallel
|
||||
results_dict = await asyncio.gather(*tasks.values())
|
||||
|
||||
# Combine results with timing
|
||||
output = {"timing": {}}
|
||||
for i, source in enumerate(tasks.keys()):
|
||||
results, source_timing = results_dict[i]
|
||||
output[source] = results
|
||||
output["timing"][f"{source}_ms"] = source_timing
|
||||
|
||||
logger.info(
|
||||
f"Parallel retrieval: vector={len(output.get('vector', []))}, "
|
||||
f"graph={len(output.get('graph', []))}, web={len(output.get('web', []))}"
|
||||
)
|
||||
|
||||
return output
|
||||
|
||||
def _reciprocal_rank_fusion(
|
||||
self,
|
||||
results_by_source: Dict[str, List],
|
||||
k: int = 60
|
||||
) -> List[Dict[str, Any]]:
|
||||
"""
|
||||
Phase 2: Merge results using Reciprocal Rank Fusion.
|
||||
|
||||
RRF formula: score = sum(1 / (k + rank)) for each source
|
||||
|
||||
Args:
|
||||
results_by_source: Results from each source
|
||||
k: RRF constant (default 60)
|
||||
|
||||
Returns:
|
||||
Merged and sorted results
|
||||
"""
|
||||
rrf_scores = {}
|
||||
|
||||
for source, results in results_by_source.items():
|
||||
for rank, result in enumerate(results, start=1):
|
||||
# Use page_id for wiki results, url hash for web results
|
||||
if result.get("page_id"):
|
||||
result_id = f"page_{result['page_id']}"
|
||||
elif result.get("url"):
|
||||
result_id = f"url_{hash(result['url'])}"
|
||||
else:
|
||||
continue # Skip results without ID
|
||||
|
||||
if result_id not in rrf_scores:
|
||||
rrf_scores[result_id] = {
|
||||
"result": result,
|
||||
"rrf_score": 0.0,
|
||||
"sources": [],
|
||||
"source_type": source
|
||||
}
|
||||
|
||||
# RRF formula: sum of 1/(k + rank) across sources
|
||||
rrf_scores[result_id]["rrf_score"] += 1 / (k + rank)
|
||||
rrf_scores[result_id]["sources"].append(source)
|
||||
|
||||
# If result appears in multiple sources, update source_type
|
||||
if len(rrf_scores[result_id]["sources"]) > 1:
|
||||
rrf_scores[result_id]["source_type"] = "+".join(
|
||||
sorted(set(rrf_scores[result_id]["sources"]))
|
||||
)
|
||||
|
||||
# Sort by RRF score descending
|
||||
sorted_results = sorted(
|
||||
rrf_scores.values(),
|
||||
key=lambda x: x["rrf_score"],
|
||||
reverse=True
|
||||
)
|
||||
|
||||
logger.info(f"RRF fusion: {len(sorted_results)} unique results from {len(results_by_source)} sources")
|
||||
|
||||
return sorted_results
|
||||
|
||||
async def _enrich_with_related_dossiers(
|
||||
self,
|
||||
results: List[Dict[str, Any]],
|
||||
user: str
|
||||
) -> List[Dict[str, Any]]:
|
||||
"""
|
||||
Phase 3: Enrich results with related documents via shared entities.
|
||||
|
||||
Args:
|
||||
results: Fused results
|
||||
user: User identifier
|
||||
|
||||
Returns:
|
||||
Results with related_dossiers added
|
||||
"""
|
||||
for result in results:
|
||||
result_data = result.get("result", {})
|
||||
page_id = result_data.get("page_id")
|
||||
|
||||
if page_id:
|
||||
try:
|
||||
related_docs = await self.graph.get_related_documents(
|
||||
page_id=page_id,
|
||||
user=user,
|
||||
limit=5
|
||||
)
|
||||
|
||||
# Convert to RelatedDossier format
|
||||
related_dossiers = []
|
||||
for doc in related_docs:
|
||||
for tag in doc.get("tags", [])[:3]: # Max 3 tags per doc
|
||||
related_dossiers.append({
|
||||
"page_id": doc["page_id"],
|
||||
"title": doc["title"],
|
||||
"path": doc["path"],
|
||||
"tag": tag,
|
||||
"shared_entities": doc["shared_entities"]
|
||||
})
|
||||
|
||||
result["related_dossiers"] = related_dossiers[:5] # Limit to 5 total
|
||||
|
||||
except Exception as e:
|
||||
logger.warning(f"Failed to get related docs for page {page_id}: {e}")
|
||||
result["related_dossiers"] = []
|
||||
else:
|
||||
result["related_dossiers"] = []
|
||||
|
||||
return results
|
||||
|
||||
async def _rerank_with_llm(
|
||||
self,
|
||||
results: List[Dict[str, Any]],
|
||||
query: str
|
||||
) -> List[Dict[str, Any]]:
|
||||
"""
|
||||
Phase 4: Re-rank results using LLM for better relevance.
|
||||
|
||||
Args:
|
||||
results: Results to re-rank (top 20)
|
||||
query: Original search query
|
||||
|
||||
Returns:
|
||||
Re-ranked results
|
||||
"""
|
||||
if len(results) <= 1:
|
||||
return results
|
||||
|
||||
try:
|
||||
# Build prompt with numbered results
|
||||
docs_text = "\n".join([
|
||||
f"{i+1}. {r['result'].get('title', 'Untitled')} - {r['result'].get('content', '')[:200]}..."
|
||||
for i, r in enumerate(results)
|
||||
])
|
||||
|
||||
prompt = f"""Given this search query and documents, rank them by relevance.
|
||||
|
||||
Query: {query}
|
||||
|
||||
Documents:
|
||||
{docs_text}
|
||||
|
||||
Return only the numbers in order of relevance (most relevant first).
|
||||
Example: 3,1,5,2,4
|
||||
|
||||
Ranking:"""
|
||||
|
||||
response = await self.ollama.generate_text(
|
||||
prompt=prompt,
|
||||
model=self.reranker_model
|
||||
)
|
||||
|
||||
# Parse response: "3,1,5,2,4" → [2, 0, 4, 1, 3] (0-indexed)
|
||||
indices_str = response.strip().split('\n')[0] # Take first line
|
||||
indices = [int(x.strip()) - 1 for x in indices_str.split(",") if x.strip().isdigit()]
|
||||
|
||||
# Reorder results according to LLM ranking
|
||||
reranked = []
|
||||
for idx in indices:
|
||||
if 0 <= idx < len(results):
|
||||
reranked.append(results[idx])
|
||||
|
||||
# Add any results that weren't in the LLM response
|
||||
for i, result in enumerate(results):
|
||||
if i not in indices and result not in reranked:
|
||||
reranked.append(result)
|
||||
|
||||
logger.info(f"LLM re-ranking: reordered {len(reranked)} results")
|
||||
return reranked
|
||||
|
||||
except Exception as e:
|
||||
logger.warning(f"LLM re-ranking failed: {e}, using RRF order")
|
||||
return results # Fallback to RRF order
|
||||
|
||||
def _format_context_for_llm(self, results: List[HybridRAGResult]) -> str:
|
||||
"""
|
||||
Phase 5: Format results into context for LLM consumption.
|
||||
|
||||
Args:
|
||||
results: Ranked results
|
||||
|
||||
Returns:
|
||||
Formatted context string
|
||||
"""
|
||||
context_parts = []
|
||||
|
||||
for i, result in enumerate(results[:10], start=1):
|
||||
# Source indicator
|
||||
source_tag = f"[{result.source_type.upper()}]"
|
||||
|
||||
# Related dossiers if available
|
||||
related = ""
|
||||
if result.related_dossiers:
|
||||
tags = ", ".join([d.tag for d in result.related_dossiers[:3]])
|
||||
related = f"\n Related research: {tags}"
|
||||
|
||||
# Build context entry
|
||||
content_preview = result.content[:300] if result.content else "(no content)"
|
||||
context_parts.append(
|
||||
f"{i}. {source_tag} {result.title}\n"
|
||||
f" {content_preview}...{related}"
|
||||
)
|
||||
|
||||
return "\n\n".join(context_parts)
|
||||
|
||||
async def _persist_search_for_librarian(
|
||||
self,
|
||||
query: str,
|
||||
user: str,
|
||||
keywords_data: Dict[str, Any],
|
||||
raw_results: Dict[str, List],
|
||||
final_results: List[Dict[str, Any]],
|
||||
timing: Dict[str, float]
|
||||
) -> Optional[str]:
|
||||
"""
|
||||
Phase 6: Store search query and results for Librarian processing.
|
||||
|
||||
Creates SearchQuery node in Neo4j with relationships to found documents
|
||||
and web results for offline knowledge consolidation.
|
||||
|
||||
Args:
|
||||
query: Search query
|
||||
user: User identifier
|
||||
keywords_data: Extracted keywords/synonyms
|
||||
raw_results: Results from each source
|
||||
final_results: Final ranked results
|
||||
timing: Performance timing
|
||||
|
||||
Returns:
|
||||
Search ID for tracking
|
||||
"""
|
||||
try:
|
||||
user_base_label = get_neo4j_user_base_label(user)
|
||||
search_id = str(uuid.uuid4())
|
||||
|
||||
# Create SearchQuery node
|
||||
create_query = f"""
|
||||
CREATE (sq:{user_base_label}_SearchQuery:SearchQuery {{
|
||||
id: $search_id,
|
||||
query: $query,
|
||||
user: $user,
|
||||
timestamp: datetime(),
|
||||
processed: false,
|
||||
total_results: $total_results,
|
||||
vector_count: $vector_count,
|
||||
graph_count: $graph_count,
|
||||
web_count: $web_count,
|
||||
keywords: $keywords,
|
||||
synonyms: $synonyms,
|
||||
timing_ms: $timing_ms
|
||||
}})
|
||||
RETURN sq.id as id
|
||||
"""
|
||||
|
||||
result = await self.graph.neo4j.execute_query(create_query, {
|
||||
"search_id": search_id,
|
||||
"query": query,
|
||||
"user": user,
|
||||
"total_results": len(final_results),
|
||||
"vector_count": len(raw_results.get("vector", [])),
|
||||
"graph_count": len(raw_results.get("graph", [])),
|
||||
"web_count": len(raw_results.get("web", [])),
|
||||
"keywords": keywords_data.get("core_keywords", []),
|
||||
"synonyms": json.dumps(keywords_data.get("synonyms", {})),
|
||||
"timing_ms": timing.get("total_ms", 0)
|
||||
})
|
||||
|
||||
# Link to found wiki documents (top 20)
|
||||
for rank, result_data in enumerate(final_results[:20], start=1):
|
||||
result = result_data.get("result", {})
|
||||
page_id = result.get("page_id")
|
||||
|
||||
if page_id:
|
||||
link_doc_query = f"""
|
||||
MATCH (sq:{user_base_label}_SearchQuery:SearchQuery {{id: $search_id}})
|
||||
MATCH (d:Document {{page_id: $page_id}})
|
||||
MERGE (sq)-[f:FOUND]->(d)
|
||||
SET f.source = $source,
|
||||
f.rank = $rank,
|
||||
f.rrf_score = $rrf_score,
|
||||
f.final_rank = $final_rank
|
||||
"""
|
||||
|
||||
await self.graph.neo4j.execute_query(link_doc_query, {
|
||||
"search_id": search_id,
|
||||
"page_id": page_id,
|
||||
"source": result_data.get("source_type", "unknown"),
|
||||
"rank": rank,
|
||||
"rrf_score": result_data.get("rrf_score", 0),
|
||||
"final_rank": result_data.get("final_rank", rank)
|
||||
})
|
||||
|
||||
# Store web results as WebResult nodes (top 10)
|
||||
web_results = [r for r in final_results[:10] if r.get("result", {}).get("url")]
|
||||
for rank, result_data in enumerate(web_results, start=1):
|
||||
result = result_data.get("result", {})
|
||||
create_web_query = f"""
|
||||
MATCH (sq:{user_base_label}_SearchQuery:SearchQuery {{id: $search_id}})
|
||||
CREATE (wr:{user_base_label}_WebResult:WebResult {{
|
||||
url: $url,
|
||||
title: $title,
|
||||
content: $content,
|
||||
search_id: $search_id,
|
||||
timestamp: datetime()
|
||||
}})
|
||||
CREATE (sq)-[:FOUND {{
|
||||
source: "web",
|
||||
rank: $rank,
|
||||
rrf_score: $rrf_score
|
||||
}}]->(wr)
|
||||
"""
|
||||
|
||||
await self.graph.neo4j.execute_query(create_web_query, {
|
||||
"search_id": search_id,
|
||||
"url": result.get("url"),
|
||||
"title": result.get("title", ""),
|
||||
"content": result.get("content", "")[:1000], # Truncate
|
||||
"rank": rank,
|
||||
"rrf_score": result_data.get("rrf_score", 0)
|
||||
})
|
||||
|
||||
logger.info(f"Persisted search {search_id} for Librarian processing")
|
||||
return search_id
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"Failed to persist search for Librarian: {e}", exc_info=True)
|
||||
return None
|
||||
|
||||
def _convert_to_result_models(self, results: List[Dict[str, Any]]) -> List[HybridRAGResult]:
|
||||
"""
|
||||
Convert internal result format to HybridRAGResult models.
|
||||
|
||||
Args:
|
||||
results: Internal result dictionaries
|
||||
|
||||
Returns:
|
||||
List of HybridRAGResult models
|
||||
"""
|
||||
models = []
|
||||
|
||||
for result_data in results:
|
||||
result = result_data.get("result", {})
|
||||
related_dossiers = result_data.get("related_dossiers", [])
|
||||
|
||||
models.append(HybridRAGResult(
|
||||
source_type=result_data.get("source_type", "unknown"),
|
||||
title=result.get("title", "Untitled"),
|
||||
content=result.get("content", ""),
|
||||
url=result.get("url"),
|
||||
page_id=result.get("page_id"),
|
||||
page_path=result.get("path"),
|
||||
rrf_score=result_data.get("rrf_score", 0),
|
||||
final_rank=result_data.get("final_rank", 0),
|
||||
sources=result_data.get("sources", []),
|
||||
related_dossiers=[RelatedDossier(**d) for d in related_dossiers],
|
||||
metadata={
|
||||
"entity_matches": result.get("entity_matches"),
|
||||
"matched_entities": result.get("matched_entities"),
|
||||
"engine": result.get("engine")
|
||||
}
|
||||
))
|
||||
|
||||
return models
|
||||
@@ -0,0 +1,711 @@
|
||||
"""
|
||||
Comprehensive tests for HybridRAG system.
|
||||
|
||||
Tests cover all 6 phases:
|
||||
- Phase 0: Query Enhancement (keyword/synonym extraction)
|
||||
- Phase 1: Parallel Retrieval (vector + graph + web)
|
||||
- Phase 2: RRF Fusion
|
||||
- Phase 3: Enrichment (related dossiers)
|
||||
- Phase 4: LLM Re-ranking
|
||||
- Phase 5: Context Formatting
|
||||
- Phase 6: Persistence (search storage)
|
||||
|
||||
Uses 'llm-tester' user to avoid contaminating production data.
|
||||
|
||||
Run with: pytest tests/test_hybrid_rag.py -v -s
|
||||
"""
|
||||
|
||||
import pytest
|
||||
import pytest_asyncio
|
||||
from typing import AsyncGenerator
|
||||
import json
|
||||
|
||||
from src.clients.neo4j_client import Neo4jClient
|
||||
from src.clients.qdrant_client import QdrantClientWrapper
|
||||
from src.clients.wikijs_client import WikiJSClient
|
||||
from src.clients.searxng_client import SearXNGClient
|
||||
from src.clients.ollama_client import OllamaClient
|
||||
from src.services.hybrid_rag_service import HybridRAGService
|
||||
from src.services.vector_service import VectorService
|
||||
from src.services.graph_service import GraphService
|
||||
from src.models.hybrid_rag import HybridRAGConfig, HybridRAGRequest
|
||||
from src.config import get_settings
|
||||
|
||||
# Test user to isolate test data
|
||||
TEST_USER = "llm-tester"
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def settings():
|
||||
"""Get application settings."""
|
||||
return get_settings()
|
||||
|
||||
|
||||
@pytest_asyncio.fixture
|
||||
async def neo4j_client(settings) -> AsyncGenerator[Neo4jClient, None]:
|
||||
"""Get connected Neo4j client."""
|
||||
client = Neo4jClient(
|
||||
uri=settings.neo4j_uri,
|
||||
user=settings.neo4j_user,
|
||||
password=settings.neo4j_password
|
||||
)
|
||||
await client.connect()
|
||||
yield client
|
||||
await client.close()
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def qdrant_client(settings) -> QdrantClientWrapper:
|
||||
"""Get Qdrant client."""
|
||||
return QdrantClientWrapper(url=settings.qdrant_url)
|
||||
|
||||
|
||||
@pytest_asyncio.fixture
|
||||
async def wiki_client(settings) -> AsyncGenerator[WikiJSClient, None]:
|
||||
"""Get Wiki.js client."""
|
||||
client = WikiJSClient(
|
||||
base_url=settings.wikijs_url,
|
||||
username=settings.wikijs_username,
|
||||
password=settings.wikijs_password
|
||||
)
|
||||
yield client
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def searxng_client(settings) -> SearXNGClient:
|
||||
"""Get SearXNG client."""
|
||||
return SearXNGClient(base_url=settings.searxng_url)
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def ollama_client(settings) -> OllamaClient:
|
||||
"""Get Ollama client."""
|
||||
return OllamaClient(base_url=settings.ollama_url)
|
||||
|
||||
|
||||
@pytest_asyncio.fixture
|
||||
async def vector_service(qdrant_client, wiki_client, ollama_client):
|
||||
"""Get VectorService instance."""
|
||||
return VectorService(qdrant_client, wiki_client, ollama_client)
|
||||
|
||||
|
||||
@pytest_asyncio.fixture
|
||||
async def graph_service(neo4j_client, wiki_client):
|
||||
"""Get GraphService instance."""
|
||||
return GraphService(neo4j_client, wiki_client)
|
||||
|
||||
|
||||
@pytest_asyncio.fixture
|
||||
async def hybrid_rag_service(
|
||||
vector_service,
|
||||
graph_service,
|
||||
searxng_client,
|
||||
ollama_client,
|
||||
settings
|
||||
):
|
||||
"""Get HybridRAGService instance."""
|
||||
return HybridRAGService(
|
||||
vector_service=vector_service,
|
||||
graph_service=graph_service,
|
||||
searxng_client=searxng_client,
|
||||
ollama_client=ollama_client,
|
||||
settings=settings
|
||||
)
|
||||
|
||||
|
||||
@pytest_asyncio.fixture
|
||||
async def test_wiki_page(wiki_client):
|
||||
"""
|
||||
Create test wiki page for llm-tester user.
|
||||
|
||||
Creates a page about Docker and Kubernetes for testing.
|
||||
"""
|
||||
from src.core.multi_tenancy import get_wikijs_namespace
|
||||
|
||||
namespace = get_wikijs_namespace(TEST_USER)
|
||||
path = f"{namespace}/testing/docker-kubernetes"
|
||||
|
||||
# Create test page
|
||||
page_data = {
|
||||
"title": "Docker and Kubernetes Testing",
|
||||
"path": path,
|
||||
"content": """# Docker and Kubernetes
|
||||
|
||||
Docker is a containerization platform that packages applications into containers.
|
||||
Kubernetes (k8s) is an orchestration platform for managing Docker containers at scale.
|
||||
|
||||
## Key Technologies
|
||||
- Docker: Container runtime
|
||||
- Kubernetes: Orchestration platform
|
||||
- Helm: Package manager for Kubernetes
|
||||
- kubectl: Command-line tool for k8s
|
||||
|
||||
## Use Cases
|
||||
Our infrastructure uses Docker containers orchestrated by Kubernetes clusters.
|
||||
We deploy microservices using Helm charts and manage them with kubectl.
|
||||
""",
|
||||
"description": "Test page for HybridRAG testing",
|
||||
"tags": ["testing", "infrastructure", "docker"]
|
||||
}
|
||||
|
||||
try:
|
||||
# Delete if exists
|
||||
existing = await wiki_client.search_pages(query="Docker and Kubernetes Testing")
|
||||
for page in existing:
|
||||
if page.get("path") == path:
|
||||
await wiki_client.delete_page(page["id"])
|
||||
|
||||
# Create new
|
||||
page = await wiki_client.create_page(**page_data)
|
||||
yield page
|
||||
|
||||
# Cleanup
|
||||
try:
|
||||
await wiki_client.delete_page(page["id"])
|
||||
except:
|
||||
pass
|
||||
except Exception as e:
|
||||
pytest.skip(f"Could not create test page: {e}")
|
||||
|
||||
|
||||
@pytest_asyncio.fixture
|
||||
async def test_graph_data(graph_service, test_wiki_page):
|
||||
"""
|
||||
Populate graph with test data for llm-tester.
|
||||
|
||||
Extracts entities from test page.
|
||||
"""
|
||||
try:
|
||||
summary = await graph_service.update_from_page(
|
||||
page_id=test_wiki_page["id"],
|
||||
user=TEST_USER
|
||||
)
|
||||
yield summary
|
||||
except Exception as e:
|
||||
pytest.skip(f"Could not populate graph: {e}")
|
||||
|
||||
|
||||
@pytest_asyncio.fixture
|
||||
async def test_vector_data(vector_service, test_wiki_page):
|
||||
"""
|
||||
Populate vector DB with test data for llm-tester.
|
||||
|
||||
Creates embeddings from test page.
|
||||
"""
|
||||
try:
|
||||
summary = await vector_service.update_from_page(
|
||||
page_id=test_wiki_page["id"],
|
||||
user=TEST_USER
|
||||
)
|
||||
yield summary
|
||||
except Exception as e:
|
||||
pytest.skip(f"Could not populate vectors: {e}")
|
||||
|
||||
|
||||
# ============================================================================
|
||||
# Unit Tests - Individual Components
|
||||
# ============================================================================
|
||||
|
||||
class TestRRFFusion:
|
||||
"""Test Reciprocal Rank Fusion algorithm."""
|
||||
|
||||
def test_rrf_single_source(self, hybrid_rag_service):
|
||||
"""Test RRF with single source."""
|
||||
results_by_source = {
|
||||
"vector": [
|
||||
{"page_id": 1, "title": "Doc 1", "content": "test"},
|
||||
{"page_id": 2, "title": "Doc 2", "content": "test"}
|
||||
]
|
||||
}
|
||||
|
||||
fused = hybrid_rag_service._reciprocal_rank_fusion(results_by_source, k=60)
|
||||
|
||||
assert len(fused) == 2
|
||||
assert fused[0]["rrf_score"] > fused[1]["rrf_score"] # Rank 1 > Rank 2
|
||||
assert fused[0]["sources"] == ["vector"]
|
||||
|
||||
def test_rrf_multiple_sources_same_doc(self, hybrid_rag_service):
|
||||
"""Test RRF with same document from multiple sources."""
|
||||
results_by_source = {
|
||||
"vector": [{"page_id": 1, "title": "Doc 1", "content": "test"}],
|
||||
"graph": [{"page_id": 1, "title": "Doc 1", "content": ""}],
|
||||
}
|
||||
|
||||
fused = hybrid_rag_service._reciprocal_rank_fusion(results_by_source, k=60)
|
||||
|
||||
assert len(fused) == 1 # Deduplicated
|
||||
assert len(fused[0]["sources"]) == 2 # Both sources
|
||||
assert "vector" in fused[0]["sources"]
|
||||
assert "graph" in fused[0]["sources"]
|
||||
# RRF score should be sum: 1/(60+1) + 1/(60+1)
|
||||
expected_score = 1/61 + 1/61
|
||||
assert abs(fused[0]["rrf_score"] - expected_score) < 0.001
|
||||
|
||||
def test_rrf_web_results(self, hybrid_rag_service):
|
||||
"""Test RRF with web results (URL-based)."""
|
||||
results_by_source = {
|
||||
"web": [
|
||||
{"url": "https://example.com/1", "title": "Web 1", "content": "test"},
|
||||
{"url": "https://example.com/2", "title": "Web 2", "content": "test"}
|
||||
]
|
||||
}
|
||||
|
||||
fused = hybrid_rag_service._reciprocal_rank_fusion(results_by_source, k=60)
|
||||
|
||||
assert len(fused) == 2
|
||||
assert fused[0]["result"]["url"] == "https://example.com/1"
|
||||
|
||||
|
||||
class TestContextFormatting:
|
||||
"""Test context formatting for LLM."""
|
||||
|
||||
def test_format_basic(self, hybrid_rag_service):
|
||||
"""Test basic context formatting."""
|
||||
from src.models.hybrid_rag import HybridRAGResult
|
||||
|
||||
results = [
|
||||
HybridRAGResult(
|
||||
source_type="vector",
|
||||
title="Test Document",
|
||||
content="This is test content for formatting",
|
||||
page_id=1,
|
||||
rrf_score=0.5,
|
||||
final_rank=1,
|
||||
sources=["vector"]
|
||||
)
|
||||
]
|
||||
|
||||
context = hybrid_rag_service._format_context_for_llm(results)
|
||||
|
||||
assert "Test Document" in context
|
||||
assert "[VECTOR]" in context
|
||||
assert "test content" in context
|
||||
|
||||
def test_format_with_related_dossiers(self, hybrid_rag_service):
|
||||
"""Test context formatting with related dossiers."""
|
||||
from src.models.hybrid_rag import HybridRAGResult, RelatedDossier
|
||||
|
||||
results = [
|
||||
HybridRAGResult(
|
||||
source_type="vector+graph",
|
||||
title="Test Document",
|
||||
content="Content",
|
||||
page_id=1,
|
||||
rrf_score=0.5,
|
||||
final_rank=1,
|
||||
sources=["vector", "graph"],
|
||||
related_dossiers=[
|
||||
RelatedDossier(
|
||||
page_id=2,
|
||||
title="Related Doc",
|
||||
path="/test/related",
|
||||
tag="infrastructure",
|
||||
shared_entities=5
|
||||
)
|
||||
]
|
||||
)
|
||||
]
|
||||
|
||||
context = hybrid_rag_service._format_context_for_llm(results)
|
||||
|
||||
assert "Related research: infrastructure" in context
|
||||
|
||||
|
||||
# ============================================================================
|
||||
# Integration Tests - Phase Testing
|
||||
# ============================================================================
|
||||
|
||||
class TestPhase0_QueryEnhancement:
|
||||
"""Test Phase 0: Query Enhancement (keyword/synonym extraction)."""
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_extract_keywords_basic(self, hybrid_rag_service):
|
||||
"""Test basic keyword extraction."""
|
||||
query = "Docker container orchestration with Kubernetes"
|
||||
|
||||
keywords_data = await hybrid_rag_service._extract_keywords_and_synonyms(query)
|
||||
|
||||
assert "core_keywords" in keywords_data
|
||||
assert "entities" in keywords_data
|
||||
assert "synonyms" in keywords_data
|
||||
assert "expansions" in keywords_data
|
||||
|
||||
# Should extract Docker and Kubernetes
|
||||
all_terms = (
|
||||
keywords_data["core_keywords"] +
|
||||
keywords_data["entities"]
|
||||
)
|
||||
assert any("docker" in term.lower() for term in all_terms)
|
||||
assert any("kubernetes" in term.lower() or "k8s" in term.lower() for term in all_terms)
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_extract_keywords_with_abbreviations(self, hybrid_rag_service):
|
||||
"""Test keyword extraction handles abbreviations."""
|
||||
query = "k8s cluster management"
|
||||
|
||||
keywords_data = await hybrid_rag_service._extract_keywords_and_synonyms(query)
|
||||
|
||||
# Should expand k8s to kubernetes
|
||||
all_data = json.dumps(keywords_data).lower()
|
||||
assert "k8s" in all_data or "kubernetes" in all_data
|
||||
|
||||
|
||||
class TestPhase1_ParallelRetrieval:
|
||||
"""Test Phase 1: Parallel Retrieval."""
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_parallel_retrieval_all_sources(
|
||||
self,
|
||||
hybrid_rag_service,
|
||||
test_wiki_page,
|
||||
test_graph_data,
|
||||
test_vector_data
|
||||
):
|
||||
"""Test parallel retrieval from all sources."""
|
||||
config = HybridRAGConfig(
|
||||
enable_vector=True,
|
||||
enable_graph=True,
|
||||
enable_web=True,
|
||||
vector_limit=5,
|
||||
graph_limit=5,
|
||||
web_limit=3
|
||||
)
|
||||
|
||||
keywords_data = {
|
||||
"core_keywords": ["docker", "kubernetes"],
|
||||
"entities": ["Docker", "Kubernetes"],
|
||||
"synonyms": {"docker": ["container"], "kubernetes": ["k8s"]},
|
||||
"expansions": {"k8s": ["kubernetes"]}
|
||||
}
|
||||
|
||||
results = await hybrid_rag_service._retrieve_parallel(
|
||||
query="docker kubernetes",
|
||||
user=TEST_USER,
|
||||
config=config,
|
||||
keywords_data=keywords_data
|
||||
)
|
||||
|
||||
assert "vector" in results
|
||||
assert "graph" in results
|
||||
assert "web" in results
|
||||
assert "timing" in results
|
||||
|
||||
# Should have timing for each source
|
||||
assert results["timing"]["vector_ms"] >= 0
|
||||
assert results["timing"]["graph_ms"] >= 0
|
||||
assert results["timing"]["web_ms"] >= 0
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_parallel_retrieval_graceful_degradation(self, hybrid_rag_service):
|
||||
"""Test graceful degradation when sources fail."""
|
||||
config = HybridRAGConfig(
|
||||
enable_vector=True,
|
||||
enable_graph=True,
|
||||
enable_web=True
|
||||
)
|
||||
|
||||
keywords_data = {"core_keywords": ["test"], "entities": [], "synonyms": {}, "expansions": {}}
|
||||
|
||||
# Even if some sources fail, should return results from working sources
|
||||
results = await hybrid_rag_service._retrieve_parallel(
|
||||
query="test query",
|
||||
user=TEST_USER,
|
||||
config=config,
|
||||
keywords_data=keywords_data
|
||||
)
|
||||
|
||||
# Should have all keys even if empty
|
||||
assert "vector" in results
|
||||
assert "graph" in results
|
||||
assert "web" in results
|
||||
|
||||
|
||||
class TestPhase3_Enrichment:
|
||||
"""Test Phase 3: Graph Enrichment."""
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_enrich_with_related_dossiers(
|
||||
self,
|
||||
hybrid_rag_service,
|
||||
graph_service,
|
||||
test_wiki_page,
|
||||
test_graph_data
|
||||
):
|
||||
"""Test enriching results with related dossiers."""
|
||||
# Create mock fused results
|
||||
fused_results = [
|
||||
{
|
||||
"result": {
|
||||
"page_id": test_wiki_page["id"],
|
||||
"title": test_wiki_page["title"],
|
||||
"content": "test"
|
||||
},
|
||||
"rrf_score": 0.5,
|
||||
"sources": ["vector"]
|
||||
}
|
||||
]
|
||||
|
||||
enriched = await hybrid_rag_service._enrich_with_related_dossiers(
|
||||
fused_results,
|
||||
user=TEST_USER
|
||||
)
|
||||
|
||||
assert len(enriched) == 1
|
||||
assert "related_dossiers" in enriched[0]
|
||||
# May or may not have related docs depending on graph state
|
||||
assert isinstance(enriched[0]["related_dossiers"], list)
|
||||
|
||||
|
||||
class TestPhase6_Persistence:
|
||||
"""Test Phase 6: Search Persistence."""
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_persist_search_creates_node(
|
||||
self,
|
||||
hybrid_rag_service,
|
||||
neo4j_client,
|
||||
test_wiki_page
|
||||
):
|
||||
"""Test that search persistence creates SearchQuery node."""
|
||||
keywords_data = {
|
||||
"core_keywords": ["docker", "kubernetes"],
|
||||
"entities": [],
|
||||
"synonyms": {},
|
||||
"expansions": {}
|
||||
}
|
||||
|
||||
raw_results = {
|
||||
"vector": [{"page_id": test_wiki_page["id"], "title": "Test", "content": "test"}],
|
||||
"graph": [],
|
||||
"web": []
|
||||
}
|
||||
|
||||
final_results = [
|
||||
{
|
||||
"result": {"page_id": test_wiki_page["id"], "title": "Test"},
|
||||
"rrf_score": 0.5,
|
||||
"final_rank": 1,
|
||||
"sources": ["vector"]
|
||||
}
|
||||
]
|
||||
|
||||
timing = {"total_ms": 1000}
|
||||
|
||||
search_id = await hybrid_rag_service._persist_search_for_librarian(
|
||||
query="test query",
|
||||
user=TEST_USER,
|
||||
keywords_data=keywords_data,
|
||||
raw_results=raw_results,
|
||||
final_results=final_results,
|
||||
timing=timing
|
||||
)
|
||||
|
||||
assert search_id is not None
|
||||
|
||||
# Verify SearchQuery node was created
|
||||
from src.core.multi_tenancy import get_neo4j_user_base_label
|
||||
user_label = get_neo4j_user_base_label(TEST_USER)
|
||||
|
||||
query = f"""
|
||||
MATCH (sq:{user_label}_SearchQuery:SearchQuery {{id: $search_id}})
|
||||
RETURN sq.query as query, sq.processed as processed
|
||||
"""
|
||||
|
||||
result = await neo4j_client.execute_query(query, {"search_id": search_id})
|
||||
assert len(result) == 1
|
||||
assert result[0]["query"] == "test query"
|
||||
assert result[0]["processed"] == False
|
||||
|
||||
# Cleanup
|
||||
cleanup_query = f"""
|
||||
MATCH (sq:{user_label}_SearchQuery:SearchQuery {{id: $search_id}})
|
||||
DETACH DELETE sq
|
||||
"""
|
||||
await neo4j_client.execute_query(cleanup_query, {"search_id": search_id})
|
||||
|
||||
|
||||
# ============================================================================
|
||||
# End-to-End Tests
|
||||
# ============================================================================
|
||||
|
||||
class TestHybridRAG_EndToEnd:
|
||||
"""End-to-end tests for complete HybridRAG flow."""
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_full_search_pipeline(
|
||||
self,
|
||||
hybrid_rag_service,
|
||||
test_wiki_page,
|
||||
test_graph_data,
|
||||
test_vector_data
|
||||
):
|
||||
"""
|
||||
Test complete HybridRAG search pipeline with all 6 phases.
|
||||
|
||||
This is the main end-to-end test that validates:
|
||||
- Phase 0: Query enhancement
|
||||
- Phase 1: Parallel retrieval
|
||||
- Phase 2: RRF fusion
|
||||
- Phase 3: Enrichment
|
||||
- Phase 4: Re-ranking
|
||||
- Phase 5: Context formatting
|
||||
- Phase 6: Persistence
|
||||
"""
|
||||
query = "How does Docker work with Kubernetes?"
|
||||
config = HybridRAGConfig(
|
||||
vector_limit=5,
|
||||
graph_limit=5,
|
||||
web_limit=3,
|
||||
enable_reranking=True,
|
||||
enable_enrichment=True,
|
||||
final_result_count=10
|
||||
)
|
||||
|
||||
# Execute full search
|
||||
response = await hybrid_rag_service.search(
|
||||
query=query,
|
||||
user=TEST_USER,
|
||||
config=config
|
||||
)
|
||||
|
||||
# Validate response structure
|
||||
assert response.query == query
|
||||
assert response.keywords is not None
|
||||
assert response.results is not None
|
||||
assert response.context is not None
|
||||
assert response.source_counts is not None
|
||||
assert response.total_results >= 0
|
||||
assert response.timing is not None
|
||||
assert response.config_used == config
|
||||
assert response.search_id is not None
|
||||
|
||||
# Validate timing breakdown
|
||||
assert response.timing.query_enhancement_ms >= 0
|
||||
assert response.timing.vector_ms >= 0
|
||||
assert response.timing.graph_ms >= 0
|
||||
assert response.timing.web_ms >= 0
|
||||
assert response.timing.fusion_ms >= 0
|
||||
assert response.timing.enrichment_ms >= 0
|
||||
assert response.timing.reranking_ms >= 0
|
||||
assert response.timing.persistence_ms >= 0
|
||||
assert response.timing.total_ms >= 0
|
||||
|
||||
# Validate keywords extraction
|
||||
assert len(response.keywords.core_keywords) > 0
|
||||
|
||||
# Validate context is formatted
|
||||
assert len(response.context) > 0
|
||||
|
||||
# Log results for inspection
|
||||
print(f"\n=== HybridRAG E2E Test Results ===")
|
||||
print(f"Query: {response.query}")
|
||||
print(f"Total Results: {response.total_results}")
|
||||
print(f"Source Counts: {response.source_counts}")
|
||||
print(f"Keywords: {response.keywords.core_keywords}")
|
||||
print(f"Total Time: {response.timing.total_ms:.0f}ms")
|
||||
print(f"Search ID: {response.search_id}")
|
||||
|
||||
if response.results:
|
||||
print(f"\nTop Result:")
|
||||
top = response.results[0]
|
||||
print(f" Title: {top.title}")
|
||||
print(f" Source: {top.source_type}")
|
||||
print(f" RRF Score: {top.rrf_score:.4f}")
|
||||
print(f" Rank: {top.final_rank}")
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_search_with_disabled_sources(
|
||||
self,
|
||||
hybrid_rag_service,
|
||||
test_wiki_page,
|
||||
test_vector_data
|
||||
):
|
||||
"""Test HybridRAG with some sources disabled."""
|
||||
config = HybridRAGConfig(
|
||||
enable_vector=True,
|
||||
enable_graph=False, # Disabled
|
||||
enable_web=False, # Disabled
|
||||
enable_reranking=False,
|
||||
final_result_count=5
|
||||
)
|
||||
|
||||
response = await hybrid_rag_service.search(
|
||||
query="docker containers",
|
||||
user=TEST_USER,
|
||||
config=config
|
||||
)
|
||||
|
||||
# Should only have vector results
|
||||
assert response.total_results >= 0
|
||||
if response.total_results > 0:
|
||||
assert all(
|
||||
"vector" in result.sources
|
||||
for result in response.results
|
||||
)
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_search_performance_target(
|
||||
self,
|
||||
hybrid_rag_service,
|
||||
test_wiki_page,
|
||||
test_graph_data,
|
||||
test_vector_data
|
||||
):
|
||||
"""Test that search completes within performance target (<3.5s)."""
|
||||
import time
|
||||
|
||||
config = HybridRAGConfig()
|
||||
|
||||
start = time.time()
|
||||
response = await hybrid_rag_service.search(
|
||||
query="kubernetes orchestration",
|
||||
user=TEST_USER,
|
||||
config=config
|
||||
)
|
||||
duration_ms = (time.time() - start) * 1000
|
||||
|
||||
print(f"\nPerformance: {duration_ms:.0f}ms (target: <3500ms)")
|
||||
|
||||
# Soft assertion - warn if exceeds target
|
||||
if duration_ms > 3500:
|
||||
print(f"WARNING: Search exceeded 3.5s target ({duration_ms:.0f}ms)")
|
||||
|
||||
|
||||
# ============================================================================
|
||||
# Cleanup Tests
|
||||
# ============================================================================
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_cleanup_test_data(neo4j_client, qdrant_client):
|
||||
"""
|
||||
Cleanup test data for llm-tester user.
|
||||
|
||||
Run this to clean up test data:
|
||||
pytest tests/test_hybrid_rag.py::test_cleanup_test_data -v -s
|
||||
"""
|
||||
from src.core.multi_tenancy import (
|
||||
get_neo4j_user_base_label,
|
||||
get_neo4j_user_label,
|
||||
get_qdrant_collection_name
|
||||
)
|
||||
|
||||
# Clean Neo4j
|
||||
user_base_label = get_neo4j_user_base_label(TEST_USER)
|
||||
user_doc_label = get_neo4j_user_label(TEST_USER)
|
||||
|
||||
# Delete all test user nodes
|
||||
delete_query = f"""
|
||||
MATCH (n)
|
||||
WHERE n:{user_base_label} OR n:{user_doc_label}
|
||||
DETACH DELETE n
|
||||
"""
|
||||
await neo4j_client.execute_query(delete_query, {})
|
||||
|
||||
# Clean Qdrant
|
||||
collection_name = get_qdrant_collection_name(TEST_USER)
|
||||
try:
|
||||
await qdrant_client.delete_collection(collection_name)
|
||||
except:
|
||||
pass
|
||||
|
||||
print(f"\n✓ Cleaned up test data for user: {TEST_USER}")
|
||||
Reference in New Issue
Block a user