Skip to main content

a3s_vec/
multi_query.rs

1//! Multi-route retrieval and deterministic reranking.
2
3use crate::error::{Error, Result};
4use crate::query::{
5    apply_diskann_query_controls, apply_fts_query_controls, apply_hnsw_query_controls,
6    apply_ivf_query_controls, apply_ivf_rabitq_query_controls, unsupported_query_controls,
7    DiskannQueryParams, FlatQueryParams, Fts, FtsQueryParams, HnswQueryParams, IvfQueryParams,
8    IvfRabitqQueryParams, SearchQuery,
9};
10use serde::{Deserialize, Serialize};
11
12/// Reranking strategy used to fuse branch results.
13#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
14pub enum RerankMethod {
15    ReciprocalRank { rank_constant: f64 },
16    Weighted { weights: Vec<f64> },
17}
18
19impl Default for RerankMethod {
20    fn default() -> Self {
21        Self::Weighted {
22            weights: Vec::new(),
23        }
24    }
25}
26
27/// A branch in a [`MultiQuery`].
28#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
29pub struct SubQuery {
30    pub field_name: Option<String>,
31    pub vector: Option<Vec<f32>>,
32    #[serde(default)]
33    pub binary_vector: Option<Vec<u8>>,
34    pub sparse_vector: Option<Vec<(u32, f32)>>,
35    pub fts: Option<Fts>,
36    pub num_candidates: i32,
37    pub params: serde_json::Map<String, serde_json::Value>,
38}
39
40impl SubQuery {
41    pub fn new() -> Result<Self> {
42        Ok(Self {
43            field_name: None,
44            vector: None,
45            binary_vector: None,
46            sparse_vector: None,
47            fts: None,
48            num_candidates: 10,
49            params: serde_json::Map::new(),
50        })
51    }
52    pub fn set_num_candidates(&mut self, n: i32) -> Result<()> {
53        if n <= 0 {
54            return Err(Error::invalid_argument("num_candidates must be positive"));
55        }
56        self.num_candidates = n;
57        Ok(())
58    }
59    pub fn num_candidates(&self) -> i32 {
60        self.num_candidates
61    }
62    pub fn set_field_name(&mut self, name: &str) -> Result<()> {
63        if name.trim().is_empty() || name.contains('\0') {
64            return Err(Error::invalid_argument("field name must be non-empty"));
65        }
66        self.field_name = Some(name.to_string());
67        Ok(())
68    }
69    pub fn set_query_vector(&mut self, data: &[f32]) -> Result<()> {
70        if data.is_empty() || !data.iter().all(|v| v.is_finite()) {
71            return Err(Error::invalid_argument(
72                "query vector must be non-empty and finite",
73            ));
74        }
75        self.vector = Some(data.to_vec());
76        self.binary_vector = None;
77        self.sparse_vector = None;
78        self.fts = None;
79        Ok(())
80    }
81    pub fn set_binary_vector(&mut self, data: &[u8]) -> Result<()> {
82        if data.is_empty() {
83            return Err(Error::invalid_argument(
84                "binary query vector must be non-empty",
85            ));
86        }
87        self.binary_vector = Some(data.to_vec());
88        self.vector = None;
89        self.sparse_vector = None;
90        self.fts = None;
91        Ok(())
92    }
93    pub fn set_sparse_vector(&mut self, indices: &[u32], values: &[f32]) -> Result<()> {
94        if indices.is_empty()
95            || indices.len() != values.len()
96            || !values.iter().all(|v| v.is_finite())
97        {
98            return Err(Error::invalid_argument("invalid sparse vector"));
99        }
100        self.sparse_vector = Some(
101            indices
102                .iter()
103                .copied()
104                .zip(values.iter().copied())
105                .collect(),
106        );
107        self.vector = None;
108        self.binary_vector = None;
109        self.fts = None;
110        Ok(())
111    }
112    pub fn set_sparse_indices(&mut self, indices: &[u32]) -> Result<()> {
113        if indices.is_empty() {
114            return Err(Error::invalid_argument("sparse indices cannot be empty"));
115        }
116        let values = self.sparse_vector.as_ref().map_or_else(
117            || vec![0.0; indices.len()],
118            |v| v.iter().map(|(_, x)| *x).collect(),
119        );
120        if values.len() != indices.len() {
121            return Err(Error::invalid_argument(
122                "sparse indices and values have different lengths",
123            ));
124        }
125        self.sparse_vector = Some(indices.iter().copied().zip(values).collect());
126        self.vector = None;
127        self.binary_vector = None;
128        self.fts = None;
129        Ok(())
130    }
131    pub fn set_sparse_values(&mut self, values: &[f32]) -> Result<()> {
132        if values.is_empty() || !values.iter().all(|v| v.is_finite()) {
133            return Err(Error::invalid_argument(
134                "sparse values cannot be empty and must be finite",
135            ));
136        }
137        let indices: Vec<u32> = if let Some(existing) = self.sparse_vector.as_ref() {
138            existing.iter().map(|(index, _)| *index).collect()
139        } else {
140            let length = u32::try_from(values.len())
141                .map_err(|_| Error::resource_exhausted("sparse vector exceeds u32 dimensions"))?;
142            (0..length).collect()
143        };
144        if indices.len() != values.len() {
145            return Err(Error::invalid_argument(
146                "sparse indices and values have different lengths",
147            ));
148        }
149        self.sparse_vector = Some(indices.into_iter().zip(values.iter().copied()).collect());
150        self.vector = None;
151        self.binary_vector = None;
152        self.fts = None;
153        Ok(())
154    }
155    pub fn set_hnsw_params(&mut self, params: HnswQueryParams) -> Result<()> {
156        apply_hnsw_query_controls(&mut self.params, params)
157    }
158    pub fn set_ivf_params(&mut self, params: IvfQueryParams) -> Result<()> {
159        apply_ivf_query_controls(&mut self.params, params)
160    }
161    pub fn set_ivf_rabitq_params(&mut self, params: IvfRabitqQueryParams) -> Result<()> {
162        apply_ivf_rabitq_query_controls(&mut self.params, params)
163    }
164    pub fn set_flat_params(&mut self, _params: FlatQueryParams) -> Result<()> {
165        unsupported_query_controls("Flat refinement")
166    }
167    pub fn set_diskann_params(&mut self, params: DiskannQueryParams) -> Result<()> {
168        apply_diskann_query_controls(&mut self.params, params)
169    }
170    pub fn set_fts_params(&mut self, params: FtsQueryParams) -> Result<()> {
171        apply_fts_query_controls(&mut self.params, params)
172    }
173    pub fn set_fts(&mut self, fts: &Fts) -> Result<()> {
174        if fts.query_string.is_none() && fts.match_string.is_none() {
175            return Err(Error::invalid_argument("FTS query has no expression"));
176        }
177        self.fts = Some(fts.clone());
178        self.vector = None;
179        self.binary_vector = None;
180        self.sparse_vector = None;
181        Ok(())
182    }
183    pub(crate) fn to_search_query(&self) -> Result<SearchQuery> {
184        let field = self
185            .field_name
186            .as_deref()
187            .ok_or_else(|| Error::invalid_argument("sub-query field_name is required"))?;
188        let route_count = usize::from(self.vector.is_some())
189            + usize::from(self.binary_vector.is_some())
190            + usize::from(self.sparse_vector.is_some())
191            + usize::from(self.fts.is_some());
192        if route_count != 1 {
193            return Err(Error::invalid_argument(
194                "sub-query must select exactly one dense, binary, sparse, or FTS route",
195            ));
196        }
197        let mut query = if let Some(vector) = &self.vector {
198            SearchQuery::new(field, vector, self.num_candidates)?
199        } else if let Some(vector) = &self.binary_vector {
200            SearchQuery::binary(field, vector, self.num_candidates)?
201        } else if let Some(sparse) = &self.sparse_vector {
202            let (i, v): (Vec<_>, Vec<_>) = sparse.iter().copied().unzip();
203            SearchQuery::sparse(field, &i, &v, self.num_candidates)?
204        } else if let Some(fts) = &self.fts {
205            SearchQuery::fts(field, fts, self.num_candidates)?
206        } else {
207            return Err(Error::invalid_argument(
208                "sub-query has no dense, binary, sparse, or FTS payload",
209            ));
210        };
211        query.params.extend(self.params.clone());
212        Ok(query)
213    }
214}
215
216/// A collection of sub-queries fused into one ranked result.
217#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
218pub struct MultiQuery {
219    pub queries: Vec<SubQuery>,
220    pub topk_value: i32,
221    pub filter: Option<String>,
222    pub include_vector_value: bool,
223    pub output_fields: Option<Vec<String>>,
224    pub rerank: RerankMethod,
225    pub normalization: Option<String>,
226}
227
228impl MultiQuery {
229    pub fn new() -> Result<Self> {
230        Ok(Self {
231            queries: Vec::new(),
232            topk_value: 10,
233            filter: None,
234            include_vector_value: false,
235            output_fields: None,
236            rerank: RerankMethod::Weighted {
237                weights: Vec::new(),
238            },
239            normalization: None,
240        })
241    }
242    pub fn add_sub_query(&mut self, sub: &SubQuery) -> Result<()> {
243        if self.queries.len() >= 1024 {
244            return Err(Error::resource_exhausted(
245                "multi-query branch limit exceeded",
246            ));
247        }
248        self.queries.push(sub.clone());
249        Ok(())
250    }
251    pub fn sub_query_count(&self) -> usize {
252        self.queries.len()
253    }
254    pub fn set_topk(&mut self, topk: i32) -> Result<()> {
255        if topk <= 0 {
256            return Err(Error::invalid_argument("topk must be positive"));
257        }
258        self.topk_value = topk;
259        Ok(())
260    }
261    pub fn topk(&self) -> i32 {
262        self.topk_value
263    }
264    pub fn set_filter(&mut self, filter: &str) -> Result<()> {
265        self.filter = (!filter.trim().is_empty()).then_some(filter.to_string());
266        Ok(())
267    }
268    pub fn set_include_vector(&mut self, include: bool) -> Result<()> {
269        self.include_vector_value = include;
270        Ok(())
271    }
272    pub fn include_vector(&self) -> bool {
273        self.include_vector_value
274    }
275    pub fn set_output_fields(&mut self, fields: &[&str]) -> Result<()> {
276        if fields
277            .iter()
278            .any(|f| f.trim().is_empty() || f.contains('\0'))
279        {
280            return Err(Error::invalid_argument("output field name is invalid"));
281        }
282        self.output_fields = Some(fields.iter().map(|v| (*v).to_string()).collect());
283        Ok(())
284    }
285    pub fn set_rerank_rrf(&mut self, rank_constant: i32) -> Result<()> {
286        if rank_constant <= 0 {
287            return Err(Error::invalid_argument(
288                "RRF rank constant must be positive",
289            ));
290        }
291        self.rerank = RerankMethod::ReciprocalRank {
292            rank_constant: f64::from(rank_constant),
293        };
294        Ok(())
295    }
296    pub fn set_rerank_weighted(&mut self, weights: &[f64]) -> Result<()> {
297        if weights.is_empty() || !weights.iter().all(|v| v.is_finite() && *v >= 0.0) {
298            return Err(Error::invalid_argument(
299                "weights must be non-empty, finite, and non-negative",
300            ));
301        }
302        self.rerank = RerankMethod::Weighted {
303            weights: weights.to_vec(),
304        };
305        Ok(())
306    }
307    pub fn set_normalization(&mut self, method: &str) -> Result<()> {
308        let normalized = method.to_ascii_lowercase();
309        if !matches!(normalized.as_str(), "none" | "minmax" | "zscore") {
310            return Err(Error::invalid_argument(
311                "normalization must be none, minmax, or zscore",
312            ));
313        }
314        self.normalization = Some(normalized);
315        Ok(())
316    }
317    pub(crate) fn effective_filter(&self) -> Option<&str> {
318        self.filter.as_deref()
319    }
320}