use std::hash::Hash;
use std::sync::Arc;
use crate::{
ScalarFunctionExpr,
expressions::{Column, LambdaVariable},
physical_expr::PhysicalExpr,
};
use arrow::{
datatypes::{DataType, Schema},
record_batch::RecordBatch,
};
use datafusion_common::{
HashMap, plan_err,
tree_node::{Transformed, TreeNode, TreeNodeRecursion, TreeNodeVisitor},
};
use datafusion_common::{HashSet, Result, internal_err};
use datafusion_expr::ColumnarValue;
#[derive(Debug, Eq, Clone)]
pub struct LambdaExpr {
params: Vec<String>,
body: Arc<dyn PhysicalExpr>,
projected_body: Arc<dyn PhysicalExpr>,
projection: Vec<usize>,
used_param_indices: Vec<usize>,
}
impl PartialEq for LambdaExpr {
fn eq(&self, other: &Self) -> bool {
self.params.eq(&other.params) && self.body.eq(&other.body)
}
}
impl Hash for LambdaExpr {
fn hash<H: std::hash::Hasher>(&self, state: &mut H) {
self.params.hash(state);
self.body.hash(state);
}
}
impl LambdaExpr {
pub fn try_new(params: Vec<String>, body: Arc<dyn PhysicalExpr>) -> Result<Self> {
if !all_unique(¶ms) {
return plan_err!(
"lambda params must be unique, got ({})",
params.join(", ")
);
}
check_async_udf(&body)?;
Ok(Self::new(params, body))
}
fn new(params: Vec<String>, body: Arc<dyn PhysicalExpr>) -> Self {
let own_params: HashSet<String> = params.iter().cloned().collect();
let mut visitor = CollectUsedVisitor {
own_params: &own_params,
used_indices: HashSet::new(),
used_param_names: HashSet::new(),
shadow_stack: Vec::new(),
};
body.visit(&mut visitor).expect("visitor is infallible");
let CollectUsedVisitor {
used_indices,
used_param_names,
..
} = visitor;
let mut projection = used_indices.into_iter().collect::<Vec<_>>();
projection.sort();
let column_index_map = projection
.iter()
.copied()
.enumerate()
.map(|(new_idx, original)| (original, new_idx))
.collect::<HashMap<_, _>>();
let projected_body = Arc::clone(&body)
.transform_down(|e| {
if let Some(column) = e.downcast_ref::<Column>() {
let original = column.index();
let projected = *column_index_map.get(&original).unwrap();
if projected != original {
return Ok(Transformed::yes(Arc::new(Column::new(
column.name(),
projected,
))));
}
} else if let Some(lambda_variable) = e.downcast_ref::<LambdaVariable>() {
let original = lambda_variable.index();
let projected = *column_index_map.get(&original).unwrap();
if projected != original {
return Ok(Transformed::yes(Arc::new(LambdaVariable::new(
projected,
Arc::clone(lambda_variable.field()),
))));
}
}
Ok(Transformed::no(e))
})
.expect("closure should be infallible")
.data;
let used_param_indices = params
.iter()
.enumerate()
.filter(|(_, name)| used_param_names.contains(*name))
.map(|(i, _)| i)
.collect();
Self {
params,
body,
projected_body,
projection,
used_param_indices,
}
}
pub fn params(&self) -> &[String] {
&self.params
}
pub fn body(&self) -> &Arc<dyn PhysicalExpr> {
&self.body
}
#[cfg(feature = "proto")]
pub fn try_from_proto(
node: &datafusion_proto_models::protobuf::PhysicalExprNode,
ctx: &datafusion_physical_expr_common::physical_expr::proto_decode::PhysicalExprDecodeCtx<'_>,
) -> Result<Arc<dyn PhysicalExpr>> {
use datafusion_physical_expr_common::expect_expr_variant;
use datafusion_proto_models::protobuf;
let lambda = expect_expr_variant!(
node,
protobuf::physical_expr_node::ExprType::Lambda,
"LambdaExpr",
);
Ok(Arc::new(LambdaExpr::try_new(
lambda.params.clone(),
ctx.decode_required_expression(lambda.body.as_deref(), "LambdaExpr", "body")?,
)?))
}
pub(crate) fn projection(&self) -> &[usize] {
&self.projection
}
pub(crate) fn projected_body(&self) -> &Arc<dyn PhysicalExpr> {
&self.projected_body
}
pub fn used_param_indices(&self) -> &[usize] {
&self.used_param_indices
}
}
struct CollectUsedVisitor<'a> {
own_params: &'a HashSet<String>,
used_indices: HashSet<usize>,
used_param_names: HashSet<String>,
shadow_stack: Vec<HashSet<String>>,
}
impl TreeNodeVisitor<'_> for CollectUsedVisitor<'_> {
type Node = Arc<dyn PhysicalExpr>;
fn f_down(&mut self, node: &Self::Node) -> Result<TreeNodeRecursion> {
if let Some(col) = node.downcast_ref::<Column>() {
self.used_indices.insert(col.index());
} else if let Some(var) = node.downcast_ref::<LambdaVariable>() {
self.used_indices.insert(var.index());
let name = var.name();
let shadowed = self.shadow_stack.iter().any(|frame| frame.contains(name));
if !shadowed && self.own_params.contains(name) {
self.used_param_names.insert(name.to_string());
}
} else if let Some(nested) = node.downcast_ref::<LambdaExpr>() {
self.shadow_stack
.push(nested.params.iter().cloned().collect());
}
Ok(TreeNodeRecursion::Continue)
}
fn f_up(&mut self, node: &Self::Node) -> Result<TreeNodeRecursion> {
if node.downcast_ref::<LambdaExpr>().is_some() {
self.shadow_stack.pop();
}
Ok(TreeNodeRecursion::Continue)
}
}
impl std::fmt::Display for LambdaExpr {
fn fmt(&self, f: &mut std::fmt::Formatter) -> std::fmt::Result {
write!(f, "({}) -> {}", self.params.join(", "), self.body)
}
}
impl PhysicalExpr for LambdaExpr {
fn data_type(&self, _input_schema: &Schema) -> Result<DataType> {
Ok(DataType::Null)
}
fn nullable(&self, _input_schema: &Schema) -> Result<bool> {
Ok(true)
}
fn evaluate(&self, _batch: &RecordBatch) -> Result<ColumnarValue> {
internal_err!("LambdaExpr::evaluate() should not be called")
}
fn children(&self) -> Vec<&Arc<dyn PhysicalExpr>> {
vec![&self.body]
}
fn with_new_children(
self: Arc<Self>,
children: Vec<Arc<dyn PhysicalExpr>>,
) -> Result<Arc<dyn PhysicalExpr>> {
let [body] = children.as_slice() else {
return internal_err!(
"LambdaExpr expects exactly 1 child, got {}",
children.len()
);
};
check_async_udf(body)?;
Ok(Arc::new(Self::new(self.params.clone(), Arc::clone(body))))
}
fn fmt_sql(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
write!(f, "({}) -> {}", self.params.join(", "), self.body)
}
#[cfg(feature = "proto")]
fn try_to_proto(
&self,
ctx: &datafusion_physical_expr_common::physical_expr::proto_encode::PhysicalExprEncodeCtx<'_>,
) -> Result<Option<datafusion_proto_models::protobuf::PhysicalExprNode>> {
use datafusion_proto_models::protobuf;
Ok(Some(protobuf::PhysicalExprNode {
expr_id: None,
expr_type: Some(protobuf::physical_expr_node::ExprType::Lambda(Box::new(
protobuf::PhysicalLambdaExprNode {
params: self.params().to_vec(),
body: Some(Box::new(ctx.encode_child(self.body())?)),
},
))),
}))
}
}
pub fn lambda(
params: impl IntoIterator<Item = impl Into<String>>,
body: Arc<dyn PhysicalExpr>,
) -> Result<Arc<dyn PhysicalExpr>> {
Ok(Arc::new(LambdaExpr::try_new(
params.into_iter().map(Into::into).collect(),
body,
)?))
}
fn all_unique(params: &[String]) -> bool {
match params.len() {
0 | 1 => true,
2 => params[0] != params[1],
_ => {
let mut set = HashSet::with_capacity(params.len());
params.iter().all(|p| set.insert(p.as_str()))
}
}
}
fn check_async_udf(body: &Arc<dyn PhysicalExpr>) -> Result<()> {
if body.exists(|expr| {
Ok(expr
.downcast_ref::<ScalarFunctionExpr>()
.is_some_and(|udf| udf.fun().as_async().is_some()))
})? {
return plan_err!(
"Async functions in lambdas aren't supported, see https://github.com/apache/datafusion/issues/22091"
);
}
Ok(())
}
#[cfg(test)]
mod tests {
use crate::expressions::{Column, LambdaVariable, NoOp, lambda::lambda};
use arrow::{
array::RecordBatch,
datatypes::{DataType, Field, Schema},
};
use std::sync::Arc;
use super::LambdaExpr;
#[test]
fn test_lambda_evaluate() {
let lambda = lambda(["a"], Arc::new(NoOp::new())).unwrap();
let batch = RecordBatch::new_empty(Arc::new(Schema::empty()));
assert!(lambda.evaluate(&batch).is_err());
}
#[test]
fn test_lambda_duplicate_name() {
assert!(lambda(["a", "a"], Arc::new(NoOp::new())).is_err());
}
#[test]
fn test_used_params_collects_only_referenced_param() {
let v_field = Arc::new(Field::new("v", DataType::Int32, true));
let body = Arc::new(LambdaVariable::new(1, Arc::clone(&v_field)));
let lambda =
LambdaExpr::try_new(vec!["k".to_string(), "v".to_string()], body).unwrap();
assert_eq!(lambda.projection(), &[1]);
assert_eq!(lambda.used_param_indices(), &[1]);
}
#[test]
fn test_used_params_all_unused() {
let body = Arc::new(NoOp::new());
let lambda =
LambdaExpr::try_new(vec!["k".to_string(), "v".to_string()], body).unwrap();
assert!(lambda.projection().is_empty());
assert!(lambda.used_param_indices().is_empty());
}
#[test]
fn test_used_params_three_params_middle_unused() {
let a_field = Arc::new(Field::new("a", DataType::Int32, true));
let c_field = Arc::new(Field::new("c", DataType::Int32, true));
let body = Arc::new(crate::expressions::BinaryExpr::new(
Arc::new(LambdaVariable::new(0, Arc::clone(&a_field))),
datafusion_expr::Operator::Plus,
Arc::new(LambdaVariable::new(2, Arc::clone(&c_field))),
));
let lambda = LambdaExpr::try_new(
vec!["a".to_string(), "b".to_string(), "c".to_string()],
body,
)
.unwrap();
assert_eq!(lambda.used_param_indices(), &[0, 2]);
}
#[test]
fn test_used_params_both_used_in_reverse_reference_order() {
let k_field = Arc::new(Field::new("k", DataType::Int32, true));
let v_field = Arc::new(Field::new("v", DataType::Int32, true));
let body = Arc::new(crate::expressions::BinaryExpr::new(
Arc::new(LambdaVariable::new(1, Arc::clone(&v_field))),
datafusion_expr::Operator::Plus,
Arc::new(LambdaVariable::new(0, Arc::clone(&k_field))),
));
let lambda =
LambdaExpr::try_new(vec!["k".to_string(), "v".to_string()], body).unwrap();
assert_eq!(lambda.projection(), &[0, 1]);
assert_eq!(lambda.used_param_indices(), &[0, 1]);
}
#[test]
fn test_used_params_handles_shadowing_inside_nested_lambda() {
let outer_k_field = Arc::new(Field::new("k", DataType::Int32, true));
let outer_v_field = Arc::new(Field::new("v", DataType::Int32, true));
let inner_v2_field = Arc::new(Field::new("v2", DataType::Int32, true));
let inner_body: Arc<dyn crate::PhysicalExpr> =
Arc::new(crate::expressions::BinaryExpr::new(
Arc::new(crate::expressions::BinaryExpr::new(
Arc::new(LambdaVariable::new(1, Arc::clone(&outer_k_field))),
datafusion_expr::Operator::Plus,
Arc::new(LambdaVariable::new(2, Arc::clone(&inner_v2_field))),
)),
datafusion_expr::Operator::Plus,
Arc::new(LambdaVariable::new(0, Arc::clone(&outer_v_field))),
));
let inner_lambda = Arc::new(
LambdaExpr::try_new(vec!["k".to_string(), "v2".to_string()], inner_body)
.unwrap(),
);
let outer_body: Arc<dyn crate::PhysicalExpr> =
Arc::new(crate::expressions::BinaryExpr::new(
Arc::new(Column::new("col", 0)),
datafusion_expr::Operator::Plus,
inner_lambda,
));
let outer_lambda =
LambdaExpr::try_new(vec!["k".to_string(), "v".to_string()], outer_body)
.unwrap();
assert_eq!(
outer_lambda.used_param_indices(),
&[1],
"only outer's `v` (index 1) should be reported as used; `k` (index 0) is \
shadowed inside the nested lambda"
);
}
}