Skip to content
Merged
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
1 change: 1 addition & 0 deletions src/lib.rs
Original file line number Diff line number Diff line change
Expand Up @@ -6,6 +6,7 @@
//! 3. Run queries with [`sql()`] instead of `SessionContext::sql`.
//!
//! See the README for the SQL syntax and result types.
mod names;
mod optimizer;
mod planner;
pub mod sql;
Expand Down
31 changes: 31 additions & 0 deletions src/names.rs
Original file line number Diff line number Diff line change
@@ -0,0 +1,31 @@
//! The internal function name and the predicates that recognise a call to it,
//! shared by the SQL rewrite, the UDF, and the plan rules.
use datafusion::{
logical_expr::Expr,
physical_expr::{PhysicalExpr, ScalarFunctionExpr, async_scalar_function::AsyncFuncExpr},
physical_plan::async_func::AsyncFuncExec,
};
use std::sync::Arc;

/// The name a `prompt_jev` call is rewritten to before planning. Private to the
/// crate: callers never write it, and it must match everywhere it appears.
pub(crate) const INTERNAL_FUNCTION: &str = "__datafusion_jev";

/// A logical call to the internal function.
pub(crate) fn is_jev(expr: &Expr) -> bool {
matches!(expr, Expr::ScalarFunction(f) if f.func.name() == INTERNAL_FUNCTION)
}

/// A physical call to the internal function.
pub(crate) fn is_jev_expr(expr: &Arc<dyn PhysicalExpr>) -> bool {
expr.downcast_ref::<ScalarFunctionExpr>()
.is_some_and(|f| f.fun().name() == INTERNAL_FUNCTION)
}

