Skip to main content

datafusion_spark/function/hash/
sha2.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 arrow::array::{ArrayRef, AsArray, BinaryArrayType, Int32Array, StringArray};
19use arrow::datatypes::{DataType, Int32Type};
20use datafusion_common::types::{
21    NativeType, logical_binary, logical_int32, logical_string,
22};
23use datafusion_common::utils::hex::{HexCase, encode_bytes};
24use datafusion_common::utils::take_function_args;
25use datafusion_common::{Result, ScalarValue, internal_err};
26use datafusion_expr::{
27    Coercion, ColumnarValue, ScalarFunctionArgs, ScalarUDFImpl, Signature,
28    TypeSignatureClass, Volatility,
29};
30use datafusion_functions::utils::make_scalar_function;
31use sha2::{self, Digest};
32use std::sync::Arc;
33
34/// Differs from DataFusion version in allowing array input for bit lengths, and
35/// also hex encoding the output.
36///
37/// <https://spark.apache.org/docs/latest/api/sql/index.html#sha2>
38#[derive(Debug, PartialEq, Eq, Hash)]
39pub struct SparkSha2 {
40    signature: Signature,
41}
42
43impl Default for SparkSha2 {
44    fn default() -> Self {
45        Self::new()
46    }
47}
48
49impl SparkSha2 {
50    pub fn new() -> Self {
51        Self {
52            signature: Signature::coercible(
53                vec![
54                    Coercion::new_implicit(
55                        TypeSignatureClass::Native(logical_binary()),
56                        vec![TypeSignatureClass::Native(logical_string())],
57                        NativeType::Binary,
58                    ),
59                    Coercion::new_implicit(
60                        TypeSignatureClass::Native(logical_int32()),
61                        vec![TypeSignatureClass::Integer],
62                        NativeType::Int32,
63                    ),
64                ],
65                Volatility::Immutable,
66            ),
67        }
68    }
69}
70
71impl ScalarUDFImpl for SparkSha2 {
72    fn name(&self) -> &str {
73        "sha2"
74    }
75
76    fn signature(&self) -> &Signature {
77        &self.signature
78    }
79
80    fn return_type(&self, _arg_types: &[DataType]) -> Result<DataType> {
81        Ok(DataType::Utf8)
82    }
83
84    fn invoke_with_args(&self, args: ScalarFunctionArgs) -> Result<ColumnarValue> {
85        let [values, bit_lengths] = take_function_args(self.name(), args.args.iter())?;
86
87        match (values, bit_lengths) {
88            (
89                ColumnarValue::Scalar(value_scalar),
90                ColumnarValue::Scalar(ScalarValue::Int32(Some(bit_length))),
91            ) => {
92                if value_scalar.is_null() {
93                    return Ok(ColumnarValue::Scalar(ScalarValue::Utf8(None)));
94                }
95
96                // Accept both Binary and Utf8 scalars (depending on coercion)
97                let bytes = match value_scalar {
98                    ScalarValue::Binary(Some(b)) => b.as_slice(),
99                    ScalarValue::LargeBinary(Some(b)) => b.as_slice(),
100                    ScalarValue::BinaryView(Some(b)) => b.as_slice(),
101                    ScalarValue::Utf8(Some(s))
102                    | ScalarValue::LargeUtf8(Some(s))
103                    | ScalarValue::Utf8View(Some(s)) => s.as_bytes(),
104                    other => {
105                        return internal_err!(
106                            "Unsupported scalar datatype for sha2: {}",
107                            other.data_type()
108                        );
109                    }
110                };
111
112                let out = match bit_length {
113                    224 => {
114                        let mut digest = sha2::Sha224::default();
115                        digest.update(bytes);
116                        Some(encode_bytes(&digest.finalize(), HexCase::Lower))
117                    }
118                    0 | 256 => {
119                        let mut digest = sha2::Sha256::default();
120                        digest.update(bytes);
121                        Some(encode_bytes(&digest.finalize(), HexCase::Lower))
122                    }
123                    384 => {
124                        let mut digest = sha2::Sha384::default();
125                        digest.update(bytes);
126                        Some(encode_bytes(&digest.finalize(), HexCase::Lower))
127                    }
128                    512 => {
129                        let mut digest = sha2::Sha512::default();
130                        digest.update(bytes);
131                        Some(encode_bytes(&digest.finalize(), HexCase::Lower))
132                    }
133                    _ => None,
134                };
135
136                Ok(ColumnarValue::Scalar(ScalarValue::Utf8(out)))
137            }
138            // Array values + scalar bit length (common case: sha2(col, 256))
139            (
140                ColumnarValue::Array(values_array),
141                ColumnarValue::Scalar(ScalarValue::Int32(Some(bit_length))),
142            ) => {
143                let output: ArrayRef = match values_array.data_type() {
144                    DataType::Binary => sha2_binary_scalar_bitlen(
145                        &values_array.as_binary::<i32>(),
146                        *bit_length,
147                    ),
148                    DataType::LargeBinary => sha2_binary_scalar_bitlen(
149                        &values_array.as_binary::<i64>(),
150                        *bit_length,
151                    ),
152                    DataType::BinaryView => sha2_binary_scalar_bitlen(
153                        &values_array.as_binary_view(),
154                        *bit_length,
155                    ),
156                    dt => return internal_err!("Unsupported datatype for sha2: {dt}"),
157                };
158                Ok(ColumnarValue::Array(output))
159            }
160            (
161                ColumnarValue::Scalar(_),
162                ColumnarValue::Scalar(ScalarValue::Int32(None)),
163            ) => Ok(ColumnarValue::Scalar(ScalarValue::Utf8(None))),
164            (
165                ColumnarValue::Array(_),
166                ColumnarValue::Scalar(ScalarValue::Int32(None)),
167            ) => Ok(ColumnarValue::Scalar(ScalarValue::Utf8(None))),
168            _ => {
169                // Fallback to existing behavior for any array/mixed cases
170                make_scalar_function(sha2_impl, vec![])(&args.args)
171            }
172        }
173    }
174}
175
176fn sha2_impl(args: &[ArrayRef]) -> Result<ArrayRef> {
177    let [values, bit_lengths] = take_function_args("sha2", args)?;
178
179    let bit_lengths = bit_lengths.as_primitive::<Int32Type>();
180    let output = match values.data_type() {
181        DataType::Binary => sha2_binary_impl(&values.as_binary::<i32>(), bit_lengths),
182        DataType::LargeBinary => {
183            sha2_binary_impl(&values.as_binary::<i64>(), bit_lengths)
184        }
185        DataType::BinaryView => sha2_binary_impl(&values.as_binary_view(), bit_lengths),
186        dt => return internal_err!("Unsupported datatype for sha2: {dt}"),
187    };
188    Ok(output)
189}
190
191fn sha2_binary_impl<'a, BinaryArrType>(
192    values: &BinaryArrType,
193    bit_lengths: &Int32Array,
194) -> ArrayRef
195where
196    BinaryArrType: BinaryArrayType<'a>,
197{
198    sha2_binary_bitlen_iter(values, bit_lengths.iter())
199}
200
201fn sha2_binary_scalar_bitlen<'a, BinaryArrType>(
202    values: &BinaryArrType,
203    bit_length: i32,
204) -> ArrayRef
205where
206    BinaryArrType: BinaryArrayType<'a>,
207{
208    sha2_binary_bitlen_iter(values, std::iter::repeat(Some(bit_length)))
209}
210
211fn sha2_binary_bitlen_iter<'a, BinaryArrType, I>(
212    values: &BinaryArrType,
213    bit_lengths: I,
214) -> ArrayRef
215where
216    BinaryArrType: BinaryArrayType<'a>,
217    I: Iterator<Item = Option<i32>>,
218{
219    let array = values
220        .iter()
221        .zip(bit_lengths)
222        .map(|(value, bit_length)| match (value, bit_length) {
223            (Some(value), Some(224)) => {
224                let mut digest = sha2::Sha224::default();
225                digest.update(value);
226                Some(encode_bytes(&digest.finalize(), HexCase::Lower))
227            }
228            (Some(value), Some(0 | 256)) => {
229                let mut digest = sha2::Sha256::default();
230                digest.update(value);
231                Some(encode_bytes(&digest.finalize(), HexCase::Lower))
232            }
233            (Some(value), Some(384)) => {
234                let mut digest = sha2::Sha384::default();
235                digest.update(value);
236                Some(encode_bytes(&digest.finalize(), HexCase::Lower))
237            }
238            (Some(value), Some(512)) => {
239                let mut digest = sha2::Sha512::default();
240                digest.update(value);
241                Some(encode_bytes(&digest.finalize(), HexCase::Lower))
242            }
243            // Unknown bit-lengths go to null, same as in Spark
244            _ => None,
245        })
246        .collect::<StringArray>();
247    Arc::new(array)
248}