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
64 changes: 43 additions & 21 deletions datafusion/spark/src/function/string/quote.rs
Original file line number Diff line number Diff line change
Expand Up @@ -15,7 +15,7 @@
// specific language governing permissions and limitations
// under the License.

use arrow::array::{ArrayRef, OffsetSizeTrait, StringArray};
use arrow::array::{Array, ArrayRef, OffsetSizeTrait, StringBuilder};
use arrow::datatypes::DataType;
use datafusion::logical_expr::{Coercion, ColumnarValue, Signature, TypeSignatureClass};
use datafusion_common::cast::{as_generic_string_array, as_string_view_array};
Expand All @@ -25,6 +25,7 @@ use datafusion_common::{Result, exec_err};
use datafusion_expr::{ScalarFunctionArgs, ScalarUDFImpl, Volatility};
use datafusion_functions::utils::make_scalar_function;

use std::fmt::Write;
use std::sync::Arc;

/// Spark-compatible `quote` expression
Expand Down Expand Up @@ -88,34 +89,55 @@ fn spark_quote_inner(arg: &[ArrayRef]) -> Result<ArrayRef> {

fn quote_array<T: OffsetSizeTrait>(array: &ArrayRef) -> Result<ArrayRef> {
let str_array = as_generic_string_array::<T>(array)?;
let result = str_array
.iter()
.map(|s| s.map(compute_quote))
.collect::<StringArray>();
Ok(Arc::new(result))
Ok(quote_impl(str_array.iter(), str_array.value_data().len()))

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

str_array.value_data().len() will over-allocate for sliced arrays.

}

fn quote_view(str_view: &ArrayRef) -> Result<ArrayRef> {
let str_array = as_string_view_array(str_view)?;
let result = str_array
.iter()
.map(|opt_str| opt_str.map(compute_quote))
.collect::<StringArray>();
Ok(Arc::new(result) as ArrayRef)
Ok(quote_impl(
str_array.iter(),
str_array.get_buffer_memory_size(),

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Looks like get_buffer_memory_size sums the buffer capacities, not their actual valid contents, so this will also over-allocate for sliced arrays.

))
}

const QUOTE_CHAR: char = '\'';
const ESCAPE_CHAR: char = '\\';
/// A literal quote in the input is emitted as this two-character escape.
const ESCAPED_QUOTE: &str = "\\'";

fn compute_quote(s: &str) -> String {
let mut quoted = String::with_capacity(s.len() + 2);
quoted.push(QUOTE_CHAR);
for c in s.chars() {
if c == QUOTE_CHAR {
quoted.push(ESCAPE_CHAR);
/// Quotes every value, writing directly into the output buffer.
///
/// `data_capacity` is a hint for the total input byte length; the output adds two
/// surrounding quotes per row plus one byte per escaped quote.
fn quote_impl<'a>(
input: impl Iterator<Item = Option<&'a str>>,
data_capacity: usize,
) -> ArrayRef {
let len = input.size_hint().0;
let mut builder = StringBuilder::with_capacity(len, data_capacity + 2 * len);
for value in input {
match value {
Some(value) => append_quoted(&mut builder, value),
None => builder.append_null(),
}
quoted.push(c);
}
quoted.push(QUOTE_CHAR);
quoted
Arc::new(builder.finish())
}

/// Appends `s` wrapped in single quotes, with any embedded quote backslash-escaped.
///
/// Writes straight into the builder's buffer — finalized by the trailing
/// `append_value("")` — so no intermediate `String` is allocated per row, and
/// copies the runs between quotes rather than one character at a time.
fn append_quoted(builder: &mut StringBuilder, s: &str) {
// `write_str` on a `GenericStringBuilder` is infallible.
let mut runs = s.split(QUOTE_CHAR);
builder.write_char(QUOTE_CHAR).unwrap();
// `split` always yields at least one run.
builder.write_str(runs.next().unwrap_or_default()).unwrap();
for run in runs {
builder.write_str(ESCAPED_QUOTE).unwrap();
builder.write_str(run).unwrap();
}
builder.write_char(QUOTE_CHAR).unwrap();
builder.append_value("");
}
58 changes: 38 additions & 20 deletions datafusion/spark/src/function/string/soundex.rs
Original file line number Diff line number Diff line change
Expand Up @@ -15,7 +15,7 @@
// specific language governing permissions and limitations
// under the License.

use arrow::array::{ArrayRef, OffsetSizeTrait, StringArray};
use arrow::array::{ArrayRef, OffsetSizeTrait, StringBuilder};
use arrow::datatypes::DataType;
use datafusion_common::cast::{as_generic_string_array, as_string_view_array};
use datafusion_common::utils::take_function_args;
Expand Down Expand Up @@ -80,21 +80,24 @@ fn spark_soundex_inner(arg: &[ArrayRef]) -> Result<ArrayRef> {
}

fn soundex_array<T: OffsetSizeTrait>(array: &ArrayRef) -> Result<ArrayRef> {
let str_array = as_generic_string_array::<T>(array)?;
let result = str_array
.iter()
.map(|s| s.map(compute_soundex))
.collect::<StringArray>();
Ok(Arc::new(result))
Ok(soundex_impl(as_generic_string_array::<T>(array)?.iter()))
}

fn soundex_view(str_view: &ArrayRef) -> Result<ArrayRef> {
let str_array = as_string_view_array(str_view)?;
let result = str_array
.iter()
.map(|opt_str| opt_str.map(compute_soundex))
.collect::<StringArray>();
Ok(Arc::new(result) as ArrayRef)
Ok(soundex_impl(as_string_view_array(str_view)?.iter()))
}

fn soundex_impl<'a>(input: impl Iterator<Item = Option<&'a str>>) -> ArrayRef {
let len = input.size_hint().0;
// A soundex code is always exactly 4 ASCII characters.
let mut builder = StringBuilder::with_capacity(len, len * SOUNDEX_LEN);
for value in input {
match value {
Some(value) => append_soundex(&mut builder, value),
None => builder.append_null(),
}
}
Arc::new(builder.finish())
}

fn classify_char(c: char) -> Option<char> {
Expand All @@ -113,20 +116,32 @@ fn is_ignored(c: char) -> bool {
matches!(c.to_ascii_uppercase(), 'H' | 'W')
}

fn compute_soundex(s: &str) -> String {
/// Length of a soundex code: an initial letter plus three digits.
const SOUNDEX_LEN: usize = 4;

/// Appends the soundex code of `s` to `builder`.
///
/// Strings that do not start with an ASCII letter are passed through unchanged.
/// Otherwise the code is built in a stack buffer, so no row allocates.
fn append_soundex(builder: &mut StringBuilder, s: &str) {
let mut chars = s.chars();

let first_char = match chars.next() {
Some(c) if c.is_ascii_alphabetic() => c.to_ascii_uppercase(),
_ => return s.to_string(),
_ => {
builder.append_value(s);
return;
}
};

let mut soundex_code = String::with_capacity(4);
soundex_code.push(first_char);
// Codes shorter than four characters are right-padded with '0'.
let mut soundex_code = [b'0'; SOUNDEX_LEN];
soundex_code[0] = first_char as u8;
let mut written = 1;
let mut last_code = classify_char(first_char);

for c in chars {
if soundex_code.len() >= 4 {
if written >= SOUNDEX_LEN {
break;
}

Expand All @@ -137,7 +152,8 @@ fn compute_soundex(s: &str) -> String {
match classify_char(c) {
Some(code) => {
if last_code != Some(code) {
soundex_code.push(code);
soundex_code[written] = code as u8;
written += 1;
}
last_code = Some(code);
}
Expand All @@ -146,5 +162,7 @@ fn compute_soundex(s: &str) -> String {
}
}
}
format!("{soundex_code:0<4}")

// SAFETY: `soundex_code` holds an ASCII letter followed by ASCII digits.
builder.append_value(unsafe { std::str::from_utf8_unchecked(&soundex_code) });
}
Loading