Skip to main content

a3s_vec/
query.rs

1//! Query payloads and per-index search controls.
2
3use crate::error::{Error, Result};
4use serde::{Deserialize, Serialize};
5use serde_json::{json, Map, Value};
6
7#[derive(Debug, Clone, Copy, PartialEq, Serialize, Deserialize)]
8pub struct HnswQueryParams {
9    pub ef: i32,
10    pub radius: f32,
11    pub is_linear: bool,
12    pub is_using_refiner: bool,
13}
14impl HnswQueryParams {
15    pub fn new(ef: i32, radius: f32, is_linear: bool, is_using_refiner: bool) -> Self {
16        Self {
17            ef,
18            radius,
19            is_linear,
20            is_using_refiner,
21        }
22    }
23    pub fn set_ef(&mut self, ef: i32) -> Result<()> {
24        if ef <= 0 {
25            return Err(Error::invalid_argument("HNSW ef must be positive"));
26        }
27        self.ef = ef;
28        Ok(())
29    }
30    pub fn ef(&self) -> i32 {
31        self.ef
32    }
33}
34
35#[derive(Debug, Clone, Copy, PartialEq, Serialize, Deserialize)]
36pub struct IvfQueryParams {
37    pub nprobe: i32,
38    pub is_using_refiner: bool,
39    pub scale_factor: f32,
40}
41impl IvfQueryParams {
42    pub fn new(nprobe: i32, is_using_refiner: bool, scale_factor: f32) -> Self {
43        Self {
44            nprobe,
45            is_using_refiner,
46            scale_factor,
47        }
48    }
49    pub fn set_nprobe(&mut self, nprobe: i32) -> Result<()> {
50        if nprobe <= 0 {
51            return Err(Error::invalid_argument("IVF nprobe must be positive"));
52        }
53        self.nprobe = nprobe;
54        Ok(())
55    }
56    pub fn nprobe(&self) -> i32 {
57        self.nprobe
58    }
59    pub fn set_scale_factor(&mut self, value: f32) -> Result<()> {
60        if !value.is_finite() || value <= 0.0 {
61            return Err(Error::invalid_argument(
62                "scale factor must be finite and positive",
63            ));
64        }
65        self.scale_factor = value;
66        Ok(())
67    }
68    pub fn scale_factor(&self) -> f32 {
69        self.scale_factor
70    }
71}
72
73#[derive(Debug, Clone, Copy, PartialEq, Serialize, Deserialize)]
74pub struct IvfRabitqQueryParams {
75    pub nprobe: i32,
76    pub radius: f32,
77    pub is_linear: bool,
78    pub is_using_refiner: bool,
79    pub scale_factor: f32,
80}
81impl IvfRabitqQueryParams {
82    pub fn new(nprobe: i32, radius: f32, is_linear: bool, is_using_refiner: bool) -> Self {
83        Self {
84            nprobe,
85            radius,
86            is_linear,
87            is_using_refiner,
88            scale_factor: 1.0,
89        }
90    }
91    pub fn set_nprobe(&mut self, nprobe: i32) -> Result<()> {
92        if nprobe <= 0 {
93            return Err(Error::invalid_argument(
94                "IVF RaBitQ nprobe must be positive",
95            ));
96        }
97        self.nprobe = nprobe;
98        Ok(())
99    }
100    pub fn nprobe(&self) -> i32 {
101        self.nprobe
102    }
103    pub fn set_scale_factor(&mut self, value: f32) -> Result<()> {
104        if !value.is_finite() || value <= 0.0 {
105            return Err(Error::invalid_argument(
106                "scale factor must be finite and positive",
107            ));
108        }
109        self.scale_factor = value;
110        Ok(())
111    }
112    pub fn scale_factor(&self) -> f32 {
113        self.scale_factor
114    }
115}
116
117#[derive(Debug, Clone, Copy, PartialEq, Serialize, Deserialize)]
118pub struct FlatQueryParams {
119    pub is_using_refiner: bool,
120    pub scale_factor: f32,
121}
122impl FlatQueryParams {
123    pub fn new(is_using_refiner: bool, scale_factor: f32) -> Self {
124        Self {
125            is_using_refiner,
126            scale_factor,
127        }
128    }
129}
130
131#[derive(Debug, Clone, Copy, PartialEq, Serialize, Deserialize)]
132pub struct DiskannQueryParams {
133    pub list_size: i32,
134}
135impl DiskannQueryParams {
136    pub fn new(list_size: i32) -> Self {
137        Self { list_size }
138    }
139    pub fn set_list_size(&mut self, value: i32) -> Result<()> {
140        if value <= 0 {
141            return Err(Error::invalid_argument(
142                "DiskANN list_size must be positive",
143            ));
144        }
145        self.list_size = value;
146        Ok(())
147    }
148    pub fn list_size(&self) -> i32 {
149        self.list_size
150    }
151}
152
153/// Full-text query controls.
154///
155/// Omitting `default_operator` preserves OR semantics. `AND` intersects all
156/// analyzed terms before BM25 scoring, which is useful for selective n-gram
157/// substring queries.
158#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
159pub struct FtsQueryParams {
160    pub default_operator: Option<String>,
161}
162
163#[derive(Debug, Clone, Copy, PartialEq, Eq)]
164pub(crate) enum FtsDefaultOperator {
165    Or,
166    And,
167}
168
169impl FtsDefaultOperator {
170    fn parse(value: &str) -> Result<Self> {
171        match value.to_ascii_lowercase().as_str() {
172            "or" => Ok(Self::Or),
173            "and" => Ok(Self::And),
174            _ => Err(Error::invalid_argument(
175                "FTS default operator must be AND or OR",
176            )),
177        }
178    }
179
180    fn as_str(self) -> &'static str {
181        match self {
182            Self::Or => "or",
183            Self::And => "and",
184        }
185    }
186}
187impl FtsQueryParams {
188    pub fn new(default_operator: Option<&str>) -> Result<Self> {
189        let mut out = Self {
190            default_operator: None,
191        };
192        if let Some(value) = default_operator {
193            out.set_default_operator(value)?;
194        }
195        Ok(out)
196    }
197    pub fn set_default_operator(&mut self, op: &str) -> Result<()> {
198        self.default_operator = Some(FtsDefaultOperator::parse(op)?.as_str().to_string());
199        Ok(())
200    }
201    pub fn default_operator(&self) -> Option<String> {
202        self.default_operator.clone()
203    }
204}
205
206#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
207pub struct Fts {
208    pub query_string: Option<String>,
209    pub match_string: Option<String>,
210}
211impl Fts {
212    pub fn new() -> Result<Self> {
213        Ok(Self {
214            query_string: None,
215            match_string: None,
216        })
217    }
218    pub fn set_query_string(&mut self, query: &str) -> Result<()> {
219        if query.trim().is_empty() {
220            return Err(Error::invalid_argument(
221                "FTS query string must not be empty",
222            ));
223        }
224        self.query_string = Some(query.to_string());
225        Ok(())
226    }
227    pub fn set_match_string(&mut self, query: &str) -> Result<()> {
228        if query.trim().is_empty() {
229            return Err(Error::invalid_argument(
230                "FTS match string must not be empty",
231            ));
232        }
233        self.match_string = Some(query.to_string());
234        Ok(())
235    }
236    pub fn query_string(&self) -> Option<String> {
237        self.query_string.clone()
238    }
239    pub fn match_string(&self) -> Option<String> {
240        self.match_string.clone()
241    }
242}
243
244/// A single dense, sparse, binary, id-based, or FTS query.
245#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
246pub struct SearchQuery {
247    pub field_name: String,
248    pub vector: Option<Vec<f32>>,
249    #[serde(default)]
250    pub binary_vector: Option<Vec<u8>>,
251    pub sparse_vector: Option<Vec<(u32, f32)>>,
252    pub id: Option<String>,
253    pub topk: i32,
254    pub filter: Option<String>,
255    pub include_vector: bool,
256    pub include_doc_id: bool,
257    pub output_fields: Option<Vec<String>>,
258    pub fts: Option<Fts>,
259    #[serde(default)]
260    pub params: Map<String, Value>,
261}
262
263/// Compatibility alias used by the zvec vocabulary.
264pub type VectorQuery = SearchQuery;
265
266impl SearchQuery {
267    pub fn new(field_name: &str, vector: &[f32], topk: i32) -> Result<Self> {
268        validate_query_header(field_name, topk)?;
269        if vector.is_empty() || !vector.iter().all(|v| v.is_finite()) {
270            return Err(Error::invalid_argument(
271                "query vector must be non-empty and finite",
272            ));
273        }
274        Ok(Self {
275            field_name: field_name.to_string(),
276            vector: Some(vector.to_vec()),
277            binary_vector: None,
278            sparse_vector: None,
279            id: None,
280            topk,
281            filter: None,
282            include_vector: false,
283            include_doc_id: false,
284            output_fields: None,
285            fts: None,
286            params: Map::new(),
287        })
288    }
289    pub fn fts(field_name: &str, fts: &Fts, topk: i32) -> Result<Self> {
290        validate_query_header(field_name, topk)?;
291        if fts.query_string.is_none() && fts.match_string.is_none() {
292            return Err(Error::invalid_argument("FTS query has no expression"));
293        }
294        Ok(Self {
295            field_name: field_name.to_string(),
296            vector: None,
297            binary_vector: None,
298            sparse_vector: None,
299            id: None,
300            topk,
301            filter: None,
302            include_vector: false,
303            include_doc_id: false,
304            output_fields: None,
305            fts: Some(fts.clone()),
306            params: Map::new(),
307        })
308    }
309    pub fn by_id(field_name: &str, id: &str, topk: i32) -> Result<Self> {
310        validate_query_header(field_name, topk)?;
311        if id.is_empty() || id.contains('\0') {
312            return Err(Error::invalid_argument(
313                "query id must be non-empty and contain no NUL byte",
314            ));
315        }
316        Ok(Self {
317            field_name: field_name.to_string(),
318            vector: None,
319            binary_vector: None,
320            sparse_vector: None,
321            id: Some(id.to_string()),
322            topk,
323            filter: None,
324            include_vector: false,
325            include_doc_id: false,
326            output_fields: None,
327            fts: None,
328            params: Map::new(),
329        })
330    }
331    pub fn sparse(field_name: &str, indices: &[u32], values: &[f32], topk: i32) -> Result<Self> {
332        validate_query_header(field_name, topk)?;
333        if indices.is_empty()
334            || indices.len() != values.len()
335            || !values.iter().all(|v| v.is_finite())
336        {
337            return Err(Error::invalid_argument(
338                "sparse query indices and values must have equal non-zero length",
339            ));
340        }
341        Ok(Self {
342            field_name: field_name.to_string(),
343            vector: None,
344            binary_vector: None,
345            sparse_vector: Some(
346                indices
347                    .iter()
348                    .copied()
349                    .zip(values.iter().copied())
350                    .collect(),
351            ),
352            id: None,
353            topk,
354            filter: None,
355            include_vector: false,
356            include_doc_id: false,
357            output_fields: None,
358            fts: None,
359            params: Map::new(),
360        })
361    }
362    /// Creates a packed Binary32 or Binary64 exact query.
363    ///
364    /// The collection schema determines the binary type and validates the byte
365    /// length. Binary search uses L2 over bit coordinates, whose squared
366    /// distance is the XOR Hamming count.
367    pub fn binary(field_name: &str, vector: &[u8], topk: i32) -> Result<Self> {
368        validate_query_header(field_name, topk)?;
369        if vector.is_empty() {
370            return Err(Error::invalid_argument(
371                "binary query vector must be non-empty",
372            ));
373        }
374        Ok(Self {
375            field_name: field_name.to_string(),
376            vector: None,
377            binary_vector: Some(vector.to_vec()),
378            sparse_vector: None,
379            id: None,
380            topk,
381            filter: None,
382            include_vector: false,
383            include_doc_id: false,
384            output_fields: None,
385            fts: None,
386            params: Map::new(),
387        })
388    }
389    pub fn builder() -> SearchQueryBuilder {
390        SearchQueryBuilder::new()
391    }
392    pub fn set_field_name(&mut self, name: &str) -> Result<()> {
393        validate_name(name)?;
394        self.field_name = name.to_string();
395        Ok(())
396    }
397    pub fn set_query_vector(&mut self, vector: &[f32]) -> Result<()> {
398        if vector.is_empty() || !vector.iter().all(|v| v.is_finite()) {
399            return Err(Error::invalid_argument(
400                "query vector must be non-empty and finite",
401            ));
402        }
403        self.vector = Some(vector.to_vec());
404        self.binary_vector = None;
405        self.sparse_vector = None;
406        self.id = None;
407        self.fts = None;
408        Ok(())
409    }
410    pub fn set_binary_vector(&mut self, vector: &[u8]) -> Result<()> {
411        if vector.is_empty() {
412            return Err(Error::invalid_argument(
413                "binary query vector must be non-empty",
414            ));
415        }
416        self.binary_vector = Some(vector.to_vec());
417        self.vector = None;
418        self.sparse_vector = None;
419        self.id = None;
420        self.fts = None;
421        Ok(())
422    }
423    pub fn set_sparse_vector(&mut self, indices: &[u32], values: &[f32]) -> Result<()> {
424        if indices.is_empty()
425            || indices.len() != values.len()
426            || !values.iter().all(|v| v.is_finite())
427        {
428            return Err(Error::invalid_argument(
429                "sparse query indices and values are invalid",
430            ));
431        }
432        self.sparse_vector = Some(
433            indices
434                .iter()
435                .copied()
436                .zip(values.iter().copied())
437                .collect(),
438        );
439        self.vector = None;
440        self.binary_vector = None;
441        self.id = None;
442        self.fts = None;
443        Ok(())
444    }
445    pub fn set_filter(&mut self, filter: &str) -> Result<()> {
446        if filter.trim().is_empty() {
447            self.filter = None;
448            return Ok(());
449        }
450        self.filter = Some(filter.to_string());
451        Ok(())
452    }
453    pub fn set_topk(&mut self, topk: i32) -> Result<()> {
454        if topk <= 0 {
455            return Err(Error::invalid_argument("topk must be positive"));
456        }
457        self.topk = topk;
458        Ok(())
459    }
460    pub fn set_include_vector(&mut self, include: bool) -> Result<()> {
461        self.include_vector = include;
462        Ok(())
463    }
464    pub fn set_include_doc_id(&mut self, include: bool) -> Result<()> {
465        self.include_doc_id = include;
466        Ok(())
467    }
468    pub fn set_output_fields(&mut self, fields: &[&str]) -> Result<()> {
469        for field in fields {
470            validate_name(field)?;
471        }
472        self.output_fields = Some(fields.iter().map(|s| (*s).to_string()).collect());
473        Ok(())
474    }
475    pub fn set_hnsw_params(&mut self, params: HnswQueryParams) -> Result<()> {
476        apply_hnsw_query_controls(&mut self.params, params)
477    }
478    pub fn set_ivf_params(&mut self, params: IvfQueryParams) -> Result<()> {
479        apply_ivf_query_controls(&mut self.params, params)
480    }
481    pub fn set_ivf_rabitq_params(&mut self, params: IvfRabitqQueryParams) -> Result<()> {
482        apply_ivf_rabitq_query_controls(&mut self.params, params)
483    }
484    pub fn set_flat_params(&mut self, _params: FlatQueryParams) -> Result<()> {
485        unsupported_query_controls("Flat refinement")
486    }
487    pub fn set_diskann_params(&mut self, params: DiskannQueryParams) -> Result<()> {
488        apply_diskann_query_controls(&mut self.params, params)
489    }
490    pub fn set_fts_params(&mut self, params: FtsQueryParams) -> Result<()> {
491        apply_fts_query_controls(&mut self.params, params)
492    }
493    pub fn set_fts(&mut self, fts: &Fts) -> Result<()> {
494        if fts.query_string.is_none() && fts.match_string.is_none() {
495            return Err(Error::invalid_argument("FTS query has no expression"));
496        }
497        self.fts = Some(fts.clone());
498        self.vector = None;
499        self.binary_vector = None;
500        self.sparse_vector = None;
501        self.id = None;
502        Ok(())
503    }
504    pub fn set_radius(&mut self, radius: f32) -> Result<()> {
505        if !radius.is_finite() {
506            return Err(Error::invalid_argument("radius must be finite"));
507        }
508        self.params.insert("radius".into(), json!(radius));
509        Ok(())
510    }
511    pub fn get_filter(&self) -> Option<&str> {
512        self.filter.as_deref()
513    }
514    pub fn has_vector(&self) -> bool {
515        self.vector.is_some() || self.binary_vector.is_some() || self.sparse_vector.is_some()
516    }
517}
518
519/// Fluent builder matching the official SDK's naming.
520///
521/// A builder can construct a dense-vector, binary-vector, or pure FTS query.
522/// The routes are mutually exclusive, as are the FTS `query_string` and
523/// `match_string` forms.
524#[derive(Debug, Clone, Default)]
525pub struct SearchQueryBuilder {
526    field_name: Option<String>,
527    vector: Option<Vec<f32>>,
528    binary_vector: Option<Vec<u8>>,
529    topk: i32,
530    filter: Option<String>,
531    include_vector: Option<bool>,
532    include_doc_id: Option<bool>,
533    output_fields: Option<Vec<String>>,
534    fts_query_string: Option<String>,
535    fts_match_string: Option<String>,
536}
537impl SearchQueryBuilder {
538    pub fn new() -> Self {
539        Self {
540            topk: 10,
541            ..Self::default()
542        }
543    }
544    pub fn field_name(mut self, name: &str) -> Self {
545        self.field_name = Some(name.to_string());
546        self
547    }
548    pub fn vector(mut self, vector: &[f32]) -> Self {
549        self.vector = Some(vector.to_vec());
550        self
551    }
552    pub fn binary_vector(mut self, vector: &[u8]) -> Self {
553        self.binary_vector = Some(vector.to_vec());
554        self
555    }
556    pub fn topk(mut self, topk: i32) -> Self {
557        self.topk = topk;
558        self
559    }
560    pub fn filter(mut self, filter: &str) -> Self {
561        self.filter = Some(filter.to_string());
562        self
563    }
564    pub fn include_vector(mut self, include: bool) -> Self {
565        self.include_vector = Some(include);
566        self
567    }
568    pub fn include_doc_id(mut self, include: bool) -> Self {
569        self.include_doc_id = Some(include);
570        self
571    }
572    pub fn output_fields(mut self, fields: &[&str]) -> Self {
573        self.output_fields = Some(fields.iter().map(|v| (*v).to_string()).collect());
574        self
575    }
576    pub fn fts_query_string(mut self, query: &str) -> Self {
577        self.fts_query_string = Some(query.to_string());
578        self
579    }
580    pub fn fts_match_string(mut self, query: &str) -> Self {
581        self.fts_match_string = Some(query.to_string());
582        self
583    }
584    /// Validates the selected route and creates the immutable query payload.
585    pub fn build(self) -> Result<SearchQuery> {
586        let field = self
587            .field_name
588            .ok_or_else(|| Error::invalid_argument("field_name is required"))?;
589        let has_query_string = self.fts_query_string.is_some();
590        let has_match_string = self.fts_match_string.is_some();
591        let route_count = usize::from(self.vector.is_some())
592            + usize::from(self.binary_vector.is_some())
593            + usize::from(has_query_string || has_match_string);
594        if route_count > 1 {
595            return Err(Error::invalid_argument(
596                "query builder cannot combine dense, binary, and FTS routes",
597            ));
598        }
599        if has_query_string && has_match_string {
600            return Err(Error::invalid_argument(
601                "query builder cannot combine FTS query_string and match_string",
602            ));
603        }
604
605        let mut query = if has_query_string || has_match_string {
606            let mut fts = Fts::new()?;
607            if let Some(value) = self.fts_query_string {
608                fts.set_query_string(&value)?;
609            }
610            if let Some(value) = self.fts_match_string {
611                fts.set_match_string(&value)?;
612            }
613            SearchQuery::fts(&field, &fts, self.topk)?
614        } else if let Some(vector) = self.vector {
615            SearchQuery::new(&field, &vector, self.topk)?
616        } else if let Some(vector) = self.binary_vector {
617            SearchQuery::binary(&field, &vector, self.topk)?
618        } else {
619            return Err(Error::invalid_argument(
620                "dense vector, binary vector, or FTS expression is required",
621            ));
622        };
623        if let Some(filter) = self.filter {
624            query.set_filter(&filter)?;
625        }
626        if let Some(value) = self.include_vector {
627            query.set_include_vector(value)?;
628        }
629        if let Some(value) = self.include_doc_id {
630            query.set_include_doc_id(value)?;
631        }
632        if let Some(fields) = self.output_fields {
633            let refs: Vec<&str> = fields.iter().map(String::as_str).collect();
634            query.set_output_fields(&refs)?;
635        }
636        Ok(query)
637    }
638}
639
640/// Group-by vector query payload.
641#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
642pub struct GroupBySearchQuery {
643    pub field_name: String,
644    pub group_by_field: String,
645    pub vector: Vec<f32>,
646    #[serde(default)]
647    pub binary_vector: Option<Vec<u8>>,
648    pub group_count: u32,
649    pub group_topk: u32,
650    pub filter: Option<String>,
651    pub include_vector: bool,
652    pub output_fields: Option<Vec<String>>,
653    pub params: Map<String, Value>,
654}
655impl GroupBySearchQuery {
656    pub fn new(
657        field_name: &str,
658        group_by_field: &str,
659        vector: &[f32],
660        group_count: u32,
661        group_topk: u32,
662    ) -> Result<Self> {
663        validate_name(field_name)?;
664        validate_name(group_by_field)?;
665        if vector.is_empty()
666            || !vector.iter().all(|v| v.is_finite())
667            || group_count == 0
668            || group_topk == 0
669        {
670            return Err(Error::invalid_argument("invalid group-by query parameters"));
671        }
672        Ok(Self {
673            field_name: field_name.to_string(),
674            group_by_field: group_by_field.to_string(),
675            vector: vector.to_vec(),
676            binary_vector: None,
677            group_count,
678            group_topk,
679            filter: None,
680            include_vector: false,
681            output_fields: None,
682            params: Map::new(),
683        })
684    }
685    /// Creates a grouped packed-binary exact query.
686    pub fn binary(
687        field_name: &str,
688        group_by_field: &str,
689        vector: &[u8],
690        group_count: u32,
691        group_topk: u32,
692    ) -> Result<Self> {
693        validate_name(field_name)?;
694        validate_name(group_by_field)?;
695        if vector.is_empty() || group_count == 0 || group_topk == 0 {
696            return Err(Error::invalid_argument("invalid group-by query parameters"));
697        }
698        Ok(Self {
699            field_name: field_name.to_string(),
700            group_by_field: group_by_field.to_string(),
701            vector: Vec::new(),
702            binary_vector: Some(vector.to_vec()),
703            group_count,
704            group_topk,
705            filter: None,
706            include_vector: false,
707            output_fields: None,
708            params: Map::new(),
709        })
710    }
711    pub fn set_filter(&mut self, filter: &str) -> Result<()> {
712        self.filter = (!filter.trim().is_empty()).then_some(filter.to_string());
713        Ok(())
714    }
715    pub fn set_include_vector(&mut self, include: bool) -> Result<()> {
716        self.include_vector = include;
717        Ok(())
718    }
719    pub fn set_output_fields(&mut self, fields: &[&str]) -> Result<()> {
720        for field in fields {
721            validate_name(field)?;
722        }
723        self.output_fields = Some(fields.iter().map(|s| (*s).to_string()).collect());
724        Ok(())
725    }
726    pub fn set_hnsw_params(&mut self, params: HnswQueryParams) -> Result<()> {
727        apply_hnsw_query_controls(&mut self.params, params)
728    }
729    pub fn set_ivf_params(&mut self, params: IvfQueryParams) -> Result<()> {
730        apply_ivf_query_controls(&mut self.params, params)
731    }
732    pub fn set_ivf_rabitq_params(&mut self, params: IvfRabitqQueryParams) -> Result<()> {
733        apply_ivf_rabitq_query_controls(&mut self.params, params)
734    }
735    pub fn set_flat_params(&mut self, _params: FlatQueryParams) -> Result<()> {
736        unsupported_query_controls("Flat refinement")
737    }
738    pub fn set_diskann_params(&mut self, params: DiskannQueryParams) -> Result<()> {
739        apply_diskann_query_controls(&mut self.params, params)
740    }
741}
742
743pub(crate) fn unsupported_query_controls(name: &str) -> Result<()> {
744    Err(Error::not_supported(format!(
745        "{name} query controls have no execution consumer"
746    )))
747}
748
749pub(crate) fn apply_hnsw_query_controls(
750    target: &mut Map<String, Value>,
751    params: HnswQueryParams,
752) -> Result<()> {
753    if params.ef <= 0 {
754        return Err(Error::invalid_argument("HNSW ef must be positive"));
755    }
756    if !params.radius.is_finite() {
757        return Err(Error::invalid_argument("HNSW radius must be finite"));
758    }
759    clear_ann_query_controls(target);
760    target.insert("type".into(), json!("hnsw"));
761    target.insert("ef".into(), json!(params.ef));
762    target.insert("is_linear".into(), json!(params.is_linear));
763    target.insert("is_using_refiner".into(), json!(params.is_using_refiner));
764    if params.radius == 0.0 {
765        target.remove("radius");
766    } else {
767        target.insert("radius".into(), json!(params.radius));
768    }
769    Ok(())
770}
771
772pub(crate) fn apply_ivf_query_controls(
773    target: &mut Map<String, Value>,
774    params: IvfQueryParams,
775) -> Result<()> {
776    if params.nprobe <= 0 {
777        return Err(Error::invalid_argument("IVF nprobe must be positive"));
778    }
779    if !params.scale_factor.is_finite() || params.scale_factor <= 0.0 {
780        return Err(Error::invalid_argument(
781            "IVF scale factor must be finite and positive",
782        ));
783    }
784    clear_ann_query_controls(target);
785    target.insert("type".into(), json!("ivf"));
786    target.insert("nprobe".into(), json!(params.nprobe));
787    target.insert("is_using_refiner".into(), json!(params.is_using_refiner));
788    target.insert("scale_factor".into(), json!(params.scale_factor));
789    Ok(())
790}
791
792pub(crate) fn apply_ivf_rabitq_query_controls(
793    target: &mut Map<String, Value>,
794    params: IvfRabitqQueryParams,
795) -> Result<()> {
796    if params.nprobe <= 0 {
797        return Err(Error::invalid_argument(
798            "IVF RaBitQ nprobe must be positive",
799        ));
800    }
801    if !params.radius.is_finite() {
802        return Err(Error::invalid_argument("IVF RaBitQ radius must be finite"));
803    }
804    if !params.scale_factor.is_finite() || params.scale_factor <= 0.0 {
805        return Err(Error::invalid_argument(
806            "IVF RaBitQ scale factor must be finite and positive",
807        ));
808    }
809    clear_ann_query_controls(target);
810    target.insert("type".into(), json!("ivf_rabitq"));
811    target.insert("nprobe".into(), json!(params.nprobe));
812    target.insert("is_linear".into(), json!(params.is_linear));
813    target.insert("is_using_refiner".into(), json!(params.is_using_refiner));
814    target.insert("scale_factor".into(), json!(params.scale_factor));
815    if params.radius == 0.0 {
816        target.remove("radius");
817    } else {
818        target.insert("radius".into(), json!(params.radius));
819    }
820    Ok(())
821}
822
823pub(crate) fn apply_diskann_query_controls(
824    target: &mut Map<String, Value>,
825    params: DiskannQueryParams,
826) -> Result<()> {
827    if params.list_size <= 0 {
828        return Err(Error::invalid_argument(
829            "DiskANN list_size must be positive",
830        ));
831    }
832    clear_ann_query_controls(target);
833    target.insert("type".into(), json!("diskann"));
834    target.insert("list_size".into(), json!(params.list_size));
835    Ok(())
836}
837
838fn clear_ann_query_controls(target: &mut Map<String, Value>) {
839    for name in [
840        "type",
841        "ef",
842        "nprobe",
843        "is_linear",
844        "is_using_refiner",
845        "scale_factor",
846        "list_size",
847    ] {
848        target.remove(name);
849    }
850}
851
852pub(crate) fn apply_fts_query_controls(
853    target: &mut Map<String, Value>,
854    params: FtsQueryParams,
855) -> Result<()> {
856    if let Some(value) = params.default_operator {
857        let operator = FtsDefaultOperator::parse(&value)?;
858        target.insert("default_operator".into(), json!(operator.as_str()));
859    } else {
860        target.remove("default_operator");
861    }
862    Ok(())
863}
864
865pub(crate) fn fts_default_operator(query: &SearchQuery) -> Result<FtsDefaultOperator> {
866    query.params.get("default_operator").map_or_else(
867        || Ok(FtsDefaultOperator::Or),
868        |value| {
869            value
870                .as_str()
871                .ok_or_else(|| {
872                    Error::invalid_argument("FTS default_operator parameter must be a string")
873                })
874                .and_then(FtsDefaultOperator::parse)
875        },
876    )
877}
878
879fn validate_name(name: &str) -> Result<()> {
880    if name.trim().is_empty() || name.contains('\0') {
881        Err(Error::invalid_argument(
882            "field name must be non-empty and contain no NUL byte",
883        ))
884    } else {
885        Ok(())
886    }
887}
888fn validate_query_header(name: &str, topk: i32) -> Result<()> {
889    validate_name(name)?;
890    if topk <= 0 {
891        return Err(Error::invalid_argument("topk must be positive"));
892    }
893    Ok(())
894}
895
896#[cfg(test)]
897mod tests {
898    use super::*;
899    use serde_json::Map;
900
901    #[test]
902    fn ann_query_controls_reject_non_positive_and_non_finite_values() {
903        let mut params = Map::new();
904        assert!(
905            apply_hnsw_query_controls(&mut params, HnswQueryParams::new(0, 0.0, false, false))
906                .is_err()
907        );
908        assert!(apply_hnsw_query_controls(
909            &mut params,
910            HnswQueryParams::new(8, f32::NAN, false, false)
911        )
912        .is_err());
913        assert!(
914            apply_hnsw_query_controls(&mut params, HnswQueryParams::new(8, 1.0, false, false))
915                .is_ok()
916        );
917        assert!(params.get("radius").is_some());
918        assert!(
919            apply_hnsw_query_controls(&mut params, HnswQueryParams::new(8, 0.0, true, true))
920                .is_ok()
921        );
922        assert!(params.get("radius").is_none());
923
924        assert!(apply_ivf_query_controls(&mut params, IvfQueryParams::new(0, false, 1.0)).is_err());
925        assert!(apply_ivf_query_controls(
926            &mut params,
927            IvfQueryParams::new(4, false, f32::INFINITY)
928        )
929        .is_err());
930        assert!(apply_ivf_query_controls(&mut params, IvfQueryParams::new(4, true, 2.0)).is_ok());
931
932        assert!(apply_ivf_rabitq_query_controls(
933            &mut params,
934            IvfRabitqQueryParams::new(0, 0.0, false, false)
935        )
936        .is_err());
937        assert!(apply_ivf_rabitq_query_controls(
938            &mut params,
939            IvfRabitqQueryParams::new(2, f32::NAN, false, false)
940        )
941        .is_err());
942        let mut bad_scale = IvfRabitqQueryParams::new(2, 0.0, false, false);
943        bad_scale.scale_factor = 0.0;
944        assert!(apply_ivf_rabitq_query_controls(&mut params, bad_scale).is_err());
945        let mut good = IvfRabitqQueryParams::new(2, 1.5, true, true);
946        good.scale_factor = 3.0;
947        assert!(apply_ivf_rabitq_query_controls(&mut params, good).is_ok());
948        assert!(params.get("radius").is_some());
949        assert!(apply_ivf_rabitq_query_controls(
950            &mut params,
951            IvfRabitqQueryParams::new(2, 0.0, false, false)
952        )
953        .is_ok());
954        assert!(params.get("radius").is_none());
955
956        assert!(apply_diskann_query_controls(&mut params, DiskannQueryParams::new(0)).is_err());
957        assert!(apply_diskann_query_controls(&mut params, DiskannQueryParams::new(64)).is_ok());
958
959        assert!(unsupported_query_controls("Flat refinement").is_err());
960        let mut query = SearchQuery::new("embedding", &[1.0, 0.0], 8).expect("query");
961        assert!(query
962            .set_flat_params(FlatQueryParams::new(false, 1.0))
963            .is_err());
964        assert!(HnswQueryParams::new(8, 0.0, false, false)
965            .set_ef(0)
966            .is_err());
967        assert!(IvfQueryParams::new(4, false, 1.0).set_nprobe(0).is_err());
968        assert!(IvfQueryParams::new(4, false, 1.0)
969            .set_scale_factor(-1.0)
970            .is_err());
971        assert!(DiskannQueryParams::new(8).set_list_size(0).is_err());
972    }
973
974    #[test]
975    fn fts_and_group_by_query_builders_cover_validation_edges() {
976        assert!(FtsQueryParams::new(Some("xor")).is_err());
977        assert!(FtsQueryParams::new(Some("AND")).is_ok());
978        assert!(FtsQueryParams::new(Some("nope")).is_err());
979        let params = FtsQueryParams::new(Some("OR")).expect("or");
980        let mut map = Map::new();
981        apply_fts_query_controls(&mut map, params).expect("apply");
982        assert_eq!(
983            map.get("default_operator").and_then(|v| v.as_str()),
984            Some("or")
985        );
986        apply_fts_query_controls(&mut map, FtsQueryParams::new(None).expect("none"))
987            .expect("clear");
988        assert!(!map.contains_key("default_operator"));
989
990        assert!(SearchQuery::new("", &[1.0], 1).is_err());
991        assert!(SearchQuery::new("embedding", &[1.0], 0).is_err());
992        assert!(SearchQuery::binary("embedding", &[], 8).is_err());
993        assert!(SearchQuery::sparse("embedding", &[0], &[1.0, 2.0], 8).is_err());
994        let mut fts = Fts::new().expect("fts");
995        assert!(SearchQuery::fts("body", &fts, 8).is_err());
996        fts.set_query_string("rust").expect("q");
997        let mut query = SearchQuery::fts("body", &fts, 8).expect("fts query");
998        query
999            .set_fts_params(FtsQueryParams::new(Some("AND")).expect("and"))
1000            .expect("set");
1001        query.set_filter("").expect("empty filter");
1002        query.set_include_vector(true).expect("include");
1003        assert!(query.set_output_fields(&["", "body"]).is_err());
1004        query.set_output_fields(&["body"]).expect("fields");
1005    }
1006}