1use 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
15pub struct MultiQuery {
21 pub(crate) handle: *mut zvec_rust_sys::zvec_multi_query_t,
22}
23
24impl MultiQuery {
25 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 pub unsafe fn as_raw(&self) -> *mut zvec_rust_sys::zvec_multi_query_t {
42 self.handle
43 }
44
45 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 pub fn sub_query_count(&self) -> usize {
56 unsafe { zvec_rust_sys::zvec_multi_query_get_sub_query_count(self.handle) }
57 }
58
59 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 pub fn topk(&self) -> i32 {
66 unsafe { zvec_rust_sys::zvec_multi_query_get_topk(self.handle) }
67 }
68
69 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 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 pub fn include_vector(&self) -> bool {
86 unsafe { zvec_rust_sys::zvec_multi_query_get_include_vector(self.handle) }
87 }
88
89 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 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 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
141unsafe impl Send for MultiQuery {}
143
144pub struct SubQuery {
148 pub(crate) handle: *mut zvec_rust_sys::zvec_sub_query_t,
149}
150
151impl SubQuery {
152 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 pub unsafe fn as_raw(&self) -> *mut zvec_rust_sys::zvec_sub_query_t {
169 self.handle
170 }
171
172 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 pub fn num_candidates(&self) -> i32 {
179 unsafe { zvec_rust_sys::zvec_sub_query_get_num_candidates(self.handle) }
180 }
181
182 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 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 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 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 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 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 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 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 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 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 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 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
322unsafe 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}