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
99 changes: 91 additions & 8 deletions datafusion/functions/src/unicode/substr.rs
Original file line number Diff line number Diff line change
Expand Up @@ -24,15 +24,16 @@ use arrow::array::{
StringArrayType, StringViewArray,
};
use arrow::buffer::{NullBuffer, ScalarBuffer};
use arrow::datatypes::DataType;
use arrow::datatypes::{DataType, Field, FieldRef};
use datafusion_common::cast::as_int64_array;
use datafusion_common::types::{
NativeType, logical_int32, logical_int64, logical_string,
};
use datafusion_common::{Result, exec_err};
use datafusion_expr::{
Coercion, ColumnarValue, Documentation, ScalarFunctionArgs, ScalarUDFImpl, Signature,
TypeSignature, TypeSignatureClass, Volatility,
Coercion, ColumnarValue, Documentation, EncodingPreservation, ReturnFieldArgs,
ScalarFunctionArgs, ScalarUDFImpl, Signature, TypeSignature, TypeSignatureClass,
Volatility,
};
use datafusion_macros::user_doc;

Expand Down Expand Up @@ -73,7 +74,8 @@ impl Default for SubstrFunc {

impl SubstrFunc {
pub fn new() -> Self {
let string = Coercion::new_exact(TypeSignatureClass::Native(logical_string()));
let string = Coercion::new_exact(TypeSignatureClass::Native(logical_string()))
.with_encoding_preservation(EncodingPreservation::dictionary());
let int64 = Coercion::new_implicit(
TypeSignatureClass::Native(logical_int64()),
vec![TypeSignatureClass::Native(logical_int32())],
Expand Down Expand Up @@ -115,8 +117,51 @@ impl ScalarUDFImpl for SubstrFunc {
Ok(arg_types[0].clone())
}

fn return_field_from_args(&self, args: ReturnFieldArgs) -> Result<FieldRef> {
// Per-row start/length materializes the dict, so the output type is the
// value type; scalar tail keeps the dict wrapper.
let input = &args.arg_fields[0];
let output_type = match input.data_type() {
DataType::Dictionary(_, value_type)
if !args.scalar_arguments.iter().skip(1).all(Option::is_some) =>
{
value_type.as_ref().clone()
}
other => other.clone(),
};
Ok(Field::new(self.name(), output_type, input.is_nullable()).into())
}

fn invoke_with_args(&self, args: ScalarFunctionArgs) -> Result<ColumnarValue> {
make_scalar_function(substr, vec![])(&args.args)
if let ColumnarValue::Array(array) = &args.args[0]
&& matches!(array.data_type(), DataType::Dictionary(_, _))
&& args.args[1..]
.iter()
.all(|v| matches!(v, ColumnarValue::Scalar(_)))
{
// Fast path: substr the N_distinct values, reuse the keys.
let dict = array.as_any_dictionary();
let values = Arc::clone(dict.values());
let mut inner = Vec::with_capacity(args.args.len());
inner.push(ColumnarValue::Array(values));
inner.extend_from_slice(&args.args[1..]);
let new_values = make_scalar_function(substr, vec![])(&inner)?
.into_array(dict.values().len())?;
return Ok(ColumnarValue::Array(dict.with_values(new_values)));
}

// Per-row start/length: materialize the dict so substr can run row-wise.
let args = if let ColumnarValue::Array(array) = &args.args[0]
&& let DataType::Dictionary(_, value_type) = array.data_type()
{
let flat = arrow::compute::kernels::cast::cast(array, value_type.as_ref())?;
let mut new_args = args.args.clone();
new_args[0] = ColumnarValue::Array(flat);
new_args
} else {
args.args
};
make_scalar_function(substr, vec![])(&args)
}

fn aliases(&self) -> &[String] {
Expand Down Expand Up @@ -354,13 +399,15 @@ mod tests {
use std::sync::Arc;

use arrow::array::{
Array, ArrayRef, AsArray, Int64Array, LargeStringArray, StringArray,
StringViewArray,
Array, ArrayRef, AsArray, DictionaryArray, Int64Array, LargeStringArray,
StringArray, StringViewArray,
};
use arrow::datatypes::DataType::{LargeUtf8, Utf8, Utf8View};
use arrow::datatypes::{DataType, Field, Int32Type};

use datafusion_common::config::ConfigOptions;
use datafusion_common::{Result, ScalarValue, exec_err};
use datafusion_expr::{ColumnarValue, ScalarUDFImpl};
use datafusion_expr::{ColumnarValue, ScalarFunctionArgs, ScalarUDFImpl};

use crate::unicode::substr::SubstrFunc;
use crate::utils::test::test_function;
Expand Down Expand Up @@ -766,4 +813,40 @@ mod tests {

Ok(())
}

#[test]
fn test_substr_dict_fast_path_preserves_encoding() -> Result<()> {
let dict_type = DataType::Dictionary(Box::new(DataType::Int32), Box::new(Utf8));
let keys = arrow::array::Int32Array::from(vec![0, 1, 0, 2, 1]);
let values = StringArray::from(vec!["silver", "copper", "golden"]);
let dict: ArrayRef = Arc::new(
DictionaryArray::<Int32Type>::try_new(keys, Arc::new(values)).unwrap(),
);

let result = SubstrFunc::new()
.invoke_with_args(ScalarFunctionArgs {
args: vec![
ColumnarValue::Array(dict),
ColumnarValue::Scalar(ScalarValue::from(1i64)),
ColumnarValue::Scalar(ScalarValue::from(3i64)),
],
arg_fields: vec![
Arc::new(Field::new("s", dict_type.clone(), true)),
Arc::new(Field::new("start", DataType::Int64, false)),
Arc::new(Field::new("len", DataType::Int64, false)),
],
number_rows: 5,
return_field: Arc::new(Field::new("r", dict_type, true)),
config_options: Arc::new(ConfigOptions::default()),
})?
.into_array(5)?;

assert!(matches!(result.data_type(), DataType::Dictionary(_, v) if **v == Utf8));
let out = result.as_dictionary::<Int32Type>();
let v = out.values().as_string::<i32>();
assert_eq!(v.len(), 3);
assert_eq!((v.value(0), v.value(1), v.value(2)), ("sil", "cop", "gol"));

Ok(())
}
}
6 changes: 6 additions & 0 deletions datafusion/sqllogictest/test_files/string/dictionary_utf8.slt
Original file line number Diff line number Diff line change
Expand Up @@ -57,6 +57,12 @@ SELECT arrow_cast('', 'Dictionary(Int32, Utf8)');
#
include ./string_query.slt.part

# Scalar start/length preserves dict encoding end-to-end.
query T
SELECT DISTINCT arrow_typeof(substr(c1, 1, 3)) FROM test_substr;
----
Dictionary(Int32, Utf8)

#
# Clean up
#
Expand Down
Loading