1use 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
12pub struct MultiQuery {
18 pub(crate) handle: *mut zvec_rust_sys::zvec_multi_query_t,
19}
20
21impl MultiQuery {
22 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 pub unsafe fn as_raw(&self) -> *mut zvec_rust_sys::zvec_multi_query_t {
39 self.handle
40 }
41
42 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 pub fn sub_query_count(&self) -> usize {
53 unsafe { zvec_rust_sys::zvec_multi_query_get_sub_query_count(self.handle) }
54 }
55
56 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 pub fn topk(&self) -> i32 {
63 unsafe { zvec_rust_sys::zvec_multi_query_get_topk(self.handle) }
64 }
65
66 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 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 pub fn include_vector(&self) -> bool {
83 unsafe { zvec_rust_sys::zvec_multi_query_get_include_vector(self.handle) }
84 }
85
86 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 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 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
138unsafe impl Send for MultiQuery {}
140
141pub struct SubQuery {
145 pub(crate) handle: *mut zvec_rust_sys::zvec_sub_query_t,
146}
147
148impl SubQuery {
149 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 pub unsafe fn as_raw(&self) -> *mut zvec_rust_sys::zvec_sub_query_t {
166 self.handle
167 }
168
169 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 pub fn num_candidates(&self) -> i32 {
176 unsafe { zvec_rust_sys::zvec_sub_query_get_num_candidates(self.handle) }
177 }
178
179 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 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 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 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 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 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 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 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 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 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
301unsafe 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}