1use crate::{Vector, VectorId};
8use serde::{Deserialize, Serialize};
9use std::collections::HashMap;
10
11#[derive(Debug, Clone, Serialize, Deserialize)]
13pub enum MetadataFilter {
14 Equals { field: String, value: FilterValue },
16 NotEquals { field: String, value: FilterValue },
18 GreaterThan { field: String, value: FilterValue },
20 GreaterThanOrEqual { field: String, value: FilterValue },
22 LessThan { field: String, value: FilterValue },
24 LessThanOrEqual { field: String, value: FilterValue },
26 In {
28 field: String,
29 values: Vec<FilterValue>,
30 },
31 NotIn {
33 field: String,
34 values: Vec<FilterValue>,
35 },
36 Contains { field: String, substring: String },
38 Regex { field: String, pattern: String },
40 Exists { field: String },
42 NotExists { field: String },
44 And(Vec<MetadataFilter>),
46 Or(Vec<MetadataFilter>),
48 Not(Box<MetadataFilter>),
50}
51
52#[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 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 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 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 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 fn parse_value(s: &str) -> FilterValue {
206 if let Ok(i) = s.parse::<i64>() {
208 return FilterValue::Integer(i);
209 }
210
211 if let Ok(f) = s.parse::<f64>() {
213 return FilterValue::Float(f);
214 }
215
216 if let Ok(b) = s.parse::<bool>() {
218 return FilterValue::Boolean(b);
219 }
220
221 if s == "null" || s.is_empty() {
223 return FilterValue::Null;
224 }
225
226 FilterValue::String(s.to_string())
228 }
229}
230
231#[derive(Debug, Clone, Serialize, Deserialize)]
233pub struct SearchFilter {
234 pub max_distance: Option<f32>,
236 pub min_distance: Option<f32>,
238 pub metadata_filter: Option<MetadataFilter>,
240 pub dimension_constraints: Option<Vec<DimensionConstraint>>,
242}
243
244#[derive(Debug, Clone, Serialize, Deserialize)]
246pub struct DimensionConstraint {
247 pub dimension: usize,
249 pub min_value: Option<f32>,
251 pub max_value: Option<f32>,
253}
254
255impl DimensionConstraint {
256 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 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 pub fn with_max_distance(mut self, max_distance: f32) -> Self {
295 self.max_distance = Some(max_distance);
296 self
297 }
298
299 pub fn with_min_distance(mut self, min_distance: f32) -> Self {
301 self.min_distance = Some(min_distance);
302 self
303 }
304
305 pub fn with_metadata_filter(mut self, filter: MetadataFilter) -> Self {
307 self.metadata_filter = Some(filter);
308 self
309 }
310
311 pub fn with_dimension_constraints(mut self, constraints: Vec<DimensionConstraint>) -> Self {
313 self.dimension_constraints = Some(constraints);
314 self
315 }
316
317 pub fn satisfies(
319 &self,
320 distance: f32,
321 vector: &Vector,
322 metadata: &HashMap<String, String>,
323 ) -> bool {
324 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 if let Some(ref filter) = self.metadata_filter {
339 if !filter.evaluate(metadata) {
340 return false;
341 }
342 }
343
344 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 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
375pub 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 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 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)); metadata.insert("category".to_string(), "pending".to_string());
651 assert!(!filter.satisfies(0.3, &vector, &metadata)); }
653}