/// The async expressions of a node. DataFusion deprecated this accessor without
/// a replacement and both physical rules need it, so the exemption lives here
/// rather than at each call site.
pub(crate) fn async_exprs(node: &AsyncFuncExec) -> &[Arc<AsyncFuncExpr>] {
#[allow(deprecated)]
node.async_exprs()
}
53 changes: 23 additions & 30 deletions src/optimizer.rs
Original file line number Diff line number Diff line change
@@ -1,22 +1,25 @@
//! Physical-plan rules around DataFusion 55's `AsyncFuncExec`.
//! Physical-plan rules around DataFusion 55's `AsyncFuncExec`, both of which
//! remove inference the query does not need.
//!
//! * [`DeduplicateJev`]: DataFusion extracts every occurrence of an async
//! function, including repeated struct field access. Keep one Jev evaluation
//! and project its result into the original slots so downstream column
//! indices remain valid.
//! * [`FilterBeforeJev`]: a `WHERE cheap AND prompt_jev(...) > x` is planned as
//! one filter above the async node, so every row is sent for inference before
//! either conjunct runs. Split it so the conjuncts that need no inference run
//! first, beneath the async node.
//! * [`FilterBeforeJev`] splits a filter that sits above the async node into
//! the conjuncts that need no inference and the rest, and runs the cheap ones
//! beneath the node. `WHERE cheap AND prompt_jev(...) > x` is otherwise
//! planned as one filter above the node, so every row is sent for inference
//! before either conjunct runs.
//! * [`DeduplicateJev`] keeps one evaluation of each distinct call and projects
//! its result into the original slots, so downstream column indices stay
//! valid. DataFusion extracts every occurrence of an async function,
//! including repeated struct field access.
use crate::names::{async_exprs, is_jev_expr};
use datafusion::{
common::{
Result,
tree_node::{Transformed, TransformedResult, TreeNode},
},
config::ConfigOptions,
physical_expr::{
PhysicalExpr, ScalarFunctionExpr, conjunction, expressions::Column, split_conjunction,
utils::collect_columns,
PhysicalExpr, async_scalar_function::AsyncFuncExpr, conjunction, expressions::Column,
split_conjunction, utils::collect_columns,
},
physical_optimizer::PhysicalOptimizerRule,
physical_plan::{
Expand All @@ -29,9 +32,11 @@ use datafusion::{
};
use std::sync::Arc;

fn is_jev_expr(expr: &Arc<dyn PhysicalExpr>) -> bool {
expr.downcast_ref::<ScalarFunctionExpr>()
.is_some_and(|f| f.fun().name() == "__datafusion_jev")
/// The async node and the expressions it evaluates, or `None` when `plan` is
/// some other node. Both rules start here.
fn async_node(plan: &Arc<dyn ExecutionPlan>) -> Option<(&AsyncFuncExec, &[Arc<AsyncFuncExpr>])> {
let node = plan.downcast_ref::<AsyncFuncExec>()?;
Some((node, async_exprs(node)))
}

#[derive(Debug)]
Expand Down Expand Up @@ -62,11 +67,9 @@ impl PhysicalOptimizerRule for FilterBeforeJev {
passthrough.push(cursor);
cursor = child;
}
let Some(node) = cursor.downcast_ref::<AsyncFuncExec>() else {
let Some((node, async_exprs)) = async_node(&cursor) else {
return Ok(Transformed::no(plan));
};
#[allow(deprecated)]
let async_exprs = node.async_exprs();
if !async_exprs.iter().any(|e| is_jev_expr(&e.func)) {
return Ok(Transformed::no(plan));
}
Expand Down Expand Up @@ -126,23 +129,13 @@ impl PhysicalOptimizerRule for DeduplicateJev {
_: &ConfigOptions,
) -> Result<Arc<dyn ExecutionPlan>> {
plan.transform_up(|plan| {
let Some(node) = plan.downcast_ref::<AsyncFuncExec>() else {
let Some((node, expressions)) = async_node(&plan) else {
return Ok(Transformed::no(plan));
};
// DF deprecated this accessor without a replacement. We need it to
// avoid duplicate paid inference while preserving the output schema.
#[allow(deprecated)]
let expressions = node.async_exprs();
let mut unique: Vec<
Arc<datafusion::physical_expr::async_scalar_function::AsyncFuncExpr>,
> = vec![];
let mut unique: Vec<Arc<AsyncFuncExpr>> = vec![];
let mut indices = vec![];
for expr in expressions {
let is_jev = expr
.func
.downcast_ref::<ScalarFunctionExpr>()
.is_some_and(|f| f.fun().name() == "__datafusion_jev");
let existing = if is_jev {
let existing = if is_jev_expr(&expr.func) {
unique.iter().position(|x| x.func.eq(&expr.func))
} else {
None
Expand Down
57 changes: 22 additions & 35 deletions src/planner.rs
Original file line number Diff line number Diff line change
@@ -1,19 +1,17 @@
//! Logical-plan adjustments so `prompt_jev` works wherever DataFusion 55 does
//! not natively plan async functions, and is never evaluated on rows a filter
//! discards.
//! Logical-plan rules that keep `prompt_jev` plannable and cheap.
//!
//! Two problems, two rules:
//!
//! * DataFusion only lifts async calls out of `Projection` and `Filter`. A call
//! in an `ORDER BY`, a window `OVER (...)`, a `GROUP BY`, or a join condition
//! reaches the synchronous invoke path and fails with "async functions should
//! not be called directly". [`HoistJev`] moves such calls into a projection
//! beneath the node and refers to them by column.
//! * DataFusion's leaf-expression pushdown moves `get_field(k, 'x')` through a
//! filter and substitutes `k`'s definition without rechecking placement, so a
//! `prompt_jev` call ends up below the filter and runs on every row.
//! [`LeafPushdownGuard`] skips those two rules for any plan that contains the
//! call, and leaves every other plan alone.
//! * [`HoistJev`] moves a call out of an `ORDER BY`, a window `OVER (...)`, a
//! `GROUP BY`, or a join condition into a projection beneath the node and
//! refers to it by column. DataFusion 55 only lifts async calls out of
//! `Projection` and `Filter`; anywhere else the call reaches the synchronous
//! invoke path and fails with "async functions should not be called
//! directly".
//! * [`LeafPushdownGuard`] skips DataFusion's two leaf-expression pushdown
//! rules for any plan that contains a call, and leaves every other plan
//! alone. Those rules move `get_field(k, 'x')` through a filter and
//! substitute `k`'s definition without rechecking placement, which would put
//! the call below the filter and run it on every row.
use crate::names::is_jev;
use datafusion::{
common::{
DFSchema, Result, plan_datafusion_err,
Expand All @@ -25,13 +23,8 @@ use datafusion::{
};
use std::sync::Arc;

pub const FUNCTION_NAME: &str = "__datafusion_jev";
const HOIST_PREFIX: &str = "__jev_hoist_";

fn is_jev(expr: &Expr) -> bool {
matches!(expr, Expr::ScalarFunction(f) if f.func.name() == FUNCTION_NAME)
}

fn contains_jev(expr: &Expr) -> bool {
expr.exists(|e| Ok(is_jev(e))).unwrap_or(false)
}
Expand Down Expand Up @@ -140,15 +133,6 @@ impl AnalyzerRule for HoistJev {
}
}

/// Collect each distinct call, project it beneath `input` under a generated
/// name, and return the new input with the call-to-column substitutions.
fn hoist_calls(
input: Arc<LogicalPlan>,
exprs: &[&Expr],
) -> Result<(LogicalPlan, Vec<(Expr, Expr)>)> {
hoist_calls_named(input, exprs, HOIST_PREFIX)
}

fn collect_calls(exprs: &[&Expr]) -> Result<Vec<Expr>> {
let mut calls: Vec<Expr> = vec![];
for expr in exprs {
Expand All @@ -163,7 +147,10 @@ fn collect_calls(exprs: &[&Expr]) -> Result<Vec<Expr>> {
Ok(calls)
}

fn hoist_calls_named(
/// Collect each distinct call, project it beneath `input` under a name made
/// from `prefix`, and return the new input with the call-to-column
/// substitutions.
fn hoist_calls(
input: Arc<LogicalPlan>,
exprs: &[&Expr],
prefix: &str,
Expand Down Expand Up @@ -213,7 +200,7 @@ fn hoist_sort(sort: Sort) -> Result<LogicalPlan> {
.map(Expr::Column)
.collect();
let keys: Vec<&Expr> = expr.iter().map(|s| &s.expr).collect();
let (projected, substitutions) = hoist_calls(input, &keys)?;
let (projected, substitutions) = hoist_calls(input, &keys, HOIST_PREFIX)?;
let expr = expr
.into_iter()
.map(|s| Ok(s.with_expr(substitute(s.expr.clone(), &substitutions)?)))
Expand All @@ -235,7 +222,7 @@ fn hoist_window(window: Window) -> Result<LogicalPlan> {
} = window;
let original: Vec<Expr> = schema.columns().into_iter().map(Expr::Column).collect();
let refs: Vec<&Expr> = window_expr.iter().collect();
let (projected, substitutions) = hoist_calls(input, &refs)?;
let (projected, substitutions) = hoist_calls(input, &refs, HOIST_PREFIX)?;
let window_expr = window_expr
.into_iter()
.map(|e| {
Expand Down Expand Up @@ -271,7 +258,7 @@ fn hoist_aggregate(agg: Aggregate) -> Result<LogicalPlan> {
..
} = agg;
let refs: Vec<&Expr> = group_expr.iter().chain(&aggr_expr).collect();
let (projected, substitutions) = hoist_calls(input, &refs)?;
let (projected, substitutions) = hoist_calls(input, &refs, HOIST_PREFIX)?;
let group_expr = group_expr
.into_iter()
.map(|e| substitute_keeping_name(e, &substitutions))
Expand Down Expand Up @@ -323,8 +310,8 @@ fn hoist_join(join: Join) -> Result<LogicalPlan> {
}
let left_refs: Vec<&Expr> = left_calls.iter().collect();
let right_refs: Vec<&Expr> = right_calls.iter().collect();
let (left, mut substitutions) = hoist_calls_named(left, &left_refs, "__jev_hoist_l")?;
let (right, right_subs) = hoist_calls_named(right, &right_refs, "__jev_hoist_r")?;
let (left, mut substitutions) = hoist_calls(left, &left_refs, "__jev_hoist_l")?;
let (right, right_subs) = hoist_calls(right, &right_refs, "__jev_hoist_r")?;
substitutions.extend(right_subs);
let on = on
.into_iter()
Expand Down
Loading
Loading