Skip to main content

zvec_rust/
query.rs

1use std::os::raw::c_void;
2
3use crate::error::{check_error, to_cstring, Error, ErrorCode, Result};
4
5/// HNSW-specific query parameters.
6pub struct HnswQueryParams {
7    pub(crate) handle: *mut zvec_rust_sys::zvec_hnsw_query_params_t,
8}
9
10impl HnswQueryParams {
11    /// Creates new HNSW query parameters.
12    pub fn new(ef: i32, radius: f32, is_linear: bool, is_using_refiner: bool) -> Self {
13        let handle = unsafe {
14            zvec_rust_sys::zvec_query_params_hnsw_create(ef, radius, is_linear, is_using_refiner)
15        };
16        HnswQueryParams { handle }
17    }
18
19    /// Sets the exploration factor.
20    pub fn set_ef(&mut self, ef: i32) -> Result<()> {
21        check_error(unsafe { zvec_rust_sys::zvec_query_params_hnsw_set_ef(self.handle, ef) })
22    }
23
24    /// Returns the exploration factor.
25    pub fn ef(&self) -> i32 {
26        unsafe { zvec_rust_sys::zvec_query_params_hnsw_get_ef(self.handle) }
27    }
28}
29
30impl Drop for HnswQueryParams {
31    fn drop(&mut self) {
32        if !self.handle.is_null() {
33            unsafe { zvec_rust_sys::zvec_query_params_hnsw_destroy(self.handle) };
34        }
35    }
36}
37
38/// IVF-specific query parameters.
39pub struct IvfQueryParams {
40    pub(crate) handle: *mut zvec_rust_sys::zvec_ivf_query_params_t,
41}
42
43impl IvfQueryParams {
44    /// Creates new IVF query parameters.
45    pub fn new(nprobe: i32, is_using_refiner: bool, scale_factor: f32) -> Self {
46        let handle = unsafe {
47            zvec_rust_sys::zvec_query_params_ivf_create(nprobe, is_using_refiner, scale_factor)
48        };
49        IvfQueryParams { handle }
50    }
51
52    /// Sets the number of probe clusters.
53    pub fn set_nprobe(&mut self, nprobe: i32) -> Result<()> {
54        check_error(unsafe { zvec_rust_sys::zvec_query_params_ivf_set_nprobe(self.handle, nprobe) })
55    }
56
57    /// Returns the number of probe clusters.
58    pub fn nprobe(&self) -> i32 {
59        unsafe { zvec_rust_sys::zvec_query_params_ivf_get_nprobe(self.handle) }
60    }
61}
62
63impl Drop for IvfQueryParams {
64    fn drop(&mut self) {
65        if !self.handle.is_null() {
66            unsafe { zvec_rust_sys::zvec_query_params_ivf_destroy(self.handle) };
67        }
68    }
69}
70
71/// IVF RaBitQ-specific query parameters.
72pub struct IvfRabitqQueryParams {
73    pub(crate) handle: *mut zvec_rust_sys::zvec_ivf_rabitq_query_params_t,
74}
75
76impl IvfRabitqQueryParams {
77    /// Creates new IVF RaBitQ query parameters.
78    pub fn new(nprobe: i32, radius: f32, is_linear: bool, is_using_refiner: bool) -> Self {
79        let handle = unsafe {
80            zvec_rust_sys::zvec_query_params_ivf_rabitq_create(
81                nprobe,
82                radius,
83                is_linear,
84                is_using_refiner,
85            )
86        };
87        IvfRabitqQueryParams { handle }
88    }
89
90    /// Sets the number of probe clusters.
91    pub fn set_nprobe(&mut self, nprobe: i32) -> Result<()> {
92        check_error(unsafe {
93            zvec_rust_sys::zvec_query_params_ivf_rabitq_set_nprobe(self.handle, nprobe)
94        })
95    }
96
97    /// Returns the number of probe clusters.
98    pub fn nprobe(&self) -> i32 {
99        unsafe { zvec_rust_sys::zvec_query_params_ivf_rabitq_get_nprobe(self.handle) }
100    }
101
102    /// Sets the candidate expansion factor used by the refiner.
103    pub fn set_scale_factor(&mut self, scale_factor: f32) -> Result<()> {
104        check_error(unsafe {
105            zvec_rust_sys::zvec_query_params_ivf_rabitq_set_scale_factor(self.handle, scale_factor)
106        })
107    }
108
109    /// Returns the candidate expansion factor used by the refiner.
110    pub fn scale_factor(&self) -> f32 {
111        unsafe { zvec_rust_sys::zvec_query_params_ivf_rabitq_get_scale_factor(self.handle) }
112    }
113}
114
115impl Drop for IvfRabitqQueryParams {
116    fn drop(&mut self) {
117        if !self.handle.is_null() {
118            unsafe { zvec_rust_sys::zvec_query_params_ivf_rabitq_destroy(self.handle) };
119        }
120    }
121}
122
123/// Flat-specific query parameters.
124pub struct FlatQueryParams {
125    pub(crate) handle: *mut zvec_rust_sys::zvec_flat_query_params_t,
126}
127
128impl FlatQueryParams {
129    /// Creates new Flat query parameters.
130    pub fn new(is_using_refiner: bool, scale_factor: f32) -> Self {
131        let handle =
132            unsafe { zvec_rust_sys::zvec_query_params_flat_create(is_using_refiner, scale_factor) };
133        FlatQueryParams { handle }
134    }
135}
136
137impl Drop for FlatQueryParams {
138    fn drop(&mut self) {
139        if !self.handle.is_null() {
140            unsafe { zvec_rust_sys::zvec_query_params_flat_destroy(self.handle) };
141        }
142    }
143}
144
145/// DiskANN-specific query parameters.
146pub struct DiskannQueryParams {
147    pub(crate) handle: *mut zvec_rust_sys::zvec_diskann_query_params_t,
148}
149
150impl DiskannQueryParams {
151    /// Creates new DiskANN query parameters.
152    ///
153    /// - `list_size`: search frontier size (default in the C library: 300)
154    pub fn new(list_size: i32) -> Self {
155        let handle = unsafe { zvec_rust_sys::zvec_query_params_diskann_create(list_size) };
156        DiskannQueryParams { handle }
157    }
158
159    /// Sets the search-time frontier size.
160    pub fn set_list_size(&mut self, list_size: i32) -> Result<()> {
161        check_error(unsafe {
162            zvec_rust_sys::zvec_query_params_diskann_set_list_size(self.handle, list_size)
163        })
164    }
165
166    /// Returns the search-time frontier size.
167    pub fn list_size(&self) -> i32 {
168        unsafe { zvec_rust_sys::zvec_query_params_diskann_get_list_size(self.handle) }
169    }
170}
171
172impl Drop for DiskannQueryParams {
173    fn drop(&mut self) {
174        if !self.handle.is_null() {
175            unsafe { zvec_rust_sys::zvec_query_params_diskann_destroy(self.handle) };
176        }
177    }
178}
179
180/// FTS-specific query parameters controlling the default boolean operator.
181pub struct FtsQueryParams {
182    pub(crate) handle: *mut zvec_rust_sys::zvec_fts_query_params_t,
183}
184
185impl FtsQueryParams {
186    /// Creates new FTS query parameters.
187    ///
188    /// `default_operator` sets the boolean operator for adjacent bare terms
189    /// ("OR" or "AND", case-insensitive). Pass `None` to use the library default.
190    pub fn new(default_operator: Option<&str>) -> Result<Self> {
191        let c_op = default_operator.map(to_cstring).transpose()?;
192        let handle = unsafe {
193            zvec_rust_sys::zvec_query_params_fts_create(
194                c_op.as_ref().map_or(std::ptr::null(), |c| c.as_ptr()),
195            )
196        };
197        if handle.is_null() {
198            return Err(Error {
199                code: ErrorCode::InternalError,
200                message: "failed to create FTS query params".into(),
201            });
202        }
203        Ok(FtsQueryParams { handle })
204    }
205
206    /// Sets the default boolean operator.
207    pub fn set_default_operator(&mut self, op: &str) -> Result<()> {
208        let c_op = to_cstring(op)?;
209        check_error(unsafe {
210            zvec_rust_sys::zvec_query_params_fts_set_default_operator(self.handle, c_op.as_ptr())
211        })
212    }
213
214    /// Returns the default boolean operator.
215    pub fn default_operator(&self) -> Option<String> {
216        unsafe {
217            let ptr = zvec_rust_sys::zvec_query_params_fts_get_default_operator(self.handle);
218            if ptr.is_null() {
219                return None;
220            }
221            Some(std::ffi::CStr::from_ptr(ptr).to_string_lossy().into_owned())
222        }
223    }
224}
225
226impl Drop for FtsQueryParams {
227    fn drop(&mut self) {
228        if !self.handle.is_null() {
229            unsafe { zvec_rust_sys::zvec_query_params_fts_destroy(self.handle) };
230        }
231    }
232}
233
234/// FTS query payload holding the query expression and match string.
235///
236/// - `query_string`: a boolean / advanced query expression
237/// - `match_string`: a natural-language match string
238pub struct Fts {
239    pub(crate) handle: *mut zvec_rust_sys::zvec_fts_t,
240}
241
242impl Fts {
243    /// Creates a new FTS query payload.
244    pub fn new() -> Result<Self> {
245        let handle = unsafe { zvec_rust_sys::zvec_fts_create() };
246        if handle.is_null() {
247            return Err(Error {
248                code: ErrorCode::InternalError,
249                message: "failed to create FTS payload".into(),
250            });
251        }
252        Ok(Fts { handle })
253    }
254
255    /// Sets the boolean / advanced query expression.
256    pub fn set_query_string(&mut self, query: &str) -> Result<()> {
257        let c = to_cstring(query)?;
258        check_error(unsafe { zvec_rust_sys::zvec_fts_set_query_string(self.handle, c.as_ptr()) })
259    }
260
261    /// Sets the natural-language match string.
262    pub fn set_match_string(&mut self, match_str: &str) -> Result<()> {
263        let c = to_cstring(match_str)?;
264        check_error(unsafe { zvec_rust_sys::zvec_fts_set_match_string(self.handle, c.as_ptr()) })
265    }
266
267    /// Returns the query expression, or `None` if not set.
268    pub fn query_string(&self) -> Option<String> {
269        unsafe {
270            let ptr = zvec_rust_sys::zvec_fts_get_query_string(self.handle);
271            if ptr.is_null() {
272                return None;
273            }
274            Some(std::ffi::CStr::from_ptr(ptr).to_string_lossy().into_owned())
275        }
276    }
277
278    /// Returns the match string, or `None` if not set.
279    pub fn match_string(&self) -> Option<String> {
280        unsafe {
281            let ptr = zvec_rust_sys::zvec_fts_get_match_string(self.handle);
282            if ptr.is_null() {
283                return None;
284            }
285            Some(std::ffi::CStr::from_ptr(ptr).to_string_lossy().into_owned())
286        }
287    }
288}
289
290impl Drop for Fts {
291    fn drop(&mut self) {
292        if !self.handle.is_null() {
293            unsafe { zvec_rust_sys::zvec_fts_destroy(self.handle) };
294        }
295    }
296}
297
298/// A vector similarity search query.
299pub struct SearchQuery {
300    pub(crate) handle: *mut zvec_rust_sys::zvec_vector_query_t,
301}
302
303impl SearchQuery {
304    /// Returns the raw FFI handle.
305    ///
306    /// # Safety
307    /// The caller must not use the handle after the `SearchQuery` is dropped.
308    pub unsafe fn as_raw(&self) -> *mut zvec_rust_sys::zvec_vector_query_t {
309        self.handle
310    }
311
312    /// Creates a `SearchQuery` from a raw FFI handle.
313    ///
314    /// # Safety
315    /// The caller must ensure the handle is valid and was created by the zvec C API.
316    /// The `SearchQuery` takes ownership and will call `zvec_vector_query_destroy` on drop.
317    pub unsafe fn from_raw(handle: *mut zvec_rust_sys::zvec_vector_query_t) -> Self {
318        SearchQuery { handle }
319    }
320
321    /// Creates a new vector query with the given field name, query vector, and topk.
322    pub fn new(field_name: &str, vector: &[f32], topk: i32) -> Result<Self> {
323        let handle = unsafe { zvec_rust_sys::zvec_vector_query_create() };
324        if handle.is_null() {
325            return Err(Error {
326                code: ErrorCode::InternalError,
327                message: "failed to create vector query".into(),
328            });
329        }
330
331        let c_field = to_cstring(field_name)?;
332        let query = SearchQuery { handle };
333
334        check_error(unsafe {
335            zvec_rust_sys::zvec_vector_query_set_field_name(query.handle, c_field.as_ptr())
336        })?;
337        check_error(unsafe {
338            zvec_rust_sys::zvec_vector_query_set_query_vector(
339                query.handle,
340                vector.as_ptr() as *const c_void,
341                std::mem::size_of_val(vector),
342            )
343        })?;
344        check_error(unsafe { zvec_rust_sys::zvec_vector_query_set_topk(query.handle, topk) })?;
345
346        Ok(query)
347    }
348
349    /// Creates a pure full-text (keyword-only) search query with no dense vector.
350    ///
351    /// [`SearchQuery::new`] and [`SearchQueryBuilder::build`] both require a query
352    /// vector, so an FTS clause can only ever be layered on top of a dense vector
353    /// (hybrid search). This constructor reaches the C API's vector-less FTS path:
354    /// once an FTS payload is attached, the underlying `zvec_vector_query` does not
355    /// require a query vector, enabling keyword-only retrieval.
356    ///
357    /// `field_name` must be the FTS-indexed field.
358    pub fn fts(field_name: &str, fts: &Fts, topk: i32) -> Result<Self> {
359        let handle = unsafe { zvec_rust_sys::zvec_vector_query_create() };
360        if handle.is_null() {
361            return Err(Error {
362                code: ErrorCode::InternalError,
363                message: "failed to create vector query".into(),
364            });
365        }
366
367        // Wrap the handle before any fallible call so `Drop` frees it on early return.
368        let mut query = SearchQuery { handle };
369        let c_field = to_cstring(field_name)?;
370
371        check_error(unsafe {
372            zvec_rust_sys::zvec_vector_query_set_field_name(query.handle, c_field.as_ptr())
373        })?;
374        check_error(unsafe { zvec_rust_sys::zvec_vector_query_set_topk(query.handle, topk) })?;
375        query.set_fts(fts)?;
376
377        Ok(query)
378    }
379
380    /// Returns a builder for constructing a search query.
381    pub fn builder() -> SearchQueryBuilder {
382        SearchQueryBuilder::new()
383    }
384
385    /// Sets the filter expression.
386    pub fn set_filter(&mut self, filter: &str) -> Result<()> {
387        let c_filter = to_cstring(filter)?;
388        check_error(unsafe {
389            zvec_rust_sys::zvec_vector_query_set_filter(self.handle, c_filter.as_ptr())
390        })
391    }
392
393    /// Sets whether to include vector data in results.
394    pub fn set_include_vector(&mut self, include: bool) -> Result<()> {
395        check_error(unsafe {
396            zvec_rust_sys::zvec_vector_query_set_include_vector(self.handle, include)
397        })
398    }
399
400    /// Sets whether to include doc ID in results.
401    pub fn set_include_doc_id(&mut self, include: bool) -> Result<()> {
402        check_error(unsafe {
403            zvec_rust_sys::zvec_vector_query_set_include_doc_id(self.handle, include)
404        })
405    }
406
407    /// Sets the output fields to include in results.
408    pub fn set_output_fields(&mut self, fields: &[&str]) -> Result<()> {
409        let c_fields: Vec<_> = fields
410            .iter()
411            .map(|f| to_cstring(f))
412            .collect::<Result<Vec<_>>>()?;
413        let c_ptrs: Vec<_> = c_fields.iter().map(|f| f.as_ptr()).collect();
414        check_error(unsafe {
415            zvec_rust_sys::zvec_vector_query_set_output_fields(
416                self.handle,
417                c_ptrs.as_ptr(),
418                c_ptrs.len(),
419            )
420        })
421    }
422
423    /// Sets HNSW query parameters (takes ownership on success).
424    pub fn set_hnsw_params(&mut self, mut params: HnswQueryParams) -> Result<()> {
425        check_error(unsafe {
426            zvec_rust_sys::zvec_vector_query_set_hnsw_params(self.handle, params.handle)
427        })?;
428        // Ownership transferred to query only on success; prevent double-free
429        params.handle = std::ptr::null_mut();
430        Ok(())
431    }
432
433    /// Sets IVF query parameters (takes ownership on success).
434    pub fn set_ivf_params(&mut self, mut params: IvfQueryParams) -> Result<()> {
435        check_error(unsafe {
436            zvec_rust_sys::zvec_vector_query_set_ivf_params(self.handle, params.handle)
437        })?;
438        params.handle = std::ptr::null_mut();
439        Ok(())
440    }
441
442    /// Sets IVF RaBitQ query parameters (takes ownership on success).
443    pub fn set_ivf_rabitq_params(&mut self, mut params: IvfRabitqQueryParams) -> Result<()> {
444        check_error(unsafe {
445            zvec_rust_sys::zvec_vector_query_set_ivf_rabitq_params(self.handle, params.handle)
446        })?;
447        params.handle = std::ptr::null_mut();
448        Ok(())
449    }
450
451    /// Sets Flat query parameters (takes ownership on success).
452    pub fn set_flat_params(&mut self, mut params: FlatQueryParams) -> Result<()> {
453        check_error(unsafe {
454            zvec_rust_sys::zvec_vector_query_set_flat_params(self.handle, params.handle)
455        })?;
456        params.handle = std::ptr::null_mut();
457        Ok(())
458    }
459
460    /// Sets DiskANN query parameters (takes ownership on success).
461    pub fn set_diskann_params(&mut self, mut params: DiskannQueryParams) -> Result<()> {
462        check_error(unsafe {
463            zvec_rust_sys::zvec_vector_query_set_diskann_params(self.handle, params.handle)
464        })?;
465        params.handle = std::ptr::null_mut();
466        Ok(())
467    }
468
469    /// Sets FTS query parameters (takes ownership on success).
470    pub fn set_fts_params(&mut self, mut params: FtsQueryParams) -> Result<()> {
471        check_error(unsafe {
472            zvec_rust_sys::zvec_vector_query_set_fts_params(self.handle, params.handle)
473        })?;
474        params.handle = std::ptr::null_mut();
475        Ok(())
476    }
477
478    /// Sets FTS payload (payload is copied, caller retains ownership).
479    pub fn set_fts(&mut self, fts: &Fts) -> Result<()> {
480        check_error(unsafe { zvec_rust_sys::zvec_vector_query_set_fts(self.handle, fts.handle) })
481    }
482}
483
484impl Drop for SearchQuery {
485    fn drop(&mut self) {
486        if !self.handle.is_null() {
487            unsafe { zvec_rust_sys::zvec_vector_query_destroy(self.handle) };
488        }
489    }
490}
491
492/// Builder for constructing a [`SearchQuery`].
493pub struct SearchQueryBuilder {
494    field_name: Option<String>,
495    vector: Option<Vec<f32>>,
496    topk: i32,
497    filter: Option<String>,
498    include_vector: Option<bool>,
499    include_doc_id: Option<bool>,
500    output_fields: Option<Vec<String>>,
501    fts_query_string: Option<String>,
502    fts_match_string: Option<String>,
503}
504
505impl SearchQueryBuilder {
506    fn new() -> Self {
507        SearchQueryBuilder {
508            field_name: None,
509            vector: None,
510            topk: 10,
511            filter: None,
512            include_vector: None,
513            include_doc_id: None,
514            output_fields: None,
515            fts_query_string: None,
516            fts_match_string: None,
517        }
518    }
519
520    /// Sets the field name to query.
521    pub fn field_name(mut self, name: &str) -> Self {
522        self.field_name = Some(name.to_string());
523        self
524    }
525
526    /// Sets the query vector.
527    pub fn vector(mut self, vector: &[f32]) -> Self {
528        self.vector = Some(vector.to_vec());
529        self
530    }
531
532    /// Sets the number of results to return.
533    pub fn topk(mut self, topk: i32) -> Self {
534        self.topk = topk;
535        self
536    }
537
538    /// Sets the filter expression.
539    pub fn filter(mut self, filter: &str) -> Self {
540        self.filter = Some(filter.to_string());
541        self
542    }
543
544    /// Sets whether to include vector data in results.
545    pub fn include_vector(mut self, include: bool) -> Self {
546        self.include_vector = Some(include);
547        self
548    }
549
550    /// Sets whether to include doc ID in results.
551    pub fn include_doc_id(mut self, include: bool) -> Self {
552        self.include_doc_id = Some(include);
553        self
554    }
555
556    /// Sets the output fields.
557    pub fn output_fields(mut self, fields: &[&str]) -> Self {
558        self.output_fields = Some(fields.iter().map(|s| s.to_string()).collect());
559        self
560    }
561
562    /// Sets the FTS boolean / advanced query expression.
563    pub fn fts_query_string(mut self, query: &str) -> Self {
564        self.fts_query_string = Some(query.to_string());
565        self
566    }
567
568    /// Sets the FTS natural-language match string.
569    pub fn fts_match_string(mut self, match_str: &str) -> Self {
570        self.fts_match_string = Some(match_str.to_string());
571        self
572    }
573
574    /// Builds the search query.
575    pub fn build(self) -> Result<SearchQuery> {
576        let field_name = self.field_name.ok_or_else(|| Error {
577            code: ErrorCode::InvalidArgument,
578            message: "field_name is required".into(),
579        })?;
580        let vector = self.vector.ok_or_else(|| Error {
581            code: ErrorCode::InvalidArgument,
582            message: "vector is required".into(),
583        })?;
584
585        let mut query = SearchQuery::new(&field_name, &vector, self.topk)?;
586
587        if let Some(filter) = &self.filter {
588            query.set_filter(filter)?;
589        }
590        if let Some(include) = self.include_vector {
591            query.set_include_vector(include)?;
592        }
593        if let Some(include) = self.include_doc_id {
594            query.set_include_doc_id(include)?;
595        }
596        if let Some(fields) = &self.output_fields {
597            let field_refs: Vec<&str> = fields.iter().map(|s| s.as_str()).collect();
598            query.set_output_fields(&field_refs)?;
599        }
600        if self.fts_query_string.is_some() || self.fts_match_string.is_some() {
601            let mut fts = Fts::new()?;
602            if let Some(qs) = &self.fts_query_string {
603                fts.set_query_string(qs)?;
604            }
605            if let Some(ms) = &self.fts_match_string {
606                fts.set_match_string(ms)?;
607            }
608            query.set_fts(&fts)?;
609        }
610
611        Ok(query)
612    }
613}
614
615/// A grouped vector similarity search query.
616pub struct GroupBySearchQuery {
617    pub(crate) handle: *mut zvec_rust_sys::zvec_group_by_vector_query_t,
618}
619
620impl GroupBySearchQuery {
621    /// Creates a new group-by search query.
622    pub fn new(
623        field_name: &str,
624        group_by_field: &str,
625        vector: &[f32],
626        group_count: u32,
627        group_topk: u32,
628    ) -> Result<Self> {
629        let handle = unsafe { zvec_rust_sys::zvec_group_by_vector_query_create() };
630        if handle.is_null() {
631            return Err(Error {
632                code: ErrorCode::InternalError,
633                message: "failed to create group by vector query".into(),
634            });
635        }
636
637        let c_field = to_cstring(field_name)?;
638        let c_group_field = to_cstring(group_by_field)?;
639
640        check_error(unsafe {
641            zvec_rust_sys::zvec_group_by_vector_query_set_field_name(handle, c_field.as_ptr())
642        })?;
643        check_error(unsafe {
644            zvec_rust_sys::zvec_group_by_vector_query_set_group_by_field_name(
645                handle,
646                c_group_field.as_ptr(),
647            )
648        })?;
649        check_error(unsafe {
650            zvec_rust_sys::zvec_group_by_vector_query_set_query_vector(
651                handle,
652                vector.as_ptr() as *const c_void,
653                std::mem::size_of_val(vector),
654            )
655        })?;
656        check_error(unsafe {
657            zvec_rust_sys::zvec_group_by_vector_query_set_group_count(handle, group_count)
658        })?;
659        check_error(unsafe {
660            zvec_rust_sys::zvec_group_by_vector_query_set_topk_per_group(handle, group_topk)
661        })?;
662
663        Ok(GroupBySearchQuery { handle })
664    }
665
666    /// Sets the filter expression.
667    pub fn set_filter(&mut self, filter: &str) -> Result<()> {
668        let c_filter = to_cstring(filter)?;
669        check_error(unsafe {
670            zvec_rust_sys::zvec_group_by_vector_query_set_filter(self.handle, c_filter.as_ptr())
671        })
672    }
673
674    /// Sets whether to include vector data in results.
675    pub fn set_include_vector(&mut self, include: bool) -> Result<()> {
676        check_error(unsafe {
677            zvec_rust_sys::zvec_group_by_vector_query_set_include_vector(self.handle, include)
678        })
679    }
680
681    /// Sets the output fields.
682    pub fn set_output_fields(&mut self, fields: &[&str]) -> Result<()> {
683        let c_fields: Vec<_> = fields
684            .iter()
685            .map(|f| to_cstring(f))
686            .collect::<Result<Vec<_>>>()?;
687        let c_ptrs: Vec<_> = c_fields.iter().map(|f| f.as_ptr()).collect();
688        check_error(unsafe {
689            zvec_rust_sys::zvec_group_by_vector_query_set_output_fields(
690                self.handle,
691                c_ptrs.as_ptr(),
692                c_ptrs.len(),
693            )
694        })
695    }
696
697    /// Sets HNSW query parameters (takes ownership on success).
698    pub fn set_hnsw_params(&mut self, mut params: HnswQueryParams) -> Result<()> {
699        check_error(unsafe {
700            zvec_rust_sys::zvec_group_by_vector_query_set_hnsw_params(self.handle, params.handle)
701        })?;
702        // Ownership transferred to query only on success; prevent double-free
703        params.handle = std::ptr::null_mut();
704        Ok(())
705    }
706
707    /// Sets IVF query parameters (takes ownership on success).
708    pub fn set_ivf_params(&mut self, mut params: IvfQueryParams) -> Result<()> {
709        check_error(unsafe {
710            zvec_rust_sys::zvec_group_by_vector_query_set_ivf_params(self.handle, params.handle)
711        })?;
712        // Ownership transferred to query only on success; prevent double-free
713        params.handle = std::ptr::null_mut();
714        Ok(())
715    }
716
717    /// Sets IVF RaBitQ query parameters (takes ownership on success).
718    pub fn set_ivf_rabitq_params(&mut self, mut params: IvfRabitqQueryParams) -> Result<()> {
719        check_error(unsafe {
720            zvec_rust_sys::zvec_group_by_vector_query_set_ivf_rabitq_params(
721                self.handle,
722                params.handle,
723            )
724        })?;
725        // Ownership transferred to query only on success; prevent double-free
726        params.handle = std::ptr::null_mut();
727        Ok(())
728    }
729
730    /// Sets Flat query parameters (takes ownership on success).
731    pub fn set_flat_params(&mut self, mut params: FlatQueryParams) -> Result<()> {
732        check_error(unsafe {
733            zvec_rust_sys::zvec_group_by_vector_query_set_flat_params(self.handle, params.handle)
734        })?;
735        // Ownership transferred to query only on success; prevent double-free
736        params.handle = std::ptr::null_mut();
737        Ok(())
738    }
739
740    /// Sets DiskANN query parameters (takes ownership on success).
741    pub fn set_diskann_params(&mut self, mut params: DiskannQueryParams) -> Result<()> {
742        check_error(unsafe {
743            zvec_rust_sys::zvec_group_by_vector_query_set_diskann_params(self.handle, params.handle)
744        })?;
745        // Ownership transferred to query only on success; prevent double-free
746        params.handle = std::ptr::null_mut();
747        Ok(())
748    }
749}
750
751impl Drop for GroupBySearchQuery {
752    fn drop(&mut self) {
753        if !self.handle.is_null() {
754            unsafe { zvec_rust_sys::zvec_group_by_vector_query_destroy(self.handle) };
755        }
756    }
757}
758
759#[cfg(test)]
760mod tests {
761    use super::*;
762
763    #[test]
764    fn test_vector_query_builder_default_values() {
765        let builder = SearchQueryBuilder::new();
766        assert!(builder.field_name.is_none());
767        assert!(builder.vector.is_none());
768        assert_eq!(builder.topk, 10);
769        assert!(builder.filter.is_none());
770        assert!(builder.include_vector.is_none());
771        assert!(builder.include_doc_id.is_none());
772        assert!(builder.output_fields.is_none());
773    }
774
775    #[test]
776    fn test_vector_query_builder_setters() {
777        let builder = SearchQueryBuilder::new()
778            .field_name("test_field")
779            .vector(&[1.0, 2.0, 3.0])
780            .topk(5)
781            .filter("age > 18")
782            .include_vector(true)
783            .include_doc_id(false)
784            .output_fields(&["name", "age"]);
785
786        assert_eq!(builder.field_name, Some("test_field".to_string()));
787        assert_eq!(builder.vector, Some(vec![1.0, 2.0, 3.0]));
788        assert_eq!(builder.topk, 5);
789        assert_eq!(builder.filter, Some("age > 18".to_string()));
790        assert_eq!(builder.include_vector, Some(true));
791        assert_eq!(builder.include_doc_id, Some(false));
792        assert_eq!(
793            builder.output_fields,
794            Some(vec!["name".to_string(), "age".to_string()])
795        );
796    }
797
798    #[test]
799    fn test_vector_query_builder_build_missing_field_name() {
800        let builder = SearchQueryBuilder::new().vector(&[1.0, 2.0, 3.0]);
801
802        let result = builder.build();
803        assert!(result.is_err());
804        if let Err(e) = result {
805            assert_eq!(e.code, ErrorCode::InvalidArgument);
806            assert!(e.message.contains("field_name is required"));
807        }
808    }
809
810    #[test]
811    fn test_vector_query_builder_build_missing_vector() {
812        let builder = SearchQueryBuilder::new().field_name("test_field");
813
814        let result = builder.build();
815        assert!(result.is_err());
816        if let Err(e) = result {
817            assert_eq!(e.code, ErrorCode::InvalidArgument);
818            assert!(e.message.contains("vector is required"));
819        }
820    }
821
822    #[test]
823    fn test_vector_query_builder_builder_method() {
824        let builder = SearchQuery::builder();
825        assert!(builder.field_name.is_none());
826        assert!(builder.vector.is_none());
827        assert_eq!(builder.topk, 10);
828    }
829
830    #[test]
831    fn test_vector_query_builder_partial_setters() {
832        let builder = SearchQueryBuilder::new()
833            .field_name("test_field")
834            .vector(&[1.0, 2.0])
835            .filter("status = 'active'");
836        assert_eq!(builder.field_name, Some("test_field".to_string()));
837        assert_eq!(builder.vector, Some(vec![1.0, 2.0]));
838        assert_eq!(builder.filter, Some("status = 'active'".to_string()));
839        assert!(builder.include_vector.is_none());
840        assert!(builder.output_fields.is_none());
841    }
842
843    #[test]
844    fn test_vector_query_builder_empty_vector() {
845        let builder = SearchQueryBuilder::new()
846            .field_name("test_field")
847            .vector(&[]);
848        assert_eq!(builder.vector, Some(vec![]));
849    }
850
851    #[test]
852    fn test_vector_query_builder_empty_output_fields() {
853        let builder = SearchQueryBuilder::new()
854            .field_name("test_field")
855            .vector(&[1.0])
856            .output_fields(&[]);
857        assert_eq!(builder.output_fields, Some(vec![]));
858    }
859
860    #[test]
861    fn test_vector_query_builder_topk_zero() {
862        let builder = SearchQueryBuilder::new()
863            .field_name("test_field")
864            .vector(&[1.0])
865            .topk(0);
866        assert_eq!(builder.topk, 0);
867    }
868
869    #[test]
870    fn test_vector_query_builder_topk_negative() {
871        let builder = SearchQueryBuilder::new()
872            .field_name("test_field")
873            .vector(&[1.0])
874            .topk(-1);
875        assert_eq!(builder.topk, -1);
876    }
877
878    #[test]
879    fn test_vector_query_builder_overwrite_field_name() {
880        let builder = SearchQueryBuilder::new()
881            .field_name("first_field")
882            .field_name("second_field");
883        assert_eq!(builder.field_name, Some("second_field".to_string()));
884    }
885
886    #[test]
887    fn test_vector_query_builder_large_vector() {
888        let large_vector: Vec<f32> = (0..1024).map(|i| i as f32).collect();
889        let builder = SearchQueryBuilder::new()
890            .field_name("test_field")
891            .vector(&large_vector);
892        assert_eq!(builder.vector.as_ref().unwrap().len(), 1024);
893    }
894
895    #[test]
896    fn test_fts_query_no_vector() {
897        // A pure FTS query must build without a query vector; the builder path
898        // (which requires `vector`) cannot express this.
899        let mut fts = Fts::new().expect("create fts payload");
900        fts.set_match_string("hello world")
901            .expect("set match string");
902
903        let query = SearchQuery::fts("content", &fts, 10);
904        assert!(
905            query.is_ok(),
906            "pure FTS query should build without a vector"
907        );
908        assert!(!unsafe { query.unwrap().as_raw() }.is_null());
909    }
910
911    #[test]
912    fn test_diskann_query_params_create_and_getters() {
913        let mut params = DiskannQueryParams::new(200);
914        assert_eq!(params.list_size(), 200);
915        params.set_list_size(300).expect("set list_size");
916        assert_eq!(params.list_size(), 300);
917    }
918
919    #[test]
920    fn test_ivf_rabitq_query_params_create_and_getters() {
921        let mut params = IvfRabitqQueryParams::new(16, 0.0, false, true);
922        assert_eq!(params.nprobe(), 16);
923        params.set_nprobe(32).expect("set nprobe");
924        assert_eq!(params.nprobe(), 32);
925        params.set_scale_factor(2.5).expect("set scale_factor");
926        assert!((params.scale_factor() - 2.5).abs() < f32::EPSILON);
927    }
928}