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};
13
14pub struct MultiQuery {
20 pub(crate) handle: *mut zvec_rust_sys::zvec_multi_query_t,
21}
22
23impl MultiQuery {
24 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 pub unsafe fn as_raw(&self) -> *mut zvec_rust_sys::zvec_multi_query_t {
41 self.handle
42 }
43
44 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 pub fn sub_query_count(&self) -> usize {
55 unsafe { zvec_rust_sys::zvec_multi_query_get_sub_query_count(self.handle) }
56 }
57
58 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 pub fn topk(&self) -> i32 {
65 unsafe { zvec_rust_sys::zvec_multi_query_get_topk(self.handle) }
66 }
67
68 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 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 pub fn include_vector(&self) -> bool {
85 unsafe { zvec_rust_sys::zvec_multi_query_get_include_vector(self.handle) }
86 }
87
88 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 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 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
140unsafe impl Send for MultiQuery {}
142
143pub struct SubQuery {
147 pub(crate) handle: *mut zvec_rust_sys::zvec_sub_query_t,
148}
149
150impl SubQuery {
151 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 pub unsafe fn as_raw(&self) -> *mut zvec_rust_sys::zvec_sub_query_t {
168 self.handle
169 }
170
171 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 pub fn num_candidates(&self) -> i32 {
178 unsafe { zvec_rust_sys::zvec_sub_query_get_num_candidates(self.handle) }
179 }
180
181 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 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 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 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 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 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 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 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 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 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 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
312unsafe 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}