1use std::str::from_utf8_unchecked;
19use std::sync::Arc;
20
21use arrow::array::{Array, ArrayAccessor, ArrayRef, StringArray, StringBuilder};
22use arrow::buffer::{Buffer, OffsetBuffer};
23use arrow::datatypes::DataType;
24use arrow::{
25 array::{as_dictionary_array, as_largestring_array, as_string_array},
26 datatypes::Int32Type,
27};
28use datafusion_common::cast::as_large_binary_array;
29use datafusion_common::cast::as_string_view_array;
30use datafusion_common::types::{NativeType, logical_int64, logical_string};
31use datafusion_common::utils::hex::{HexCase, ToHex, encode_bytes_into};
32use datafusion_common::utils::take_function_args;
33use datafusion_common::{
34 DataFusionError,
35 cast::{as_binary_array, as_fixed_size_binary_array, as_int64_array},
36 exec_datafusion_err, exec_err,
37};
38use datafusion_expr::{
39 Coercion, ColumnarValue, EncodingPreservation, ScalarFunctionArgs, ScalarUDFImpl,
40 Signature, TypeSignature, TypeSignatureClass, Volatility,
41};
42#[derive(Debug, PartialEq, Eq, Hash)]
44pub struct SparkHex {
45 signature: Signature,
46 aliases: Vec<String>,
47}
48
49impl Default for SparkHex {
50 fn default() -> Self {
51 Self::new()
52 }
53}
54
55impl SparkHex {
56 pub fn new() -> Self {
57 let int64 = Coercion::new_implicit(
58 TypeSignatureClass::Native(logical_int64()),
59 vec![TypeSignatureClass::Numeric],
60 NativeType::Int64,
61 );
62
63 let string = Coercion::new_exact(TypeSignatureClass::Native(logical_string()));
64
65 let binary = Coercion::new_exact(TypeSignatureClass::Binary)
66 .with_encoding_preservation(EncodingPreservation::dictionary());
67
68 let variants = vec![
69 TypeSignature::Coercible(vec![int64]),
71 TypeSignature::Coercible(vec![string]),
73 TypeSignature::Coercible(vec![binary]),
75 ];
76
77 Self {
78 signature: Signature::one_of(variants, Volatility::Immutable),
79 aliases: vec![],
80 }
81 }
82}
83
84impl ScalarUDFImpl for SparkHex {
85 fn name(&self) -> &str {
86 "hex"
87 }
88
89 fn signature(&self) -> &Signature {
90 &self.signature
91 }
92
93 fn return_type(&self, arg_types: &[DataType]) -> datafusion_common::Result<DataType> {
94 Ok(match &arg_types[0] {
95 DataType::Dictionary(key_type, _) => {
96 DataType::Dictionary(key_type.clone(), Box::new(DataType::Utf8))
97 }
98 _ => DataType::Utf8,
99 })
100 }
101
102 fn invoke_with_args(
103 &self,
104 args: ScalarFunctionArgs,
105 ) -> datafusion_common::Result<ColumnarValue> {
106 spark_hex(&args.args)
107 }
108
109 fn aliases(&self) -> &[String] {
110 &self.aliases
111 }
112}
113
114#[inline]
115fn append_hex_bytes(
116 values: &mut Vec<u8>,
117 bytes: &[u8],
118 case: HexCase,
119) -> Result<i32, DataFusionError> {
120 let additional = bytes
121 .len()
122 .checked_mul(2)
123 .ok_or_else(|| exec_datafusion_err!("hex output size overflow"))?;
124 values.try_reserve(additional).map_err(|e| {
125 exec_datafusion_err!("failed to reserve {additional} bytes for hex output: {e}")
126 })?;
127 encode_bytes_into(bytes, case, values);
128 i32::try_from(values.len())
129 .map_err(|_| exec_datafusion_err!("hex output exceeds i32 offset range"))
130}
131
132fn hex_encode_bytes<'a, A, T>(
134 array: &A,
135 lowercase: bool,
136) -> Result<ArrayRef, DataFusionError>
137where
138 A: ArrayAccessor<Item = &'a T>,
139 T: AsRef<[u8]> + ?Sized + 'a,
140{
141 let case = if lowercase {
142 HexCase::Lower
143 } else {
144 HexCase::Upper
145 };
146 let len = array.len();
147 let nulls = array.nulls().cloned();
148
149 let mut values: Vec<u8> = Vec::with_capacity(len * 64);
154 let mut offsets: Vec<i32> = Vec::with_capacity(len + 1);
155 offsets.push(0);
156
157 if let Some(ref nulls) = nulls {
158 for i in 0..len {
159 if nulls.is_valid(i) {
160 let bytes = unsafe { array.value_unchecked(i) }.as_ref();
162 offsets.push(append_hex_bytes(&mut values, bytes, case)?);
163 } else {
164 offsets.push(i32::try_from(values.len()).map_err(|_| {
165 exec_datafusion_err!("hex output exceeds i32 offset range")
166 })?);
167 }
168 }
169 } else {
170 for i in 0..len {
171 let bytes = unsafe { array.value_unchecked(i) }.as_ref();
173 offsets.push(append_hex_bytes(&mut values, bytes, case)?);
174 }
175 }
176
177 let array = unsafe {
182 StringArray::new_unchecked(
183 OffsetBuffer::new(offsets.into()),
184 Buffer::from_vec(values),
185 nulls,
186 )
187 };
188 Ok(Arc::new(array))
189}
190
191fn hex_encode_int64(
193 iter: impl Iterator<Item = Option<i64>>,
194 len: usize,
195) -> Result<ArrayRef, DataFusionError> {
196 let mut builder = StringBuilder::with_capacity(len, len * 16);
197
198 for v in iter {
199 if let Some(num) = v {
200 let mut temp = [0u8; 16];
201 let slice = num.write_hex(HexCase::Upper, &mut temp);
202 unsafe {
204 builder.append_value(from_utf8_unchecked(slice));
205 }
206 } else {
207 builder.append_null();
208 }
209 }
210
211 Ok(Arc::new(builder.finish()))
212}
213
214pub fn spark_hex(args: &[ColumnarValue]) -> Result<ColumnarValue, DataFusionError> {
216 compute_hex(args, false)
217}
218
219pub fn spark_sha2_hex(args: &[ColumnarValue]) -> Result<ColumnarValue, DataFusionError> {
221 compute_hex(args, true)
222}
223
224pub fn compute_hex(
225 args: &[ColumnarValue],
226 lowercase: bool,
227) -> Result<ColumnarValue, DataFusionError> {
228 let input = match take_function_args("hex", args)? {
229 [ColumnarValue::Scalar(value)] => ColumnarValue::Array(value.to_array()?),
230 [ColumnarValue::Array(arr)] => ColumnarValue::Array(Arc::clone(arr)),
231 };
232
233 match &input {
234 ColumnarValue::Array(array) => match array.data_type() {
235 DataType::Int64 => {
236 let array = as_int64_array(array)?;
237 Ok(ColumnarValue::Array(hex_encode_int64(
238 array.iter(),
239 array.len(),
240 )?))
241 }
242 DataType::Utf8 => {
243 let array = as_string_array(array);
244 Ok(ColumnarValue::Array(hex_encode_bytes(&array, lowercase)?))
245 }
246 DataType::Utf8View => {
247 let array = as_string_view_array(array)?;
248 Ok(ColumnarValue::Array(hex_encode_bytes(&array, lowercase)?))
249 }
250 DataType::LargeUtf8 => {
251 let array = as_largestring_array(array);
252 Ok(ColumnarValue::Array(hex_encode_bytes(&array, lowercase)?))
253 }
254 DataType::Binary => {
255 let array = as_binary_array(array)?;
256 Ok(ColumnarValue::Array(hex_encode_bytes(&array, lowercase)?))
257 }
258 DataType::LargeBinary => {
259 let array = as_large_binary_array(array)?;
260 Ok(ColumnarValue::Array(hex_encode_bytes(&array, lowercase)?))
261 }
262 DataType::FixedSizeBinary(_) => {
263 let array = as_fixed_size_binary_array(array)?;
264 Ok(ColumnarValue::Array(hex_encode_bytes(&array, lowercase)?))
265 }
266 DataType::Dictionary(key_type, _) => {
267 if **key_type != DataType::Int32 {
268 return exec_err!(
269 "hex only supports Int32 dictionary keys, get: {}",
270 key_type
271 );
272 }
273
274 let dict = as_dictionary_array::<Int32Type>(&array);
275 let dict_values = dict.values();
276
277 let encoded_values = match dict_values.data_type() {
278 DataType::Int64 => {
279 let arr = as_int64_array(dict_values)?;
280 hex_encode_int64(arr.iter(), arr.len())?
281 }
282 DataType::Utf8 => {
283 let arr = as_string_array(dict_values);
284 hex_encode_bytes(&arr, lowercase)?
285 }
286 DataType::LargeUtf8 => {
287 let arr = as_largestring_array(dict_values);
288 hex_encode_bytes(&arr, lowercase)?
289 }
290 DataType::Utf8View => {
291 let arr = as_string_view_array(dict_values)?;
292 hex_encode_bytes(&arr, lowercase)?
293 }
294 DataType::Binary => {
295 let arr = as_binary_array(dict_values)?;
296 hex_encode_bytes(&arr, lowercase)?
297 }
298 DataType::LargeBinary => {
299 let arr = as_large_binary_array(dict_values)?;
300 hex_encode_bytes(&arr, lowercase)?
301 }
302 DataType::FixedSizeBinary(_) => {
303 let arr = as_fixed_size_binary_array(dict_values)?;
304 hex_encode_bytes(&arr, lowercase)?
305 }
306 _ => {
307 return exec_err!(
308 "hex got an unexpected argument type: {}",
309 dict_values.data_type()
310 );
311 }
312 };
313
314 let new_dict = dict.with_values(encoded_values);
315 Ok(ColumnarValue::Array(Arc::new(new_dict)))
316 }
317 _ => exec_err!("hex got an unexpected argument type: {}", array.data_type()),
318 },
319 _ => exec_err!("native hex does not support scalar values at this time"),
320 }
321}
322
323#[cfg(test)]
324mod test {
325 use std::sync::Arc;
326
327 use arrow::array::{
328 Array, BinaryArray, DictionaryArray, Int32Array, Int64Array, StringArray,
329 };
330 use arrow::{
331 array::{
332 BinaryDictionaryBuilder, PrimitiveDictionaryBuilder, StringDictionaryBuilder,
333 as_string_array,
334 },
335 datatypes::{Int32Type, Int64Type},
336 };
337 use datafusion_common::cast::as_dictionary_array;
338 use datafusion_expr::ColumnarValue;
339
340 #[test]
341 fn test_dictionary_hex_utf8() {
342 let mut input_builder = StringDictionaryBuilder::<Int32Type>::new();
343 input_builder.append_value("hi");
344 input_builder.append_value("bye");
345 input_builder.append_null();
346 input_builder.append_value("rust");
347 let input = input_builder.finish();
348
349 let mut expected_builder = StringDictionaryBuilder::<Int32Type>::new();
350 expected_builder.append_value("6869");
351 expected_builder.append_value("627965");
352 expected_builder.append_null();
353 expected_builder.append_value("72757374");
354 let expected = expected_builder.finish();
355
356 let columnar_value = ColumnarValue::Array(Arc::new(input));
357 let result = super::spark_hex(&[columnar_value]).unwrap();
358
359 let result = match result {
360 ColumnarValue::Array(array) => array,
361 _ => panic!("Expected array"),
362 };
363
364 let result = as_dictionary_array(&result).unwrap();
365
366 assert_eq!(result, &expected);
367 }
368
369 #[test]
370 fn test_dictionary_hex_int64() {
371 let mut input_builder = PrimitiveDictionaryBuilder::<Int32Type, Int64Type>::new();
372 input_builder.append_value(1);
373 input_builder.append_value(2);
374 input_builder.append_null();
375 input_builder.append_value(3);
376 let input = input_builder.finish();
377
378 let mut expected_builder = StringDictionaryBuilder::<Int32Type>::new();
379 expected_builder.append_value("1");
380 expected_builder.append_value("2");
381 expected_builder.append_null();
382 expected_builder.append_value("3");
383 let expected = expected_builder.finish();
384
385 let columnar_value = ColumnarValue::Array(Arc::new(input));
386 let result = super::spark_hex(&[columnar_value]).unwrap();
387
388 let result = match result {
389 ColumnarValue::Array(array) => array,
390 _ => panic!("Expected array"),
391 };
392
393 let result = as_dictionary_array(&result).unwrap();
394
395 assert_eq!(result, &expected);
396 }
397
398 #[test]
399 fn test_dictionary_hex_binary() {
400 let mut input_builder = BinaryDictionaryBuilder::<Int32Type>::new();
401 input_builder.append_value("1");
402 input_builder.append_value("j");
403 input_builder.append_null();
404 input_builder.append_value("3");
405 let input = input_builder.finish();
406
407 let mut expected_builder = StringDictionaryBuilder::<Int32Type>::new();
408 expected_builder.append_value("31");
409 expected_builder.append_value("6A");
410 expected_builder.append_null();
411 expected_builder.append_value("33");
412 let expected = expected_builder.finish();
413
414 let columnar_value = ColumnarValue::Array(Arc::new(input));
415 let result = super::spark_hex(&[columnar_value]).unwrap();
416
417 let result = match result {
418 ColumnarValue::Array(array) => array,
419 _ => panic!("Expected array"),
420 };
421
422 let result = as_dictionary_array(&result).unwrap();
423
424 assert_eq!(result, &expected);
425 }
426
427 #[test]
428 fn test_hex_int64() {
429 let cases = vec![
430 (0_i64, "0"),
431 (1, "1"),
432 (15, "F"),
433 (16, "10"),
434 (255, "FF"),
435 (256, "100"),
436 (1234, "4D2"),
437 (i64::MAX, "7FFFFFFFFFFFFFFF"),
438 (i64::MIN, "8000000000000000"),
439 (-1, "FFFFFFFFFFFFFFFF"),
440 ];
441
442 let arr =
443 super::hex_encode_int64(cases.iter().map(|(n, _)| Some(*n)), cases.len())
444 .unwrap();
445 let arr = as_string_array(&arr);
446 for (i, (num, expected)) in cases.iter().enumerate() {
447 assert_eq!(*expected, arr.value(i), "hex({num})");
448 }
449 }
450
451 #[test]
452 fn test_hex_encode_bytes_lowercase() {
453 let input = StringArray::from(vec![Some("hi"), Some("bye"), None, Some("rust")]);
458 let input_ref = &input;
459 let result = super::hex_encode_bytes(&input_ref, true).unwrap();
460 let result = as_string_array(&result);
461
462 let expected =
463 StringArray::from(vec![Some("6869"), Some("627965"), None, Some("72757374")]);
464 assert_eq!(result, &expected);
465 }
466
467 #[test]
468 fn test_spark_hex_binary_round_trip_all_bytes() {
469 let payload: Vec<u8> = (0u8..=255).collect();
472 let bin_array = BinaryArray::from(vec![Some(payload.as_slice())]);
473
474 let result =
475 super::spark_hex(&[ColumnarValue::Array(Arc::new(bin_array))]).unwrap();
476 let array = match result {
477 ColumnarValue::Array(array) => array,
478 _ => panic!("Expected array"),
479 };
480 let strings = as_string_array(&array);
481 let mut expected = String::with_capacity(512);
482 for byte in 0u8..=255 {
483 use std::fmt::Write;
484 write!(expected, "{byte:02X}").unwrap();
485 }
486 assert_eq!(strings.value(0), expected);
487 }
488
489 #[test]
490 fn test_spark_hex_binary_no_nulls() {
491 let input = BinaryArray::from(vec![
492 b"".as_slice(),
493 b"\x00\x7f\x80\xff".as_slice(),
494 b"DataFusion".as_slice(),
495 ]);
496
497 let result = super::spark_hex(&[ColumnarValue::Array(Arc::new(input))]).unwrap();
498 let array = match result {
499 ColumnarValue::Array(array) => array,
500 _ => panic!("Expected array"),
501 };
502 let strings = as_string_array(&array);
503
504 assert_eq!(strings.nulls(), None);
505 assert_eq!(
506 strings,
507 &StringArray::from(vec!["", "007F80FF", "44617461467573696F6E"])
508 );
509 }
510
511 #[test]
512 fn test_spark_hex_binary_reuses_input_nulls() {
513 let input = BinaryArray::from(vec![
514 Some(b"skip".as_slice()),
515 None,
516 Some(b"\x00\xff".as_slice()),
517 Some(b"hex".as_slice()),
518 None,
519 ])
520 .slice(1, 4);
521 let input_nulls = input.nulls().unwrap().clone();
522
523 let result = super::spark_hex(&[ColumnarValue::Array(Arc::new(input))]).unwrap();
524 let array = match result {
525 ColumnarValue::Array(array) => array,
526 _ => panic!("Expected array"),
527 };
528 let strings = as_string_array(&array);
529 let output_nulls = strings.nulls().unwrap();
530
531 assert_eq!(output_nulls, &input_nulls);
532 assert!(output_nulls.inner().ptr_eq(input_nulls.inner()));
533 assert_eq!(
534 strings,
535 &StringArray::from(vec![None, Some("00FF"), Some("686578"), None])
536 );
537 }
538
539 #[test]
540 fn test_spark_hex_int64() {
541 let int_array = Int64Array::from(vec![Some(1), Some(2), None, Some(3)]);
542 let columnar_value = ColumnarValue::Array(Arc::new(int_array));
543
544 let result = super::spark_hex(&[columnar_value]).unwrap();
545 let result = match result {
546 ColumnarValue::Array(array) => array,
547 _ => panic!("Expected array"),
548 };
549
550 let string_array = as_string_array(&result);
551 let expected_array = StringArray::from(vec![
552 Some("1".to_string()),
553 Some("2".to_string()),
554 None,
555 Some("3".to_string()),
556 ]);
557
558 assert_eq!(string_array, &expected_array);
559 }
560
561 #[test]
562 fn test_dict_values_null() {
563 let keys = Int32Array::from(vec![Some(0), None, Some(1)]);
564 let vals = Int64Array::from(vec![Some(32), None]);
565 let dict = DictionaryArray::new(keys, Arc::new(vals));
567
568 let columnar_value = ColumnarValue::Array(Arc::new(dict));
569 let result = super::spark_hex(&[columnar_value]).unwrap();
570
571 let result = match result {
572 ColumnarValue::Array(array) => array,
573 _ => panic!("Expected array"),
574 };
575
576 let result = as_dictionary_array(&result).unwrap();
577
578 let keys = Int32Array::from(vec![Some(0), None, Some(1)]);
579 let vals = StringArray::from(vec![Some("20"), None]);
580 let expected = DictionaryArray::new(keys, Arc::new(vals));
581
582 assert_eq!(&expected, result);
583 }
584
585 #[test]
586 fn test_dict_binary_values_null() {
587 let keys = Int32Array::from(vec![Some(0), None, Some(1)]);
588 let vals = BinaryArray::from(vec![Some(b"hi".as_slice()), None]);
589 let dict = DictionaryArray::new(keys, Arc::new(vals));
591
592 let result = super::spark_hex(&[ColumnarValue::Array(Arc::new(dict))]).unwrap();
593 let result = match result {
594 ColumnarValue::Array(array) => array,
595 _ => panic!("Expected array"),
596 };
597 let result = as_dictionary_array(&result).unwrap();
598
599 let keys = Int32Array::from(vec![Some(0), None, Some(1)]);
600 let vals = StringArray::from(vec![Some("6869"), None]);
601 let expected = DictionaryArray::new(keys, Arc::new(vals));
602
603 assert_eq!(&expected, result);
604 }
605}