Skip to main content

zvec_rust/
multi_query.rs

1//! Multi-query and sub-query types for combined vector searches.
2//!
3//! A [`MultiQuery`] aggregates multiple [`SubQuery`] instances and reranks the
4//! results using either RRF (Reciprocal Rank Fusion) or weighted strategies.
5
6use std::os::raw::c_void;
7
8use crate::error::{check_error, to_cstring, Error, ErrorCode, Result};
9use crate::query::Fts;
10use crate::query::{
11    DiskannQueryParams, FlatQueryParams, FtsQueryParams, HnswQueryParams, IvfQueryParams,
12    IvfRabitqQueryParams,
13};
14
15/// A multi-query operation combining multiple [`SubQuery`] objects.
16///
17/// Use [`MultiQuery::new`] to construct, then [`MultiQuery::add_sub_query`] to
18/// attach individual sub-queries. Configure top-k, filters, and rerank strategy
19/// before passing to [`crate::Collection::multi_query`].
20pub struct MultiQuery {
21    pub(crate) handle: *mut zvec_rust_sys::zvec_multi_query_t,
22}
23
24impl MultiQuery {
25    /// Creates a new multi-query.
26    pub fn new() -> Result<Self> {
27        let handle = unsafe { zvec_rust_sys::zvec_multi_query_create() };
28        if handle.is_null() {
29            return Err(Error {
30                code: ErrorCode::InternalError,
31                message: "failed to create multi-query".into(),
32            });
33        }
34        Ok(MultiQuery { handle })
35    }
36
37    /// Returns the raw FFI handle.
38    ///
39    /// # Safety
40    /// The caller must not use the handle after the `MultiQuery` is dropped.
41    pub unsafe fn as_raw(&self) -> *mut zvec_rust_sys::zvec_multi_query_t {
42        self.handle
43    }
44
45    /// Adds a sub-query to this multi-query.
46    ///
47    /// The sub-query is copied internally; the caller retains ownership.
48    pub fn add_sub_query(&mut self, sub: &SubQuery) -> Result<()> {
49        check_error(unsafe {
50            zvec_rust_sys::zvec_multi_query_add_sub_query(self.handle, sub.handle)
51        })
52    }
53
54    /// Returns the number of sub-queries currently registered.
55    pub fn sub_query_count(&self) -> usize {
56        unsafe { zvec_rust_sys::zvec_multi_query_get_sub_query_count(self.handle) }
57    }
58
59    /// Sets the top-k parameter for the merged result.
60    pub fn set_topk(&mut self, topk: i32) -> Result<()> {
61        check_error(unsafe { zvec_rust_sys::zvec_multi_query_set_topk(self.handle, topk) })
62    }
63
64    /// Returns the configured top-k value.
65    pub fn topk(&self) -> i32 {
66        unsafe { zvec_rust_sys::zvec_multi_query_get_topk(self.handle) }
67    }
68
69    /// Sets the filter expression applied to the final merged result.
70    pub fn set_filter(&mut self, filter: &str) -> Result<()> {
71        let c_filter = to_cstring(filter)?;
72        check_error(unsafe {
73            zvec_rust_sys::zvec_multi_query_set_filter(self.handle, c_filter.as_ptr())
74        })
75    }
76
77    /// Sets whether vector data is returned in the result documents.
78    pub fn set_include_vector(&mut self, include: bool) -> Result<()> {
79        check_error(unsafe {
80            zvec_rust_sys::zvec_multi_query_set_include_vector(self.handle, include)
81        })
82    }
83
84    /// Returns whether vector data is included in the result documents.
85    pub fn include_vector(&self) -> bool {
86        unsafe { zvec_rust_sys::zvec_multi_query_get_include_vector(self.handle) }
87    }
88
89    /// Sets the output field whitelist.
90    pub fn set_output_fields(&mut self, fields: &[&str]) -> Result<()> {
91        if fields.is_empty() {
92            return Ok(());
93        }
94        let c_fields: Vec<_> = fields
95            .iter()
96            .map(|f| to_cstring(f))
97            .collect::<Result<Vec<_>>>()?;
98        let c_ptrs: Vec<_> = c_fields.iter().map(|s| s.as_ptr()).collect();
99        check_error(unsafe {
100            zvec_rust_sys::zvec_multi_query_set_output_fields(
101                self.handle,
102                c_ptrs.as_ptr(),
103                c_ptrs.len(),
104            )
105        })
106    }
107
108    /// Configures Reciprocal Rank Fusion (RRF) as the rerank strategy.
109    pub fn set_rerank_rrf(&mut self, rank_constant: i32) -> Result<()> {
110        check_error(unsafe {
111            zvec_rust_sys::zvec_multi_query_set_rerank_rrf(self.handle, rank_constant)
112        })
113    }
114
115    /// Configures a weighted rerank strategy with the given per-sub-query weights.
116    pub fn set_rerank_weighted(&mut self, weights: &[f64]) -> Result<()> {
117        if weights.is_empty() {
118            return Err(Error {
119                code: ErrorCode::InvalidArgument,
120                message: "weights cannot be empty".into(),
121            });
122        }
123        check_error(unsafe {
124            zvec_rust_sys::zvec_multi_query_set_rerank_weighted(
125                self.handle,
126                weights.as_ptr(),
127                weights.len(),
128            )
129        })
130    }
131}
132
133impl Drop for MultiQuery {
134    fn drop(&mut self) {
135        if !self.handle.is_null() {
136            unsafe { zvec_rust_sys::zvec_multi_query_destroy(self.handle) };
137        }
138    }
139}
140
141// Safety: MultiQuery owns its handle exclusively.
142unsafe impl Send for MultiQuery {}
143
144/// A sub-query inside a [`MultiQuery`].
145///
146/// Each sub-query targets a single vector or sparse-vector field.
147pub struct SubQuery {
148    pub(crate) handle: *mut zvec_rust_sys::zvec_sub_query_t,
149}
150
151impl SubQuery {
152    /// Creates a new sub-query.
153    pub fn new() -> Result<Self> {
154        let handle = unsafe { zvec_rust_sys::zvec_sub_query_create() };
155        if handle.is_null() {
156            return Err(Error {
157                code: ErrorCode::InternalError,
158                message: "failed to create sub-query".into(),
159            });
160        }
161        Ok(SubQuery { handle })
162    }
163
164    /// Returns the raw FFI handle.
165    ///
166    /// # Safety
167    /// The caller must not use the handle after the `SubQuery` is dropped.
168    pub unsafe fn as_raw(&self) -> *mut zvec_rust_sys::zvec_sub_query_t {
169        self.handle
170    }
171
172    /// Sets the number of candidates to retrieve before reranking.
173    pub fn set_num_candidates(&mut self, n: i32) -> Result<()> {
174        check_error(unsafe { zvec_rust_sys::zvec_sub_query_set_num_candidates(self.handle, n) })
175    }
176
177    /// Returns the number of candidates.
178    pub fn num_candidates(&self) -> i32 {
179        unsafe { zvec_rust_sys::zvec_sub_query_get_num_candidates(self.handle) }
180    }
181
182    /// Sets the target field name.
183    pub fn set_field_name(&mut self, name: &str) -> Result<()> {
184        let c_name = to_cstring(name)?;
185        check_error(unsafe {
186            zvec_rust_sys::zvec_sub_query_set_field_name(self.handle, c_name.as_ptr())
187        })
188    }
189
190    /// Sets the dense query vector (f32).
191    pub fn set_query_vector(&mut self, data: &[f32]) -> Result<()> {
192        if data.is_empty() {
193            return Err(Error {
194                code: ErrorCode::InvalidArgument,
195                message: "query vector cannot be empty".into(),
196            });
197        }
198        let bytes = std::mem::size_of_val(data);
199        check_error(unsafe {
200            zvec_rust_sys::zvec_sub_query_set_query_vector(
201                self.handle,
202                data.as_ptr() as *const c_void,
203                bytes,
204            )
205        })
206    }
207
208    /// Sets the sparse vector (indices + values, equal length).
209    pub fn set_sparse_vector(&mut self, indices: &[u32], values: &[f32]) -> Result<()> {
210        if indices.len() != values.len() {
211            return Err(Error {
212                code: ErrorCode::InvalidArgument,
213                message: "indices and values must have the same length".into(),
214            });
215        }
216        if indices.is_empty() {
217            return Err(Error {
218                code: ErrorCode::InvalidArgument,
219                message: "sparse vector cannot be empty".into(),
220            });
221        }
222        check_error(unsafe {
223            zvec_rust_sys::zvec_sub_query_set_sparse_vector(
224                self.handle,
225                indices.as_ptr(),
226                values.as_ptr(),
227                indices.len(),
228            )
229        })
230    }
231
232    /// Sets only the sparse-vector indices.
233    pub fn set_sparse_indices(&mut self, indices: &[u32]) -> Result<()> {
234        check_error(unsafe {
235            zvec_rust_sys::zvec_sub_query_set_sparse_indices(
236                self.handle,
237                indices.as_ptr(),
238                indices.len(),
239            )
240        })
241    }
242
243    /// Sets only the sparse-vector values.
244    pub fn set_sparse_values(&mut self, values: &[f32]) -> Result<()> {
245        check_error(unsafe {
246            zvec_rust_sys::zvec_sub_query_set_sparse_values(
247                self.handle,
248                values.as_ptr(),
249                values.len(),
250            )
251        })
252    }
253
254    /// Sets HNSW query parameters (takes ownership on success).
255    pub fn set_hnsw_params(&mut self, mut params: HnswQueryParams) -> Result<()> {
256        check_error(unsafe {
257            zvec_rust_sys::zvec_sub_query_set_hnsw_params(self.handle, params.handle)
258        })?;
259        params.handle = std::ptr::null_mut();
260        Ok(())
261    }
262
263    /// Sets IVF query parameters (takes ownership on success).
264    pub fn set_ivf_params(&mut self, mut params: IvfQueryParams) -> Result<()> {
265        check_error(unsafe {
266            zvec_rust_sys::zvec_sub_query_set_ivf_params(self.handle, params.handle)
267        })?;
268        params.handle = std::ptr::null_mut();
269        Ok(())
270    }
271
272    /// Sets IVF RaBitQ query parameters (takes ownership on success).
273    pub fn set_ivf_rabitq_params(&mut self, mut params: IvfRabitqQueryParams) -> Result<()> {
274        check_error(unsafe {
275            zvec_rust_sys::zvec_sub_query_set_ivf_rabitq_params(self.handle, params.handle)
276        })?;
277        params.handle = std::ptr::null_mut();
278        Ok(())
279    }
280
281    /// Sets Flat query parameters (takes ownership on success).
282    pub fn set_flat_params(&mut self, mut params: FlatQueryParams) -> Result<()> {
283        check_error(unsafe {
284            zvec_rust_sys::zvec_sub_query_set_flat_params(self.handle, params.handle)
285        })?;
286        params.handle = std::ptr::null_mut();
287        Ok(())
288    }
289
290    /// Sets DiskANN query parameters (takes ownership on success).
291    pub fn set_diskann_params(&mut self, mut params: DiskannQueryParams) -> Result<()> {
292        check_error(unsafe {
293            zvec_rust_sys::zvec_sub_query_set_diskann_params(self.handle, params.handle)
294        })?;
295        params.handle = std::ptr::null_mut();
296        Ok(())
297    }
298
299    /// Sets FTS query parameters (takes ownership on success).
300    pub fn set_fts_params(&mut self, mut params: FtsQueryParams) -> Result<()> {
301        check_error(unsafe {
302            zvec_rust_sys::zvec_sub_query_set_fts_params(self.handle, params.handle)
303        })?;
304        params.handle = std::ptr::null_mut();
305        Ok(())
306    }
307
308    /// Sets FTS payload (payload is copied, caller retains ownership).
309    pub fn set_fts(&mut self, fts: &Fts) -> Result<()> {
310        check_error(unsafe { zvec_rust_sys::zvec_sub_query_set_fts(self.handle, fts.handle) })
311    }
312}
313
314impl Drop for SubQuery {
315    fn drop(&mut self) {
316        if !self.handle.is_null() {
317            unsafe { zvec_rust_sys::zvec_sub_query_destroy(self.handle) };
318        }
319    }
320}
321
322// Safety: SubQuery owns its handle exclusively.
323unsafe impl Send for SubQuery {}
324
325#[cfg(test)]
326mod tests {
327    use super::*;
328
329    #[test]
330    fn create_and_drop_multi_query() {
331        let mq = MultiQuery::new().expect("create multi-query");
332        assert_eq!(mq.sub_query_count(), 0);
333    }
334
335    #[test]
336    fn create_and_drop_sub_query() {
337        let _sq = SubQuery::new().expect("create sub-query");
338    }
339
340    #[test]
341    fn multi_query_basic_setters() {
342        let mut mq = MultiQuery::new().expect("create multi-query");
343        mq.set_topk(20).expect("set topk");
344        assert_eq!(mq.topk(), 20);
345        mq.set_include_vector(true).expect("set include_vector");
346        assert!(mq.include_vector());
347    }
348
349    #[test]
350    fn add_sub_query_increments_count() {
351        let mut mq = MultiQuery::new().expect("create multi-query");
352        let mut sq = SubQuery::new().expect("create sub-query");
353        sq.set_field_name("vec").expect("set field name");
354        sq.set_query_vector(&[0.1, 0.2, 0.3, 0.4])
355            .expect("set query vector");
356        sq.set_num_candidates(50).expect("set num candidates");
357        mq.add_sub_query(&sq).expect("add sub-query");
358        assert_eq!(mq.sub_query_count(), 1);
359    }
360
361    #[test]
362    fn rerank_weighted_rejects_empty() {
363        let mut mq = MultiQuery::new().expect("create multi-query");
364        let err = mq.set_rerank_weighted(&[]).unwrap_err();
365        assert_eq!(err.code, ErrorCode::InvalidArgument);
366    }
367
368    #[test]
369    fn sub_query_set_fts() {
370        let mut sub = SubQuery::new().expect("create sub-query");
371        sub.set_field_name("content").expect("set field name");
372
373        let mut fts = Fts::new().expect("create fts payload");
374        fts.set_match_string("hello world")
375            .expect("set match string");
376
377        sub.set_fts(&fts).expect("set fts payload");
378    }
379
380    #[test]
381    fn sub_query_set_fts_params() {
382        let mut sub = SubQuery::new().expect("create sub-query");
383        sub.set_field_name("content").expect("set field name");
384
385        let params = FtsQueryParams::new(Some("AND")).expect("create fts params");
386        sub.set_fts_params(params).expect("set fts params");
387    }
388
389    #[test]
390    fn sub_query_set_ivf_rabitq_params() {
391        let mut sub = SubQuery::new().expect("create sub-query");
392        sub.set_field_name("embedding").expect("set field name");
393        sub.set_query_vector(&[0.1, 0.2, 0.3, 0.4])
394            .expect("set query vector");
395
396        let params = IvfRabitqQueryParams::new(16, 0.0, false, false);
397        sub.set_ivf_rabitq_params(params)
398            .expect("set ivf rabitq params");
399    }
400
401    #[test]
402    fn sub_query_set_diskann_params() {
403        let mut sub = SubQuery::new().expect("create sub-query");
404        sub.set_field_name("embedding").expect("set field name");
405        sub.set_query_vector(&[0.1, 0.2, 0.3, 0.4])
406            .expect("set query vector");
407
408        let params = DiskannQueryParams::new(200);
409        sub.set_diskann_params(params).expect("set diskann params");
410    }
411}