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
256 changes: 153 additions & 103 deletions datafusion/catalog/src/information_schema.rs
Original file line number Diff line number Diff line change
Expand Up @@ -32,17 +32,19 @@ use async_trait::async_trait;
use datafusion_common::DataFusionError;
use datafusion_common::config::{ConfigEntry, ConfigOptions};
use datafusion_common::error::Result;
use datafusion_common::types::NativeType;
use datafusion_common::types::{LogicalType, NativeType};
use datafusion_execution::TaskContext;
use datafusion_execution::runtime_env::RuntimeEnv;
use datafusion_expr::function::WindowUDFFieldArgs;
use datafusion_expr::type_coercion::functions::fields_with_udf;
use datafusion_expr::{
AggregateUDF, ReturnFieldArgs, ScalarUDF, Signature, TypeSignature, WindowUDF,
};
use datafusion_expr::{TableType, Volatility};
use datafusion_physical_plan::SendableRecordBatchStream;
use datafusion_physical_plan::stream::RecordBatchStreamAdapter;
use datafusion_physical_plan::streaming::PartitionStream;
use itertools::Itertools;
use std::collections::{BTreeSet, HashMap, HashSet};
use std::fmt::Debug;
use std::sync::Arc;
Expand Down Expand Up @@ -454,118 +456,99 @@ impl InformationSchemaConfig {
}
}

/// get the arguments and return types of a UDF
/// returns a tuple of (arg_types, return_type)
/// Origins used to enumerate the physical types a native type can take
const RESOLVE_CAST_SOURCES: [DataType; 2] = [DataType::Null, DataType::LargeUtf8];

/// Build argument fields for `information_schema` to provide possible return types
fn resolve_informational_fields(idx: usize, t: &NativeType) -> Vec<FieldRef> {
// Since native types map to several physical types, resolve it against
// ambiguous types to get canonical `DataType`s for the native type.
// Skip origins the type has no cast from (e.g. `Struct` from `LargeUtf8`)
RESOLVE_CAST_SOURCES
.iter()
.filter_map(|source| t.default_cast_for(source).ok())
.unique()
.map(|dt| Arc::new(Field::new(format!("arg_{idx}"), dt, true)))
.collect()
}

/// Function information schema is a set of tuples - argument types and an optional return type
type FunctionInformationSchema = BTreeSet<(Vec<String>, Option<String>)>;

