Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
22 changes: 10 additions & 12 deletions datafusion/spark/src/function/url/parse_url.rs
Original file line number Diff line number Diff line change
Expand Up @@ -18,8 +18,8 @@
use std::sync::Arc;

use arrow::array::{
Array, ArrayRef, GenericStringBuilder, LargeStringArray, StringArray,
StringArrayType, StringViewArray,
Array, ArrayRef, AsArray, LargeStringArray, StringArray, StringArrayType,
StringViewArray, new_null_array,
};
use arrow::datatypes::DataType;
use datafusion_common::cast::{
Expand Down Expand Up @@ -272,20 +272,18 @@ pub fn spark_handled_parse_url(
),
}
} else {
// The 'key' argument is omitted, assume all values are null
// Create 'null' string array for 'key' argument
let mut builder: GenericStringBuilder<i32> = GenericStringBuilder::new();
for _ in 0..args[0].len() {
builder.append_null();
}
let key = builder.finish();
// The 'key' argument is omitted, assume all values are null.
// `new_null_array` allocates the null array outright, rather than
// appending one null per row through a builder.
let key_array = new_null_array(&DataType::Utf8, args[0].len());
let key = key_array.as_string::<i32>();

match (url.data_type(), part.data_type()) {
(DataType::Utf8, DataType::Utf8) => {
process_parse_url::<_, _, _, StringArray>(
as_string_array(url)?,
as_string_array(part)?,
&key,
key,
handler_err,
false,
)
Expand All @@ -294,7 +292,7 @@ pub fn spark_handled_parse_url(
process_parse_url::<_, _, _, StringViewArray>(
as_string_view_array(url)?,
as_string_view_array(part)?,
&key,
key,
handler_err,
false,
)
Expand All @@ -303,7 +301,7 @@ pub fn spark_handled_parse_url(
process_parse_url::<_, _, _, LargeStringArray>(
as_large_string_array(url)?,
as_large_string_array(part)?,
&key,
key,
handler_err,
false,
)
Expand Down
9 changes: 4 additions & 5 deletions datafusion/spark/src/function/url/try_url_decode.rs
Original file line number Diff line number Diff line change
Expand Up @@ -24,7 +24,9 @@ use datafusion_expr::{
};
use datafusion_functions::utils::make_scalar_function;

use crate::function::url::url_decode::{UrlDecode, spark_handled_url_decode};
use crate::function::url::url_decode::{
OnDecodeError, UrlDecode, spark_handled_url_decode,
};

#[derive(Debug, PartialEq, Eq, Hash)]
pub struct TryUrlDecode {
Expand Down Expand Up @@ -67,10 +69,7 @@ impl ScalarUDFImpl for TryUrlDecode {
}

fn spark_try_url_decode(args: &[ArrayRef]) -> Result<ArrayRef> {
spark_handled_url_decode(args, |x| match x {
Err(_) => Ok(None),
result => result,
})
spark_handled_url_decode(args, OnDecodeError::Null)
}

#[cfg(test)]
Expand Down
100 changes: 70 additions & 30 deletions datafusion/spark/src/function/url/url_decode.rs
Original file line number Diff line number Diff line change
Expand Up @@ -18,7 +18,9 @@
use std::borrow::Cow;
use std::sync::Arc;

use arrow::array::{ArrayRef, LargeStringArray, StringArray, StringViewArray};
use arrow::array::{
Array, ArrayRef, LargeStringBuilder, StringBuilder, StringViewBuilder,
};
use arrow::datatypes::DataType;
use datafusion_common::cast::{
as_large_string_array, as_string_array, as_string_view_array,
Expand Down Expand Up @@ -61,18 +63,25 @@ impl UrlDecode {
///
/// # Returns
///
/// * `Ok(String)` - The decoded string
/// * `Ok(Cow<str>)` - The decoded string, borrowed from `value` when there
/// was nothing to rewrite and owned otherwise
/// * `Err(DataFusionError)` - If the input is malformed or contains invalid UTF-8
///
fn decode(value: &str) -> Result<String> {
fn decode(value: &str) -> Result<Cow<'_, str>> {
// Check if the string has valid percent encoding
Self::validate_percent_encoding(value)?;

let replaced = Self::replace_plus(value.as_bytes());
percent_decode(&replaced)
.decode_utf8()
.map_err(|e| exec_datafusion_err!("Invalid UTF-8 sequence: {e}"))
.map(|parsed| parsed.into_owned())
match Self::replace_plus(value.as_bytes()) {
// No '+' was rewritten, so the decode can borrow from `value` itself.
Cow::Borrowed(bytes) => percent_decode(bytes)
.decode_utf8()
.map_err(|e| exec_datafusion_err!("Invalid UTF-8 sequence: {e}")),
// Rewriting '+' already allocated, so owning the decoded form here
// costs nothing beyond what has been spent.
Cow::Owned(bytes) => percent_decode(&bytes)
.decode_utf8()
.map(|decoded| Cow::Owned(decoded.into_owned()))
.map_err(|e| exec_datafusion_err!("Invalid UTF-8 sequence: {e}")),
}
}

/// Replace b'+' with b' '
Expand Down Expand Up @@ -155,6 +164,15 @@ impl ScalarUDFImpl for UrlDecode {
}
}

/// How [`spark_handled_url_decode`] reacts to a malformed input value.
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub enum OnDecodeError {
/// Propagate the error, as `url_decode` does.
Fail,
/// Return NULL for that row, as `try_url_decode` does.
Null,
}

/// Core implementation of URL decoding function.
///
/// # Arguments
Expand All @@ -165,38 +183,59 @@ impl ScalarUDFImpl for UrlDecode {
///
/// * `Ok(ArrayRef)` - A new array of the same type containing decoded strings
/// * `Err(DataFusionError)` - If validation fails or invalid arguments are provided
///
fn spark_url_decode(args: &[ArrayRef]) -> Result<ArrayRef> {
spark_handled_url_decode(args, |x| x)
spark_handled_url_decode(args, OnDecodeError::Fail)
}

pub fn spark_handled_url_decode(
args: &[ArrayRef],
err_handle_fn: impl Fn(Result<Option<String>>) -> Result<Option<String>>,
on_error: OnDecodeError,
) -> Result<ArrayRef> {
if args.len() != 1 {
return exec_err!("`url_decode` expects 1 argument");
}

// Decoded values go straight into the builder, so a row that needs no
// unescaping is copied once rather than materialised as its own `String`.
macro_rules! decode_all {
($array:expr, $builder:expr) => {{
let array = $array;
let mut builder = $builder;
for value in array.iter() {
let Some(value) = value else {
builder.append_null();
continue;
};
match UrlDecode::decode(value) {
Ok(decoded) => builder.append_value(&decoded),
Err(e) => match on_error {
OnDecodeError::Fail => return Err(e),
OnDecodeError::Null => builder.append_null(),
},
}
}
Ok(Arc::new(builder.finish()) as ArrayRef)
}};
}

match &args[0].data_type() {
DataType::Utf8 => as_string_array(&args[0])?
.iter()
.map(|x| x.map(UrlDecode::decode).transpose())
.map(&err_handle_fn)
.collect::<Result<StringArray>>()
.map(|array| Arc::new(array) as ArrayRef),
DataType::LargeUtf8 => as_large_string_array(&args[0])?
.iter()
.map(|x| x.map(UrlDecode::decode).transpose())
.map(&err_handle_fn)
.collect::<Result<LargeStringArray>>()
.map(|array| Arc::new(array) as ArrayRef),
DataType::Utf8View => as_string_view_array(&args[0])?
.iter()
.map(|x| x.map(UrlDecode::decode).transpose())
.map(&err_handle_fn)
.collect::<Result<StringViewArray>>()
.map(|array| Arc::new(array) as ArrayRef),
DataType::Utf8 => {
let array = as_string_array(&args[0])?;
let builder =
StringBuilder::with_capacity(array.len(), array.value_data().len());
decode_all!(array, builder)
}
DataType::LargeUtf8 => {
let array = as_large_string_array(&args[0])?;
let builder =
LargeStringBuilder::with_capacity(array.len(), array.value_data().len());
decode_all!(array, builder)
}
DataType::Utf8View => {
let array = as_string_view_array(&args[0])?;
let builder = StringViewBuilder::with_capacity(array.len());
decode_all!(array, builder)
}
other => exec_err!("`url_decode`: Expr must be STRING, got {other:?}"),
}
}
Expand All @@ -205,6 +244,7 @@ pub fn spark_handled_url_decode(
mod tests {

use super::*;
use arrow::array::StringArray;

#[test]
fn test_decode() -> Result<()> {
Expand Down
71 changes: 41 additions & 30 deletions datafusion/spark/src/function/url/url_encode.rs
Original file line number Diff line number Diff line change
Expand Up @@ -17,7 +17,9 @@

use std::sync::Arc;

use arrow::array::{ArrayRef, LargeStringArray, StringArray, StringViewArray};
use arrow::array::{
Array, ArrayRef, LargeStringBuilder, StringBuilder, StringViewBuilder,
};
use arrow::datatypes::DataType;
use datafusion_common::cast::{
as_large_string_array, as_string_array, as_string_view_array,
Expand Down Expand Up @@ -46,20 +48,6 @@ impl UrlEncode {
signature: Signature::string(1, Volatility::Immutable),
}
}

/// Encode a string to application/x-www-form-urlencoded format.
///
/// # Arguments
///
/// * `value` - The string to encode
///
/// # Returns
///
/// * `Ok(String)` - The encoded string
///
fn encode(value: &str) -> Result<String> {
Ok(byte_serialize(value.as_bytes()).collect::<String>())
}
}

impl ScalarUDFImpl for UrlEncode {
Expand Down Expand Up @@ -105,22 +93,45 @@ fn spark_url_encode(args: &[ArrayRef]) -> Result<ArrayRef> {
return exec_err!("`url_encode` expects 1 argument");
}

// The percent-encoded form of each value is assembled in a single scratch buffer
// reused across rows, rather than allocating a `String` per row.
macro_rules! encode_all {
($array:expr, $builder:expr) => {{
let array = $array;
let mut builder = $builder;
let mut encoded = String::new();
for value in array.iter() {
match value {
Some(value) => {
encoded.clear();
encoded.extend(byte_serialize(value.as_bytes()));
builder.append_value(&encoded);
}
None => builder.append_null(),
}
}
Ok(Arc::new(builder.finish()) as ArrayRef)
}};
}

match &args[0].data_type() {
DataType::Utf8 => as_string_array(&args[0])?
.iter()
.map(|x| x.map(UrlEncode::encode).transpose())
.collect::<Result<StringArray>>()
.map(|array| Arc::new(array) as ArrayRef),
DataType::LargeUtf8 => as_large_string_array(&args[0])?
.iter()
.map(|x| x.map(UrlEncode::encode).transpose())
.collect::<Result<LargeStringArray>>()
.map(|array| Arc::new(array) as ArrayRef),
DataType::Utf8View => as_string_view_array(&args[0])?
.iter()
.map(|x| x.map(UrlEncode::encode).transpose())
.collect::<Result<StringViewArray>>()
.map(|array| Arc::new(array) as ArrayRef),
DataType::Utf8 => {
let array = as_string_array(&args[0])?;
let builder =
StringBuilder::with_capacity(array.len(), array.value_data().len());
encode_all!(array, builder)
}
DataType::LargeUtf8 => {
let array = as_large_string_array(&args[0])?;
let builder =
LargeStringBuilder::with_capacity(array.len(), array.value_data().len());
encode_all!(array, builder)
}
DataType::Utf8View => {
let array = as_string_view_array(&args[0])?;
let builder = StringViewBuilder::with_capacity(array.len());
encode_all!(array, builder)
}
other => exec_err!("`url_encode`: Expr must be STRING, got {other:?}"),
}
}
Loading