Skip to main content

oxirs_vec/
filtered_search.rs

1//! Filtered search capabilities for vector indices
2//!
3//! This module provides advanced filtering capabilities for vector search,
4//! allowing searches to be constrained by metadata predicates, value ranges,
5//! and complex logical conditions.
6
7use crate::{Vector, VectorId};
8use serde::{Deserialize, Serialize};
9use std::collections::HashMap;
10
11/// Metadata filter for search operations
12#[derive(Debug, Clone, Serialize, Deserialize)]
13pub enum MetadataFilter {
14    /// Exact match on a metadata field
15    Equals { field: String, value: FilterValue },
16    /// Field value is not equal to the given value
17    NotEquals { field: String, value: FilterValue },
18    /// Field value is greater than the given value
19    GreaterThan { field: String, value: FilterValue },
20    /// Field value is greater than or equal to the given value
21    GreaterThanOrEqual { field: String, value: FilterValue },
22    /// Field value is less than the given value
23    LessThan { field: String, value: FilterValue },
24    /// Field value is less than or equal to the given value
25    LessThanOrEqual { field: String, value: FilterValue },
26    /// Field value is in the given set
27    In {
28        field: String,
29        values: Vec<FilterValue>,
30    },
31    /// Field value is not in the given set
32    NotIn {
33        field: String,
34        values: Vec<FilterValue>,
35    },
36    /// Field value contains the given substring
37    Contains { field: String, substring: String },
38    /// Field value matches the given regex pattern
39    Regex { field: String, pattern: String },
40    /// Field exists (has any value)
41    Exists { field: String },
42    /// Field does not exist or is null
43    NotExists { field: String },
44    /// Logical AND of multiple filters
45    And(Vec<MetadataFilter>),
46    /// Logical OR of multiple filters
47    Or(Vec<MetadataFilter>),
48    /// Logical NOT of a filter
49    Not(Box<MetadataFilter>),
50}
51
52/// Value type for filter predicates
53#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
54pub enum FilterValue {
55    String(String),
56    Integer(i64),
57    Float(f64),
58    Boolean(bool),
59    Null,
60}
61
62impl FilterValue {
63    /// Numeric view of a value, coercing across the numeric kinds (Integer /
64    /// Float, and numeric-looking Strings) so that e.g. a stored `Integer(75)`
65    /// can be meaningfully compared against a filter literal `Float(50.0)`.
66    fn as_numeric(&self) -> Option<f64> {
67        match self {
68            FilterValue::Integer(i) => Some(*i as f64),
69            FilterValue::Float(f) => Some(*f),
70            FilterValue::String(s) => s.parse::<f64>().ok(),
71            _ => None,
72        }
73    }
74
75    /// Compare two filter values.
76    ///
77    /// Numeric values compare across types (Integer vs. Float, and
78    /// numeric-looking Strings) by promoting both sides to `f64`; this fixes
79    /// `GreaterThan`/`LessThan`/`*OrEqual` predicates that would otherwise
80    /// silently return `Ordering::Equal` for cross-type numeric comparisons
81    /// (e.g. stored `"75"` -> `Integer(75)` against filter `Float(50.0)`).
82    fn compare(&self, other: &FilterValue) -> std::cmp::Ordering {
83        match (self, other) {
84            (FilterValue::String(a), FilterValue::String(b)) => a.cmp(b),
85            (FilterValue::Integer(a), FilterValue::Integer(b)) => a.cmp(b),
86            (FilterValue::Float(a), FilterValue::Float(b)) => {
87                a.partial_cmp(b).unwrap_or(std::cmp::Ordering::Equal)
88            }
89            (FilterValue::Boolean(a), FilterValue::Boolean(b)) => a.cmp(b),
90            _ => {
91                // Cross-type: fall back to numeric coercion when both sides are
92                // numeric (or numeric-looking strings). If either side is not
93                // numeric, the values are genuinely incomparable -> Equal (the
94                // pre-existing conservative default).
95                match (self.as_numeric(), other.as_numeric()) {
96                    (Some(a), Some(b)) => a.partial_cmp(&b).unwrap_or(std::cmp::Ordering::Equal),
97                    _ => std::cmp::Ordering::Equal,
98                }
99            }
100        }
101    }
102}
103
104impl MetadataFilter {
105    /// Evaluate the filter against a metadata map
106    pub fn evaluate(&self, metadata: &HashMap<String, String>) -> bool {
107        match self {
108            MetadataFilter::Equals { field, value } => {
109                if let Some(field_value) = metadata.get(field) {
110                    let parsed_value = Self::parse_value(field_value);
111                    &parsed_value == value
112                } else {
113                    false
114                }
115            }
116            MetadataFilter::NotEquals { field, value } => {
117                if let Some(field_value) = metadata.get(field) {
118                    let parsed_value = Self::parse_value(field_value);
119                    &parsed_value != value
120                } else {
121                    true
122                }
123            }
124            MetadataFilter::GreaterThan { field, value } => {
125                if let Some(field_value) = metadata.get(field) {
126                    let parsed_value = Self::parse_value(field_value);
127                    parsed_value.compare(value) == std::cmp::Ordering::Greater
128                } else {
129                    false
130                }
131            }
132            MetadataFilter::GreaterThanOrEqual { field, value } => {
133                if let Some(field_value) = metadata.get(field) {
134                    let parsed_value = Self::parse_value(field_value);
135                    matches!(
136                        parsed_value.compare(value),
137                        std::cmp::Ordering::Greater | std::cmp::Ordering::Equal
138                    )
139                } else {
140                    false
141                }
142            }
143            MetadataFilter::LessThan { field, value } => {
144                if let Some(field_value) = metadata.get(field) {
145                    let parsed_value = Self::parse_value(field_value);
146                    parsed_value.compare(value) == std::cmp::Ordering::Less
147                } else {
148                    false
149                }
150            }
151            MetadataFilter::LessThanOrEqual { field, value } => {
152                if let Some(field_value) = metadata.get(field) {
153                    let parsed_value = Self::parse_value(field_value);
154                    matches!(
155                        parsed_value.compare(value),
156                        std::cmp::Ordering::Less | std::cmp::Ordering::Equal
157                    )
158                } else {
159                    false
160                }
161            }
162            MetadataFilter::In { field, values } => {
163                if let Some(field_value) = metadata.get(field) {
164                    let parsed_value = Self::parse_value(field_value);
165                    values.contains(&parsed_value)
166                } else {
167                    false
168                }
169            }
170            MetadataFilter::NotIn { field, values } => {
171                if let Some(field_value) = metadata.get(field) {
172                    let parsed_value = Self::parse_value(field_value);
173                    !values.contains(&parsed_value)
174                } else {
175                    true
176                }
177            }
178            MetadataFilter::Contains { field, substring } => {
179                if let Some(field_value) = metadata.get(field) {
180                    field_value.contains(substring)
181                } else {
182                    false
183                }
184            }
185            MetadataFilter::Regex { field, pattern } => {
186                if let Some(field_value) = metadata.get(field) {
187                    if let Ok(regex) = regex::Regex::new(pattern) {
188                        regex.is_match(field_value)
189                    } else {
190                        false
191                    }
192                } else {
193                    false
194                }
195            }
196            MetadataFilter::Exists { field } => metadata.contains_key(field),
197            MetadataFilter::NotExists { field } => !metadata.contains_key(field),
198            MetadataFilter::And(filters) => filters.iter().all(|f| f.evaluate(metadata)),
199            MetadataFilter::Or(filters) => filters.iter().any(|f| f.evaluate(metadata)),
200            MetadataFilter::Not(filter) => !filter.evaluate(metadata),
201        }
202    }
203
204    /// Parse a string value into a FilterValue
205    fn parse_value(s: &str) -> FilterValue {
206        // Try to parse as integer
207        if let Ok(i) = s.parse::<i64>() {
208            return FilterValue::Integer(i);
209        }
210
211        // Try to parse as float
212        if let Ok(f) = s.parse::<f64>() {
213            return FilterValue::Float(f);
214        }
215
216        // Try to parse as boolean
217        if let Ok(b) = s.parse::<bool>() {
218            return FilterValue::Boolean(b);
219        }
220
221        // Check for null
222        if s == "null" || s.is_empty() {
223            return FilterValue::Null;
224        }
225
226        // Default to string
227        FilterValue::String(s.to_string())
228    }
229}
230
231/// Search filter combining distance and metadata constraints
232#[derive(Debug, Clone, Serialize, Deserialize)]
233pub struct SearchFilter {
234    /// Maximum distance threshold
235    pub max_distance: Option<f32>,
236    /// Minimum distance threshold
237    pub min_distance: Option<f32>,
238    /// Metadata filter predicates
239    pub metadata_filter: Option<MetadataFilter>,
240    /// Vector dimension constraints
241    pub dimension_constraints: Option<Vec<DimensionConstraint>>,
242}
243
244/// Constraint on specific vector dimensions
245#[derive(Debug, Clone, Serialize, Deserialize)]
246pub struct DimensionConstraint {
247    /// Dimension index
248    pub dimension: usize,
249    /// Minimum value for this dimension
250    pub min_value: Option<f32>,
251    /// Maximum value for this dimension
252    pub max_value: Option<f32>,
253}
254
255impl DimensionConstraint {
256    /// Check if a vector satisfies this dimension constraint
257    pub fn satisfies(&self, vector: &Vector) -> bool {
258        let values = vector.as_f32();
259
260        if self.dimension >= values.len() {
261            return false;
262        }
263
264        let value = values[self.dimension];
265
266        if let Some(min) = self.min_value {
267            if value < min {
268                return false;
269            }
270        }
271
272        if let Some(max) = self.max_value {
273            if value > max {
274                return false;
275            }
276        }
277
278        true
279    }
280}
281
282impl SearchFilter {
283    /// Create a new empty search filter
284    pub fn new() -> Self {
285        Self {
286            max_distance: None,
287            min_distance: None,
288            metadata_filter: None,
289            dimension_constraints: None,
290        }
291    }
292
293    /// Set maximum distance threshold
294    pub fn with_max_distance(mut self, max_distance: f32) -> Self {
295        self.max_distance = Some(max_distance);
296        self
297    }
298
299    /// Set minimum distance threshold
300    pub fn with_min_distance(mut self, min_distance: f32) -> Self {
301        self.min_distance = Some(min_distance);
302        self
303    }
304
305    /// Set metadata filter
306    pub fn with_metadata_filter(mut self, filter: MetadataFilter) -> Self {
307        self.metadata_filter = Some(filter);
308        self
309    }
310
311    /// Set dimension constraints
312    pub fn with_dimension_constraints(mut self, constraints: Vec<DimensionConstraint>) -> Self {
313        self.dimension_constraints = Some(constraints);
314        self
315    }
316
317    /// Check if a search result satisfies this filter
318    pub fn satisfies(
319        &self,
320        distance: f32,
321        vector: &Vector,
322        metadata: &HashMap<String, String>,
323    ) -> bool {
324        // Check distance constraints
325        if let Some(max) = self.max_distance {
326            if distance > max {
327                return false;
328            }
329        }
330
331        if let Some(min) = self.min_distance {
332            if distance < min {
333                return false;
334            }
335        }
336
337        // Check metadata filter
338        if let Some(ref filter) = self.metadata_filter {
339            if !filter.evaluate(metadata) {
340                return false;
341            }
342        }
343
344        // Check dimension constraints
345        if let Some(ref constraints) = self.dimension_constraints {
346            for constraint in constraints {
347                if !constraint.satisfies(vector) {
348                    return false;
349                }
350            }
351        }
352
353        true
354    }
355
356    /// Filter a list of search results
357    pub fn filter_results(
358        &self,
359        results: Vec<(VectorId, f32, Vector, HashMap<String, String>)>,
360    ) -> Vec<(VectorId, f32)> {
361        results
362            .into_iter()
363            .filter(|(_, distance, vector, metadata)| self.satisfies(*distance, vector, metadata))
364            .map(|(id, distance, _, _)| (id, distance))
365            .collect()
366    }
367}
368
369impl Default for SearchFilter {
370    fn default() -> Self {
371        Self::new()
372    }
373}
374
375/// Builder for complex filter expressions
376pub struct FilterBuilder {
377    filters: Vec<MetadataFilter>,
378}
379
380impl FilterBuilder {
381    pub fn new() -> Self {
382        Self {
383            filters: Vec::new(),
384        }
385    }
386
387    pub fn equals(mut self, field: impl Into<String>, value: FilterValue) -> Self {
388        self.filters.push(MetadataFilter::Equals {
389            field: field.into(),
390            value,
391        });
392        self
393    }
394
395    pub fn not_equals(mut self, field: impl Into<String>, value: FilterValue) -> Self {
396        self.filters.push(MetadataFilter::NotEquals {
397            field: field.into(),
398            value,
399        });
400        self
401    }
402
403    pub fn greater_than(mut self, field: impl Into<String>, value: FilterValue) -> Self {
404        self.filters.push(MetadataFilter::GreaterThan {
405            field: field.into(),
406            value,
407        });
408        self
409    }
410
411    pub fn less_than(mut self, field: impl Into<String>, value: FilterValue) -> Self {
412        self.filters.push(MetadataFilter::LessThan {
413            field: field.into(),
414            value,
415        });
416        self
417    }
418
419    pub fn contains(mut self, field: impl Into<String>, substring: impl Into<String>) -> Self {
420        self.filters.push(MetadataFilter::Contains {
421            field: field.into(),
422            substring: substring.into(),
423        });
424        self
425    }
426
427    pub fn regex(mut self, field: impl Into<String>, pattern: impl Into<String>) -> Self {
428        self.filters.push(MetadataFilter::Regex {
429            field: field.into(),
430            pattern: pattern.into(),
431        });
432        self
433    }
434
435    pub fn exists(mut self, field: impl Into<String>) -> Self {
436        self.filters.push(MetadataFilter::Exists {
437            field: field.into(),
438        });
439        self
440    }
441
442    pub fn build_and(self) -> MetadataFilter {
443        if self.filters.len() == 1 {
444            self.filters
445                .into_iter()
446                .next()
447                .expect("filters validated to have exactly one element")
448        } else {
449            MetadataFilter::And(self.filters)
450        }
451    }
452
453    pub fn build_or(self) -> MetadataFilter {
454        if self.filters.len() == 1 {
455            self.filters
456                .into_iter()
457                .next()
458                .expect("filters validated to have exactly one element")
459        } else {
460            MetadataFilter::Or(self.filters)
461        }
462    }
463}
464
465impl Default for FilterBuilder {
466    fn default() -> Self {
467        Self::new()
468    }
469}
470
471#[cfg(test)]
472mod tests {
473    use super::*;
474
475    #[test]
476    fn test_equals_filter() {
477        let filter = MetadataFilter::Equals {
478            field: "category".to_string(),
479            value: FilterValue::String("news".to_string()),
480        };
481
482        let mut metadata = HashMap::new();
483        metadata.insert("category".to_string(), "news".to_string());
484
485        assert!(filter.evaluate(&metadata));
486
487        metadata.insert("category".to_string(), "sports".to_string());
488        assert!(!filter.evaluate(&metadata));
489    }
490
491    #[test]
492    fn test_greater_than_filter() {
493        let filter = MetadataFilter::GreaterThan {
494            field: "score".to_string(),
495            value: FilterValue::Integer(50),
496        };
497
498        let mut metadata = HashMap::new();
499        metadata.insert("score".to_string(), "75".to_string());
500        assert!(filter.evaluate(&metadata));
501
502        metadata.insert("score".to_string(), "25".to_string());
503        assert!(!filter.evaluate(&metadata));
504    }
505
506    #[test]
507    fn regression_cross_type_numeric_comparison() {
508        // Stored "75" auto-detects as Integer(75); the filter literal is a Float.
509        // Before the fix, compare() fell through to Ordering::Equal so the
510        // predicate silently returned false even though 75 > 50.
511        let mut metadata = HashMap::new();
512        metadata.insert("score".to_string(), "75".to_string());
513
514        let gt = MetadataFilter::GreaterThan {
515            field: "score".to_string(),
516            value: FilterValue::Float(50.0),
517        };
518        assert!(gt.evaluate(&metadata), "75 > 50.0 must be true");
519
520        let lt = MetadataFilter::LessThan {
521            field: "score".to_string(),
522            value: FilterValue::Float(50.0),
523        };
524        assert!(!lt.evaluate(&metadata), "75 < 50.0 must be false");
525
526        // Float stored value vs Integer literal, the other direction.
527        metadata.insert("score".to_string(), "12.5".to_string());
528        let ge = MetadataFilter::GreaterThanOrEqual {
529            field: "score".to_string(),
530            value: FilterValue::Integer(12),
531        };
532        assert!(ge.evaluate(&metadata), "12.5 >= 12 must be true");
533
534        let le = MetadataFilter::LessThanOrEqual {
535            field: "score".to_string(),
536            value: FilterValue::Integer(12),
537        };
538        assert!(!le.evaluate(&metadata), "12.5 <= 12 must be false");
539    }
540
541    #[test]
542    fn test_and_filter() {
543        let filter = MetadataFilter::And(vec![
544            MetadataFilter::Equals {
545                field: "status".to_string(),
546                value: FilterValue::String("active".to_string()),
547            },
548            MetadataFilter::GreaterThan {
549                field: "priority".to_string(),
550                value: FilterValue::Integer(5),
551            },
552        ]);
553
554        let mut metadata = HashMap::new();
555        metadata.insert("status".to_string(), "active".to_string());
556        metadata.insert("priority".to_string(), "8".to_string());
557        assert!(filter.evaluate(&metadata));
558
559        metadata.insert("priority".to_string(), "3".to_string());
560        assert!(!filter.evaluate(&metadata));
561    }
562
563    #[test]
564    fn test_or_filter() {
565        let filter = MetadataFilter::Or(vec![
566            MetadataFilter::Equals {
567                field: "type".to_string(),
568                value: FilterValue::String("urgent".to_string()),
569            },
570            MetadataFilter::Equals {
571                field: "type".to_string(),
572                value: FilterValue::String("critical".to_string()),
573            },
574        ]);
575
576        let mut metadata = HashMap::new();
577        metadata.insert("type".to_string(), "urgent".to_string());
578        assert!(filter.evaluate(&metadata));
579
580        metadata.insert("type".to_string(), "critical".to_string());
581        assert!(filter.evaluate(&metadata));
582
583        metadata.insert("type".to_string(), "normal".to_string());
584        assert!(!filter.evaluate(&metadata));
585    }
586
587    #[test]
588    fn test_contains_filter() {
589        let filter = MetadataFilter::Contains {
590            field: "description".to_string(),
591            substring: "important".to_string(),
592        };
593
594        let mut metadata = HashMap::new();
595        metadata.insert(
596            "description".to_string(),
597            "This is an important message".to_string(),
598        );
599        assert!(filter.evaluate(&metadata));
600
601        metadata.insert("description".to_string(), "Regular message".to_string());
602        assert!(!filter.evaluate(&metadata));
603    }
604
605    #[test]
606    fn test_filter_builder() {
607        let filter = FilterBuilder::new()
608            .equals("category", FilterValue::String("tech".to_string()))
609            .greater_than("score", FilterValue::Integer(70))
610            .build_and();
611
612        let mut metadata = HashMap::new();
613        metadata.insert("category".to_string(), "tech".to_string());
614        metadata.insert("score".to_string(), "85".to_string());
615        assert!(filter.evaluate(&metadata));
616    }
617
618    #[test]
619    fn test_dimension_constraint() {
620        let constraint = DimensionConstraint {
621            dimension: 0,
622            min_value: Some(0.0),
623            max_value: Some(1.0),
624        };
625
626        let vec1 = Vector::new(vec![0.5, 0.3, 0.7]);
627        assert!(constraint.satisfies(&vec1));
628
629        let vec2 = Vector::new(vec![1.5, 0.3, 0.7]);
630        assert!(!constraint.satisfies(&vec2));
631    }
632
633    #[test]
634    fn test_search_filter() {
635        let filter = SearchFilter::new()
636            .with_max_distance(0.5)
637            .with_metadata_filter(MetadataFilter::Equals {
638                field: "category".to_string(),
639                value: FilterValue::String("approved".to_string()),
640            });
641
642        let mut metadata = HashMap::new();
643        metadata.insert("category".to_string(), "approved".to_string());
644
645        let vector = Vector::new(vec![1.0, 2.0, 3.0]);
646
647        assert!(filter.satisfies(0.3, &vector, &metadata));
648        assert!(!filter.satisfies(0.7, &vector, &metadata)); // distance too high
649
650        metadata.insert("category".to_string(), "pending".to_string());
651        assert!(!filter.satisfies(0.3, &vector, &metadata)); // metadata doesn't match
652    }
653}