Skip to main content

reddb_server/storage/query/executors/
vector.rs

1//! Vector Query Executor
2//!
3//! Executes VECTOR SEARCH queries using HNSW approximate nearest neighbor search.
4//! Supports metadata filtering, multiple distance metrics, and cross-references.
5
6use std::collections::HashMap;
7use std::sync::Arc;
8
9use crate::storage::engine::distance::{distance, DistanceMetric};
10use crate::storage::engine::hnsw::{HnswConfig, HnswIndex};
11use crate::storage::engine::unified_index::UnifiedIndex;
12use crate::storage::engine::vector_metadata::{MetadataFilter, MetadataValue};
13use crate::storage::engine::vector_store::VectorStore;
14use crate::storage::query::ast::{QueryExpr, VectorQuery, VectorSource};
15use crate::storage::query::sql_lowering::effective_vector_filter;
16use crate::storage::query::unified::{
17    ExecutionError, QueryStats, UnifiedRecord, UnifiedResult, VectorSearchResult,
18};
19use crate::storage::schema::Value;
20
21/// Vector query executor using HNSW index
22pub struct VectorExecutor {
23    /// Vector store for segment management
24    vector_store: Arc<VectorStore>,
25    /// Cross-reference index for linking vectors to nodes/rows
26    unified_index: Option<Arc<UnifiedIndex>>,
27}
28
29impl VectorExecutor {
30    /// Create a new vector executor
31    pub fn new(vector_store: Arc<VectorStore>) -> Self {
32        Self {
33            vector_store,
34            unified_index: None,
35        }
36    }
37
38    /// Create with cross-reference support
39    pub fn with_unified_index(mut self, index: Arc<UnifiedIndex>) -> Self {
40        self.unified_index = Some(index);
41        self
42    }
43
44    /// Execute a vector search query
45    pub fn execute(&self, query: &VectorQuery) -> Result<UnifiedResult, ExecutionError> {
46        let start = std::time::Instant::now();
47
48        // Resolve the query vector
49        let query_vector = self.resolve_vector_source(&query.query_vector)?;
50
51        // Get the collection
52        let collection = self.vector_store.get(&query.collection).ok_or_else(|| {
53            ExecutionError::new(format!("Vector collection not found: {}", query.collection))
54        })?;
55
56        // Search the vector store with filter
57        let search_results = collection.search_with_filter(
58            &query_vector,
59            query.k,
60            effective_vector_filter(query).as_ref(),
61        );
62
63        // Build result
64        let mut result = UnifiedResult::with_columns(vec![
65            "id".to_string(),
66            "distance".to_string(),
67            "collection".to_string(),
68        ]);
69
70        if query.include_vectors {
71            result.columns.push("vector".to_string());
72        }
73        if query.include_metadata {
74            result.columns.push("metadata".to_string());
75        }
76
77        // Convert search results to unified records
78        for sr in search_results {
79            // Apply threshold filter if specified
80            if let Some(threshold) = query.threshold {
81                if sr.distance > threshold {
82                    continue;
83                }
84            }
85
86            let mut record = UnifiedRecord::new();
87
88            // Build vector search result
89            let mut vsr = VectorSearchResult::new(sr.id, &query.collection, sr.distance);
90
91            // Include vector data if requested and available
92            if query.include_vectors {
93                if let Some(vec_data) = sr.vector {
94                    vsr = vsr.with_vector(vec_data);
95                }
96            }
97
98            // Include metadata if requested and available
99            if query.include_metadata {
100                if let Some(ref meta_entry) = sr.metadata {
101                    // Convert MetadataEntry to HashMap<String, Value>
102                    let mut meta_map: HashMap<String, Value> = HashMap::new();
103                    for (k, v) in &meta_entry.strings {
104                        meta_map.insert(k.clone(), Value::text(v.clone()));
105                    }
106                    for (k, v) in &meta_entry.integers {
107                        meta_map.insert(k.clone(), Value::Integer(*v));
108                    }
109                    for (k, v) in &meta_entry.floats {
110                        meta_map.insert(k.clone(), Value::Float(*v));
111                    }
112                    for (k, v) in &meta_entry.bools {
113                        meta_map.insert(k.clone(), Value::Boolean(*v));
114                    }
115                    vsr = vsr.with_metadata(meta_map);
116                }
117            }
118
119            // Add cross-references if available
120            if let Some(ref unified) = self.unified_index {
121                // Check for linked node
122                if let Some(node_id) = unified.get_vector_node(&query.collection, sr.id) {
123                    vsr = vsr.with_linked_node(node_id);
124                }
125
126                // Check for linked row
127                if let Some(row_key) = unified.get_vector_row(&query.collection, sr.id) {
128                    vsr = vsr.with_linked_row(&row_key.table, row_key.row_id);
129                }
130            }
131
132            // Add basic values to record
133            record.set_arc(Arc::from("id"), Value::Integer(sr.id as i64));
134            record.set_arc(Arc::from("distance"), Value::Float(sr.distance as f64));
135            record.set_arc(
136                Arc::from("collection"),
137                Value::text(query.collection.clone()),
138            );
139
140            record.vector_results.push(vsr);
141            result.push(record);
142        }
143
144        // Update stats
145        result.stats = QueryStats {
146            nodes_scanned: 0,
147            edges_scanned: 0,
148            rows_scanned: result.len() as u64,
149            exec_time_us: start.elapsed().as_micros() as u64,
150            ..Default::default()
151        };
152
153        Ok(result)
154    }
155
156    /// Resolve vector source to actual vector data
157    fn resolve_vector_source(&self, source: &VectorSource) -> Result<Vec<f32>, ExecutionError> {
158        match source {
159            VectorSource::Literal(vec) => Ok(vec.clone()),
160
161            VectorSource::Text(text) => {
162                // Text embedding would require an embedding model
163                // For now, return an error indicating this needs external embedding
164                Err(ExecutionError::new(format!(
165                    "Text embedding not yet implemented. Provide a literal vector or use an embedding service for: '{}'",
166                    text
167                )))
168            }
169
170            VectorSource::Reference {
171                collection,
172                vector_id,
173            } => {
174                if let Some(coll) = self.vector_store.get(collection) {
175                    coll.get(*vector_id).cloned().ok_or_else(|| {
176                        ExecutionError::new(format!(
177                            "Reference vector not found: {}:{}",
178                            collection, vector_id
179                        ))
180                    })
181                } else {
182                    Err(ExecutionError::new(format!(
183                        "Vector collection not found: {}",
184                        collection
185                    )))
186                }
187            }
188
189            VectorSource::Subquery(expr) => self.resolve_subquery_vector(expr.as_ref()),
190        }
191    }
192
193    fn resolve_subquery_vector(&self, expr: &QueryExpr) -> Result<Vec<f32>, ExecutionError> {
194        match expr {
195            QueryExpr::Vector(query) => {
196                let result = self.execute(query)?;
197                let (collection, vector_id) =
198                    vector_subquery_reference(&result.records, &query.collection)?;
199                self.resolve_vector_source(&VectorSource::Reference {
200                    collection,
201                    vector_id,
202                })
203            }
204            other => Err(ExecutionError::new(format!(
205                "Vector subqueries currently support only nested VECTOR SEARCH expressions, got {}",
206                query_expr_name(other)
207            ))),
208        }
209    }
210}
211
212/// Convert MetadataValue to Value for unified results
213fn metadata_value_to_value(mv: MetadataValue) -> Value {
214    match mv {
215        MetadataValue::String(s) => Value::text(s),
216        MetadataValue::Integer(i) => Value::Integer(i),
217        MetadataValue::Float(f) => Value::Float(f),
218        MetadataValue::Bool(b) => Value::Boolean(b),
219        MetadataValue::Null => Value::Null,
220    }
221}
222
223// ============================================================================
224// In-Memory Executor for Testing
225// ============================================================================
226
227/// Simple in-memory vector executor for testing without full VectorStore
228pub struct InMemoryVectorExecutor {
229    /// Vectors indexed by (collection, id)
230    vectors: HashMap<(String, u64), Vec<f32>>,
231    /// Metadata indexed by (collection, id)
232    metadata: HashMap<(String, u64), HashMap<String, MetadataValue>>,
233    /// HNSW indexes by collection
234    indexes: HashMap<String, HnswIndex>,
235    /// Cross-reference index
236    unified_index: Option<Arc<UnifiedIndex>>,
237}
238
239impl InMemoryVectorExecutor {
240    /// Create a new in-memory executor
241    pub fn new() -> Self {
242        Self {
243            vectors: HashMap::new(),
244            metadata: HashMap::new(),
245            indexes: HashMap::new(),
246            unified_index: None,
247        }
248    }
249
250    /// Add cross-reference support
251    pub fn with_unified_index(mut self, index: Arc<UnifiedIndex>) -> Self {
252        self.unified_index = Some(index);
253        self
254    }
255
256    /// Add a vector to a collection
257    pub fn add_vector(
258        &mut self,
259        collection: &str,
260        id: u64,
261        vector: Vec<f32>,
262        meta: Option<HashMap<String, MetadataValue>>,
263    ) {
264        let dim = vector.len();
265
266        // Store vector
267        self.vectors
268            .insert((collection.to_string(), id), vector.clone());
269
270        // Store metadata
271        if let Some(m) = meta {
272            self.metadata.insert((collection.to_string(), id), m);
273        }
274
275        // Add to HNSW index
276        let index = self
277            .indexes
278            .entry(collection.to_string())
279            .or_insert_with(|| {
280                let config = HnswConfig {
281                    m: 16,
282                    m_max0: 32,
283                    ef_construction: 200,
284                    ef_search: 50,
285                    ml: 1.0 / (16.0_f64).ln(),
286                    metric: DistanceMetric::L2,
287                };
288                HnswIndex::new(dim, config)
289            });
290
291        index.insert_with_id(id, vector.clone());
292    }
293
294    /// Execute a vector query
295    pub fn execute(&self, query: &VectorQuery) -> Result<UnifiedResult, ExecutionError> {
296        let start = std::time::Instant::now();
297
298        // Resolve query vector
299        let query_vector = match &query.query_vector {
300            VectorSource::Literal(v) => v.clone(),
301            VectorSource::Reference {
302                collection,
303                vector_id,
304            } => self
305                .vectors
306                .get(&(collection.clone(), *vector_id))
307                .cloned()
308                .ok_or_else(|| ExecutionError::new("Reference vector not found"))?,
309            VectorSource::Text(t) => {
310                return Err(ExecutionError::new(format!(
311                    "Text embedding not implemented: '{}'",
312                    t
313                )));
314            }
315            VectorSource::Subquery(expr) => self.resolve_subquery_vector(expr.as_ref())?,
316        };
317
318        let metric = query.metric.unwrap_or(DistanceMetric::L2);
319
320        // Get or create result
321        let mut result = UnifiedResult::with_columns(vec![
322            "id".to_string(),
323            "distance".to_string(),
324            "collection".to_string(),
325        ]);
326
327        // Search using HNSW if available, otherwise brute force
328        let search_results: Vec<(u64, f32)> =
329            if let Some(index) = self.indexes.get(&query.collection) {
330                // HNSW search returns DistanceResult with id and distance
331                let mut results: Vec<_> = index
332                    .search(&query_vector, query.k)
333                    .into_iter()
334                    .map(|r| (r.id, r.distance))
335                    .collect();
336                results.sort_by(|a, b| {
337                    match a.1.partial_cmp(&b.1).unwrap_or(std::cmp::Ordering::Equal) {
338                        std::cmp::Ordering::Equal => a.0.cmp(&b.0),
339                        ordering => ordering,
340                    }
341                });
342                results
343            } else {
344                // Brute force search
345                self.brute_force_search(&query.collection, &query_vector, query.k, metric)
346            };
347
348        for (vector_id, dist) in search_results {
349            // Apply threshold
350            if let Some(threshold) = query.threshold {
351                if dist > threshold {
352                    continue;
353                }
354            }
355
356            // Apply metadata filter
357            if let Some(ref filter) = query.filter {
358                let key = (query.collection.clone(), vector_id);
359                if let Some(meta) = self.metadata.get(&key) {
360                    if !evaluate_filter(filter, meta) {
361                        continue;
362                    }
363                } else {
364                    continue; // No metadata, filter fails
365                }
366            }
367
368            let mut record = UnifiedRecord::new();
369            let mut vsr = VectorSearchResult::new(vector_id, &query.collection, dist);
370
371            if query.include_vectors {
372                if let Some(vec) = self.vectors.get(&(query.collection.clone(), vector_id)) {
373                    vsr = vsr.with_vector(vec.clone());
374                }
375            }
376
377            if query.include_metadata {
378                if let Some(meta) = self.metadata.get(&(query.collection.clone(), vector_id)) {
379                    let meta_map: HashMap<String, Value> = meta
380                        .iter()
381                        .map(|(k, v)| (k.clone(), metadata_value_to_value(v.clone())))
382                        .collect();
383                    vsr = vsr.with_metadata(meta_map);
384                }
385            }
386
387            // Add cross-references
388            if let Some(ref unified) = self.unified_index {
389                if let Some(node_id) = unified.get_vector_node(&query.collection, vector_id) {
390                    vsr = vsr.with_linked_node(node_id);
391                }
392
393                if let Some(row_key) = unified.get_vector_row(&query.collection, vector_id) {
394                    vsr = vsr.with_linked_row(&row_key.table, row_key.row_id);
395                }
396            }
397
398            record.set_arc(Arc::from("id"), Value::Integer(vector_id as i64));
399            record.set_arc(Arc::from("distance"), Value::Float(dist as f64));
400            record.set_arc(
401                Arc::from("collection"),
402                Value::text(query.collection.clone()),
403            );
404            record.vector_results.push(vsr);
405            result.push(record);
406        }
407
408        result.stats = QueryStats {
409            nodes_scanned: 0,
410            edges_scanned: 0,
411            rows_scanned: self.vectors.len() as u64,
412            exec_time_us: start.elapsed().as_micros() as u64,
413            ..Default::default()
414        };
415
416        Ok(result)
417    }
418
419    fn resolve_subquery_vector(&self, expr: &QueryExpr) -> Result<Vec<f32>, ExecutionError> {
420        match expr {
421            QueryExpr::Vector(query) => {
422                let result = self.execute(query)?;
423                let (collection, vector_id) =
424                    vector_subquery_reference(&result.records, &query.collection)?;
425                self.vectors
426                    .get(&(collection, vector_id))
427                    .cloned()
428                    .ok_or_else(|| ExecutionError::new("Subquery reference vector not found"))
429            }
430            other => Err(ExecutionError::new(format!(
431                "Vector subqueries currently support only nested VECTOR SEARCH expressions, got {}",
432                query_expr_name(other)
433            ))),
434        }
435    }
436
437    /// Brute force search when no index is available
438    fn brute_force_search(
439        &self,
440        collection: &str,
441        query: &[f32],
442        k: usize,
443        metric: DistanceMetric,
444    ) -> Vec<(u64, f32)> {
445        let mut results: Vec<(u64, f32)> = self
446            .vectors
447            .iter()
448            .filter(|((c, _), _)| c == collection)
449            .map(|((_, id), vec)| {
450                let dist = distance(query, vec, metric);
451                (*id, dist)
452            })
453            .collect();
454
455        results.sort_by(
456            |a, b| match a.1.partial_cmp(&b.1).unwrap_or(std::cmp::Ordering::Equal) {
457                std::cmp::Ordering::Equal => a.0.cmp(&b.0),
458                ordering => ordering,
459            },
460        );
461        results.truncate(k);
462        results
463    }
464}
465
466impl Default for InMemoryVectorExecutor {
467    fn default() -> Self {
468        Self::new()
469    }
470}
471
472/// Evaluate a metadata filter against metadata values
473fn evaluate_filter(filter: &MetadataFilter, metadata: &HashMap<String, MetadataValue>) -> bool {
474    match filter {
475        MetadataFilter::Eq(field, value) => metadata
476            .get(field)
477            .map(|candidate| candidate.matches_eq(value))
478            .unwrap_or(false),
479        MetadataFilter::Ne(field, value) => metadata
480            .get(field)
481            .map(|candidate| !candidate.matches_eq(value))
482            .unwrap_or(true),
483        MetadataFilter::Lt(field, value) => metadata
484            .get(field)
485            .and_then(|candidate| candidate.compare(value))
486            .map(|ord| ord == std::cmp::Ordering::Less)
487            .unwrap_or(false),
488        MetadataFilter::Lte(field, value) => metadata
489            .get(field)
490            .and_then(|candidate| candidate.compare(value))
491            .map(|ord| ord != std::cmp::Ordering::Greater)
492            .unwrap_or(false),
493        MetadataFilter::Gt(field, value) => metadata
494            .get(field)
495            .and_then(|candidate| candidate.compare(value))
496            .map(|ord| ord == std::cmp::Ordering::Greater)
497            .unwrap_or(false),
498        MetadataFilter::Gte(field, value) => metadata
499            .get(field)
500            .and_then(|candidate| candidate.compare(value))
501            .map(|ord| ord != std::cmp::Ordering::Less)
502            .unwrap_or(false),
503        MetadataFilter::In(field, values) => metadata
504            .get(field)
505            .map(|candidate| values.iter().any(|value| candidate.matches_eq(value)))
506            .unwrap_or(false),
507        MetadataFilter::NotIn(field, values) => metadata
508            .get(field)
509            .map(|candidate| !values.iter().any(|value| candidate.matches_eq(value)))
510            .unwrap_or(true),
511        MetadataFilter::Contains(field, substring) => {
512            if let Some(MetadataValue::String(s)) = metadata.get(field) {
513                s.contains(substring)
514            } else {
515                false
516            }
517        }
518        MetadataFilter::And(filters) => filters.iter().all(|f| evaluate_filter(f, metadata)),
519        MetadataFilter::Or(filters) => filters.iter().any(|f| evaluate_filter(f, metadata)),
520        MetadataFilter::Not(inner) => !evaluate_filter(inner, metadata),
521        MetadataFilter::StartsWith(field, prefix) => {
522            if let Some(MetadataValue::String(s)) = metadata.get(field) {
523                s.starts_with(prefix)
524            } else {
525                false
526            }
527        }
528        MetadataFilter::EndsWith(field, suffix) => {
529            if let Some(MetadataValue::String(s)) = metadata.get(field) {
530                s.ends_with(suffix)
531            } else {
532                false
533            }
534        }
535        MetadataFilter::GeoRadius { .. } => false,
536        MetadataFilter::Exists(field) => metadata.contains_key(field),
537        MetadataFilter::NotExists(field) => !metadata.contains_key(field),
538    }
539}
540
541fn vector_subquery_reference(
542    records: &[UnifiedRecord],
543    default_collection: &str,
544) -> Result<(String, u64), ExecutionError> {
545    let record = records
546        .first()
547        .ok_or_else(|| ExecutionError::new("Vector subquery returned no rows"))?;
548
549    let collection: String = match record.get("collection") {
550        Some(Value::Text(collection)) => collection.to_string(),
551        _ => default_collection.to_string(),
552    };
553
554    let vector_id = match record.get("id") {
555        Some(Value::Integer(id)) if *id >= 0 => *id as u64,
556        Some(Value::UnsignedInteger(id)) => *id,
557        other => {
558            return Err(ExecutionError::new(format!(
559                "Vector subquery must expose an integer id column, got {other:?}"
560            )));
561        }
562    };
563
564    Ok((collection, vector_id))
565}
566
567fn query_expr_name(expr: &QueryExpr) -> &'static str {
568    match expr {
569        QueryExpr::Table(_) => "table",
570        QueryExpr::Graph(_) => "graph",
571        QueryExpr::Join(_) => "join",
572        QueryExpr::Path(_) => "path",
573        QueryExpr::Vector(_) => "vector",
574        QueryExpr::Hybrid(_) => "hybrid",
575        QueryExpr::Insert(_) => "insert",
576        QueryExpr::Update(_) => "update",
577        QueryExpr::Delete(_) => "delete",
578        QueryExpr::CreateTable(_) => "create_table",
579        QueryExpr::CreateCollection(_) => "create_collection",
580        QueryExpr::CreateVector(_) => "create_vector",
581        QueryExpr::DropTable(_) => "drop_table",
582        QueryExpr::DropGraph(_) => "drop_graph",
583        QueryExpr::DropVector(_) => "drop_vector",
584        QueryExpr::DropDocument(_) => "drop_document",
585        QueryExpr::DropKv(_) => "drop_kv",
586        QueryExpr::DropCollection(_) => "drop_collection",
587        QueryExpr::Truncate(_) => "truncate",
588        QueryExpr::AlterTable(_) => "alter_table",
589        QueryExpr::CreateVcsRef(_) => "create_vcs_ref",
590        QueryExpr::DropVcsRef(_) => "drop_vcs_ref",
591        QueryExpr::ForkStore(_) => "fork_store",
592        QueryExpr::PromoteFork(_) => "promote_fork",
593        QueryExpr::DropFork(_) => "drop_fork",
594        QueryExpr::VcsCommand(_) => "vcs_command",
595        QueryExpr::GraphCommand(_) => "graph_command",
596        QueryExpr::SearchCommand(_) => "search_command",
597        QueryExpr::Ask(_) => "ask",
598        QueryExpr::CreateIndex(_) => "create_index",
599        QueryExpr::DropIndex(_) => "drop_index",
600        QueryExpr::ProbabilisticCommand(_) => "probabilistic_command",
601        QueryExpr::CreateTimeSeries(_) => "create_timeseries",
602        QueryExpr::CreateMetric(_) => "create_metric",
603        QueryExpr::AlterMetric(_) => "alter_metric",
604        QueryExpr::CreateSlo(_) => "create_slo",
605        QueryExpr::DropTimeSeries(_) => "drop_timeseries",
606        QueryExpr::CreateQueue(_) => "create_queue",
607        QueryExpr::AlterQueue(_) => "alter_queue",
608        QueryExpr::DropQueue(_) => "drop_queue",
609        QueryExpr::QueueSelect(_) => "queue_select",
610        QueryExpr::QueueCommand(_) => "queue_command",
611        QueryExpr::KvCommand(_) => "kv_command",
612        QueryExpr::ConfigCommand(_) => "config_command",
613        QueryExpr::CreateTree(_) => "create_tree",
614        QueryExpr::DropTree(_) => "drop_tree",
615        QueryExpr::TreeCommand(_) => "tree_command",
616        QueryExpr::SetConfig { .. } => "set_config",
617        QueryExpr::ShowConfig { .. } => "show_config",
618        QueryExpr::Scrub { .. } => "scrub",
619        QueryExpr::SetSecret { .. } => "set_secret",
620        QueryExpr::DeleteSecret { .. } => "delete_secret",
621        QueryExpr::SetKv { .. } => "set_kv",
622        QueryExpr::DeleteKv { .. } => "delete_kv",
623        QueryExpr::ShowSecrets { .. } => "show_secrets",
624        QueryExpr::SetTenant(_) => "set_tenant",
625        QueryExpr::ShowTenant => "show_tenant",
626        QueryExpr::ExplainAlter(_) => "explain_alter",
627        QueryExpr::TransactionControl(_) => "transaction_control",
628        QueryExpr::MaintenanceCommand(_) => "maintenance_command",
629        QueryExpr::CreateSchema(_) => "create_schema",
630        QueryExpr::DropSchema(_) => "drop_schema",
631        QueryExpr::CreateSequence(_) => "create_sequence",
632        QueryExpr::DropSequence(_) => "drop_sequence",
633        QueryExpr::CopyFrom(_) => "copy_from",
634        QueryExpr::CreateView(_) => "create_view",
635        QueryExpr::DropView(_) => "drop_view",
636        QueryExpr::RefreshMaterializedView(_) => "refresh_materialized_view",
637        QueryExpr::CreatePolicy(_) => "create_policy",
638        QueryExpr::DropPolicy(_) => "drop_policy",
639        QueryExpr::CreateServer(_) => "create_server",
640        QueryExpr::DropServer(_) => "drop_server",
641        QueryExpr::CreateForeignTable(_) => "create_foreign_table",
642        QueryExpr::DropForeignTable(_) => "drop_foreign_table",
643        QueryExpr::Grant(_) => "grant",
644        QueryExpr::Revoke(_) => "revoke",
645        QueryExpr::AlterUser(_) => "alter_user",
646        QueryExpr::CreateUser(_) => "create_user",
647        QueryExpr::CreateIamPolicy { .. } => "create_iam_policy",
648        QueryExpr::DropIamPolicy { .. } => "drop_iam_policy",
649        QueryExpr::AttachPolicy { .. } => "attach_policy",
650        QueryExpr::DetachPolicy { .. } => "detach_policy",
651        QueryExpr::ShowPolicies { .. } => "show_policies",
652        QueryExpr::ShowEffectivePermissions { .. } => "show_effective_permissions",
653        QueryExpr::RankOf(_) => "rank_of",
654        QueryExpr::ApproxRankOf(_) => "approx_rank_of",
655        QueryExpr::RankRange(_) => "rank_range",
656        QueryExpr::SimulatePolicy { .. } => "simulate_policy",
657        QueryExpr::LintPolicy { .. } => "lint_policy",
658        QueryExpr::MigratePolicyMode { .. } => "migrate_policy_mode",
659        QueryExpr::CreateMigration(_) => "create_migration",
660        QueryExpr::ApplyMigration(_) => "apply_migration",
661        QueryExpr::RollbackMigration(_) => "rollback_migration",
662        QueryExpr::ExplainMigration(_) => "explain_migration",
663        QueryExpr::Explain(_) => "explain",
664        QueryExpr::EventsBackfill(_) => "events_backfill",
665        QueryExpr::EventsBackfillStatus { .. } => "events_backfill_status",
666    }
667}
668
669// ============================================================================
670// Tests
671// ============================================================================
672
673#[cfg(test)]
674mod tests {
675    use super::*;
676
677    #[test]
678    fn test_in_memory_vector_search() {
679        let mut executor = InMemoryVectorExecutor::new();
680
681        // Add some vectors
682        executor.add_vector("test", 1, vec![1.0, 0.0, 0.0], None);
683        executor.add_vector("test", 2, vec![0.0, 1.0, 0.0], None);
684        executor.add_vector("test", 3, vec![0.0, 0.0, 1.0], None);
685        executor.add_vector("test", 4, vec![0.9, 0.1, 0.0], None);
686
687        let query = VectorQuery {
688            alias: None,
689            collection: "test".to_string(),
690            query_vector: VectorSource::Literal(vec![1.0, 0.0, 0.0]),
691            k: 2,
692            filter: None,
693            metric: Some(DistanceMetric::L2),
694            include_vectors: false,
695            include_metadata: false,
696            threshold: None,
697        };
698
699        let result = executor.execute(&query).unwrap();
700        assert_eq!(result.len(), 2);
701
702        // First result should be vector 1 (exact match)
703        let first = &result.records[0];
704        assert_eq!(first.get("id"), Some(&Value::Integer(1)));
705    }
706
707    #[test]
708    fn test_vector_search_with_metadata_filter() {
709        let mut executor = InMemoryVectorExecutor::new();
710
711        let mut meta1 = HashMap::new();
712        meta1.insert("type".to_string(), MetadataValue::String("cve".to_string()));
713        meta1.insert("severity".to_string(), MetadataValue::Integer(9));
714
715        let mut meta2 = HashMap::new();
716        meta2.insert("type".to_string(), MetadataValue::String("cve".to_string()));
717        meta2.insert("severity".to_string(), MetadataValue::Integer(5));
718
719        let mut meta3 = HashMap::new();
720        meta3.insert(
721            "type".to_string(),
722            MetadataValue::String("advisory".to_string()),
723        );
724        meta3.insert("severity".to_string(), MetadataValue::Integer(8));
725
726        executor.add_vector("vulns", 1, vec![1.0, 0.0], Some(meta1));
727        executor.add_vector("vulns", 2, vec![0.9, 0.1], Some(meta2));
728        executor.add_vector("vulns", 3, vec![0.8, 0.2], Some(meta3));
729
730        // Search with filter: type = 'cve' AND severity >= 7
731        let query = VectorQuery {
732            alias: None,
733            collection: "vulns".to_string(),
734            query_vector: VectorSource::Literal(vec![1.0, 0.0]),
735            k: 10,
736            filter: Some(MetadataFilter::And(vec![
737                MetadataFilter::Eq("type".to_string(), MetadataValue::String("cve".to_string())),
738                MetadataFilter::Gte("severity".to_string(), MetadataValue::Integer(7)),
739            ])),
740            metric: Some(DistanceMetric::L2),
741            include_vectors: false,
742            include_metadata: true,
743            threshold: None,
744        };
745
746        let result = executor.execute(&query).unwrap();
747
748        // Only vector 1 matches (type=cve, severity=9)
749        assert_eq!(result.len(), 1);
750        assert_eq!(result.records[0].get("id"), Some(&Value::Integer(1)));
751    }
752
753    #[test]
754    fn test_vector_search_with_threshold() {
755        let mut executor = InMemoryVectorExecutor::new();
756
757        executor.add_vector("test", 1, vec![1.0, 0.0], None);
758        executor.add_vector("test", 2, vec![0.0, 1.0], None); // Far from query
759
760        let query = VectorQuery {
761            alias: None,
762            collection: "test".to_string(),
763            query_vector: VectorSource::Literal(vec![1.0, 0.0]),
764            k: 10,
765            filter: None,
766            metric: Some(DistanceMetric::L2),
767            include_vectors: false,
768            include_metadata: false,
769            threshold: Some(0.5), // Only include close matches
770        };
771
772        let result = executor.execute(&query).unwrap();
773
774        // Only vector 1 is within threshold
775        assert_eq!(result.len(), 1);
776    }
777
778    #[test]
779    fn test_vector_search_include_vectors() {
780        let mut executor = InMemoryVectorExecutor::new();
781
782        executor.add_vector("test", 1, vec![1.0, 2.0, 3.0], None);
783
784        let query = VectorQuery {
785            alias: None,
786            collection: "test".to_string(),
787            query_vector: VectorSource::Literal(vec![1.0, 2.0, 3.0]),
788            k: 1,
789            filter: None,
790            metric: Some(DistanceMetric::L2),
791            include_vectors: true,
792            include_metadata: false,
793            threshold: None,
794        };
795
796        let result = executor.execute(&query).unwrap();
797        assert_eq!(result.len(), 1);
798
799        let vsr = &result.records[0].vector_results[0];
800        assert!(vsr.vector.is_some());
801        assert_eq!(vsr.vector.as_ref().unwrap(), &vec![1.0, 2.0, 3.0]);
802    }
803
804    #[test]
805    fn test_vector_executor_reference_source() {
806        let mut store = VectorStore::new();
807        let collection = store.create_collection("refs", 2);
808        let ref_id = collection.insert(vec![1.0, 0.0], None).unwrap();
809        collection.insert(vec![0.0, 1.0], None).unwrap();
810
811        let executor = VectorExecutor::new(Arc::new(store));
812        let query = VectorQuery {
813            alias: None,
814            collection: "refs".to_string(),
815            query_vector: VectorSource::Reference {
816                collection: "refs".to_string(),
817                vector_id: ref_id,
818            },
819            k: 1,
820            filter: None,
821            metric: Some(DistanceMetric::L2),
822            include_vectors: false,
823            include_metadata: false,
824            threshold: None,
825        };
826
827        let result = executor.execute(&query).unwrap();
828        assert_eq!(result.len(), 1);
829        assert_eq!(result.records[0].get("id"), Some(&Value::Integer(0)));
830    }
831
832    #[test]
833    fn test_vector_executor_subquery_source() {
834        let mut store = VectorStore::new();
835        let collection = store.create_collection("refs", 2);
836        collection.insert(vec![1.0, 0.0], None).unwrap();
837        collection.insert(vec![0.0, 1.0], None).unwrap();
838
839        let executor = VectorExecutor::new(Arc::new(store));
840        let inner = VectorQuery {
841            alias: None,
842            collection: "refs".to_string(),
843            query_vector: VectorSource::Literal(vec![1.0, 0.0]),
844            k: 1,
845            filter: None,
846            metric: Some(DistanceMetric::L2),
847            include_vectors: false,
848            include_metadata: false,
849            threshold: None,
850        };
851        let query = VectorQuery {
852            alias: None,
853            collection: "refs".to_string(),
854            query_vector: VectorSource::Subquery(Box::new(QueryExpr::Vector(inner))),
855            k: 1,
856            filter: None,
857            metric: Some(DistanceMetric::L2),
858            include_vectors: false,
859            include_metadata: false,
860            threshold: None,
861        };
862
863        let result = executor.execute(&query).unwrap();
864        assert_eq!(result.len(), 1);
865        assert_eq!(result.records[0].get("id"), Some(&Value::Integer(0)));
866    }
867}