use std::{ffi::c_char, marker::PhantomData, ptr, slice, str};
use crate::{
error::{Error, Result},
ffi::{
duckdb_destroy_logical_type, duckdb_get_type_id, duckdb_string_t, duckdb_string_t_data,
duckdb_string_t_length, duckdb_validity_row_is_valid, duckdb_validity_set_row_invalid,
duckdb_vector, duckdb_vector_assign_string_element_len,
duckdb_vector_ensure_validity_writable, duckdb_vector_get_column_type,
duckdb_vector_get_data, duckdb_vector_get_validity,
},
types::{
numeric::{hugeint_from_i128, i128_from_hugeint, u128_from_uhugeint, uhugeint_from_u128},
value::DuckValue,
},
};
use super::UdfResult;
pub struct VectorRef<'a> {
ptr: duckdb_vector,
_marker: PhantomData<&'a ()>,
}
impl<'a> VectorRef<'a> {
pub(crate) unsafe fn new(ptr: duckdb_vector) -> Self {
Self { ptr, _marker: PhantomData }
}
pub fn is_null(
&self,
row: usize,
) -> bool {
let validity = unsafe { duckdb_vector_get_validity(self.ptr) };
!unsafe { duckdb_validity_row_is_valid(validity, row as u64) }
}
pub fn get<T: ScalarArg<'a>>(
&self,
row: usize,
) -> UdfResult<T> {
T::read(self, row)
}
pub fn as_duck_value(
&self,
row: usize,
) -> Result<DuckValue> {
let mut lt = unsafe { duckdb_vector_get_column_type(self.ptr) };
let type_id = unsafe { duckdb_get_type_id(lt) };
unsafe { duckdb_destroy_logical_type(&mut lt) };
DuckValue::from_duckdb_vec(self.ptr, type_id, row as u64).map_err(Error::ConversionError)
}
pub(crate) fn raw(&self) -> duckdb_vector {
self.ptr
}
}
pub struct VectorMut<'a> {
ptr: duckdb_vector,
_marker: PhantomData<&'a mut ()>,
}
impl<'a> VectorMut<'a> {
pub(crate) unsafe fn new(ptr: duckdb_vector) -> Self {
Self { ptr, _marker: PhantomData }
}
pub fn set_null(
&mut self,
row: usize,
) {
unsafe { duckdb_vector_ensure_validity_writable(self.ptr) };
let validity = unsafe { duckdb_vector_get_validity(self.ptr) };
unsafe { duckdb_validity_set_row_invalid(validity, row as u64) };
}
pub fn set<T: ScalarRet>(
&mut self,
row: usize,
value: T,
) -> UdfResult<()> {
value.write(self, row)
}
pub(crate) fn raw(&mut self) -> duckdb_vector {
self.ptr
}
}
pub trait ScalarArg<'v>: Sized {
fn read(
v: &VectorRef<'v>,
row: usize,
) -> UdfResult<Self>;
}
pub trait ScalarRet {
fn write(
self,
v: &mut VectorMut<'_>,
row: usize,
) -> UdfResult<()>;
}
macro_rules! impl_scalar_fixed {
($($rust_type:ty),+ $(,)?) => {
$(
impl<'v> ScalarArg<'v> for $rust_type {
fn read(v: &VectorRef<'v>, row: usize) -> UdfResult<Self> {
let data = unsafe { duckdb_vector_get_data(v.raw()) as *const $rust_type };
Ok(unsafe { *data.add(row) })
}
}
impl ScalarRet for $rust_type {
fn write(self, v: &mut VectorMut<'_>, row: usize) -> UdfResult<()> {
let data = unsafe { duckdb_vector_get_data(v.raw()) as *mut $rust_type };
unsafe { ptr::write(data.add(row), self) };
Ok(())
}
}
)+
};
}
impl_scalar_fixed!(bool, i8, i16, i32, i64, u8, u16, u32, u64, f32, f64);
impl<'v> ScalarArg<'v> for i128 {
fn read(
v: &VectorRef<'v>,
row: usize,
) -> UdfResult<Self> {
let data = unsafe { duckdb_vector_get_data(v.raw()) as *const crate::ffi::duckdb_hugeint };
Ok(i128_from_hugeint(unsafe { *data.add(row) }))
}
}
impl ScalarRet for i128 {
fn write(
self,
v: &mut VectorMut<'_>,
row: usize,
) -> UdfResult<()> {
let data = unsafe { duckdb_vector_get_data(v.raw()) as *mut crate::ffi::duckdb_hugeint };
unsafe { ptr::write(data.add(row), hugeint_from_i128(self)) };
Ok(())
}
}
impl<'v> ScalarArg<'v> for u128 {
fn read(
v: &VectorRef<'v>,
row: usize,
) -> UdfResult<Self> {
let data = unsafe { duckdb_vector_get_data(v.raw()) as *const crate::ffi::duckdb_uhugeint };
Ok(u128_from_uhugeint(unsafe { *data.add(row) }))
}
}
impl ScalarRet for u128 {
fn write(
self,
v: &mut VectorMut<'_>,
row: usize,
) -> UdfResult<()> {
let data = unsafe { duckdb_vector_get_data(v.raw()) as *mut crate::ffi::duckdb_uhugeint };
unsafe { ptr::write(data.add(row), uhugeint_from_u128(self)) };
Ok(())
}
}
fn read_str_bytes<'v>(
v: &VectorRef<'v>,
row: usize,
) -> &'v [u8] {
let data = unsafe { duckdb_vector_get_data(v.raw()) as *mut duckdb_string_t };
let str_ptr = unsafe { data.add(row) };
let c_ptr = unsafe { duckdb_string_t_data(str_ptr) };
let len = unsafe { duckdb_string_t_length(*str_ptr) } as usize;
unsafe { slice::from_raw_parts(c_ptr.cast::<u8>(), len) }
}
fn write_str_bytes(
v: &mut VectorMut<'_>,
row: usize,
bytes: &[u8],
) {
unsafe {
duckdb_vector_assign_string_element_len(
v.raw(),
row as u64,
bytes.as_ptr() as *const c_char,
bytes.len() as u64,
)
};
}
impl<'v> ScalarArg<'v> for &'v str {
fn read(
v: &VectorRef<'v>,
row: usize,
) -> UdfResult<Self> {
str::from_utf8(read_str_bytes(v, row)).map_err(|e| Box::new(e) as _)
}
}
impl<'v> ScalarArg<'v> for String {
fn read(
v: &VectorRef<'v>,
row: usize,
) -> UdfResult<Self> {
<&str as ScalarArg<'v>>::read(v, row).map(str::to_owned)
}
}
impl ScalarRet for &str {
fn write(
self,
v: &mut VectorMut<'_>,
row: usize,
) -> UdfResult<()> {
write_str_bytes(v, row, self.as_bytes());
Ok(())
}
}
impl ScalarRet for String {
fn write(
self,
v: &mut VectorMut<'_>,
row: usize,
) -> UdfResult<()> {
write_str_bytes(v, row, self.as_bytes());
Ok(())
}
}
impl<'v, T: ScalarArg<'v>> ScalarArg<'v> for Option<T> {
fn read(
v: &VectorRef<'v>,
row: usize,
) -> UdfResult<Self> {
if v.is_null(row) {
Ok(None)
} else {
T::read(v, row).map(Some)
}
}
}
impl<T: ScalarRet> ScalarRet for Option<T> {
fn write(
self,
v: &mut VectorMut<'_>,
row: usize,
) -> UdfResult<()> {
match self {
Some(value) => value.write(v, row),
None => {
v.set_null(row);
Ok(())
},
}
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::udf::data_chunk::DataChunkHandle;
use crate::udf::logical_type::LogicalType;
#[test]
fn i32_round_trips() {
let types = [LogicalType::of::<i32>().unwrap()];
let mut chunk = DataChunkHandle::new(&types).unwrap();
{
let mut vec = chunk.vector_mut(0).unwrap();
vec.set(0, 42i32).unwrap();
}
let vec = chunk.vector(0).unwrap();
let got: i32 = vec.get(0).unwrap();
assert_eq!(got, 42);
}
#[test]
fn i128_round_trips_negative() {
let types = [LogicalType::of::<i128>().unwrap()];
let mut chunk = DataChunkHandle::new(&types).unwrap();
let value: i128 = -170_141_183_460_469_231_731_687_303_715_884_105_000;
{
let mut vec = chunk.vector_mut(0).unwrap();
vec.set(0, value).unwrap();
}
let vec = chunk.vector(0).unwrap();
let got: i128 = vec.get(0).unwrap();
assert_eq!(got, value);
}
#[test]
fn short_string_round_trips() {
let types = [LogicalType::of::<String>().unwrap()];
let mut chunk = DataChunkHandle::new(&types).unwrap();
{
let mut vec = chunk.vector_mut(0).unwrap();
vec.set(0, "hi").unwrap();
}
let vec = chunk.vector(0).unwrap();
let got: String = vec.get(0).unwrap();
assert_eq!(got, "hi");
let got_ref: &str = vec.get(0).unwrap();
assert_eq!(got_ref, "hi");
}
#[test]
fn long_string_round_trips() {
let types = [LogicalType::of::<String>().unwrap()];
let mut chunk = DataChunkHandle::new(&types).unwrap();
let long = "x".repeat(64);
{
let mut vec = chunk.vector_mut(0).unwrap();
vec.set(0, long.as_str()).unwrap();
}
let vec = chunk.vector(0).unwrap();
let got: String = vec.get(0).unwrap();
assert_eq!(got, long);
}
#[test]
fn null_round_trips_via_option() {
let types = [LogicalType::of::<i32>().unwrap()];
let mut chunk = DataChunkHandle::new(&types).unwrap();
{
let mut vec = chunk.vector_mut(0).unwrap();
vec.set(0, None::<i32>).unwrap();
}
let vec = chunk.vector(0).unwrap();
assert!(vec.is_null(0));
let got: Option<i32> = vec.get(0).unwrap();
assert_eq!(got, None);
}
#[test]
fn some_round_trips_via_option() {
let types = [LogicalType::of::<i32>().unwrap()];
let mut chunk = DataChunkHandle::new(&types).unwrap();
{
let mut vec = chunk.vector_mut(0).unwrap();
vec.set(0, Some(7i32)).unwrap();
}
let vec = chunk.vector(0).unwrap();
assert!(!vec.is_null(0));
let got: Option<i32> = vec.get(0).unwrap();
assert_eq!(got, Some(7));
}
}