datafusion_spark/function/hash/
sha2.rs1use 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#[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 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 (
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 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 _ => None,
245 })
246 .collect::<StringArray>();
247 Arc::new(array)
248}