/// Get the arguments and return types of a function from its signature
fn get_args_and_return_types(
signature: &Signature,
return_field: impl Fn(&[FieldRef]) -> Result<FieldRef>,
) -> Result<FunctionInformationSchema> {
let arg_types = signature.type_signature.get_representative_types();
if arg_types.is_empty() {
// Edge case if function doesn't have arguments
return Ok(BTreeSet::from([(vec![], None)]));
}
arg_types
.into_iter()
.map(|arg_types| {
// Get possible types for each input arg
let arg_fields = arg_types
.iter()
.enumerate()
.map(|(i, t)| resolve_informational_fields(i, t))
.collect::<Vec<_>>();
// Build combinations of arg types with the return type
let return_types = arg_fields

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.

generate_series/range now report only (String, Date, Interval) → List(String), but the query returns List(Date32). Base also listed List(Date), but only because its Date64 example happened to hit (_, Some(Date64)) in Range::return_type. That match is lost now that Date resolves only to Date32. The underlying cause predates this PR: the return type is computed on arguments that haven't been coerced, which the get_representative_types docs say callers must do. Fine to handle in a follow-up.

select arrow_typeof(generate_series('2020-01-01', DATE '2020-01-03', INTERVAL '1 day'));
-- List(Date32)
select data_type from information_schema.parameters
where specific_name = 'generate_series' and parameter_mode = 'OUT';
-- has List(String) for (String, Date, Interval), no List(Date)

Fix: coerce the arguments the same way the planner does before asking for the return type. I tried this locally: it fixes this row and also corrects sum(Int32) (NULL → Int64), trunc(Int64) (NULL → Float64), median(Int64) (Int64 → Float64) and date_trunc(String, Date) (Date → Timestamp(ns)), all checked against arrow_typeof. It does change the existing date_trunc expectations in information_schema.slt.

 use datafusion_expr::function::WindowUDFFieldArgs;
+use datafusion_expr::type_coercion::functions::fields_with_udf;
     get_args_and_return_types(udf.signature(), |arg_fields| {
+        let arg_fields = &fields_with_udf(arg_fields, udf.as_ref())?;
         let scalar_arguments = &vec![None; arg_fields.len()];
     get_args_and_return_types(udaf.signature(), |arg_fields| {
-        udaf.return_field(arg_fields)
+        udaf.return_field(&fields_with_udf(arg_fields, udaf.as_ref())?)
     })
     get_args_and_return_types(udwf.signature(), |arg_fields| {
+        let arg_fields = &fields_with_udf(arg_fields, udwf.as_ref())?;
         udwf.field(WindowUDFFieldArgs::new(arg_fields, udwf.name()))

Copy link
Copy Markdown
Contributor Author

Choose a reason for hiding this comment

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

Added, tests are updated to check date_trunc and a new generate_series case

.into_iter()
.multi_cartesian_product()
.filter_map(|arg_fields| return_field(&arg_fields).ok())
.map(|f| Some(remove_native_type_prefix(&f.data_type().into())))
.collect::<BTreeSet<_>>();
let return_types = if return_types.is_empty() {
// Indicate `None` if the return type cannot be represented from a signature,
BTreeSet::from([None])
} else {
return_types
};
let arg_types = arg_types
.iter()
.map(remove_native_type_prefix)
.collect::<Vec<_>>();
let tuples = return_types
.into_iter()
.map(move |return_type| (arg_types.clone(), return_type));
Ok(tuples)
})
.flatten_ok()
.collect()
}

fn get_udf_args_and_return_types(
udf: &Arc<ScalarUDF>,
) -> Result<BTreeSet<(Vec<String>, Option<String>)>> {
let signature = udf.signature();
let arg_types = signature.type_signature.get_example_types();
if arg_types.is_empty() {
Ok(vec![(vec![], None)].into_iter().collect::<BTreeSet<_>>())
} else {
Ok(arg_types
.into_iter()
.map(|arg_types| {
let arg_fields: Vec<FieldRef> = arg_types
.iter()
.enumerate()
.map(|(i, t)| {
Arc::new(Field::new(format!("arg_{i}"), t.clone(), true))
})
.collect();
let scalar_arguments = vec![None; arg_fields.len()];
let return_type = udf
.return_field_from_args(ReturnFieldArgs {
arg_fields: &arg_fields,
scalar_arguments: &scalar_arguments,
})
.map(|f| {
remove_native_type_prefix(&NativeType::from(
f.data_type().clone(),
))
})
.ok();
let arg_types = arg_types
.into_iter()
.map(|t| remove_native_type_prefix(&NativeType::from(t)))
.collect::<Vec<_>>();
(arg_types, return_type)
})
.collect::<BTreeSet<_>>())
}
) -> Result<FunctionInformationSchema> {
get_args_and_return_types(udf.signature(), |arg_fields| {
let arg_fields = &fields_with_udf(arg_fields, udf.as_ref())?;
let scalar_arguments = &vec![None; arg_fields.len()];
udf.return_field_from_args(ReturnFieldArgs {
arg_fields,
scalar_arguments,
})
})
}

fn get_udaf_args_and_return_types(
udaf: &Arc<AggregateUDF>,
) -> Result<BTreeSet<(Vec<String>, Option<String>)>> {
let signature = udaf.signature();
let arg_types = signature.type_signature.get_example_types();
if arg_types.is_empty() {
Ok(vec![(vec![], None)].into_iter().collect::<BTreeSet<_>>())
} else {
Ok(arg_types
.into_iter()
.map(|arg_types| {
let arg_fields: Vec<FieldRef> = arg_types
.iter()
.enumerate()
.map(|(i, t)| {
Arc::new(Field::new(format!("arg_{i}"), t.clone(), true))
})
.collect();
let return_type = udaf
.return_field(&arg_fields)
.map(|f| {
remove_native_type_prefix(&NativeType::from(
f.data_type().clone(),
))
})
.ok();
let arg_types = arg_types
.into_iter()
.map(|t| remove_native_type_prefix(&NativeType::from(t)))
.collect::<Vec<_>>();
(arg_types, return_type)
})
.collect::<BTreeSet<_>>())
}
) -> Result<FunctionInformationSchema> {
get_args_and_return_types(udaf.signature(), |arg_fields| {
let arg_fields = &fields_with_udf(arg_fields, udaf.as_ref())?;
udaf.return_field(arg_fields)
})
}

fn get_udwf_args_and_return_types(
udwf: &Arc<WindowUDF>,
) -> Result<BTreeSet<(Vec<String>, Option<String>)>> {
let signature = udwf.signature();
let arg_types = signature.type_signature.get_example_types();
if arg_types.is_empty() {
Ok(vec![(vec![], None)].into_iter().collect::<BTreeSet<_>>())
} else {
Ok(arg_types
.into_iter()
.map(|arg_types| {
let arg_fields: Vec<FieldRef> = arg_types
.iter()
.enumerate()
.map(|(i, t)| {
Arc::new(Field::new(format!("arg_{i}"), t.clone(), true))
})
.collect();
let return_type = udwf
.field(WindowUDFFieldArgs::new(&arg_fields, udwf.name()))
.map(|f| {
remove_native_type_prefix(&NativeType::from(
f.data_type().clone(),
))
})
.ok();
let arg_types = arg_types
.into_iter()
.map(|t| remove_native_type_prefix(&NativeType::from(t)))
.collect::<Vec<_>>();
(arg_types, return_type)
})
.collect::<BTreeSet<_>>())
}
) -> Result<FunctionInformationSchema> {
get_args_and_return_types(udwf.signature(), |arg_fields| {
let arg_fields = &fields_with_udf(arg_fields, udwf.as_ref())?;
udwf.field(WindowUDFFieldArgs::new(arg_fields, udwf.name()))
})
}

#[inline]
Expand Down Expand Up @@ -1510,6 +1493,9 @@ mod tests {
use super::*;
use crate::CatalogProvider;
use arrow::array::Array;
use arrow::datatypes::Fields;
use datafusion_common::ScalarValue;
use datafusion_expr::{ColumnarValue, ScalarFunctionArgs, ScalarUDFImpl};

#[test]
fn schemata_builder_emits_canonical_schema_and_rows() {
Expand Down Expand Up @@ -1580,6 +1566,70 @@ mod tests {
assert_eq!("BASE TABLE", builder.table_types.finish().value(0));
}

// UDF
#[derive(Debug, PartialEq, Eq, Hash)]
struct TestScalarUDF {
signature: Signature,
}
impl ScalarUDFImpl for TestScalarUDF {
fn name(&self) -> &str {
"TestScalarUDF"
}

fn signature(&self) -> &Signature {
&self.signature
}

fn return_type(&self, arg_types: &[DataType]) -> Result<DataType> {
Ok(arg_types.last().unwrap().clone())
}

fn invoke_with_args(&self, _args: ScalarFunctionArgs) -> Result<ColumnarValue> {
Ok(ColumnarValue::Scalar(ScalarValue::from("a")))
}
}

#[test]
fn test_get_udf_args_and_return_types() -> Result<()> {
// heterogeneous arguments to test mixed arguments retrieval
let signature = Signature::exact(
[
vec![DataType::Int32; 6],
vec![DataType::Float32; 6],
vec![DataType::Utf8; 1],
]
.concat(),
Volatility::Stable,
);
let udf = Arc::new(ScalarUDF::from(TestScalarUDF { signature }));
let result = get_udf_args_and_return_types(&udf)?;
assert_eq!(result.len(), 1);
let (args, ret) = result.iter().next().unwrap();
assert_eq!(
*args,
[
vec![String::from("Int32"); 6],
vec![String::from("Float32"); 6],
vec![String::from("String"); 1]
]
.concat()
);
assert_eq!(*ret, Some(String::from("String")));

Ok(())
}

#[test]
fn test_get_udf_args_and_return_types_nested() -> Result<()> {
let struct_type =
DataType::Struct(Fields::from(vec![Field::new("a", DataType::Int32, true)]));
let signature = Signature::exact(vec![struct_type], Volatility::Stable);
let udf = Arc::new(ScalarUDF::from(TestScalarUDF { signature }));
let result = get_udf_args_and_return_types(&udf)?;
assert_eq!(result.len(), 1);
Ok(())
}

#[derive(Debug)]
struct Fixture;

Expand Down
Loading
Loading