1use std::os::raw::c_void;
2
3use crate::error::{check_error, to_cstring, Error, ErrorCode, Result};
4
5pub struct HnswQueryParams {
7 pub(crate) handle: *mut zvec_rust_sys::zvec_hnsw_query_params_t,
8}
9
10impl HnswQueryParams {
11 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 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 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
38pub struct IvfQueryParams {
40 pub(crate) handle: *mut zvec_rust_sys::zvec_ivf_query_params_t,
41}
42
43impl IvfQueryParams {
44 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 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 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
71pub struct FlatQueryParams {
73 pub(crate) handle: *mut zvec_rust_sys::zvec_flat_query_params_t,
74}
75
76impl FlatQueryParams {
77 pub fn new(is_using_refiner: bool, scale_factor: f32) -> Self {
79 let handle =
80 unsafe { zvec_rust_sys::zvec_query_params_flat_create(is_using_refiner, scale_factor) };
81 FlatQueryParams { handle }
82 }
83}
84
85impl Drop for FlatQueryParams {
86 fn drop(&mut self) {
87 if !self.handle.is_null() {
88 unsafe { zvec_rust_sys::zvec_query_params_flat_destroy(self.handle) };
89 }
90 }
91}
92
93pub struct FtsQueryParams {
95 pub(crate) handle: *mut zvec_rust_sys::zvec_fts_query_params_t,
96}
97
98impl FtsQueryParams {
99 pub fn new(default_operator: Option<&str>) -> Result<Self> {
104 let c_op = default_operator.map(to_cstring).transpose()?;
105 let handle = unsafe {
106 zvec_rust_sys::zvec_query_params_fts_create(
107 c_op.as_ref().map_or(std::ptr::null(), |c| c.as_ptr()),
108 )
109 };
110 if handle.is_null() {
111 return Err(Error {
112 code: ErrorCode::InternalError,
113 message: "failed to create FTS query params".into(),
114 });
115 }
116 Ok(FtsQueryParams { handle })
117 }
118
119 pub fn set_default_operator(&mut self, op: &str) -> Result<()> {
121 let c_op = to_cstring(op)?;
122 check_error(unsafe {
123 zvec_rust_sys::zvec_query_params_fts_set_default_operator(self.handle, c_op.as_ptr())
124 })
125 }
126
127 pub fn default_operator(&self) -> Option<String> {
129 unsafe {
130 let ptr = zvec_rust_sys::zvec_query_params_fts_get_default_operator(self.handle);
131 if ptr.is_null() {
132 return None;
133 }
134 Some(std::ffi::CStr::from_ptr(ptr).to_string_lossy().into_owned())
135 }
136 }
137}
138
139impl Drop for FtsQueryParams {
140 fn drop(&mut self) {
141 if !self.handle.is_null() {
142 unsafe { zvec_rust_sys::zvec_query_params_fts_destroy(self.handle) };
143 }
144 }
145}
146
147pub struct Fts {
152 pub(crate) handle: *mut zvec_rust_sys::zvec_fts_t,
153}
154
155impl Fts {
156 pub fn new() -> Result<Self> {
158 let handle = unsafe { zvec_rust_sys::zvec_fts_create() };
159 if handle.is_null() {
160 return Err(Error {
161 code: ErrorCode::InternalError,
162 message: "failed to create FTS payload".into(),
163 });
164 }
165 Ok(Fts { handle })
166 }
167
168 pub fn set_query_string(&mut self, query: &str) -> Result<()> {
170 let c = to_cstring(query)?;
171 check_error(unsafe { zvec_rust_sys::zvec_fts_set_query_string(self.handle, c.as_ptr()) })
172 }
173
174 pub fn set_match_string(&mut self, match_str: &str) -> Result<()> {
176 let c = to_cstring(match_str)?;
177 check_error(unsafe { zvec_rust_sys::zvec_fts_set_match_string(self.handle, c.as_ptr()) })
178 }
179
180 pub fn query_string(&self) -> Option<String> {
182 unsafe {
183 let ptr = zvec_rust_sys::zvec_fts_get_query_string(self.handle);
184 if ptr.is_null() {
185 return None;
186 }
187 Some(std::ffi::CStr::from_ptr(ptr).to_string_lossy().into_owned())
188 }
189 }
190
191 pub fn match_string(&self) -> Option<String> {
193 unsafe {
194 let ptr = zvec_rust_sys::zvec_fts_get_match_string(self.handle);
195 if ptr.is_null() {
196 return None;
197 }
198 Some(std::ffi::CStr::from_ptr(ptr).to_string_lossy().into_owned())
199 }
200 }
201}
202
203impl Drop for Fts {
204 fn drop(&mut self) {
205 if !self.handle.is_null() {
206 unsafe { zvec_rust_sys::zvec_fts_destroy(self.handle) };
207 }
208 }
209}
210
211pub struct SearchQuery {
213 pub(crate) handle: *mut zvec_rust_sys::zvec_vector_query_t,
214}
215
216impl SearchQuery {
217 pub unsafe fn as_raw(&self) -> *mut zvec_rust_sys::zvec_vector_query_t {
222 self.handle
223 }
224
225 pub unsafe fn from_raw(handle: *mut zvec_rust_sys::zvec_vector_query_t) -> Self {
231 SearchQuery { handle }
232 }
233
234 pub fn new(field_name: &str, vector: &[f32], topk: i32) -> Result<Self> {
236 let handle = unsafe { zvec_rust_sys::zvec_vector_query_create() };
237 if handle.is_null() {
238 return Err(Error {
239 code: ErrorCode::InternalError,
240 message: "failed to create vector query".into(),
241 });
242 }
243
244 let c_field = to_cstring(field_name)?;
245 let query = SearchQuery { handle };
246
247 check_error(unsafe {
248 zvec_rust_sys::zvec_vector_query_set_field_name(query.handle, c_field.as_ptr())
249 })?;
250 check_error(unsafe {
251 zvec_rust_sys::zvec_vector_query_set_query_vector(
252 query.handle,
253 vector.as_ptr() as *const c_void,
254 std::mem::size_of_val(vector),
255 )
256 })?;
257 check_error(unsafe { zvec_rust_sys::zvec_vector_query_set_topk(query.handle, topk) })?;
258
259 Ok(query)
260 }
261
262 pub fn fts(field_name: &str, fts: &Fts, topk: i32) -> Result<Self> {
272 let handle = unsafe { zvec_rust_sys::zvec_vector_query_create() };
273 if handle.is_null() {
274 return Err(Error {
275 code: ErrorCode::InternalError,
276 message: "failed to create vector query".into(),
277 });
278 }
279
280 let mut query = SearchQuery { handle };
282 let c_field = to_cstring(field_name)?;
283
284 check_error(unsafe {
285 zvec_rust_sys::zvec_vector_query_set_field_name(query.handle, c_field.as_ptr())
286 })?;
287 check_error(unsafe { zvec_rust_sys::zvec_vector_query_set_topk(query.handle, topk) })?;
288 query.set_fts(fts)?;
289
290 Ok(query)
291 }
292
293 pub fn builder() -> SearchQueryBuilder {
295 SearchQueryBuilder::new()
296 }
297
298 pub fn set_filter(&mut self, filter: &str) -> Result<()> {
300 let c_filter = to_cstring(filter)?;
301 check_error(unsafe {
302 zvec_rust_sys::zvec_vector_query_set_filter(self.handle, c_filter.as_ptr())
303 })
304 }
305
306 pub fn set_include_vector(&mut self, include: bool) -> Result<()> {
308 check_error(unsafe {
309 zvec_rust_sys::zvec_vector_query_set_include_vector(self.handle, include)
310 })
311 }
312
313 pub fn set_include_doc_id(&mut self, include: bool) -> Result<()> {
315 check_error(unsafe {
316 zvec_rust_sys::zvec_vector_query_set_include_doc_id(self.handle, include)
317 })
318 }
319
320 pub fn set_output_fields(&mut self, fields: &[&str]) -> Result<()> {
322 let c_fields: Vec<_> = fields
323 .iter()
324 .map(|f| to_cstring(f))
325 .collect::<Result<Vec<_>>>()?;
326 let c_ptrs: Vec<_> = c_fields.iter().map(|f| f.as_ptr()).collect();
327 check_error(unsafe {
328 zvec_rust_sys::zvec_vector_query_set_output_fields(
329 self.handle,
330 c_ptrs.as_ptr(),
331 c_ptrs.len(),
332 )
333 })
334 }
335
336 pub fn set_hnsw_params(&mut self, mut params: HnswQueryParams) -> Result<()> {
338 check_error(unsafe {
339 zvec_rust_sys::zvec_vector_query_set_hnsw_params(self.handle, params.handle)
340 })?;
341 params.handle = std::ptr::null_mut();
343 Ok(())
344 }
345
346 pub fn set_ivf_params(&mut self, mut params: IvfQueryParams) -> Result<()> {
348 check_error(unsafe {
349 zvec_rust_sys::zvec_vector_query_set_ivf_params(self.handle, params.handle)
350 })?;
351 params.handle = std::ptr::null_mut();
352 Ok(())
353 }
354
355 pub fn set_flat_params(&mut self, mut params: FlatQueryParams) -> Result<()> {
357 check_error(unsafe {
358 zvec_rust_sys::zvec_vector_query_set_flat_params(self.handle, params.handle)
359 })?;
360 params.handle = std::ptr::null_mut();
361 Ok(())
362 }
363
364 pub fn set_fts_params(&mut self, mut params: FtsQueryParams) -> Result<()> {
366 check_error(unsafe {
367 zvec_rust_sys::zvec_vector_query_set_fts_params(self.handle, params.handle)
368 })?;
369 params.handle = std::ptr::null_mut();
370 Ok(())
371 }
372
373 pub fn set_fts(&mut self, fts: &Fts) -> Result<()> {
375 check_error(unsafe { zvec_rust_sys::zvec_vector_query_set_fts(self.handle, fts.handle) })
376 }
377}
378
379impl Drop for SearchQuery {
380 fn drop(&mut self) {
381 if !self.handle.is_null() {
382 unsafe { zvec_rust_sys::zvec_vector_query_destroy(self.handle) };
383 }
384 }
385}
386
387pub struct SearchQueryBuilder {
389 field_name: Option<String>,
390 vector: Option<Vec<f32>>,
391 topk: i32,
392 filter: Option<String>,
393 include_vector: Option<bool>,
394 include_doc_id: Option<bool>,
395 output_fields: Option<Vec<String>>,
396 fts_query_string: Option<String>,
397 fts_match_string: Option<String>,
398}
399
400impl SearchQueryBuilder {
401 fn new() -> Self {
402 SearchQueryBuilder {
403 field_name: None,
404 vector: None,
405 topk: 10,
406 filter: None,
407 include_vector: None,
408 include_doc_id: None,
409 output_fields: None,
410 fts_query_string: None,
411 fts_match_string: None,
412 }
413 }
414
415 pub fn field_name(mut self, name: &str) -> Self {
417 self.field_name = Some(name.to_string());
418 self
419 }
420
421 pub fn vector(mut self, vector: &[f32]) -> Self {
423 self.vector = Some(vector.to_vec());
424 self
425 }
426
427 pub fn topk(mut self, topk: i32) -> Self {
429 self.topk = topk;
430 self
431 }
432
433 pub fn filter(mut self, filter: &str) -> Self {
435 self.filter = Some(filter.to_string());
436 self
437 }
438
439 pub fn include_vector(mut self, include: bool) -> Self {
441 self.include_vector = Some(include);
442 self
443 }
444
445 pub fn include_doc_id(mut self, include: bool) -> Self {
447 self.include_doc_id = Some(include);
448 self
449 }
450
451 pub fn output_fields(mut self, fields: &[&str]) -> Self {
453 self.output_fields = Some(fields.iter().map(|s| s.to_string()).collect());
454 self
455 }
456
457 pub fn fts_query_string(mut self, query: &str) -> Self {
459 self.fts_query_string = Some(query.to_string());
460 self
461 }
462
463 pub fn fts_match_string(mut self, match_str: &str) -> Self {
465 self.fts_match_string = Some(match_str.to_string());
466 self
467 }
468
469 pub fn build(self) -> Result<SearchQuery> {
471 let field_name = self.field_name.ok_or_else(|| Error {
472 code: ErrorCode::InvalidArgument,
473 message: "field_name is required".into(),
474 })?;
475 let vector = self.vector.ok_or_else(|| Error {
476 code: ErrorCode::InvalidArgument,
477 message: "vector is required".into(),
478 })?;
479
480 let mut query = SearchQuery::new(&field_name, &vector, self.topk)?;
481
482 if let Some(filter) = &self.filter {
483 query.set_filter(filter)?;
484 }
485 if let Some(include) = self.include_vector {
486 query.set_include_vector(include)?;
487 }
488 if let Some(include) = self.include_doc_id {
489 query.set_include_doc_id(include)?;
490 }
491 if let Some(fields) = &self.output_fields {
492 let field_refs: Vec<&str> = fields.iter().map(|s| s.as_str()).collect();
493 query.set_output_fields(&field_refs)?;
494 }
495 if self.fts_query_string.is_some() || self.fts_match_string.is_some() {
496 let mut fts = Fts::new()?;
497 if let Some(qs) = &self.fts_query_string {
498 fts.set_query_string(qs)?;
499 }
500 if let Some(ms) = &self.fts_match_string {
501 fts.set_match_string(ms)?;
502 }
503 query.set_fts(&fts)?;
504 }
505
506 Ok(query)
507 }
508}
509
510pub struct GroupBySearchQuery {
512 pub(crate) handle: *mut zvec_rust_sys::zvec_group_by_vector_query_t,
513}
514
515impl GroupBySearchQuery {
516 pub fn new(
518 field_name: &str,
519 group_by_field: &str,
520 vector: &[f32],
521 group_count: u32,
522 group_topk: u32,
523 ) -> Result<Self> {
524 let handle = unsafe { zvec_rust_sys::zvec_group_by_vector_query_create() };
525 if handle.is_null() {
526 return Err(Error {
527 code: ErrorCode::InternalError,
528 message: "failed to create group by vector query".into(),
529 });
530 }
531
532 let c_field = to_cstring(field_name)?;
533 let c_group_field = to_cstring(group_by_field)?;
534
535 check_error(unsafe {
536 zvec_rust_sys::zvec_group_by_vector_query_set_field_name(handle, c_field.as_ptr())
537 })?;
538 check_error(unsafe {
539 zvec_rust_sys::zvec_group_by_vector_query_set_group_by_field_name(
540 handle,
541 c_group_field.as_ptr(),
542 )
543 })?;
544 check_error(unsafe {
545 zvec_rust_sys::zvec_group_by_vector_query_set_query_vector(
546 handle,
547 vector.as_ptr() as *const c_void,
548 std::mem::size_of_val(vector),
549 )
550 })?;
551 check_error(unsafe {
552 zvec_rust_sys::zvec_group_by_vector_query_set_group_count(handle, group_count)
553 })?;
554 check_error(unsafe {
555 zvec_rust_sys::zvec_group_by_vector_query_set_group_topk(handle, group_topk)
556 })?;
557
558 Ok(GroupBySearchQuery { handle })
559 }
560
561 pub fn set_filter(&mut self, filter: &str) -> Result<()> {
563 let c_filter = to_cstring(filter)?;
564 check_error(unsafe {
565 zvec_rust_sys::zvec_group_by_vector_query_set_filter(self.handle, c_filter.as_ptr())
566 })
567 }
568
569 pub fn set_include_vector(&mut self, include: bool) -> Result<()> {
571 check_error(unsafe {
572 zvec_rust_sys::zvec_group_by_vector_query_set_include_vector(self.handle, include)
573 })
574 }
575
576 pub fn set_output_fields(&mut self, fields: &[&str]) -> Result<()> {
578 let c_fields: Vec<_> = fields
579 .iter()
580 .map(|f| to_cstring(f))
581 .collect::<Result<Vec<_>>>()?;
582 let c_ptrs: Vec<_> = c_fields.iter().map(|f| f.as_ptr()).collect();
583 check_error(unsafe {
584 zvec_rust_sys::zvec_group_by_vector_query_set_output_fields(
585 self.handle,
586 c_ptrs.as_ptr(),
587 c_ptrs.len(),
588 )
589 })
590 }
591
592 pub fn set_hnsw_params(&mut self, mut params: HnswQueryParams) -> Result<()> {
594 check_error(unsafe {
595 zvec_rust_sys::zvec_group_by_vector_query_set_hnsw_params(self.handle, params.handle)
596 })?;
597 params.handle = std::ptr::null_mut();
599 Ok(())
600 }
601
602 pub fn set_ivf_params(&mut self, mut params: IvfQueryParams) -> Result<()> {
604 check_error(unsafe {
605 zvec_rust_sys::zvec_group_by_vector_query_set_ivf_params(self.handle, params.handle)
606 })?;
607 params.handle = std::ptr::null_mut();
609 Ok(())
610 }
611
612 pub fn set_flat_params(&mut self, mut params: FlatQueryParams) -> Result<()> {
614 check_error(unsafe {
615 zvec_rust_sys::zvec_group_by_vector_query_set_flat_params(self.handle, params.handle)
616 })?;
617 params.handle = std::ptr::null_mut();
619 Ok(())
620 }
621}
622
623impl Drop for GroupBySearchQuery {
624 fn drop(&mut self) {
625 if !self.handle.is_null() {
626 unsafe { zvec_rust_sys::zvec_group_by_vector_query_destroy(self.handle) };
627 }
628 }
629}
630
631#[cfg(test)]
632mod tests {
633 use super::*;
634
635 #[test]
636 fn test_vector_query_builder_default_values() {
637 let builder = SearchQueryBuilder::new();
638 assert!(builder.field_name.is_none());
639 assert!(builder.vector.is_none());
640 assert_eq!(builder.topk, 10);
641 assert!(builder.filter.is_none());
642 assert!(builder.include_vector.is_none());
643 assert!(builder.include_doc_id.is_none());
644 assert!(builder.output_fields.is_none());
645 }
646
647 #[test]
648 fn test_vector_query_builder_setters() {
649 let builder = SearchQueryBuilder::new()
650 .field_name("test_field")
651 .vector(&[1.0, 2.0, 3.0])
652 .topk(5)
653 .filter("age > 18")
654 .include_vector(true)
655 .include_doc_id(false)
656 .output_fields(&["name", "age"]);
657
658 assert_eq!(builder.field_name, Some("test_field".to_string()));
659 assert_eq!(builder.vector, Some(vec![1.0, 2.0, 3.0]));
660 assert_eq!(builder.topk, 5);
661 assert_eq!(builder.filter, Some("age > 18".to_string()));
662 assert_eq!(builder.include_vector, Some(true));
663 assert_eq!(builder.include_doc_id, Some(false));
664 assert_eq!(
665 builder.output_fields,
666 Some(vec!["name".to_string(), "age".to_string()])
667 );
668 }
669
670 #[test]
671 fn test_vector_query_builder_build_missing_field_name() {
672 let builder = SearchQueryBuilder::new().vector(&[1.0, 2.0, 3.0]);
673
674 let result = builder.build();
675 assert!(result.is_err());
676 if let Err(e) = result {
677 assert_eq!(e.code, ErrorCode::InvalidArgument);
678 assert!(e.message.contains("field_name is required"));
679 }
680 }
681
682 #[test]
683 fn test_vector_query_builder_build_missing_vector() {
684 let builder = SearchQueryBuilder::new().field_name("test_field");
685
686 let result = builder.build();
687 assert!(result.is_err());
688 if let Err(e) = result {
689 assert_eq!(e.code, ErrorCode::InvalidArgument);
690 assert!(e.message.contains("vector is required"));
691 }
692 }
693
694 #[test]
695 fn test_vector_query_builder_builder_method() {
696 let builder = SearchQuery::builder();
697 assert!(builder.field_name.is_none());
698 assert!(builder.vector.is_none());
699 assert_eq!(builder.topk, 10);
700 }
701
702 #[test]
703 fn test_vector_query_builder_partial_setters() {
704 let builder = SearchQueryBuilder::new()
705 .field_name("test_field")
706 .vector(&[1.0, 2.0])
707 .filter("status = 'active'");
708 assert_eq!(builder.field_name, Some("test_field".to_string()));
709 assert_eq!(builder.vector, Some(vec![1.0, 2.0]));
710 assert_eq!(builder.filter, Some("status = 'active'".to_string()));
711 assert!(builder.include_vector.is_none());
712 assert!(builder.output_fields.is_none());
713 }
714
715 #[test]
716 fn test_vector_query_builder_empty_vector() {
717 let builder = SearchQueryBuilder::new()
718 .field_name("test_field")
719 .vector(&[]);
720 assert_eq!(builder.vector, Some(vec![]));
721 }
722
723 #[test]
724 fn test_vector_query_builder_empty_output_fields() {
725 let builder = SearchQueryBuilder::new()
726 .field_name("test_field")
727 .vector(&[1.0])
728 .output_fields(&[]);
729 assert_eq!(builder.output_fields, Some(vec![]));
730 }
731
732 #[test]
733 fn test_vector_query_builder_topk_zero() {
734 let builder = SearchQueryBuilder::new()
735 .field_name("test_field")
736 .vector(&[1.0])
737 .topk(0);
738 assert_eq!(builder.topk, 0);
739 }
740
741 #[test]
742 fn test_vector_query_builder_topk_negative() {
743 let builder = SearchQueryBuilder::new()
744 .field_name("test_field")
745 .vector(&[1.0])
746 .topk(-1);
747 assert_eq!(builder.topk, -1);
748 }
749
750 #[test]
751 fn test_vector_query_builder_overwrite_field_name() {
752 let builder = SearchQueryBuilder::new()
753 .field_name("first_field")
754 .field_name("second_field");
755 assert_eq!(builder.field_name, Some("second_field".to_string()));
756 }
757
758 #[test]
759 fn test_vector_query_builder_large_vector() {
760 let large_vector: Vec<f32> = (0..1024).map(|i| i as f32).collect();
761 let builder = SearchQueryBuilder::new()
762 .field_name("test_field")
763 .vector(&large_vector);
764 assert_eq!(builder.vector.as_ref().unwrap().len(), 1024);
765 }
766
767 #[test]
768 fn test_fts_query_no_vector() {
769 let mut fts = Fts::new().expect("create fts payload");
772 fts.set_match_string("hello world")
773 .expect("set match string");
774
775 let query = SearchQuery::fts("content", &fts, 10);
776 assert!(
777 query.is_ok(),
778 "pure FTS query should build without a vector"
779 );
780 assert!(!unsafe { query.unwrap().as_raw() }.is_null());
781 }
782}