Skip to main content

datafusion_spark/function/math/
hex.rs

1// Licensed to the Apache Software Foundation (ASF) under one
2// or more contributor license agreements.  See the NOTICE file
3// distributed with this work for additional information
4// regarding copyright ownership.  The ASF licenses this file
5// to you under the Apache License, Version 2.0 (the
6// "License"); you may not use this file except in compliance
7// with the License.  You may obtain a copy of the License at
8//
9//   http://www.apache.org/licenses/LICENSE-2.0
10//
11// Unless required by applicable law or agreed to in writing,
12// software distributed under the License is distributed on an
13// "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY
14// KIND, either express or implied.  See the License for the
15// specific language governing permissions and limitations
16// under the License.
17
18use 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/// <https://spark.apache.org/docs/latest/api/sql/index.html#hex>
43#[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            // accepts numeric types
70            TypeSignature::Coercible(vec![int64]),
71            // accepts string types (Utf8, Utf8View, LargeUtf8)
72            TypeSignature::Coercible(vec![string]),
73            // accepts binary types (Binary, FixedSizeBinary, LargeBinary)
74            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
132/// Generic hex encoding for byte array types
133fn 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    // Write hex digits directly into one growing value buffer, tracking offsets
150    // ourselves. Each input byte becomes exactly two output bytes, so there is
151    // no per-row `String`/`StringBuilder` copy — the hex digits are written once
152    // into the final buffer.
153    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                // SAFETY: `i` is in bounds and the validity buffer marks it valid.
161                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            // SAFETY: `i` is in bounds and no null buffer means every value is valid.
172            let bytes = unsafe { array.value_unchecked(i) }.as_ref();
173            offsets.push(append_hex_bytes(&mut values, bytes, case)?);
174        }
175    }
176
177    // SAFETY: the value buffer contains only ASCII hex digits (valid UTF-8) and
178    // the offsets are monotonically increasing and end at `values.len()`, so the
179    // array invariants hold. This mirrors the previous `from_utf8_unchecked`
180    // path and avoids a redundant UTF-8 validation pass over the whole buffer.
181    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
191/// Generic hex encoding for int64 type
192fn 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            // SAFETY: slice contains only ASCII hex digests, which are valid UTF-8
203            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
214/// Spark-compatible `hex` function
215pub fn spark_hex(args: &[ColumnarValue]) -> Result<ColumnarValue, DataFusionError> {
216    compute_hex(args, false)
217}
218
219/// Spark-compatible `sha2` function
220pub 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        // Every in-repo caller of `hex_encode_bytes` goes through `spark_hex`,
454        // which always passes `lowercase = false`. The `lowercase = true` path
455        // is reachable only via `spark_sha2_hex`, which has no in-workspace
456        // caller, so it otherwise has no coverage. Drive it directly here.
457        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        // Single-row binary input containing every byte value, encoded in
470        // a single column. Catches per-byte regressions in the bytes path.
471        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        // [32, null, null]
566        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        // [b"hi", null, null]
590        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}