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