1use 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#[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#[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#[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}