use std::ffi::CStr;
use std::os::raw::c_void;
use crate::error::{check_error, to_cstring, Error, ErrorCode, Result};
use crate::types::DataType;
pub struct Doc {
pub(crate) handle: *mut zvec_rust_sys::zvec_doc_t,
owned: bool,
}
impl Doc {
pub unsafe fn as_raw(&self) -> *mut zvec_rust_sys::zvec_doc_t {
self.handle
}
pub unsafe fn from_raw(handle: *mut zvec_rust_sys::zvec_doc_t) -> Self {
Doc {
handle,
owned: true,
}
}
pub fn new() -> Result<Self> {
let handle = unsafe { zvec_rust_sys::zvec_doc_create() };
if handle.is_null() {
return Err(Error {
code: ErrorCode::InternalError,
message: "failed to create document".into(),
});
}
Ok(Doc {
handle,
owned: true,
})
}
#[allow(dead_code)]
pub(crate) fn from_borrowed(handle: *mut zvec_rust_sys::zvec_doc_t) -> Self {
Doc {
handle,
owned: false,
}
}
pub fn set_pk(&mut self, pk: &str) {
let c_pk = to_cstring(pk).expect("pk must not contain null bytes");
unsafe { zvec_rust_sys::zvec_doc_set_pk(self.handle, c_pk.as_ptr()) };
}
pub fn get_pk(&self) -> Option<&str> {
unsafe {
let ptr = zvec_rust_sys::zvec_doc_get_pk_pointer(self.handle);
if ptr.is_null() {
None
} else {
CStr::from_ptr(ptr).to_str().ok()
}
}
}
pub fn get_score(&self) -> f32 {
unsafe { zvec_rust_sys::zvec_doc_get_score(self.handle) }
}
#[allow(dead_code)]
pub(crate) fn get_doc_id(&self) -> u64 {
unsafe { zvec_rust_sys::zvec_doc_get_doc_id(self.handle) }
}
pub fn field_count(&self) -> usize {
unsafe { zvec_rust_sys::zvec_doc_get_field_count(self.handle) }
}
pub fn is_empty(&self) -> bool {
unsafe { zvec_rust_sys::zvec_doc_is_empty(self.handle) }
}
pub fn has_field(&self, name: &str) -> bool {
let c_name = match to_cstring(name) {
Ok(s) => s,
Err(_) => return false,
};
unsafe { zvec_rust_sys::zvec_doc_has_field(self.handle, c_name.as_ptr()) }
}
pub fn is_field_null(&self, name: &str) -> bool {
let c_name = match to_cstring(name) {
Ok(s) => s,
Err(_) => return false,
};
unsafe { zvec_rust_sys::zvec_doc_is_field_null(self.handle, c_name.as_ptr()) }
}
pub fn add_string(&mut self, name: &str, value: &str) -> Result<()> {
let c_name = to_cstring(name)?;
let c_value = to_cstring(value)?;
let bytes = c_value.as_bytes();
check_error(unsafe {
zvec_rust_sys::zvec_doc_add_field_by_value(
self.handle,
c_name.as_ptr(),
DataType::String as u32,
bytes.as_ptr() as *const c_void,
bytes.len(),
)
})
}
pub fn add_bool(&mut self, name: &str, value: bool) -> Result<()> {
let c_name = to_cstring(name)?;
check_error(unsafe {
zvec_rust_sys::zvec_doc_add_field_by_value(
self.handle,
c_name.as_ptr(),
DataType::Bool as u32,
&value as *const bool as *const c_void,
std::mem::size_of::<bool>(),
)
})
}
pub fn add_i32(&mut self, name: &str, value: i32) -> Result<()> {
let c_name = to_cstring(name)?;
check_error(unsafe {
zvec_rust_sys::zvec_doc_add_field_by_value(
self.handle,
c_name.as_ptr(),
DataType::Int32 as u32,
&value as *const i32 as *const c_void,
std::mem::size_of::<i32>(),
)
})
}
pub fn add_i64(&mut self, name: &str, value: i64) -> Result<()> {
let c_name = to_cstring(name)?;
check_error(unsafe {
zvec_rust_sys::zvec_doc_add_field_by_value(
self.handle,
c_name.as_ptr(),
DataType::Int64 as u32,
&value as *const i64 as *const c_void,
std::mem::size_of::<i64>(),
)
})
}
pub fn add_u32(&mut self, name: &str, value: u32) -> Result<()> {
let c_name = to_cstring(name)?;
check_error(unsafe {
zvec_rust_sys::zvec_doc_add_field_by_value(
self.handle,
c_name.as_ptr(),
DataType::Uint32 as u32,
&value as *const u32 as *const c_void,
std::mem::size_of::<u32>(),
)
})
}
pub fn add_u64(&mut self, name: &str, value: u64) -> Result<()> {
let c_name = to_cstring(name)?;
check_error(unsafe {
zvec_rust_sys::zvec_doc_add_field_by_value(
self.handle,
c_name.as_ptr(),
DataType::Uint64 as u32,
&value as *const u64 as *const c_void,
std::mem::size_of::<u64>(),
)
})
}
pub fn add_f32(&mut self, name: &str, value: f32) -> Result<()> {
let c_name = to_cstring(name)?;
check_error(unsafe {
zvec_rust_sys::zvec_doc_add_field_by_value(
self.handle,
c_name.as_ptr(),
DataType::Float as u32,
&value as *const f32 as *const c_void,
std::mem::size_of::<f32>(),
)
})
}
pub fn add_f64(&mut self, name: &str, value: f64) -> Result<()> {
let c_name = to_cstring(name)?;
check_error(unsafe {
zvec_rust_sys::zvec_doc_add_field_by_value(
self.handle,
c_name.as_ptr(),
DataType::Double as u32,
&value as *const f64 as *const c_void,
std::mem::size_of::<f64>(),
)
})
}
pub fn add_vector_f32(&mut self, name: &str, vector: &[f32]) -> Result<()> {
let c_name = to_cstring(name)?;
check_error(unsafe {
zvec_rust_sys::zvec_doc_add_field_by_value(
self.handle,
c_name.as_ptr(),
DataType::VectorFp32 as u32,
vector.as_ptr() as *const c_void,
std::mem::size_of_val(vector),
)
})
}
pub fn add_vector_f64(&mut self, name: &str, vector: &[f64]) -> Result<()> {
let c_name = to_cstring(name)?;
check_error(unsafe {
zvec_rust_sys::zvec_doc_add_field_by_value(
self.handle,
c_name.as_ptr(),
DataType::VectorFp64 as u32,
vector.as_ptr() as *const c_void,
std::mem::size_of_val(vector),
)
})
}
pub fn add_binary(&mut self, name: &str, value: &[u8]) -> Result<()> {
let c_name = to_cstring(name)?;
check_error(unsafe {
zvec_rust_sys::zvec_doc_add_field_by_value(
self.handle,
c_name.as_ptr(),
DataType::Binary as u32,
value.as_ptr() as *const c_void,
value.len(),
)
})
}
pub fn add_vector_i8(&mut self, name: &str, vector: &[i8]) -> Result<()> {
let c_name = to_cstring(name)?;
check_error(unsafe {
zvec_rust_sys::zvec_doc_add_field_by_value(
self.handle,
c_name.as_ptr(),
DataType::VectorInt8 as u32,
vector.as_ptr() as *const c_void,
std::mem::size_of_val(vector),
)
})
}
pub fn add_vector_i16(&mut self, name: &str, vector: &[i16]) -> Result<()> {
let c_name = to_cstring(name)?;
check_error(unsafe {
zvec_rust_sys::zvec_doc_add_field_by_value(
self.handle,
c_name.as_ptr(),
DataType::VectorInt16 as u32,
vector.as_ptr() as *const c_void,
std::mem::size_of_val(vector),
)
})
}
fn add_typed_array<T>(&mut self, name: &str, data_type: DataType, values: &[T]) -> Result<()> {
let c_name = to_cstring(name)?;
check_error(unsafe {
zvec_rust_sys::zvec_doc_add_field_by_value(
self.handle,
c_name.as_ptr(),
data_type as u32,
values.as_ptr() as *const c_void,
std::mem::size_of_val(values),
)
})
}
pub fn add_array_i32(&mut self, name: &str, values: &[i32]) -> Result<()> {
self.add_typed_array(name, DataType::ArrayInt32, values)
}
pub fn add_array_i64(&mut self, name: &str, values: &[i64]) -> Result<()> {
self.add_typed_array(name, DataType::ArrayInt64, values)
}
pub fn add_array_u32(&mut self, name: &str, values: &[u32]) -> Result<()> {
self.add_typed_array(name, DataType::ArrayUint32, values)
}
pub fn add_array_u64(&mut self, name: &str, values: &[u64]) -> Result<()> {
self.add_typed_array(name, DataType::ArrayUint64, values)
}
pub fn add_array_f32(&mut self, name: &str, values: &[f32]) -> Result<()> {
self.add_typed_array(name, DataType::ArrayFloat, values)
}
pub fn add_array_f64(&mut self, name: &str, values: &[f64]) -> Result<()> {
self.add_typed_array(name, DataType::ArrayDouble, values)
}
pub fn add_array_bool(&mut self, name: &str, values: &[bool]) -> Result<()> {
self.add_typed_array(name, DataType::ArrayBool, values)
}
pub fn set_field_null(&mut self, name: &str) -> Result<()> {
let c_name = to_cstring(name)?;
check_error(unsafe { zvec_rust_sys::zvec_doc_set_field_null(self.handle, c_name.as_ptr()) })
}
pub fn remove_field(&mut self, name: &str) -> Result<()> {
let c_name = to_cstring(name)?;
check_error(unsafe { zvec_rust_sys::zvec_doc_remove_field(self.handle, c_name.as_ptr()) })
}
fn get_basic_field<T: Copy + Default>(
&self,
name: &str,
data_type: DataType,
) -> Result<Option<T>> {
if !self.has_field(name) || self.is_field_null(name) {
return Ok(None);
}
let c_name = to_cstring(name)?;
let mut value: T = T::default();
check_error(unsafe {
zvec_rust_sys::zvec_doc_get_field_value_basic(
self.handle,
c_name.as_ptr(),
data_type as u32,
&mut value as *mut T as *mut c_void,
std::mem::size_of::<T>(),
)
})?;
Ok(Some(value))
}
fn get_pointer_field(
&self,
name: &str,
data_type: DataType,
) -> Result<Option<(*const c_void, usize)>> {
let c_name = to_cstring(name)?;
let mut value_ptr: *const c_void = std::ptr::null();
let mut value_size: usize = 0;
check_error(unsafe {
zvec_rust_sys::zvec_doc_get_field_value_pointer(
self.handle,
c_name.as_ptr(),
data_type as u32,
&mut value_ptr,
&mut value_size,
)
})?;
if value_ptr.is_null() || value_size == 0 {
return Ok(None);
}
Ok(Some((value_ptr, value_size)))
}
fn get_typed_vec<T: Copy>(&self, name: &str, data_type: DataType) -> Result<Option<Vec<T>>> {
let Some((ptr, size)) = self.get_pointer_field(name, data_type)? else {
return Ok(None);
};
let elem_size = std::mem::size_of::<T>();
if elem_size > 1 && size % elem_size != 0 {
return Err(Error {
code: ErrorCode::InternalError,
message: format!(
"data size {} is not aligned to element size {}",
size, elem_size
),
});
}
let count = size / elem_size;
let slice = unsafe { std::slice::from_raw_parts(ptr as *const T, count) };
Ok(Some(slice.to_vec()))
}
pub fn get_string(&self, name: &str) -> Result<Option<String>> {
let Some((ptr, _size)) = self.get_pointer_field(name, DataType::String)? else {
return Ok(None);
};
unsafe {
let cstr = CStr::from_ptr(ptr as *const std::os::raw::c_char);
Ok(Some(cstr.to_string_lossy().into_owned()))
}
}
pub fn get_bool(&self, name: &str) -> Result<Option<bool>> {
self.get_basic_field(name, DataType::Bool)
}
pub fn get_i32(&self, name: &str) -> Result<Option<i32>> {
self.get_basic_field(name, DataType::Int32)
}
pub fn get_i64(&self, name: &str) -> Result<Option<i64>> {
self.get_basic_field(name, DataType::Int64)
}
pub fn get_u32(&self, name: &str) -> Result<Option<u32>> {
self.get_basic_field(name, DataType::Uint32)
}
pub fn get_u64(&self, name: &str) -> Result<Option<u64>> {
self.get_basic_field(name, DataType::Uint64)
}
pub fn get_f32(&self, name: &str) -> Result<Option<f32>> {
self.get_basic_field(name, DataType::Float)
}
pub fn get_f64(&self, name: &str) -> Result<Option<f64>> {
self.get_basic_field(name, DataType::Double)
}
pub fn get_vector_f32(&self, name: &str) -> Result<Option<Vec<f32>>> {
self.get_typed_vec(name, DataType::VectorFp32)
}
pub fn get_vector_f64(&self, name: &str) -> Result<Option<Vec<f64>>> {
self.get_typed_vec(name, DataType::VectorFp64)
}
pub fn get_binary(&self, name: &str) -> Result<Option<Vec<u8>>> {
self.get_typed_vec(name, DataType::Binary)
}
pub fn get_vector_i8(&self, name: &str) -> Result<Option<Vec<i8>>> {
self.get_typed_vec(name, DataType::VectorInt8)
}
pub fn get_vector_i16(&self, name: &str) -> Result<Option<Vec<i16>>> {
self.get_typed_vec(name, DataType::VectorInt16)
}
pub fn get_array_i32(&self, name: &str) -> Result<Option<Vec<i32>>> {
self.get_typed_vec(name, DataType::ArrayInt32)
}
pub fn get_array_i64(&self, name: &str) -> Result<Option<Vec<i64>>> {
self.get_typed_vec(name, DataType::ArrayInt64)
}
pub fn get_array_u32(&self, name: &str) -> Result<Option<Vec<u32>>> {
self.get_typed_vec(name, DataType::ArrayUint32)
}
pub fn get_array_u64(&self, name: &str) -> Result<Option<Vec<u64>>> {
self.get_typed_vec(name, DataType::ArrayUint64)
}
pub fn get_array_f32(&self, name: &str) -> Result<Option<Vec<f32>>> {
self.get_typed_vec(name, DataType::ArrayFloat)
}
pub fn get_array_f64(&self, name: &str) -> Result<Option<Vec<f64>>> {
self.get_typed_vec(name, DataType::ArrayDouble)
}
pub fn get_array_bool(&self, name: &str) -> Result<Option<Vec<bool>>> {
self.get_typed_vec(name, DataType::ArrayBool)
}
pub fn clear(&mut self) {
unsafe { zvec_rust_sys::zvec_doc_clear(self.handle) };
}
}
impl Drop for Doc {
fn drop(&mut self) {
if self.owned && !self.handle.is_null() {
unsafe { zvec_rust_sys::zvec_doc_destroy(self.handle) };
}
}
}
pub fn free_docs(docs: Vec<Doc>) {
drop(docs);
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn from_borrowed_does_not_own() {
let doc = Doc::from_borrowed(std::ptr::null_mut());
assert!(!doc.owned);
assert!(doc.handle.is_null());
}
#[test]
fn from_raw_takes_ownership() {
let doc = unsafe { Doc::from_raw(std::ptr::null_mut()) };
assert!(doc.owned);
}
}