Files
poimen-memory/crates/mem-cli/src/hybrid_query_worker.rs
T

388 lines
13 KiB
Rust

use crate::query_optimizer::{QueryContext, QueryOptimizer, RRFConfig, RRFFusion, SearchStrategy};
use crate::opensearch_client::OpenSearchClient;
use anyhow::Result;
use mem_llm::EmbeddingsClient;
use mem_store::VectorStore;
use serde::{Deserialize, Serialize};
use std::sync::Arc;
use std::time::Instant;
/// Hybrid Query Result with score breakdown
#[derive(Clone, Debug, Serialize, Deserialize)]
pub struct HybridQueryResult {
pub id: String,
pub text: String,
pub source: String,
pub level: String,
pub breadcrumb: Vec<String>,
// Scoring breakdown
pub final_score: f32,
pub semantic_score: Option<f32>, // From pgvector
pub lexical_score: Option<f32>, // From OpenSearch
pub fusion_method: String, // "rrf" or "weighted_linear"
pub rank: usize,
pub retrieval_engine: String, // "semantic_only", "lexical_only", or "hybrid"
}
/// Hybrid Query Response
#[derive(Clone, Debug, Serialize, Deserialize)]
pub struct HybridQueryResponse {
pub query: String,
pub project: String,
pub search_strategy: String,
pub strategy_confidence: f32,
pub results: Vec<HybridQueryResult>,
pub metrics: QueryMetrics,
}
/// Query execution metrics
#[derive(Clone, Debug, Serialize, Deserialize)]
pub struct QueryMetrics {
pub total_time_ms: u128,
pub semantic_time_ms: Option<u128>,
pub lexical_time_ms: Option<u128>,
pub fusion_time_ms: u128,
pub semantic_results_count: Option<usize>,
pub lexical_results_count: Option<usize>,
pub final_results_count: usize,
}
/// Hybrid Query Worker: orchestrates parallel retrieval
pub struct HybridQueryWorker {
optimizer: Arc<QueryOptimizer>,
vector_store: Arc<VectorStore>,
embeddings: Arc<EmbeddingsClient>,
opensearch: Option<Arc<OpenSearchClient>>,
rrf_config: RRFConfig,
}
impl HybridQueryWorker {
pub fn new(
vector_store: Arc<VectorStore>,
embeddings: Arc<EmbeddingsClient>,
opensearch: Option<Arc<OpenSearchClient>>,
) -> Self {
Self {
optimizer: Arc::new(QueryOptimizer::new()),
vector_store,
embeddings,
opensearch,
rrf_config: RRFConfig::default(),
}
}
/// Main entry point: hybrid query with full orchestration
pub async fn query(
&self,
project: &str,
question: &str,
limit: i64,
jwt_token: &str,
) -> Result<HybridQueryResponse> {
let start = Instant::now();
// Stage 1: Optimize query
let mut query_ctx = self.optimizer.optimize_query(question).await?;
// Stage 2: Generate embedding
query_ctx.embedding = Some(self.embeddings.embed_one(question).await?);
// Stage 3: Execute retrieval based on strategy
let (semantic_results, lexical_results, metrics) = match &query_ctx.search_strategy {
SearchStrategy::Hybrid => {
self.retrieve_hybrid(
project,
&query_ctx,
limit,
jwt_token,
)
.await?
}
SearchStrategy::SemanticOnly => {
let sem_results = self.retrieve_semantic(project, &query_ctx, limit).await?;
(Some(sem_results), None, QueryMetrics::default())
}
SearchStrategy::LexicalOnly => {
let lex_results = self.retrieve_lexical(project, &query_ctx, limit, jwt_token).await?;
(None, Some(lex_results), QueryMetrics::default())
}
SearchStrategy::LexicalFirst => {
self.retrieve_cascading(
project,
&query_ctx,
limit,
jwt_token,
)
.await?
}
};
// Stage 4: Fuse results
let fusion_start = Instant::now();
let fused = self.fuse_results(semantic_results, lexical_results)?;
let fusion_time_ms = fusion_start.elapsed().as_millis();
// Stage 5: Build response
let results = self.build_results(fused, &query_ctx).await?;
let mut metrics = metrics;
metrics.total_time_ms = start.elapsed().as_millis();
metrics.fusion_time_ms = fusion_time_ms;
metrics.final_results_count = results.len();
Ok(HybridQueryResponse {
query: question.to_string(),
project: project.to_string(),
search_strategy: format!("{:?}", query_ctx.search_strategy),
strategy_confidence: query_ctx.confidence,
results,
metrics,
})
}
/// Hybrid retrieval: parallel pgvector + OpenSearch
async fn retrieve_hybrid(
&self,
project: &str,
query_ctx: &QueryContext,
limit: i64,
jwt_token: &str,
) -> Result<(Option<Vec<(String, f32)>>, Option<Vec<(String, f32)>>, QueryMetrics)> {
let embedding = query_ctx
.embedding
.as_ref()
.ok_or_else(|| anyhow::anyhow!("no embedding generated"))?;
// Parallel execution
let semantic_fut = self.retrieve_semantic(project, query_ctx, limit);
let lexical_fut = self.retrieve_lexical(project, query_ctx, limit, jwt_token);
let sem_start = Instant::now();
let (semantic_results, lexical_results) = tokio::try_join!(semantic_fut, lexical_fut)?;
let sem_time = sem_start.elapsed().as_millis();
let metrics = QueryMetrics {
semantic_time_ms: Some(sem_time),
lexical_time_ms: Some(sem_time), // Parallel, so roughly same
semantic_results_count: Some(semantic_results.len()),
lexical_results_count: Some(lexical_results.len()),
..Default::default()
};
Ok((Some(semantic_results), Some(lexical_results), metrics))
}
/// Cascading retrieval: lexical → semantic
async fn retrieve_cascading(
&self,
project: &str,
query_ctx: &QueryContext,
limit: i64,
jwt_token: &str,
) -> Result<(Option<Vec<(String, f32)>>, Option<Vec<(String, f32)>>, QueryMetrics)> {
// Stage 1: Lexical search (narrow down)
let lex_start = Instant::now();
let lexical_results = self.retrieve_lexical(project, query_ctx, limit * 4, jwt_token).await?;
let lex_time = lex_start.elapsed().as_millis();
// Extract chunk IDs from lexical results
let chunk_ids: Vec<String> = lexical_results.iter().map(|(id, _)| id.clone()).collect();
// Stage 2: Semantic rerank (on narrowed set)
let sem_start = Instant::now();
let semantic_results = self
.retrieve_semantic_with_ids(project, query_ctx, limit, &chunk_ids)
.await?;
let sem_time = sem_start.elapsed().as_millis();
let metrics = QueryMetrics {
lexical_time_ms: Some(lex_time),
semantic_time_ms: Some(sem_time),
lexical_results_count: Some(lexical_results.len()),
semantic_results_count: Some(semantic_results.len()),
..Default::default()
};
Ok((Some(semantic_results), Some(lexical_results), metrics))
}
/// Retrieve from pgvector (semantic search)
async fn retrieve_semantic(
&self,
project: &str,
query_ctx: &QueryContext,
limit: i64,
) -> Result<Vec<(String, f32)>> {
let embedding = query_ctx
.embedding
.as_ref()
.ok_or_else(|| anyhow::anyhow!("no embedding generated"))?;
// Query pgvector with filters
let results = self
.vector_store
.search(embedding, project, limit, None)
.await?;
// Convert to (id, score) tuples
let scored: Vec<(String, f32)> = results
.into_iter()
.map(|(id, score, _)| (id, score))
.collect();
Ok(scored)
}
/// Retrieve from pgvector with specific chunk IDs (for cascading)
async fn retrieve_semantic_with_ids(
&self,
project: &str,
query_ctx: &QueryContext,
limit: i64,
chunk_ids: &[String],
) -> Result<Vec<(String, f32)>> {
let embedding = query_ctx
.embedding
.as_ref()
.ok_or_else(|| anyhow::anyhow!("no embedding generated"))?;
// Query pgvector filtered by chunk IDs
let results = self
.vector_store
.search_with_ids(embedding, project, limit, chunk_ids)
.await?;
let scored: Vec<(String, f32)> = results
.into_iter()
.map(|(id, score, _)| (id, score))
.collect();
Ok(scored)
}
/// Retrieve from OpenSearch (lexical search)
async fn retrieve_lexical(
&self,
project: &str,
query_ctx: &QueryContext,
limit: i64,
jwt_token: &str,
) -> Result<Vec<(String, f32)>> {
let opensearch = self
.opensearch
.as_ref()
.ok_or_else(|| anyhow::anyhow!("OpenSearch not configured"))?;
// Query OpenSearch with JWT auth
let results = opensearch
.lexical_search(&query_ctx.normalized_query, limit as usize, jwt_token)
.await?;
// Convert to (id, score) tuples
let scored: Vec<(String, f32)> = results
.into_iter()
.map(|(id, score, _, _, _)| (id, score))
.collect();
Ok(scored)
}
/// Fuse semantic and lexical results using RRF
fn fuse_results(
&self,
semantic: Option<Vec<(String, f32)>>,
lexical: Option<Vec<(String, f32)>>,
) -> Result<Vec<(String, f32)>> {
match (semantic, lexical) {
(Some(sem), Some(lex)) => {
// Use RRF for fusion
let fusion = RRFFusion::new(self.rrf_config.clone());
Ok(fusion.fuse(sem, lex))
}
(Some(sem), None) => {
// Semantic only: return top-k
let mut results = sem;
results.sort_by(|a, b| b.1.partial_cmp(&a.1).unwrap());
results.truncate(self.rrf_config.final_k);
Ok(results)
}
(None, Some(lex)) => {
// Lexical only: return top-k
let mut results = lex;
results.sort_by(|a, b| b.1.partial_cmp(&a.1).unwrap());
results.truncate(self.rrf_config.final_k);
Ok(results)
}
(None, None) => Err(anyhow::anyhow!("no results from either engine")),
}
}
/// Build response with enriched metadata
async fn build_results(
&self,
fused: Vec<(String, f32)>,
query_ctx: &QueryContext,
) -> Result<Vec<HybridQueryResult>> {
let mut results = Vec::new();
for (rank, (id, score)) in fused.into_iter().enumerate() {
// Fetch full chunk metadata from database
let chunk = self.vector_store.get_chunk(&id).await?;
results.push(HybridQueryResult {
id: id.clone(),
text: chunk.text,
source: chunk.source,
level: chunk.level.unwrap_or_default(),
breadcrumb: chunk.breadcrumb.unwrap_or_default(),
final_score: score,
semantic_score: None, // Would need to track separately
lexical_score: None, // Would need to track separately
fusion_method: "rrf".to_string(),
rank: rank + 1,
retrieval_engine: format!("{:?}", query_ctx.search_strategy),
});
}
Ok(results)
}
}
impl Default for QueryMetrics {
fn default() -> Self {
Self {
total_time_ms: 0,
semantic_time_ms: None,
lexical_time_ms: None,
fusion_time_ms: 0,
semantic_results_count: None,
lexical_results_count: None,
final_results_count: 0,
}
}
}
#[cfg(test)]
mod tests {
use super::*;
// These tests require mock implementations of VectorStore and EmbeddingsClient
// Placeholder tests for structure verification
#[test]
fn test_hybrid_response_structure() {
let resp = HybridQueryResponse {
query: "test".to_string(),
project: "poimen".to_string(),
search_strategy: "Hybrid".to_string(),
strategy_confidence: 0.95,
results: vec![],
metrics: QueryMetrics::default(),
};
assert_eq!(resp.query, "test");
assert_eq!(resp.strategy_confidence, 0.95);
}
}