1use std::{ffi::c_char, marker::PhantomData, ptr, slice, str};
11
12use crate::{
13 error::{Error, Result},
14 ffi::{
15 duckdb_destroy_logical_type, duckdb_get_type_id, duckdb_string_t, duckdb_string_t_data,
16 duckdb_string_t_length, duckdb_validity_row_is_valid, duckdb_validity_set_row_invalid,
17 duckdb_vector, duckdb_vector_assign_string_element_len,
18 duckdb_vector_ensure_validity_writable, duckdb_vector_get_column_type,
19 duckdb_vector_get_data, duckdb_vector_get_validity,
20 },
21 types::{
22 numeric::{hugeint_from_i128, i128_from_hugeint, u128_from_uhugeint, uhugeint_from_u128},
23 value::DuckValue,
24 },
25};
26
27use super::UdfResult;
28
29pub struct VectorRef<'a> {
34 ptr: duckdb_vector,
35 _marker: PhantomData<&'a ()>,
36}
37
38impl<'a> VectorRef<'a> {
39 pub(crate) unsafe fn new(ptr: duckdb_vector) -> Self {
46 Self { ptr, _marker: PhantomData }
47 }
48
49 pub fn is_null(
51 &self,
52 row: usize,
53 ) -> bool {
54 let validity = unsafe { duckdb_vector_get_validity(self.ptr) };
56 !unsafe { duckdb_validity_row_is_valid(validity, row as u64) }
61 }
62
63 pub fn get<T: ScalarArg<'a>>(
69 &self,
70 row: usize,
71 ) -> UdfResult<T> {
72 T::read(self, row)
73 }
74
75 pub fn as_duck_value(
82 &self,
83 row: usize,
84 ) -> Result<DuckValue> {
85 let mut lt = unsafe { duckdb_vector_get_column_type(self.ptr) };
88 let type_id = unsafe { duckdb_get_type_id(lt) };
90 unsafe { duckdb_destroy_logical_type(&mut lt) };
92 DuckValue::from_duckdb_vec(self.ptr, type_id, row as u64).map_err(Error::ConversionError)
93 }
94
95 pub(crate) fn raw(&self) -> duckdb_vector {
97 self.ptr
98 }
99}
100
101pub struct VectorMut<'a> {
106 ptr: duckdb_vector,
107 _marker: PhantomData<&'a mut ()>,
108}
109
110impl<'a> VectorMut<'a> {
111 pub(crate) unsafe fn new(ptr: duckdb_vector) -> Self {
119 Self { ptr, _marker: PhantomData }
120 }
121
122 pub fn set_null(
124 &mut self,
125 row: usize,
126 ) {
127 unsafe { duckdb_vector_ensure_validity_writable(self.ptr) };
129 let validity = unsafe { duckdb_vector_get_validity(self.ptr) };
132 unsafe { duckdb_validity_set_row_invalid(validity, row as u64) };
135 }
136
137 pub fn set<T: ScalarRet>(
143 &mut self,
144 row: usize,
145 value: T,
146 ) -> UdfResult<()> {
147 value.write(self, row)
148 }
149
150 pub(crate) fn raw(&mut self) -> duckdb_vector {
152 self.ptr
153 }
154}
155
156pub trait ScalarArg<'v>: Sized {
158 fn read(
164 v: &VectorRef<'v>,
165 row: usize,
166 ) -> UdfResult<Self>;
167}
168
169pub trait ScalarRet {
171 fn write(
177 self,
178 v: &mut VectorMut<'_>,
179 row: usize,
180 ) -> UdfResult<()>;
181}
182
183macro_rules! impl_scalar_fixed {
189 ($($rust_type:ty),+ $(,)?) => {
190 $(
191 impl<'v> ScalarArg<'v> for $rust_type {
192 fn read(v: &VectorRef<'v>, row: usize) -> UdfResult<Self> {
193 let data = unsafe { duckdb_vector_get_data(v.raw()) as *const $rust_type };
196 Ok(unsafe { *data.add(row) })
199 }
200 }
201 impl ScalarRet for $rust_type {
202 fn write(self, v: &mut VectorMut<'_>, row: usize) -> UdfResult<()> {
203 let data = unsafe { duckdb_vector_get_data(v.raw()) as *mut $rust_type };
206 unsafe { ptr::write(data.add(row), self) };
210 Ok(())
211 }
212 }
213 )+
214 };
215}
216
217impl_scalar_fixed!(bool, i8, i16, i32, i64, u8, u16, u32, u64, f32, f64);
218
219impl<'v> ScalarArg<'v> for i128 {
223 fn read(
224 v: &VectorRef<'v>,
225 row: usize,
226 ) -> UdfResult<Self> {
227 let data = unsafe { duckdb_vector_get_data(v.raw()) as *const crate::ffi::duckdb_hugeint };
229 Ok(i128_from_hugeint(unsafe { *data.add(row) }))
231 }
232}
233impl ScalarRet for i128 {
234 fn write(
235 self,
236 v: &mut VectorMut<'_>,
237 row: usize,
238 ) -> UdfResult<()> {
239 let data = unsafe { duckdb_vector_get_data(v.raw()) as *mut crate::ffi::duckdb_hugeint };
241 unsafe { ptr::write(data.add(row), hugeint_from_i128(self)) };
243 Ok(())
244 }
245}
246impl<'v> ScalarArg<'v> for u128 {
247 fn read(
248 v: &VectorRef<'v>,
249 row: usize,
250 ) -> UdfResult<Self> {
251 let data = unsafe { duckdb_vector_get_data(v.raw()) as *const crate::ffi::duckdb_uhugeint };
253 Ok(u128_from_uhugeint(unsafe { *data.add(row) }))
255 }
256}
257impl ScalarRet for u128 {
258 fn write(
259 self,
260 v: &mut VectorMut<'_>,
261 row: usize,
262 ) -> UdfResult<()> {
263 let data = unsafe { duckdb_vector_get_data(v.raw()) as *mut crate::ffi::duckdb_uhugeint };
265 unsafe { ptr::write(data.add(row), uhugeint_from_u128(self)) };
267 Ok(())
268 }
269}
270
271fn read_str_bytes<'v>(
275 v: &VectorRef<'v>,
276 row: usize,
277) -> &'v [u8] {
278 let data = unsafe { duckdb_vector_get_data(v.raw()) as *mut duckdb_string_t };
280 let str_ptr = unsafe { data.add(row) };
282 let c_ptr = unsafe { duckdb_string_t_data(str_ptr) };
286 let len = unsafe { duckdb_string_t_length(*str_ptr) } as usize;
288 unsafe { slice::from_raw_parts(c_ptr.cast::<u8>(), len) }
290}
291
292fn write_str_bytes(
293 v: &mut VectorMut<'_>,
294 row: usize,
295 bytes: &[u8],
296) {
297 unsafe {
300 duckdb_vector_assign_string_element_len(
301 v.raw(),
302 row as u64,
303 bytes.as_ptr() as *const c_char,
304 bytes.len() as u64,
305 )
306 };
307}
308
309impl<'v> ScalarArg<'v> for &'v str {
310 fn read(
311 v: &VectorRef<'v>,
312 row: usize,
313 ) -> UdfResult<Self> {
314 str::from_utf8(read_str_bytes(v, row)).map_err(|e| Box::new(e) as _)
315 }
316}
317impl<'v> ScalarArg<'v> for String {
318 fn read(
319 v: &VectorRef<'v>,
320 row: usize,
321 ) -> UdfResult<Self> {
322 <&str as ScalarArg<'v>>::read(v, row).map(str::to_owned)
323 }
324}
325impl ScalarRet for &str {
326 fn write(
327 self,
328 v: &mut VectorMut<'_>,
329 row: usize,
330 ) -> UdfResult<()> {
331 write_str_bytes(v, row, self.as_bytes());
332 Ok(())
333 }
334}
335impl ScalarRet for String {
336 fn write(
337 self,
338 v: &mut VectorMut<'_>,
339 row: usize,
340 ) -> UdfResult<()> {
341 write_str_bytes(v, row, self.as_bytes());
342 Ok(())
343 }
344}
345
346impl<'v, T: ScalarArg<'v>> ScalarArg<'v> for Option<T> {
347 fn read(
348 v: &VectorRef<'v>,
349 row: usize,
350 ) -> UdfResult<Self> {
351 if v.is_null(row) {
352 Ok(None)
353 } else {
354 T::read(v, row).map(Some)
355 }
356 }
357}
358impl<T: ScalarRet> ScalarRet for Option<T> {
359 fn write(
360 self,
361 v: &mut VectorMut<'_>,
362 row: usize,
363 ) -> UdfResult<()> {
364 match self {
365 Some(value) => value.write(v, row),
366 None => {
367 v.set_null(row);
368 Ok(())
369 },
370 }
371 }
372}
373
374#[cfg(test)]
375mod tests {
376 use super::*;
377 use crate::types::LogicalType;
378 use crate::udf::data_chunk::DataChunkHandle;
379
380 #[test]
381 fn i32_round_trips() {
382 let types = [LogicalType::of::<i32>().unwrap()];
383 let mut chunk = DataChunkHandle::new(&types).unwrap();
384 {
385 let mut vec = chunk.vector_mut(0).unwrap();
386 vec.set(0, 42i32).unwrap();
387 }
388 let vec = chunk.vector(0).unwrap();
389 let got: i32 = vec.get(0).unwrap();
390 assert_eq!(got, 42);
391 }
392
393 #[test]
394 fn i128_round_trips_negative() {
395 let types = [LogicalType::of::<i128>().unwrap()];
396 let mut chunk = DataChunkHandle::new(&types).unwrap();
397 let value: i128 = -170_141_183_460_469_231_731_687_303_715_884_105_000;
398 {
399 let mut vec = chunk.vector_mut(0).unwrap();
400 vec.set(0, value).unwrap();
401 }
402 let vec = chunk.vector(0).unwrap();
403 let got: i128 = vec.get(0).unwrap();
404 assert_eq!(got, value);
405 }
406
407 #[test]
408 fn short_string_round_trips() {
409 let types = [LogicalType::of::<String>().unwrap()];
410 let mut chunk = DataChunkHandle::new(&types).unwrap();
411 {
412 let mut vec = chunk.vector_mut(0).unwrap();
413 vec.set(0, "hi").unwrap();
414 }
415 let vec = chunk.vector(0).unwrap();
416 let got: String = vec.get(0).unwrap();
417 assert_eq!(got, "hi");
418 let got_ref: &str = vec.get(0).unwrap();
419 assert_eq!(got_ref, "hi");
420 }
421
422 #[test]
423 fn long_string_round_trips() {
424 let types = [LogicalType::of::<String>().unwrap()];
425 let mut chunk = DataChunkHandle::new(&types).unwrap();
426 let long = "x".repeat(64);
427 {
428 let mut vec = chunk.vector_mut(0).unwrap();
429 vec.set(0, long.as_str()).unwrap();
430 }
431 let vec = chunk.vector(0).unwrap();
432 let got: String = vec.get(0).unwrap();
433 assert_eq!(got, long);
434 }
435
436 #[test]
437 fn null_round_trips_via_option() {
438 let types = [LogicalType::of::<i32>().unwrap()];
439 let mut chunk = DataChunkHandle::new(&types).unwrap();
440 {
441 let mut vec = chunk.vector_mut(0).unwrap();
442 vec.set(0, None::<i32>).unwrap();
443 }
444 let vec = chunk.vector(0).unwrap();
445 assert!(vec.is_null(0));
446 let got: Option<i32> = vec.get(0).unwrap();
447 assert_eq!(got, None);
448 }
449
450 #[test]
451 fn some_round_trips_via_option() {
452 let types = [LogicalType::of::<i32>().unwrap()];
453 let mut chunk = DataChunkHandle::new(&types).unwrap();
454 {
455 let mut vec = chunk.vector_mut(0).unwrap();
456 vec.set(0, Some(7i32)).unwrap();
457 }
458 let vec = chunk.vector(0).unwrap();
459 assert!(!vec.is_null(0));
460 let got: Option<i32> = vec.get(0).unwrap();
461 assert_eq!(got, Some(7));
462 }
463}