use std::sync::Arc;
use arrow::datatypes::{DataType, Field};
use datafusion_common::{
Result, ScalarValue,
tree_node::{Transformed, TreeNode, TreeNodeRecursion},
};
use datafusion_expr::ScalarUDFImpl;
use datafusion_functions::core::file_row_index::FileRowIndexFunc;
use datafusion_functions::core::input_file_name::InputFileNameFunc;
use datafusion_physical_expr::ScalarFunctionExpr;
use datafusion_physical_expr::expressions::{CastExpr, Column, Literal};
use datafusion_physical_expr::projection::{ProjectionExpr, ProjectionExprs};
use datafusion_physical_expr_common::physical_expr::PhysicalExpr;
pub fn expr_references_scalar_udf<T: ScalarUDFImpl>(
expr: &Arc<dyn PhysicalExpr>,
) -> bool {
let mut found = false;
expr.apply(|node| {
if ScalarFunctionExpr::try_downcast_func::<T>(node.as_ref()).is_some() {
found = true;
return Ok(TreeNodeRecursion::Stop);
}
Ok(TreeNodeRecursion::Continue)
})
.expect("Infallible traversal of PhysicalExpr tree failed");
found
}
fn rewrite_scalar_udf<T, F>(
expr: Arc<dyn PhysicalExpr>,
mut replacement: F,
) -> Result<Arc<dyn PhysicalExpr>>
where
T: ScalarUDFImpl,
F: FnMut(&ScalarFunctionExpr) -> Result<Arc<dyn PhysicalExpr>>,
{
expr.transform_up(|node| {
if let Some(scalar_fn) = ScalarFunctionExpr::try_downcast_func::<T>(node.as_ref())
{
Ok(Transformed::yes(replacement(scalar_fn)?))
} else {
Ok(Transformed::no(node))
}
})
.map(|transformed| transformed.data)
}
pub fn rewrite_file_row_index_expr(
expr: Arc<dyn PhysicalExpr>,
row_index_name: &str,
row_index_idx: usize,
) -> Result<Arc<dyn PhysicalExpr>> {
rewrite_scalar_udf::<FileRowIndexFunc, _>(expr, |_| {
let source = Arc::new(Column::new(row_index_name, row_index_idx));
let target_field = Arc::new(Field::new("file_row_index", DataType::Int64, true));
Ok(Arc::new(CastExpr::new_with_target_field(
source,
target_field,
None,
)))
})
}
pub fn rewrite_file_row_index_projection(
base_projection: &ProjectionExprs,
projection: &ProjectionExprs,
row_index_col: &Column,
) -> Result<ProjectionExprs> {
let mut base_exprs = base_projection.as_ref().to_vec();
let row_index_projection_idx =
base_projection.projected_column_position(row_index_col);
if row_index_projection_idx.is_none() {
base_exprs.push(ProjectionExpr {
expr: Arc::new(row_index_col.clone()),
alias: row_index_col.name().to_owned(),
});
}
let rewritten_projection = projection.clone().try_map_exprs(|expr| {
rewrite_file_row_index_expr(
expr,
row_index_col.name(),
row_index_projection_idx.unwrap_or(base_exprs.len() - 1),
)
})?;
ProjectionExprs::new(base_exprs).try_merge(&rewritten_projection)
}
pub fn rewrite_input_file_name_in_projection(
projection: ProjectionExprs,
file_name: &str,
) -> Result<ProjectionExprs> {
if !projection
.iter()
.any(|p| expr_references_scalar_udf::<InputFileNameFunc>(&p.expr))
{
return Ok(projection);
}
let file_name_lit =
Arc::new(Literal::new(ScalarValue::Utf8(Some(file_name.to_string()))))
as Arc<dyn PhysicalExpr>;
projection.try_map_exprs(|expr| {
rewrite_scalar_udf::<InputFileNameFunc, _>(expr, |_| {
Ok(Arc::clone(&file_name_lit))
})
})
}
#[cfg(test)]
mod tests {
use super::*;
use arrow::datatypes::Schema;
use datafusion_common::config::ConfigOptions;
use datafusion_expr::{Operator, ScalarUDF};
use datafusion_physical_expr::expressions;
use std::collections::HashMap;
fn file_row_index_expr() -> Arc<dyn PhysicalExpr> {
Arc::new(ScalarFunctionExpr::new(
"file_row_index",
Arc::new(ScalarUDF::from(FileRowIndexFunc::new())),
vec![],
Arc::new(Field::new("file_row_index", DataType::Int64, true)),
Arc::new(ConfigOptions::default()),
))
}
fn input_file_name_expr() -> Arc<dyn PhysicalExpr> {
Arc::new(ScalarFunctionExpr::new(
"input_file_name",
Arc::new(ScalarUDF::from(InputFileNameFunc::new())),
vec![],
Arc::new(Field::new("input_file_name", DataType::Utf8, true)),
Arc::new(ConfigOptions::default()),
))
}
#[test]
fn test_rewrite_scalar_udf_replaces_nested_typed_udf() -> Result<()> {
let expr = Arc::new(expressions::BinaryExpr::new(
file_row_index_expr(),
Operator::Plus,
expressions::lit(ScalarValue::Int64(Some(1))),
)) as Arc<dyn PhysicalExpr>;
let rewritten = rewrite_scalar_udf::<FileRowIndexFunc, _>(expr, |_| {
Ok(expressions::lit(ScalarValue::Int64(Some(7))))
})?;
let binary = rewritten
.downcast_ref::<expressions::BinaryExpr>()
.expect("rewritten expression should remain binary");
assert_eq!(binary.op(), &Operator::Plus);
let left = binary
.left()
.downcast_ref::<Literal>()
.expect("left side should be rewritten to a literal");
assert_eq!(left.value(), &ScalarValue::Int64(Some(7)));
let right = binary
.right()
.downcast_ref::<Literal>()
.expect("right side should remain the original literal");
assert_eq!(right.value(), &ScalarValue::Int64(Some(1)));
Ok(())
}
#[test]
fn test_rewrite_input_file_name_in_projection() -> Result<()> {
let file_name = "part=west/data.parquet";
let projection = ProjectionExprs::new([
ProjectionExpr::new(input_file_name_expr(), "file_name"),
ProjectionExpr::new(
Arc::new(expressions::BinaryExpr::new(
input_file_name_expr(),
Operator::Eq,
expressions::lit(ScalarValue::Utf8(Some(file_name.to_string()))),
)),
"matches_file",
),
]);
let rewritten = rewrite_input_file_name_in_projection(projection, file_name)?;
let rewritten = rewritten.as_ref();
assert_eq!(rewritten[0].alias, "file_name");
assert_eq!(rewritten[1].alias, "matches_file");
let file_name_lit = rewritten[0]
.expr
.downcast_ref::<Literal>()
.expect("input_file_name should rewrite to a literal");
assert_eq!(
file_name_lit.value(),
&ScalarValue::Utf8(Some(file_name.to_string()))
);
let binary = rewritten[1]
.expr
.downcast_ref::<expressions::BinaryExpr>()
.expect("nested expression should remain binary");
assert_eq!(binary.op(), &Operator::Eq);
let left = binary
.left()
.downcast_ref::<Literal>()
.expect("nested input_file_name should rewrite to a literal");
assert_eq!(
left.value(),
&ScalarValue::Utf8(Some(file_name.to_string()))
);
let right = binary
.right()
.downcast_ref::<Literal>()
.expect("comparison literal should remain unchanged");
assert_eq!(
right.value(),
&ScalarValue::Utf8(Some(file_name.to_string()))
);
Ok(())
}
#[test]
fn test_rewrite_file_row_index_expr_to_source_column() -> Result<()> {
let expr = rewrite_file_row_index_expr(
file_row_index_expr(),
"__datafusion_file_row_index",
2,
)?;
let cast_expr = expr
.downcast_ref::<CastExpr>()
.expect("file row index expression should be a cast");
assert_eq!(cast_expr.cast_type(), &DataType::Int64);
let target_field = cast_expr.target_field();
assert_eq!(target_field.name(), "file_row_index");
assert_eq!(target_field.data_type(), &DataType::Int64);
assert!(target_field.is_nullable());
assert!(target_field.metadata().is_empty());
let source = cast_expr
.expr()
.downcast_ref::<Column>()
.expect("source column");
assert_eq!(source.name(), "__datafusion_file_row_index");
assert_eq!(source.index(), 2);
let input_schema = Schema::new(vec![
Field::new("value", DataType::Int64, true),
Field::new("__datafusion_file_row_index", DataType::Int64, false)
.with_metadata(HashMap::from([(
"source".to_string(),
"virtual".to_string(),
)])),
]);
let return_field = expr.return_field(&input_schema)?;
assert_eq!(return_field.name(), "file_row_index");
assert_eq!(return_field.data_type(), &DataType::Int64);
assert!(return_field.is_nullable());
assert!(return_field.metadata().is_empty());
Ok(())
}
}