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