From f7c27ffdf3528f2fdfd0eb704cdf8faa225f9b12 Mon Sep 17 00:00:00 2001 From: Eddie A Tejeda <669988+eddietejeda@users.noreply.github.com> Date: Wed, 23 Sep 2026 17:14:20 -0700 Subject: [PATCH] fix: plan prompt_jev in sort, window, group by, and joins; keep filters first DataFusion 55 lifts async calls only out of Projection and Filter, so a prompt_jev in ORDER BY, OVER (...), GROUP BY, or a join condition hit the synchronous path and failed with "async functions should not be called directly". A new analyzer rule hoists such calls into a projection beneath the node and refers to them by column, aliasing rewritten expressions so the plan above keeps its column names. A call in a join condition that spans both sides is refused with a clear planning error. DataFusion's leaf-expression pushdown also moved get_field(k, 'x') through a filter and substituted k's definition without rechecking placement, so a prompt_jev call ran on every row before the filter. The two pushdown rules are now wrapped to skip plans that contain the call; other plans keep the optimization. --- Cargo.lock | 24 ++-- README.md | 5 + src/lib.rs | 17 +++ src/planner.rs | 353 +++++++++++++++++++++++++++++++++++++++++++++++++ tests/sql.rs | 114 ++++++++++++++++ 5 files changed, 501 insertions(+), 12 deletions(-) create mode 100644 src/planner.rs diff --git a/Cargo.lock b/Cargo.lock index 5681f50..6b665d4 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -876,6 +876,18 @@ dependencies = [ "datafusion-physical-expr-common", ] +[[package]] +name = "datafusion-jev" +version = "0.1.0" +dependencies = [ + "async-trait", + "datafusion", + "futures", + "serde", + "serde_json", + "tokio", +] + [[package]] name = "datafusion-macros" version = "55.1.0" @@ -1317,18 +1329,6 @@ version = "0.4.3" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "7f24254aa9a54b5c858eaee2f5bccdb46aaf0e486a595ed5fd8f86ba55232a70" -[[package]] -name = "datafusion-jev" -version = "0.1.0" -dependencies = [ - "async-trait", - "datafusion", - "futures", - "serde", - "serde_json", - "tokio", -] - [[package]] name = "http" version = "1.5.0" diff --git a/README.md b/README.md index da42b36..d90fac4 100644 --- a/README.md +++ b/README.md @@ -145,6 +145,11 @@ each row must be judged in complete isolation. `prompt_jev('hello', ...)` costs about one request per 256 rows scanned. - **Repeated calls** with identical arguments in one query are evaluated once. Reading several fields from one result does not repeat the request. +- **Anywhere in a query.** `prompt_jev` works in `SELECT`, `WHERE`, `ORDER BY`, + `GROUP BY`, `HAVING`, window `OVER (...)` clauses, and join conditions. A call + in a join condition must use columns from one side of the join only. +- **Filters run first.** Rows removed by `WHERE` are never sent to the service, + including when the call sits inside a CTE or subquery. - **Limits.** Each row's text may be up to 64 KiB after JSON encoding. Longer text fails the query. Requests are capped at 256 KiB. - **Outages.** Each request is tried three times with a 30-second timeout. If all diff --git a/src/lib.rs b/src/lib.rs index e1bb452..e068aa0 100644 --- a/src/lib.rs +++ b/src/lib.rs @@ -7,6 +7,7 @@ //! //! See the README for the SQL syntax and result types. mod optimizer; +mod planner; pub mod sql; mod udf; use async_trait::async_trait; @@ -43,7 +44,23 @@ pub fn register(ctx: &SessionContext, client: Arc) { ctx.register_udf(udf::function(client)); let state_ref = ctx.state_ref(); let mut state = state_ref.write(); + // Keep DataFusion's optimizer list as is, except that the two leaf-pushdown + // rules are skipped for plans calling prompt_jev (see `planner`). + let optimizer_rules = state + .optimizer() + .rules + .iter() + .map(|rule| { + if planner::LeafPushdownGuard::GUARDED.contains(&rule.name()) { + planner::LeafPushdownGuard::wrap(Arc::clone(rule)) + } else { + Arc::clone(rule) + } + }) + .collect(); *state = datafusion::execution::SessionStateBuilder::new_from_existing(state.clone()) + .with_analyzer_rule(Arc::new(planner::HoistJev)) + .with_optimizer_rules(optimizer_rules) .with_physical_optimizer_rule(Arc::new(optimizer::DeduplicateJev)) .build(); } diff --git a/src/planner.rs b/src/planner.rs new file mode 100644 index 0000000..04b8cde --- /dev/null +++ b/src/planner.rs @@ -0,0 +1,353 @@ +//! 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. +//! +//! 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. +use datafusion::{ + common::{ + DFSchema, Result, plan_datafusion_err, + tree_node::{Transformed, TreeNode, TreeNodeRecursion}, + }, + config::ConfigOptions, + logical_expr::{Aggregate, Expr, Join, LogicalPlan, LogicalPlanBuilder, Sort, Window, col}, + optimizer::{AnalyzerRule, ApplyOrder, OptimizerConfig, OptimizerRule}, +}; +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) +} + +fn plan_contains_jev(plan: &LogicalPlan) -> bool { + let mut found = false; + let _ = plan.apply_with_subqueries(|p| { + if p.expressions().iter().any(contains_jev) { + found = true; + return Ok(TreeNodeRecursion::Stop); + } + Ok(TreeNodeRecursion::Continue) + }); + found +} + +/// Runs a wrapped optimizer rule only for plans that do not call `prompt_jev`. +#[derive(Debug)] +pub struct LeafPushdownGuard { + inner: Arc, +} + +impl LeafPushdownGuard { + /// Rules that move expressions toward leaves and, in DataFusion 55, carry + /// an async call along with them. + pub const GUARDED: &[&str] = &["extract_leaf_expressions", "push_down_leaf_projections"]; + + pub fn wrap( + inner: Arc, + ) -> Arc { + Arc::new(Self { inner }) + } +} + +impl OptimizerRule for LeafPushdownGuard { + fn name(&self) -> &str { + self.inner.name() + } + fn apply_order(&self) -> Option { + // The guard needs the whole plan to decide, so it drives the traversal + // itself and replays the inner rule's own order below. + None + } + fn supports_rewrite(&self) -> bool { + true + } + fn rewrite( + &self, + plan: LogicalPlan, + config: &dyn OptimizerConfig, + ) -> Result> { + if plan_contains_jev(&plan) { + return Ok(Transformed::no(plan)); + } + match self.inner.apply_order() { + Some(ApplyOrder::TopDown) => { + plan.transform_down_with_subqueries(|p| self.inner.rewrite(p, config)) + } + Some(ApplyOrder::BottomUp) => { + plan.transform_up_with_subqueries(|p| self.inner.rewrite(p, config)) + } + None => self.inner.rewrite(plan, config), + } + } +} + +/// Moves `prompt_jev` calls out of sort keys and window expressions into a +/// projection beneath, so DataFusion plans them as ordinary async projections. +#[derive(Debug)] +pub struct HoistJev; + +impl AnalyzerRule for HoistJev { + fn name(&self) -> &str { + "datafusion_jev_hoist" + } + fn analyze(&self, plan: LogicalPlan, _: &ConfigOptions) -> Result { + plan.transform_up_with_subqueries(|p| match p { + LogicalPlan::Sort(sort) if sort.expr.iter().any(|s| contains_jev(&s.expr)) => { + hoist_sort(sort).map(Transformed::yes) + } + LogicalPlan::Window(window) if window.window_expr.iter().any(contains_jev) => { + hoist_window(window).map(Transformed::yes) + } + LogicalPlan::Aggregate(agg) + if agg + .group_expr + .iter() + .chain(&agg.aggr_expr) + .any(contains_jev) => + { + hoist_aggregate(agg).map(Transformed::yes) + } + LogicalPlan::Join(join) + if join + .on + .iter() + .flat_map(|(l, r)| [l, r]) + .chain(&join.filter) + .any(contains_jev) => + { + hoist_join(join).map(Transformed::yes) + } + other => Ok(Transformed::no(other)), + }) + .map(|t| t.data) + } +} + +/// 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, + exprs: &[&Expr], +) -> Result<(LogicalPlan, Vec<(Expr, Expr)>)> { + hoist_calls_named(input, exprs, HOIST_PREFIX) +} + +fn collect_calls(exprs: &[&Expr]) -> Result> { + let mut calls: Vec = vec![]; + for expr in exprs { + expr.apply(|e| { + if is_jev(e) && !calls.contains(e) { + calls.push(e.clone()); + return Ok(TreeNodeRecursion::Jump); + } + Ok(TreeNodeRecursion::Continue) + })?; + } + Ok(calls) +} + +fn hoist_calls_named( + input: Arc, + exprs: &[&Expr], + prefix: &str, +) -> Result<(LogicalPlan, Vec<(Expr, Expr)>)> { + let calls = collect_calls(exprs)?; + if calls.is_empty() { + return Ok((Arc::unwrap_or_clone(input), vec![])); + } + let mut projection: Vec = input + .schema() + .columns() + .into_iter() + .map(Expr::Column) + .collect(); + let mut substitutions = vec![]; + for (i, call) in calls.into_iter().enumerate() { + let name = format!("{prefix}{i}"); + projection.push(call.clone().alias(&name)); + substitutions.push((call, col(&name))); + } + let projected = LogicalPlanBuilder::from(Arc::unwrap_or_clone(input)) + .project(projection)? + .build()?; + Ok((projected, substitutions)) +} + +fn substitute(expr: Expr, substitutions: &[(Expr, Expr)]) -> Result { + expr.transform_down(|e| { + if let Some((_, replacement)) = substitutions.iter().find(|(call, _)| *call == e) { + return Ok(Transformed::new( + replacement.clone(), + true, + TreeNodeRecursion::Jump, + )); + } + Ok(Transformed::no(e)) + }) + .map(|t| t.data) +} + +fn hoist_sort(sort: Sort) -> Result { + let Sort { expr, input, fetch } = sort; + let original: Vec = input + .schema() + .columns() + .into_iter() + .map(Expr::Column) + .collect(); + let keys: Vec<&Expr> = expr.iter().map(|s| &s.expr).collect(); + let (projected, substitutions) = hoist_calls(input, &keys)?; + let expr = expr + .into_iter() + .map(|s| Ok(s.with_expr(substitute(s.expr.clone(), &substitutions)?))) + .collect::>>()?; + let builder = LogicalPlanBuilder::from(projected); + let builder = match fetch { + Some(n) => builder.sort_with_limit(expr, Some(n))?, + None => builder.sort(expr)?, + }; + // Drop the hoisted columns again so the plan above sees the schema it planned for. + builder.project(original)?.build() +} + +fn hoist_window(window: Window) -> Result { + let Window { + input, + window_expr, + schema, + } = window; + let original: Vec = schema.columns().into_iter().map(Expr::Column).collect(); + let refs: Vec<&Expr> = window_expr.iter().collect(); + let (projected, substitutions) = hoist_calls(input, &refs)?; + let window_expr = window_expr + .into_iter() + .map(|e| { + // Rewriting the expression changes its generated name; keep the + // original so the projection above still finds its column. + let name = e.schema_name().to_string(); + Ok(substitute(e, &substitutions)?.alias(name)) + }) + .collect::>>()?; + LogicalPlanBuilder::from(projected) + .window(window_expr)? + .project(original)? + .build() +} + +/// Substitute and, when that changed the expression, keep its original output +/// name so the plan above still finds the column. +fn substitute_keeping_name(expr: Expr, substitutions: &[(Expr, Expr)]) -> Result { + let name = expr.schema_name().to_string(); + let rewritten = substitute(expr, substitutions)?; + Ok(if rewritten.schema_name().to_string() == name { + rewritten + } else { + rewritten.alias(name) + }) +} + +fn hoist_aggregate(agg: Aggregate) -> Result { + let Aggregate { + input, + group_expr, + aggr_expr, + .. + } = agg; + let refs: Vec<&Expr> = group_expr.iter().chain(&aggr_expr).collect(); + let (projected, substitutions) = hoist_calls(input, &refs)?; + let group_expr = group_expr + .into_iter() + .map(|e| substitute_keeping_name(e, &substitutions)) + .collect::>>()?; + let aggr_expr = aggr_expr + .into_iter() + .map(|e| substitute_keeping_name(e, &substitutions)) + .collect::>>()?; + LogicalPlanBuilder::from(projected) + .aggregate(group_expr, aggr_expr)? + .build() +} + +fn belongs_to(expr: &Expr, schema: &DFSchema) -> bool { + expr.column_refs() + .into_iter() + .all(|c| schema.index_of_column(c).is_ok()) +} + +fn hoist_join(join: Join) -> Result { + let Join { + left, + right, + on, + filter, + join_type, + join_constraint, + schema, + null_equality, + null_aware, + } = join; + let original: Vec = schema.columns().into_iter().map(Expr::Column).collect(); + // Every call must be computable on one side. A call over both sides has + // no input to be projected from; the caller can compute it in a CTE. + let mut all: Vec<&Expr> = on.iter().flat_map(|(l, r)| [l, r]).collect(); + all.extend(&filter); + let (mut left_calls, mut right_calls) = (vec![], vec![]); + for call in collect_calls(&all)? { + if belongs_to(&call, left.schema()) { + left_calls.push(call); + } else if belongs_to(&call, right.schema()) { + right_calls.push(call); + } else { + return Err(plan_datafusion_err!( + "prompt_jev in a join condition must reference columns from one side only; \ + compute it in a CTE or subquery first" + )); + } + } + 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")?; + substitutions.extend(right_subs); + let on = on + .into_iter() + .map(|(l, r)| { + Ok(( + substitute(l, &substitutions)?, + substitute(r, &substitutions)?, + )) + }) + .collect::>>()?; + let filter = filter.map(|f| substitute(f, &substitutions)).transpose()?; + let join = Join::try_new( + Arc::new(left), + Arc::new(right), + on, + filter, + join_type, + join_constraint, + null_equality, + null_aware, + )?; + // The join now carries the hoisted columns; project back to what was planned. + LogicalPlanBuilder::from(LogicalPlan::Join(join)) + .project(original)? + .build() +} diff --git a/tests/sql.rs b/tests/sql.rs index bcacfcc..313117c 100644 --- a/tests/sql.rs +++ b/tests/sql.rs @@ -664,3 +664,117 @@ async fn strict_dialect_gets_a_clear_error() { .to_string(); assert!(!plain.contains("prompt_jev"), "{plain}"); } + +#[tokio::test] +async fn order_by_prompt_jev_is_hoisted_and_planned() { + let mock = Arc::new(Mock::default()); + let result = run( + mock.clone(), + "SELECT id FROM (VALUES (1,'a'),(2,'b'),(3,'c')) t(id, body) \ + ORDER BY prompt_jev(body, 'Urgent?') DESC, id LIMIT 2", + ) + .await + .unwrap(); + assert_eq!(result.iter().map(|b| b.num_rows()).sum::(), 2); + assert_eq!(mock.calls.lock().unwrap().len(), 1, "one batched request"); +} + +#[tokio::test] +async fn window_over_prompt_jev_is_hoisted_and_planned() { + let mock = Arc::new(Mock::default()); + let result = run( + mock.clone(), + "SELECT id, rank() OVER (ORDER BY prompt_jev(body, 'Urgent?') DESC) AS rk \ + FROM (VALUES (1,'a'),(2,'b'),(3,'c')) t(id, body) ORDER BY rk, id", + ) + .await + .unwrap(); + let text = display(&result); + assert_eq!(result.iter().map(|b| b.num_rows()).sum::(), 3); + assert!(text.contains("rk"), "{text}"); + assert_eq!(mock.calls.lock().unwrap().len(), 1); +} + +#[tokio::test] +async fn filter_runs_before_inference_when_a_field_is_read_through_a_cte() { + // DataFusion's leaf pushdown would otherwise move get_field(k, 'choice') and + // the call it wraps beneath the filter, sending every row to the provider. + let mock = Arc::new(Mock::default()); + let result = run( + mock.clone(), + "WITH c AS (SELECT id, prompt_jev(body, 'Topic?', choice := ['a','b']) AS k \ + FROM (VALUES (1,'keep'),(2,'drop'),(3,'drop')) t(id, body) WHERE id = 1) \ + SELECT id, k.choice FROM c", + ) + .await + .unwrap(); + assert_eq!(result.iter().map(|b| b.num_rows()).sum::(), 1); + let calls = mock.calls.lock().unwrap(); + assert_eq!(calls.len(), 1); + assert_eq!( + calls[0]["questions"].as_object().unwrap().len(), + 1, + "only the row that passed the filter is asked about: {}", + calls[0] + ); +} + +#[tokio::test] +async fn leaf_pushdown_still_runs_for_queries_without_prompt_jev() { + // The guard must not disable DataFusion's optimization for ordinary queries. + let ctx = SessionContext::new(); + datafusion_jev::register(&ctx, Arc::new(Mock::default())); + let df = datafusion_jev::sql( + &ctx, + "EXPLAIN VERBOSE SELECT s.x FROM (SELECT struct(id AS x) AS s FROM (VALUES (1),(2)) t(id) WHERE id = 1) q", + ) + .await + .unwrap(); + let text = display(&df.collect().await.unwrap()); + assert!(text.contains("__datafusion_extracted"), "{text}"); +} + +#[tokio::test] +async fn group_by_prompt_jev_is_planned() { + let mock = Arc::new(Mock::default()); + let result = run( + mock.clone(), + "SELECT round(prompt_jev(body, 'Urgent?'), 1) AS p, count(*) AS n \ + FROM (VALUES (1,'a'),(2,'b'),(3,'c')) t(id, body) GROUP BY 1 ORDER BY 1", + ) + .await + .unwrap(); + assert!(result.iter().map(|b| b.num_rows()).sum::() >= 1); + assert_eq!(mock.calls.lock().unwrap().len(), 1); +} + +#[tokio::test] +async fn one_sided_join_condition_with_prompt_jev_is_hoisted() { + let mock = Arc::new(Mock::default()); + let result = run( + mock.clone(), + "SELECT a.id FROM (VALUES (1,'x'),(2,'y')) a(id, body) JOIN (VALUES (1),(2)) b(id) \ + ON a.id = b.id AND prompt_jev(a.body, 'Urgent?') > 0.5 ORDER BY a.id", + ) + .await + .unwrap(); + assert_eq!(result.iter().map(|b| b.num_rows()).sum::(), 2); + assert_eq!(mock.calls.lock().unwrap().len(), 1); +} + +#[tokio::test] +async fn two_sided_join_condition_is_a_clear_planning_error() { + let mock = Arc::new(Mock::default()); + let err = run( + mock.clone(), + "SELECT a.id FROM (VALUES (1,'x')) a(id, body) JOIN (VALUES (1,'y')) b(id, body) \ + ON a.id = b.id AND prompt_jev(a.body || b.body, 'Same?') > 0.5", + ) + .await + .err() + .unwrap() + .to_string(); + assert!(err.contains("one side only"), "{err}"); + assert!(!err.contains("called directly"), "{err}"); + assert!(mock.calls.lock().unwrap().is_empty()); +}