1use std::ffi::CStr;
2use std::os::raw::c_void;
3
4use crate::error::{check_error, to_cstring, Error, ErrorCode, Result};
5use crate::types::DataType;
6
7pub struct Doc {
12 pub(crate) handle: *mut zvec_rust_sys::zvec_doc_t,
13 owned: bool,
14}
15
16impl Doc {
17 pub unsafe fn as_raw(&self) -> *mut zvec_rust_sys::zvec_doc_t {
22 self.handle
23 }
24
25 pub unsafe fn from_raw(handle: *mut zvec_rust_sys::zvec_doc_t) -> Self {
31 Doc {
32 handle,
33 owned: true,
34 }
35 }
36
37 pub fn new() -> Result<Self> {
39 let handle = unsafe { zvec_rust_sys::zvec_doc_create() };
40 if handle.is_null() {
41 return Err(Error {
42 code: ErrorCode::InternalError,
43 message: "failed to create document".into(),
44 });
45 }
46 Ok(Doc {
47 handle,
48 owned: true,
49 })
50 }
51
52 #[allow(dead_code)]
54 pub(crate) fn from_borrowed(handle: *mut zvec_rust_sys::zvec_doc_t) -> Self {
55 Doc {
56 handle,
57 owned: false,
58 }
59 }
60
61 pub fn set_pk(&mut self, pk: &str) {
63 let c_pk = to_cstring(pk).expect("pk must not contain null bytes");
64 unsafe { zvec_rust_sys::zvec_doc_set_pk(self.handle, c_pk.as_ptr()) };
65 }
66
67 pub fn get_pk(&self) -> Option<&str> {
69 unsafe {
70 let ptr = zvec_rust_sys::zvec_doc_get_pk_pointer(self.handle);
71 if ptr.is_null() {
72 None
73 } else {
74 CStr::from_ptr(ptr).to_str().ok()
75 }
76 }
77 }
78
79 pub fn get_score(&self) -> f32 {
81 unsafe { zvec_rust_sys::zvec_doc_get_score(self.handle) }
82 }
83
84 #[allow(dead_code)]
85 pub(crate) fn get_doc_id(&self) -> u64 {
86 unsafe { zvec_rust_sys::zvec_doc_get_doc_id(self.handle) }
87 }
88
89 pub fn field_count(&self) -> usize {
91 unsafe { zvec_rust_sys::zvec_doc_get_field_count(self.handle) }
92 }
93
94 pub fn is_empty(&self) -> bool {
96 unsafe { zvec_rust_sys::zvec_doc_is_empty(self.handle) }
97 }
98
99 pub fn has_field(&self, name: &str) -> bool {
101 let c_name = match to_cstring(name) {
102 Ok(s) => s,
103 Err(_) => return false,
104 };
105 unsafe { zvec_rust_sys::zvec_doc_has_field(self.handle, c_name.as_ptr()) }
106 }
107
108 pub fn is_field_null(&self, name: &str) -> bool {
110 let c_name = match to_cstring(name) {
111 Ok(s) => s,
112 Err(_) => return false,
113 };
114 unsafe { zvec_rust_sys::zvec_doc_is_field_null(self.handle, c_name.as_ptr()) }
115 }
116
117 pub fn add_string(&mut self, name: &str, value: &str) -> Result<()> {
123 let c_name = to_cstring(name)?;
124 let c_value = to_cstring(value)?;
125 let bytes = c_value.as_bytes();
130 check_error(unsafe {
131 zvec_rust_sys::zvec_doc_add_field_by_value(
132 self.handle,
133 c_name.as_ptr(),
134 DataType::String as u32,
135 bytes.as_ptr() as *const c_void,
136 bytes.len(),
137 )
138 })
139 }
140
141 pub fn add_bool(&mut self, name: &str, value: bool) -> Result<()> {
143 let c_name = to_cstring(name)?;
144 check_error(unsafe {
145 zvec_rust_sys::zvec_doc_add_field_by_value(
146 self.handle,
147 c_name.as_ptr(),
148 DataType::Bool as u32,
149 &value as *const bool as *const c_void,
150 std::mem::size_of::<bool>(),
151 )
152 })
153 }
154
155 pub fn add_i32(&mut self, name: &str, value: i32) -> Result<()> {
157 let c_name = to_cstring(name)?;
158 check_error(unsafe {
159 zvec_rust_sys::zvec_doc_add_field_by_value(
160 self.handle,
161 c_name.as_ptr(),
162 DataType::Int32 as u32,
163 &value as *const i32 as *const c_void,
164 std::mem::size_of::<i32>(),
165 )
166 })
167 }
168
169 pub fn add_i64(&mut self, name: &str, value: i64) -> Result<()> {
171 let c_name = to_cstring(name)?;
172 check_error(unsafe {
173 zvec_rust_sys::zvec_doc_add_field_by_value(
174 self.handle,
175 c_name.as_ptr(),
176 DataType::Int64 as u32,
177 &value as *const i64 as *const c_void,
178 std::mem::size_of::<i64>(),
179 )
180 })
181 }
182
183 pub fn add_u32(&mut self, name: &str, value: u32) -> Result<()> {
185 let c_name = to_cstring(name)?;
186 check_error(unsafe {
187 zvec_rust_sys::zvec_doc_add_field_by_value(
188 self.handle,
189 c_name.as_ptr(),
190 DataType::Uint32 as u32,
191 &value as *const u32 as *const c_void,
192 std::mem::size_of::<u32>(),
193 )
194 })
195 }
196
197 pub fn add_u64(&mut self, name: &str, value: u64) -> Result<()> {
199 let c_name = to_cstring(name)?;
200 check_error(unsafe {
201 zvec_rust_sys::zvec_doc_add_field_by_value(
202 self.handle,
203 c_name.as_ptr(),
204 DataType::Uint64 as u32,
205 &value as *const u64 as *const c_void,
206 std::mem::size_of::<u64>(),
207 )
208 })
209 }
210
211 pub fn add_f32(&mut self, name: &str, value: f32) -> Result<()> {
213 let c_name = to_cstring(name)?;
214 check_error(unsafe {
215 zvec_rust_sys::zvec_doc_add_field_by_value(
216 self.handle,
217 c_name.as_ptr(),
218 DataType::Float as u32,
219 &value as *const f32 as *const c_void,
220 std::mem::size_of::<f32>(),
221 )
222 })
223 }
224
225 pub fn add_f64(&mut self, name: &str, value: f64) -> Result<()> {
227 let c_name = to_cstring(name)?;
228 check_error(unsafe {
229 zvec_rust_sys::zvec_doc_add_field_by_value(
230 self.handle,
231 c_name.as_ptr(),
232 DataType::Double as u32,
233 &value as *const f64 as *const c_void,
234 std::mem::size_of::<f64>(),
235 )
236 })
237 }
238
239 pub fn add_vector_f32(&mut self, name: &str, vector: &[f32]) -> Result<()> {
241 let c_name = to_cstring(name)?;
242 check_error(unsafe {
243 zvec_rust_sys::zvec_doc_add_field_by_value(
244 self.handle,
245 c_name.as_ptr(),
246 DataType::VectorFp32 as u32,
247 vector.as_ptr() as *const c_void,
248 std::mem::size_of_val(vector),
249 )
250 })
251 }
252
253 pub fn add_vector_f64(&mut self, name: &str, vector: &[f64]) -> Result<()> {
255 let c_name = to_cstring(name)?;
256 check_error(unsafe {
257 zvec_rust_sys::zvec_doc_add_field_by_value(
258 self.handle,
259 c_name.as_ptr(),
260 DataType::VectorFp64 as u32,
261 vector.as_ptr() as *const c_void,
262 std::mem::size_of_val(vector),
263 )
264 })
265 }
266
267 pub fn add_binary(&mut self, name: &str, value: &[u8]) -> Result<()> {
269 let c_name = to_cstring(name)?;
270 check_error(unsafe {
271 zvec_rust_sys::zvec_doc_add_field_by_value(
272 self.handle,
273 c_name.as_ptr(),
274 DataType::Binary as u32,
275 value.as_ptr() as *const c_void,
276 value.len(),
277 )
278 })
279 }
280
281 pub fn add_vector_i8(&mut self, name: &str, vector: &[i8]) -> Result<()> {
283 let c_name = to_cstring(name)?;
284 check_error(unsafe {
285 zvec_rust_sys::zvec_doc_add_field_by_value(
286 self.handle,
287 c_name.as_ptr(),
288 DataType::VectorInt8 as u32,
289 vector.as_ptr() as *const c_void,
290 std::mem::size_of_val(vector),
291 )
292 })
293 }
294
295 pub fn add_vector_i16(&mut self, name: &str, vector: &[i16]) -> Result<()> {
297 let c_name = to_cstring(name)?;
298 check_error(unsafe {
299 zvec_rust_sys::zvec_doc_add_field_by_value(
300 self.handle,
301 c_name.as_ptr(),
302 DataType::VectorInt16 as u32,
303 vector.as_ptr() as *const c_void,
304 std::mem::size_of_val(vector),
305 )
306 })
307 }
308
309 fn add_typed_array<T>(&mut self, name: &str, data_type: DataType, values: &[T]) -> Result<()> {
314 let c_name = to_cstring(name)?;
315 check_error(unsafe {
316 zvec_rust_sys::zvec_doc_add_field_by_value(
317 self.handle,
318 c_name.as_ptr(),
319 data_type as u32,
320 values.as_ptr() as *const c_void,
321 std::mem::size_of_val(values),
322 )
323 })
324 }
325
326 pub fn add_array_i32(&mut self, name: &str, values: &[i32]) -> Result<()> {
328 self.add_typed_array(name, DataType::ArrayInt32, values)
329 }
330
331 pub fn add_array_i64(&mut self, name: &str, values: &[i64]) -> Result<()> {
333 self.add_typed_array(name, DataType::ArrayInt64, values)
334 }
335
336 pub fn add_array_u32(&mut self, name: &str, values: &[u32]) -> Result<()> {
338 self.add_typed_array(name, DataType::ArrayUint32, values)
339 }
340
341 pub fn add_array_u64(&mut self, name: &str, values: &[u64]) -> Result<()> {
343 self.add_typed_array(name, DataType::ArrayUint64, values)
344 }
345
346 pub fn add_array_f32(&mut self, name: &str, values: &[f32]) -> Result<()> {
348 self.add_typed_array(name, DataType::ArrayFloat, values)
349 }
350
351 pub fn add_array_f64(&mut self, name: &str, values: &[f64]) -> Result<()> {
353 self.add_typed_array(name, DataType::ArrayDouble, values)
354 }
355
356 pub fn add_array_bool(&mut self, name: &str, values: &[bool]) -> Result<()> {
358 self.add_typed_array(name, DataType::ArrayBool, values)
359 }
360
361 pub fn set_field_null(&mut self, name: &str) -> Result<()> {
363 let c_name = to_cstring(name)?;
364 check_error(unsafe { zvec_rust_sys::zvec_doc_set_field_null(self.handle, c_name.as_ptr()) })
365 }
366
367 pub fn remove_field(&mut self, name: &str) -> Result<()> {
369 let c_name = to_cstring(name)?;
370 check_error(unsafe { zvec_rust_sys::zvec_doc_remove_field(self.handle, c_name.as_ptr()) })
371 }
372
373 fn get_basic_field<T: Copy + Default>(
378 &self,
379 name: &str,
380 data_type: DataType,
381 ) -> Result<Option<T>> {
382 if !self.has_field(name) || self.is_field_null(name) {
383 return Ok(None);
384 }
385 let c_name = to_cstring(name)?;
386 let mut value: T = T::default();
387 check_error(unsafe {
388 zvec_rust_sys::zvec_doc_get_field_value_basic(
389 self.handle,
390 c_name.as_ptr(),
391 data_type as u32,
392 &mut value as *mut T as *mut c_void,
393 std::mem::size_of::<T>(),
394 )
395 })?;
396 Ok(Some(value))
397 }
398
399 fn get_pointer_field(
400 &self,
401 name: &str,
402 data_type: DataType,
403 ) -> Result<Option<(*const c_void, usize)>> {
404 let c_name = to_cstring(name)?;
405 let mut value_ptr: *const c_void = std::ptr::null();
406 let mut value_size: usize = 0;
407 check_error(unsafe {
408 zvec_rust_sys::zvec_doc_get_field_value_pointer(
409 self.handle,
410 c_name.as_ptr(),
411 data_type as u32,
412 &mut value_ptr,
413 &mut value_size,
414 )
415 })?;
416 if value_ptr.is_null() || value_size == 0 {
417 return Ok(None);
418 }
419 Ok(Some((value_ptr, value_size)))
420 }
421
422 fn get_typed_vec<T: Copy>(&self, name: &str, data_type: DataType) -> Result<Option<Vec<T>>> {
423 let Some((ptr, size)) = self.get_pointer_field(name, data_type)? else {
424 return Ok(None);
425 };
426 let elem_size = std::mem::size_of::<T>();
427 if elem_size > 1 && size % elem_size != 0 {
428 return Err(Error {
429 code: ErrorCode::InternalError,
430 message: format!(
431 "data size {} is not aligned to element size {}",
432 size, elem_size
433 ),
434 });
435 }
436 let count = size / elem_size;
437 let slice = unsafe { std::slice::from_raw_parts(ptr as *const T, count) };
438 Ok(Some(slice.to_vec()))
439 }
440
441 pub fn get_string(&self, name: &str) -> Result<Option<String>> {
443 let Some((ptr, _size)) = self.get_pointer_field(name, DataType::String)? else {
444 return Ok(None);
445 };
446 unsafe {
447 let cstr = CStr::from_ptr(ptr as *const std::os::raw::c_char);
448 Ok(Some(cstr.to_string_lossy().into_owned()))
449 }
450 }
451
452 pub fn get_bool(&self, name: &str) -> Result<Option<bool>> {
454 self.get_basic_field(name, DataType::Bool)
455 }
456
457 pub fn get_i32(&self, name: &str) -> Result<Option<i32>> {
459 self.get_basic_field(name, DataType::Int32)
460 }
461
462 pub fn get_i64(&self, name: &str) -> Result<Option<i64>> {
464 self.get_basic_field(name, DataType::Int64)
465 }
466
467 pub fn get_u32(&self, name: &str) -> Result<Option<u32>> {
469 self.get_basic_field(name, DataType::Uint32)
470 }
471
472 pub fn get_u64(&self, name: &str) -> Result<Option<u64>> {
474 self.get_basic_field(name, DataType::Uint64)
475 }
476
477 pub fn get_f32(&self, name: &str) -> Result<Option<f32>> {
479 self.get_basic_field(name, DataType::Float)
480 }
481
482 pub fn get_f64(&self, name: &str) -> Result<Option<f64>> {
484 self.get_basic_field(name, DataType::Double)
485 }
486
487 pub fn get_vector_f32(&self, name: &str) -> Result<Option<Vec<f32>>> {
489 self.get_typed_vec(name, DataType::VectorFp32)
490 }
491
492 pub fn get_vector_f64(&self, name: &str) -> Result<Option<Vec<f64>>> {
494 self.get_typed_vec(name, DataType::VectorFp64)
495 }
496
497 pub fn get_binary(&self, name: &str) -> Result<Option<Vec<u8>>> {
499 self.get_typed_vec(name, DataType::Binary)
500 }
501
502 pub fn get_vector_i8(&self, name: &str) -> Result<Option<Vec<i8>>> {
504 self.get_typed_vec(name, DataType::VectorInt8)
505 }
506
507 pub fn get_vector_i16(&self, name: &str) -> Result<Option<Vec<i16>>> {
509 self.get_typed_vec(name, DataType::VectorInt16)
510 }
511
512 pub fn get_array_i32(&self, name: &str) -> Result<Option<Vec<i32>>> {
518 self.get_typed_vec(name, DataType::ArrayInt32)
519 }
520
521 pub fn get_array_i64(&self, name: &str) -> Result<Option<Vec<i64>>> {
523 self.get_typed_vec(name, DataType::ArrayInt64)
524 }
525
526 pub fn get_array_u32(&self, name: &str) -> Result<Option<Vec<u32>>> {
528 self.get_typed_vec(name, DataType::ArrayUint32)
529 }
530
531 pub fn get_array_u64(&self, name: &str) -> Result<Option<Vec<u64>>> {
533 self.get_typed_vec(name, DataType::ArrayUint64)
534 }
535
536 pub fn get_array_f32(&self, name: &str) -> Result<Option<Vec<f32>>> {
538 self.get_typed_vec(name, DataType::ArrayFloat)
539 }
540
541 pub fn get_array_f64(&self, name: &str) -> Result<Option<Vec<f64>>> {
543 self.get_typed_vec(name, DataType::ArrayDouble)
544 }
545
546 pub fn get_array_bool(&self, name: &str) -> Result<Option<Vec<bool>>> {
548 self.get_typed_vec(name, DataType::ArrayBool)
549 }
550
551 pub fn clear(&mut self) {
553 unsafe { zvec_rust_sys::zvec_doc_clear(self.handle) };
554 }
555}
556
557impl Drop for Doc {
558 fn drop(&mut self) {
559 if self.owned && !self.handle.is_null() {
560 unsafe { zvec_rust_sys::zvec_doc_destroy(self.handle) };
561 }
562 }
563}
564
565pub fn free_docs(docs: Vec<Doc>) {
567 drop(docs);
569}
570
571#[cfg(test)]
572mod tests {
573 use super::*;
574
575 #[test]
576 fn from_borrowed_does_not_own() {
577 let doc = Doc::from_borrowed(std::ptr::null_mut());
579 assert!(!doc.owned);
580 assert!(doc.handle.is_null());
581 }
582
583 #[test]
584 fn from_raw_takes_ownership() {
585 let doc = unsafe { Doc::from_raw(std::ptr::null_mut()) };
586 assert!(doc.owned);
587 }
589}