From 2f94af2a7965235c75014dd81bb9d6be0fee3b90 Mon Sep 17 00:00:00 2001 From: Eddie A Tejeda <669988+eddietejeda@users.noreply.github.com> Date: Thu, 24 Sep 2026 09:03:32 -0700 Subject: [PATCH 1/2] refactor: share the internal name and split the long functions No behaviour, error-message, or public API change. - new src/names.rs: INTERNAL_FUNCTION plus the is_jev / is_jev_expr predicates and an async_exprs wrapper that carries the single #[allow(deprecated)]; sql.rs, udf.rs, planner.rs and optimizer.rs now use them instead of repeating the literal. - optimizer.rs: both rules start from one async_node helper; module doc now describes both rules in code order. - planner.rs: hoist_calls_named folded into hoist_calls with a prefix parameter; module doc rewritten as a description of the two rules. - sql.rs: rewrite_statement split into parse_call and parse_batch_size; criteria() match arms reformatted. - udf.rs: plan_batches (with a Batches struct) and probabilities() extracted from invoke_async_with_args and answer(). - tests/sql.rs split into tests/sql/{main,common,syntax,execution, planning}.rs; the same 36 tests, unchanged. --- src/lib.rs | 1 + src/names.rs | 31 ++ src/optimizer.rs | 53 ++- src/planner.rs | 57 +-- src/sql.rs | 209 +++++----- src/udf.rs | 147 ++++--- tests/sql.rs | 893 ----------------------------------------- tests/sql/common.rs | 112 ++++++ tests/sql/execution.rs | 448 +++++++++++++++++++++ tests/sql/main.rs | 5 + tests/sql/planning.rs | 217 ++++++++++ tests/sql/syntax.rs | 126 ++++++ 12 files changed, 1181 insertions(+), 1118 deletions(-) create mode 100644 src/names.rs delete mode 100644 tests/sql.rs create mode 100644 tests/sql/common.rs create mode 100644 tests/sql/execution.rs create mode 100644 tests/sql/main.rs create mode 100644 tests/sql/planning.rs create mode 100644 tests/sql/syntax.rs diff --git a/src/lib.rs b/src/lib.rs index cd2b5f5..a806db1 100644 --- a/src/lib.rs +++ b/src/lib.rs @@ -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; diff --git a/src/names.rs b/src/names.rs new file mode 100644 index 0000000..e36ab64 --- /dev/null +++ b/src/names.rs @@ -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) -> bool { + expr.downcast_ref::() + .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] { + #[allow(deprecated)] + node.async_exprs() +} diff --git a/src/optimizer.rs b/src/optimizer.rs index b6f2369..c3af4ba 100644 --- a/src/optimizer.rs +++ b/src/optimizer.rs @@ -1,13 +1,16 @@ -//! 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, @@ -15,8 +18,8 @@ use datafusion::{ }, 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::{ @@ -29,9 +32,11 @@ use datafusion::{ }; use std::sync::Arc; -fn is_jev_expr(expr: &Arc) -> bool { - expr.downcast_ref::() - .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) -> Option<(&AsyncFuncExec, &[Arc])> { + let node = plan.downcast_ref::()?; + Some((node, async_exprs(node))) } #[derive(Debug)] @@ -62,11 +67,9 @@ impl PhysicalOptimizerRule for FilterBeforeJev { passthrough.push(cursor); cursor = child; } - let Some(node) = cursor.downcast_ref::() 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)); } @@ -126,23 +129,13 @@ impl PhysicalOptimizerRule for DeduplicateJev { _: &ConfigOptions, ) -> Result> { plan.transform_up(|plan| { - let Some(node) = plan.downcast_ref::() 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, - > = vec![]; + let mut unique: Vec> = vec![]; let mut indices = vec![]; for expr in expressions { - let is_jev = expr - .func - .downcast_ref::() - .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 diff --git a/src/planner.rs b/src/planner.rs index 04b8cde..3d7b836 100644 --- a/src/planner.rs +++ b/src/planner.rs @@ -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, @@ -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) } @@ -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, - 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 { @@ -163,7 +147,10 @@ fn collect_calls(exprs: &[&Expr]) -> Result> { 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, exprs: &[&Expr], prefix: &str, @@ -213,7 +200,7 @@ fn hoist_sort(sort: Sort) -> Result { .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)?))) @@ -235,7 +222,7 @@ fn hoist_window(window: Window) -> Result { } = 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 (projected, substitutions) = hoist_calls(input, &refs, HOIST_PREFIX)?; let window_expr = window_expr .into_iter() .map(|e| { @@ -271,7 +258,7 @@ fn hoist_aggregate(agg: Aggregate) -> Result { .. } = 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)) @@ -323,8 +310,8 @@ fn hoist_join(join: Join) -> Result { } 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() diff --git a/src/sql.rs b/src/sql.rs index 54d32c4..0fdf3d9 100644 --- a/src/sql.rs +++ b/src/sql.rs @@ -1,4 +1,5 @@ //! Lower the public SQL syntax to a validated, constant configuration. +use crate::names::INTERNAL_FUNCTION; use datafusion::{ common::{Result, plan_datafusion_err}, sql::sqlparser::ast::{ @@ -120,7 +121,9 @@ fn criteria(e: &Expr) -> Result> { return Err(plan_datafusion_err!("duplicate criterion field")); } match key { - "label" => label = Some(string(&field.value)?), + "label" => { + label = Some(string(&field.value)?); + } "description" => { let null = matches!( field.value.as_ref(), @@ -130,7 +133,9 @@ fn criteria(e: &Expr) -> Result> { description = Some(string(&field.value)?); } } - _ => return Err(plan_datafusion_err!("unknown criterion field: {key}")), + _ => { + return Err(plan_datafusion_err!("unknown criterion field: {key}")); + } } } Ok(Criterion { @@ -140,6 +145,107 @@ fn criteria(e: &Expr) -> Result> { }) .collect() } +fn parse_batch_size(e: &Expr) -> Result { + match e { + Expr::Value(v) => match &v.value { + Value::Number(n, _) => n.parse().ok(), + _ => None, + }, + _ => None, + } + .ok_or_else(|| plan_datafusion_err!("batch_size must be a constant integer")) +} +/// Read one `prompt_jev` call: the input expression and the question it asks. +/// The question is not validated here; the caller does that before rewriting. +fn parse_call(f: &mut ast::Function) -> Result<(Expr, Question)> { + if f.filter.is_some() + || f.over.is_some() + || f.null_treatment.is_some() + || !f.within_group.is_empty() + || !matches!(f.parameters, FunctionArguments::None) + { + return Err(plan_datafusion_err!( + "prompt_jev does not accept aggregate/window modifiers" + )); + } + let FunctionArguments::List(args) = &mut f.args else { + return Err(plan_datafusion_err!( + "prompt_jev requires input and instructions" + )); + }; + if args.duplicate_treatment.is_some() || !args.clauses.is_empty() { + return Err(plan_datafusion_err!("invalid prompt_jev modifiers")); + } + let mut positional = Vec::new(); + let mut named = HashSet::new(); + let mut q = Question { + instructions: String::new(), + kind: "noul".into(), + criteria: vec![], + batch_size: 32, + }; + let mut mode = false; + for arg in &args.args { + match arg { + FunctionArg::Unnamed(FunctionArgExpr::Expr(e)) if named.is_empty() => { + positional.push(e.clone()) + } + FunctionArg::Named { + name, + arg: FunctionArgExpr::Expr(e), + .. + } => { + let name = if name.quote_style.is_some() { + name.value.clone() + } else { + name.value.to_lowercase() + }; + if !named.insert(name.clone()) { + return Err(plan_datafusion_err!( + "duplicate prompt_jev argument: {name}" + )); + } + match name.as_str() { + "choice" | "score" | "noul" => { + if mode { + return Err(plan_datafusion_err!( + "choice, score, and noul are mutually exclusive" + )); + } + mode = true; + q.kind = name; + q.criteria = criteria(e)?; + if q.kind == "noul" && q.criteria.is_empty() { + return Err(plan_datafusion_err!( + "explicit noul criteria require true and false" + )); + } + } + "batch_size" => { + q.batch_size = parse_batch_size(e)?; + } + _ => { + return Err(plan_datafusion_err!( + "unsupported prompt_jev argument: {name}" + )); + } + } + } + _ => { + return Err(plan_datafusion_err!( + "prompt_jev requires two positional arguments followed by named options" + )); + } + } + } + if positional.len() != 2 { + return Err(plan_datafusion_err!( + "prompt_jev requires input and constant instructions" + )); + } + q.instructions = string(&positional[1])?; + Ok((positional.remove(0), q)) +} /// Call after parsing and before DataFusion plans the statement. Idempotent. pub fn rewrite_statement(statement: &mut ast::Statement) -> Result<()> { let flow = ast::visit_expressions_mut(statement, |expr| { @@ -158,106 +264,17 @@ pub fn rewrite_statement(statement: &mut ast::Statement) -> Result<()> { return ControlFlow::Continue(()); } let result = (|| { - if f.filter.is_some() - || f.over.is_some() - || f.null_treatment.is_some() - || !f.within_group.is_empty() - || !matches!(f.parameters, FunctionArguments::None) - { - return Err(plan_datafusion_err!( - "prompt_jev does not accept aggregate/window modifiers" - )); - } + let (input, q) = parse_call(f)?; + q.validate()?; + let config = serde_json::to_string(&q).map_err(|e| plan_datafusion_err!("{e}"))?; let FunctionArguments::List(args) = &mut f.args else { return Err(plan_datafusion_err!( "prompt_jev requires input and instructions" )); }; - if args.duplicate_treatment.is_some() || !args.clauses.is_empty() { - return Err(plan_datafusion_err!("invalid prompt_jev modifiers")); - } - let mut positional = Vec::new(); - let mut named = HashSet::new(); - let mut q = Question { - instructions: String::new(), - kind: "noul".into(), - criteria: vec![], - batch_size: 32, - }; - let mut mode = false; - for arg in &args.args { - match arg { - FunctionArg::Unnamed(FunctionArgExpr::Expr(e)) if named.is_empty() => { - positional.push(e.clone()) - } - FunctionArg::Named { - name, - arg: FunctionArgExpr::Expr(e), - .. - } => { - let name = if name.quote_style.is_some() { - name.value.clone() - } else { - name.value.to_lowercase() - }; - if !named.insert(name.clone()) { - return Err(plan_datafusion_err!( - "duplicate prompt_jev argument: {name}" - )); - } - match name.as_str() { - "choice" | "score" | "noul" => { - if mode { - return Err(plan_datafusion_err!( - "choice, score, and noul are mutually exclusive" - )); - } - mode = true; - q.kind = name; - q.criteria = criteria(e)?; - if q.kind == "noul" && q.criteria.is_empty() { - return Err(plan_datafusion_err!( - "explicit noul criteria require true and false" - )); - } - } - "batch_size" => { - q.batch_size = match e { - Expr::Value(v) => match &v.value { - Value::Number(n, _) => n.parse().ok(), - _ => None, - }, - _ => None, - } - .ok_or_else(|| { - plan_datafusion_err!("batch_size must be a constant integer") - })?; - } - _ => { - return Err(plan_datafusion_err!( - "unsupported prompt_jev argument: {name}" - )); - } - } - } - _ => { - return Err(plan_datafusion_err!( - "prompt_jev requires two positional arguments followed by named options" - )); - } - } - } - if positional.len() != 2 { - return Err(plan_datafusion_err!( - "prompt_jev requires input and constant instructions" - )); - } - q.instructions = string(&positional[1])?; - q.validate()?; - let config = serde_json::to_string(&q).map_err(|e| plan_datafusion_err!("{e}"))?; - f.name = ObjectName::from(vec![Ident::new("__datafusion_jev")]); + f.name = ObjectName::from(vec![Ident::new(INTERNAL_FUNCTION)]); args.args = vec![ - FunctionArg::Unnamed(FunctionArgExpr::Expr(positional.remove(0))), + FunctionArg::Unnamed(FunctionArgExpr::Expr(input)), FunctionArg::Unnamed(FunctionArgExpr::Expr(Expr::Value( Value::SingleQuotedString(config).into(), ))), diff --git a/src/udf.rs b/src/udf.rs index 1b0f30a..3bf6fb0 100644 --- a/src/udf.rs +++ b/src/udf.rs @@ -1,4 +1,4 @@ -use crate::{JevClient, RequestError, sql::Question}; +use crate::{JevClient, RequestError, names::INTERNAL_FUNCTION, sql::Question}; use async_trait::async_trait; use datafusion::{ arrow::{ @@ -118,7 +118,7 @@ fn is_text(t: &DataType) -> bool { } impl ScalarUDFImpl for Jev { fn name(&self) -> &str { - "__datafusion_jev" + INTERNAL_FUNCTION } fn signature(&self) -> &Signature { &self.signature @@ -156,22 +156,9 @@ fn number(value: &Value, upper: f64) -> Result { .filter(|v| v.is_finite() && *v >= 0.0 && *v <= upper) .ok_or_else(|| exec_datafusion_err!("Jev returned an invalid numeric answer")) } -fn answer(q: &Question, value: &Value) -> Result { - if value["type"].as_str() != Some(q.kind.as_str()) { - return Err(exec_datafusion_err!( - "Jev returned an unexpected answer type" - )); - } - if q.kind == "noul" { - return Ok(ScalarValue::Float64(Some(number(&value["noul"], 1.0)?))); - } - let fs = fields(q); - let DataType::List(item) = fs[1].data_type() else { - unreachable!() - }; - let DataType::Struct(pfields) = item.data_type() else { - unreachable!() - }; +/// One struct entry per criterion, in the caller's order, and the sum of their +/// probabilities. +fn probabilities(q: &Question, value: &Value, pfields: &Fields) -> Result<(Vec, f64)> { let probs = value["probabilities"] .as_object() .ok_or_else(|| exec_datafusion_err!("Jev omitted probabilities"))?; @@ -198,6 +185,25 @@ fn answer(q: &Question, value: &Value) -> Result { values.push(ScalarValue::Float64(Some(probability))); entries.push(scalar_struct(pfields.clone(), values)?); } + Ok((entries, sum)) +} +fn answer(q: &Question, value: &Value) -> Result { + if value["type"].as_str() != Some(q.kind.as_str()) { + return Err(exec_datafusion_err!( + "Jev returned an unexpected answer type" + )); + } + if q.kind == "noul" { + return Ok(ScalarValue::Float64(Some(number(&value["noul"], 1.0)?))); + } + let fs = fields(q); + let DataType::List(item) = fs[1].data_type() else { + unreachable!() + }; + let DataType::Struct(pfields) = item.data_type() else { + unreachable!() + }; + let (entries, sum) = probabilities(q, value, pfields)?; // Providers round each probability independently, so the error a valid // answer can accumulate grows with the number of criteria. let tolerance = (0.01 + 0.001 * q.criteria.len() as f64).min(0.3); @@ -288,6 +294,64 @@ fn request_body(q: &Question, rows: &[(usize, String)]) -> Value { .collect(); json!({"model": "jev-latest", "state": state, "questions": questions}) } +/// The requests to make for one input array. +struct Batches { + /// Each request: the distinct texts it asks about, as (text index, text). + batches: Vec>, + /// For each distinct text, the rows of the input that held it. + rows_for: Vec>, +} +/// Ask about each distinct text once, and group the texts into requests that +/// respect `batch_size` and the encoded-size budget. Repeated and constant +/// inputs are common, and the answer is copied back to every row that shared +/// the text. +fn plan_batches(input: &ArrayRef, batch_size: usize) -> Result { + let mut unique: Vec<(String, usize)> = vec![]; + let mut rows_for: Vec> = vec![]; + let mut seen: HashMap = HashMap::new(); + for i in 0..input.len() { + if input.data_type() == &DataType::Null || input.is_null(i) { + continue; + } + let s = match ScalarValue::try_from_array(input, i)? { + ScalarValue::Utf8(Some(s)) + | ScalarValue::Utf8View(Some(s)) + | ScalarValue::LargeUtf8(Some(s)) => s, + _ => return Err(exec_datafusion_err!("prompt_jev input must be text")), + }; + if let Some(&u) = seen.get(&s) { + rows_for[u].push(i); + continue; + } + // Measure what actually travels: JSON escaping can expand text severalfold. + let encoded = serde_json::to_string(&s) + .map_err(|_| exec_datafusion_err!("invalid Jev request"))? + .len(); + if encoded > 64 * 1024 { + return Err(exec_datafusion_err!( + "prompt_jev input exceeds 64 KiB once JSON-encoded" + )); + } + seen.insert(s.clone(), unique.len()); + unique.push((s, encoded)); + rows_for.push(vec![i]); + } + let mut batches = vec![]; + let mut batch = vec![]; + let mut bytes = 0; + for (u, (s, encoded)) in unique.into_iter().enumerate() { + if !batch.is_empty() && (batch.len() >= batch_size || bytes + encoded > 32 * 1024) { + batches.push(std::mem::take(&mut batch)); + bytes = 0; + } + bytes += encoded; + batch.push((u, s)); + } + if !batch.is_empty() { + batches.push(batch); + } + Ok(Batches { batches, rows_for }) +} impl Jev { async fn batch( &self, @@ -373,52 +437,7 @@ impl AsyncScalarUDFImpl for Jev { if matches!(input.data_type(), DataType::Dictionary(_, _)) { input = datafusion::arrow::compute::cast(&input, &DataType::Utf8)?; } - // Ask about each distinct text once. Repeated and constant inputs are - // common, and the answer is copied back to every row that shared the text. - let mut unique: Vec<(String, usize)> = vec![]; - let mut rows_for: Vec> = vec![]; - let mut seen: HashMap = HashMap::new(); - for i in 0..len { - if input.data_type() == &DataType::Null || input.is_null(i) { - continue; - } - let s = match ScalarValue::try_from_array(&input, i)? { - ScalarValue::Utf8(Some(s)) - | ScalarValue::Utf8View(Some(s)) - | ScalarValue::LargeUtf8(Some(s)) => s, - _ => return Err(exec_datafusion_err!("prompt_jev input must be text")), - }; - if let Some(&u) = seen.get(&s) { - rows_for[u].push(i); - continue; - } - // Measure what actually travels: JSON escaping can expand text severalfold. - let encoded = serde_json::to_string(&s) - .map_err(|_| exec_datafusion_err!("invalid Jev request"))? - .len(); - if encoded > 64 * 1024 { - return Err(exec_datafusion_err!( - "prompt_jev input exceeds 64 KiB once JSON-encoded" - )); - } - seen.insert(s.clone(), unique.len()); - unique.push((s, encoded)); - rows_for.push(vec![i]); - } - let mut batches = vec![]; - let mut batch = vec![]; - let mut bytes = 0; - for (u, (s, encoded)) in unique.into_iter().enumerate() { - if !batch.is_empty() && (batch.len() >= q.batch_size || bytes + encoded > 32 * 1024) { - batches.push(std::mem::take(&mut batch)); - bytes = 0; - } - bytes += encoded; - batch.push((u, s)); - } - if !batch.is_empty() { - batches.push(batch); - } + let Batches { batches, rows_for } = plan_batches(&input, q.batch_size)?; let mut out = vec![ScalarValue::try_from(&datatype(&q))?; len]; let mut pending = stream::iter(batches.into_iter().map(|rows| self.batch(&q, rows))).buffer_unordered(8); diff --git a/tests/sql.rs b/tests/sql.rs deleted file mode 100644 index 9b84ce8..0000000 --- a/tests/sql.rs +++ /dev/null @@ -1,893 +0,0 @@ -use async_trait::async_trait; -use datafusion::{ - arrow::{ - array::{ - Array, ArrayRef, DictionaryArray, Float64Array, ListArray, RecordBatch, StringArray, - StructArray, - }, - datatypes::{Field, Int32Type, Schema, UInt32Type}, - }, - common::Result, - datasource::MemTable, - prelude::*, -}; -use datafusion_jev::{JevClient, RequestError}; -use serde_json::{Value, json}; -use std::sync::{Arc, Mutex}; - -#[derive(Debug, Default)] -struct Mock { - calls: Mutex>, - unavailable: bool, - fatal: bool, - malformed: bool, - /// When set, every choice probability takes this value instead of a - /// one-hot distribution, so a test can control the probability sum. - probability: Option, -} -#[async_trait] -impl JevClient for Mock { - async fn request(&self, body: Value) -> std::result::Result { - self.calls.lock().unwrap().push(body.clone()); - if self.unavailable { - return Err(RequestError::Unavailable); - } - if self.fatal { - return Err(RequestError::Fatal("HTTP 422".into())); - } - if self.malformed { - return Ok(json!({"answers":{}})); - } - let mut answers = serde_json::Map::new(); - for (key, q) in body["questions"].as_object().unwrap() { - let value = match q["type"].as_str().unwrap() { - "noul" => json!({"type":"noul","noul":0.9}), - "choice" => { - let labels: Vec<_> = q["criteria"].as_object().unwrap().keys().collect(); - let probabilities: serde_json::Map = labels - .iter() - .enumerate() - .map(|(i, k)| { - let p = self.probability.unwrap_or(if i == 0 { 1.0 } else { 0.0 }); - ((*k).clone(), json!(p)) - }) - .collect(); - json!({"type":"choice","choice":labels[0],"confidence":0.95,"probabilities":probabilities}) - } - "score" => { - let n = q["criteria"].as_array().unwrap().len(); - let probabilities: serde_json::Map = (0..n) - .map(|i| { - ( - i.to_string(), - json!(if i == 0 { - 0.25 - } else if i == n - 1 { - 0.75 - } else { - 0.0 - }), - ) - }) - .collect(); - json!({"type":"score","score":0.75*(n-1) as f64,"confidence":0.7,"probabilities":probabilities}) - } - _ => unreachable!(), - }; - answers.insert(key.clone(), value); - } - Ok(json!({"answers":answers})) - } -} -async fn run(client: Arc, sql: &str) -> Result> { - plan_and_run(SessionContext::new(), client, sql).await -} -/// Run `sql` against a single-column table `t(body)` built from `body`. -async fn run_table( - client: Arc, - sql: &str, - body: ArrayRef, -) -> Result> { - let ctx = SessionContext::new(); - let schema = Arc::new(Schema::new(vec![Field::new( - "body", - body.data_type().clone(), - true, - )])); - let batch = RecordBatch::try_new(schema.clone(), vec![body])?; - ctx.register_table("t", Arc::new(MemTable::try_new(schema, vec![vec![batch]])?))?; - plan_and_run(ctx, client, sql).await -} -async fn plan_and_run( - ctx: SessionContext, - client: Arc, - sql: &str, -) -> Result> { - datafusion_jev::register(&ctx, client); - datafusion_jev::sql(&ctx, sql).await?.collect().await -} -fn display(batches: &[RecordBatch]) -> String { - datafusion::arrow::util::pretty::pretty_format_batches(batches) - .unwrap() - .to_string() -} -#[tokio::test] -async fn exact_user_syntax_and_typed_fields() { - let mock = Arc::new(Mock::default()); - let result = run( - mock.clone(), - r#" - WITH classified AS ( - SELECT conversation_id, prompt_jev(transcript, 'Identify the customer''s main complaint', choice := [ - {label: 'billing', description: 'Payments, invoices, and refunds'}, - {label: 'technical', description: 'Errors, outages, and integrations'}, - {label: 'sales', description: 'Pricing and upgrades'}, - {label: 'account', description: 'Cancellations and account administration'} - ]) AS classification - FROM (VALUES (1, 'Please refund this payment')) t(conversation_id, transcript) - ) SELECT conversation_id, classification.choice, classification.confidence, classification.probabilities FROM classified - "#, - ) - .await - .unwrap(); - let text = display(&result); - assert!(text.contains("0.95"), "{text}"); - assert!(text.contains("account"), "{text}"); - assert_eq!( - mock.calls.lock().unwrap().len(), - 1, - "reading fields must not repeat inference" - ); -} -#[tokio::test] -async fn batching_and_null_input() { - let mock = Arc::new(Mock::default()); - let result = run( - mock.clone(), - "SELECT id, prompt_jev(body, 'Urgent?', batch_size := 2) AS p \ - FROM (VALUES (1,'a'), (2,NULL), (3,'b'), (4,'c')) t(id,body)", - ) - .await - .unwrap(); - assert_eq!(result.iter().map(|b| b.num_rows()).sum::(), 4); - let calls = mock.calls.lock().unwrap(); - assert_eq!(calls.len(), 2); - assert_eq!( - calls - .iter() - .map(|v| v["questions"].as_object().unwrap().len()) - .sum::(), - 3 - ); - assert!(display(&result).contains("0.9")); -} -#[tokio::test] -async fn score_and_noul_types() { - let result = run( - Arc::new(Mock::default()), - "SELECT prompt_jev('bad', 'Severity?', score := ['low','medium','high']) AS s, \ - prompt_jev('bad','Urgent?') AS p", - ) - .await - .unwrap(); - let text = display(&result); - assert!(text.contains("1.5"), "{text}"); - assert!(text.contains("0.9"), "{text}"); - assert!(text.contains("medium"), "{text}"); -} -#[tokio::test] -async fn invalid_arguments_never_call_provider() { - for sql in [ - "SELECT prompt_jev('a','q',choice := ['a'])", - "SELECT prompt_jev('a','q',choice := ['a','a'])", - "SELECT prompt_jev('a','q',choice := ['a','b'], score := ['x','y'])", - "SELECT prompt_jev('a','q',batch_size := 0)", - "SELECT prompt_jev('a','q',batch_size := 65)", - "SELECT prompt_jev('a',body) FROM (VALUES ('q')) t(body)", - "SELECT prompt_jev('a','q',choice := body) FROM (VALUES ('q')) t(body)", - "SELECT prompt_jev('a','q',noul := ['yes','no'])", - "SELECT prompt_jev('a','q',model := 'foo')", - "SELECT prompt_jev(42,'q')", - "SELECT prompt_jev('a','q',noul := [])", - ] { - let mock = Arc::new(Mock::default()); - assert!(run(mock.clone(), sql).await.is_err(), "{sql}"); - assert!(mock.calls.lock().unwrap().is_empty()); - } -} -#[tokio::test] -async fn unavailable_is_retried_then_null() { - let mock = Arc::new(Mock { - unavailable: true, - ..Default::default() - }); - let result = run(mock.clone(), "SELECT prompt_jev('a','q') AS p") - .await - .unwrap(); - assert!(result[0].column(0).is_null(0)); - assert_eq!(mock.calls.lock().unwrap().len(), 3); -} -#[tokio::test] -async fn provider_validation_and_malformed_answers_fail() { - for mock in [ - Mock { - fatal: true, - ..Default::default() - }, - Mock { - malformed: true, - ..Default::default() - }, - ] { - let mock = Arc::new(mock); - assert!( - run(mock.clone(), "SELECT prompt_jev('a','q') AS p") - .await - .is_err() - ); - assert_eq!(mock.calls.lock().unwrap().len(), 1); - } -} -#[tokio::test] -async fn nulls_and_empty_results_make_no_requests() { - let mock = Arc::new(Mock::default()); - run(mock.clone(), "SELECT prompt_jev(NULL,'q') AS p") - .await - .unwrap(); - run( - mock.clone(), - "SELECT prompt_jev(body,'q') AS p FROM (VALUES ('a')) t(body) WHERE false", - ) - .await - .unwrap(); - assert!(mock.calls.lock().unwrap().is_empty()); -} -#[tokio::test] -async fn semantic_filter_and_aggregate() { - let mock = Arc::new(Mock::default()); - let result = run( - mock, - "SELECT count(*) AS n FROM (VALUES ('a'),('b')) t(body) \ - WHERE prompt_jev(body,'Urgent?') > 0.8", - ) - .await - .unwrap(); - assert!(display(&result).contains("2")); -} - -#[derive(Debug)] -struct Hanging { - started: Arc, - dropped: Arc, -} -struct DropSignal(Arc); -impl Drop for DropSignal { - fn drop(&mut self) { - self.0.notify_one(); - } -} -#[async_trait] -impl JevClient for Hanging { - async fn request(&self, _: Value) -> std::result::Result { - let _signal = DropSignal(self.dropped.clone()); - self.started.notify_one(); - std::future::pending().await - } -} -#[tokio::test] -async fn cancelling_query_drops_inflight_request() { - let started = Arc::new(tokio::sync::Notify::new()); - let dropped = Arc::new(tokio::sync::Notify::new()); - let client = Arc::new(Hanging { - started: started.clone(), - dropped: dropped.clone(), - }); - let task = tokio::spawn(async move { run(client, "SELECT prompt_jev('a','q') AS p").await }); - tokio::time::timeout(std::time::Duration::from_secs(5), started.notified()) - .await - .unwrap(); - task.abort(); - let _ = task.await; - tokio::time::timeout(std::time::Duration::from_secs(5), dropped.notified()) - .await - .unwrap(); -} -#[tokio::test] -async fn isolated_batches_keep_state_as_single_text() { - let mock = Arc::new(Mock::default()); - run( - mock.clone(), - "SELECT prompt_jev(body,'q',batch_size := 1) FROM (VALUES ('alpha'),('beta')) t(body)", - ) - .await - .unwrap(); - let calls = mock.calls.lock().unwrap(); - assert_eq!(calls.len(), 2); - assert!(calls.iter().all(|c| c["state"].is_string())); -} -#[tokio::test] -async fn explain_does_not_run_inference() { - let mock = Arc::new(Mock::default()); - run(mock.clone(), "EXPLAIN SELECT prompt_jev('a','q') AS p") - .await - .unwrap(); - assert!(mock.calls.lock().unwrap().is_empty()); -} - -#[derive(Debug, Default)] -struct Concurrent { - mock: Mock, - active: std::sync::atomic::AtomicUsize, - peak: std::sync::atomic::AtomicUsize, -} -#[async_trait] -impl JevClient for Concurrent { - async fn request(&self, body: Value) -> std::result::Result { - use std::sync::atomic::Ordering::SeqCst; - let active = self.active.fetch_add(1, SeqCst) + 1; - self.peak.fetch_max(active, SeqCst); - tokio::time::sleep(std::time::Duration::from_millis(10)).await; - let result = self.mock.request(body).await; - self.active.fetch_sub(1, SeqCst); - result - } -} -#[tokio::test] -async fn inference_concurrency_is_bounded() { - let mock = Arc::new(Concurrent::default()); - let values = (0..64) - .map(|i| format!("('row {i}')")) - .collect::>() - .join(","); - let sql = format!("SELECT prompt_jev(body,'q',batch_size := 1) FROM (VALUES {values}) t(body)"); - run(mock.clone(), &sql).await.unwrap(); - let peak = mock.peak.load(std::sync::atomic::Ordering::SeqCst); - assert!(peak > 1 && peak <= 8, "peak requests: {peak}"); - assert_eq!(mock.mock.calls.lock().unwrap().len(), 64); -} - -#[tokio::test] -async fn identical_inputs_are_requested_once() { - let mock = Arc::new(Mock::default()); - let result = run( - mock.clone(), - "SELECT prompt_jev(body,'q') AS p FROM (VALUES ('a'),('a'),('b'),('a')) t(body)", - ) - .await - .unwrap(); - let calls = mock.calls.lock().unwrap(); - assert_eq!(calls.len(), 1, "repeated text must be asked about once"); - assert_eq!(calls[0]["questions"].as_object().unwrap().len(), 2); - let column = result[0].column_by_name("p").unwrap(); - assert_eq!(result.iter().map(|b| b.num_rows()).sum::(), 4); - assert_eq!(column.null_count(), 0, "every row keeps an answer"); -} -#[tokio::test] -async fn constant_input_is_requested_once() { - let mock = Arc::new(Mock::default()); - let result = run( - mock.clone(), - "SELECT prompt_jev('x','q') AS p FROM (VALUES (1),(2),(3)) t(id)", - ) - .await - .unwrap(); - assert_eq!(result.iter().map(|b| b.num_rows()).sum::(), 3); - let calls = mock.calls.lock().unwrap(); - assert_eq!(calls.len(), 1); - assert_eq!(calls[0]["questions"].as_object().unwrap().len(), 1); -} -fn many_labels(n: usize) -> String { - (0..n) - .map(|i| format!("'c{i}'")) - .collect::>() - .join(",") -} -#[tokio::test] -async fn probability_sum_tolerance_scales_with_criteria_count() { - // 200 criteria rounded to three decimals can drift far from one. - let mock = Arc::new(Mock { - probability: Some(0.004), - ..Default::default() - }); - let sql = format!( - "SELECT prompt_jev('a','q',choice := [{}]) AS p", - many_labels(200) - ); - run(mock, &sql).await.expect("rounding drift is tolerated"); - // Garbage is still rejected: 200 * 0.0075 = 1.5. - let mock = Arc::new(Mock { - probability: Some(0.0075), - ..Default::default() - }); - let sql = format!( - "SELECT prompt_jev('a','q',choice := [{}]) AS p", - many_labels(200) - ); - assert!(run(mock, &sql).await.is_err()); - // A small question keeps the tight tolerance. - let mock = Arc::new(Mock { - probability: Some(0.4), - ..Default::default() - }); - assert!( - run(mock, "SELECT prompt_jev('a','q',choice := ['x','y']) AS p") - .await - .is_err() - ); -} -#[tokio::test] -async fn oversized_encoded_input_is_rejected_without_batch_size_advice() { - let mock = Arc::new(Mock::default()); - // ~60k control characters: six bytes each once JSON-encoded. - let body = Arc::new(StringArray::from(vec!["\u{1}".repeat(60_000)])) as ArrayRef; - let error = run_table( - mock.clone(), - "SELECT prompt_jev(body,'q',batch_size := 1) AS p FROM t", - body, - ) - .await - .unwrap_err() - .to_string(); - assert!(error.contains("input exceeds"), "{error}"); - assert!(!error.contains("batch_size"), "{error}"); - assert!(mock.calls.lock().unwrap().is_empty()); -} -#[tokio::test] -async fn dictionary_encoded_input_is_accepted() { - let mock = Arc::new(Mock::default()); - let body = Arc::new( - vec!["alpha", "beta", "alpha"] - .into_iter() - .collect::>(), - ) as ArrayRef; - let result = run_table( - mock.clone(), - "SELECT prompt_jev(body,'q') AS p FROM t", - body, - ) - .await - .unwrap(); - assert_eq!(result.iter().map(|b| b.num_rows()).sum::(), 3); - assert_eq!(result[0].column_by_name("p").unwrap().null_count(), 0); - let calls = mock.calls.lock().unwrap(); - assert_eq!(calls.len(), 1); - assert_eq!(calls[0]["questions"].as_object().unwrap().len(), 2); -} -#[tokio::test] -async fn batches_split_on_encoded_size() { - let mock = Arc::new(Mock::default()); - // Three distinct ~20 KiB rows: the 32 KiB batch budget allows only one each. - let body = Arc::new(StringArray::from( - (0..3) - .map(|i| format!("{i}{}", "x".repeat(20_000))) - .collect::>(), - )) as ArrayRef; - run_table( - mock.clone(), - "SELECT prompt_jev(body,'q') AS p FROM t", - body, - ) - .await - .unwrap(); - let calls = mock.calls.lock().unwrap(); - assert_eq!( - calls.len(), - 3, - "default batch_size must still split by size" - ); - assert!(calls.iter().all(|c| c["state"].is_string())); -} -#[tokio::test] -async fn repeated_expressions_share_one_request_and_keep_columns() { - let mock = Arc::new(Mock::default()); - let result = run( - mock.clone(), - "WITH c AS (SELECT id, \ - prompt_jev(body,'a',choice := ['alpha','beta']) AS x, \ - prompt_jev(body,'b') AS y, \ - prompt_jev(body,'a',choice := ['alpha','beta']) AS z, \ - body FROM (VALUES (1,'m'),(2,'n')) t(id,body)) \ - SELECT id, x.choice AS xc, y, z.choice AS zc, body FROM c", - ) - .await - .unwrap(); - assert_eq!( - mock.calls.lock().unwrap().len(), - 2, - "the repeated expression must not be evaluated twice" - ); - let names: Vec<_> = result[0] - .schema() - .fields() - .iter() - .map(|f| f.name().clone()) - .collect(); - assert_eq!(names, ["id", "xc", "y", "zc", "body"]); - let text = display(&result); - for expected in ["alpha", "0.9", "| m ", "| n "] { - assert!(text.contains(expected), "{expected} missing from {text}"); - } - let xc = result[0].column_by_name("xc").unwrap(); - let zc = result[0].column_by_name("zc").unwrap(); - assert_eq!(xc.as_ref(), zc.as_ref(), "both slots hold the same answer"); -} -/// The struct field `name` of the `probabilities` list entries, as (value, probability). -fn probabilities(batch: &RecordBatch, name: &str) -> (Vec, Vec, Option>) { - let column = batch.column_by_name(name).unwrap(); - let outer = column.as_any().downcast_ref::().unwrap(); - let list = outer - .column_by_name("probabilities") - .unwrap() - .as_any() - .downcast_ref::() - .unwrap(); - let entries = list.value(0); - let entries = entries.as_any().downcast_ref::().unwrap(); - let values = entries - .column_by_name("value") - .unwrap() - .as_any() - .downcast_ref::() - .unwrap(); - let probs = entries - .column_by_name("probability") - .unwrap() - .as_any() - .downcast_ref::() - .unwrap(); - let index = entries.column_by_name("index").map(|c| { - c.as_any() - .downcast_ref::>() - .unwrap() - .values() - .to_vec() - }); - ( - values.iter().map(|v| v.unwrap().to_owned()).collect(), - probs.values().to_vec(), - index, - ) -} -#[tokio::test] -async fn probabilities_follow_caller_order() { - let result = run( - Arc::new(Mock::default()), - "SELECT prompt_jev('a','q',choice := ['billing','technical','sales']) AS c, \ - prompt_jev('a','q',score := ['low','medium','high']) AS s", - ) - .await - .unwrap(); - let (labels, probs, index) = probabilities(&result[0], "c"); - assert_eq!(labels, ["billing", "technical", "sales"]); - assert_eq!(probs, [1.0, 0.0, 0.0]); - assert_eq!(index, None, "choice carries no index"); - let (labels, probs, index) = probabilities(&result[0], "s"); - assert_eq!(labels, ["low", "medium", "high"]); - assert_eq!(probs, [0.25, 0.0, 0.75]); - assert_eq!(index, Some(vec![0, 1, 2])); -} -#[tokio::test] -async fn score_criteria_carry_their_labels() { - let mock = Arc::new(Mock::default()); - run( - mock.clone(), - "SELECT prompt_jev('a','q',score := [\ - {label: 'low', description: 'minor annoyance'}, \ - {label: 'high', description: 'outage'}]) AS s", - ) - .await - .unwrap(); - let calls = mock.calls.lock().unwrap(); - let criteria = calls[0]["questions"]["row_0"]["criteria"] - .as_array() - .unwrap(); - assert_eq!( - criteria, - &vec![json!("low: minor annoyance"), json!("high: outage")] - ); -} - -#[derive(Debug, Default)] -struct HangsOnce { - mock: Mock, - calls: std::sync::atomic::AtomicUsize, -} -#[async_trait] -impl JevClient for HangsOnce { - async fn request(&self, body: Value) -> std::result::Result { - if self.calls.fetch_add(1, std::sync::atomic::Ordering::SeqCst) == 0 { - std::future::pending::<()>().await; - } - self.mock.request(body).await - } -} -#[tokio::test(start_paused = true)] -async fn timed_out_request_is_retried() { - let client = Arc::new(HangsOnce::default()); - let result = run(client.clone(), "SELECT prompt_jev('a','q') AS p") - .await - .unwrap(); - assert!(display(&result).contains("0.9")); - assert_eq!(client.calls.load(std::sync::atomic::Ordering::SeqCst), 2); -} - -#[tokio::test] -async fn drop_in_sql_handles_ddl_set_and_options() { - let mock = Arc::new(Mock::default()); - let ctx = SessionContext::new(); - datafusion_jev::register(&ctx, mock.clone()); - datafusion_jev::sql(&ctx, "CREATE TABLE t AS VALUES ('a'), ('b')") - .await - .unwrap() - .collect() - .await - .unwrap(); - datafusion_jev::sql(&ctx, "SET datafusion.execution.batch_size = 1024") - .await - .unwrap() - .collect() - .await - .unwrap(); - let result = datafusion_jev::sql(&ctx, "SELECT prompt_jev(column1, 'q') AS p FROM t") - .await - .unwrap() - .collect() - .await - .unwrap(); - assert_eq!(result.iter().map(|b| b.num_rows()).sum::(), 2); - assert_eq!(mock.calls.lock().unwrap().len(), 1); - let forbidden = datafusion_jev::sql_with_options( - &ctx, - "CREATE TABLE u AS VALUES (1)", - datafusion::execution::context::SQLOptions::new().with_allow_ddl(false), - ) - .await; - assert!(forbidden.is_err()); -} -#[tokio::test] -async fn strict_dialect_gets_a_clear_error() { - let ctx = SessionContext::new_with_config( - SessionConfig::new().set_str("datafusion.sql_parser.dialect", "postgresql"), - ); - datafusion_jev::register(&ctx, Arc::new(Mock::default())); - let err = datafusion_jev::sql(&ctx, "SELECT prompt_jev('a','q', choice := ['x','y'])") - .await - .err() - .unwrap() - .to_string(); - assert!(err.contains("generic (default) or duckdb"), "{err}"); - assert!(err.contains("postgresql"), "{err}"); - let plain = datafusion_jev::sql(&ctx, "SELECT 1 +") - .await - .err() - .unwrap() - .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()); -} - -#[tokio::test] -async fn cheap_conjuncts_filter_rows_before_inference() { - // `WHERE id = 1 AND prompt_jev(...) > 0.5` must not send rows 2 and 3. - let mock = Arc::new(Mock::default()); - let result = run( - mock.clone(), - "SELECT id FROM (VALUES (1,'keep'),(2,'drop'),(3,'drop')) t(id, body) \ - WHERE id = 1 AND prompt_jev(body, 'Urgent?') > 0.5", - ) - .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 passing the cheap predicate is asked about: {}", - calls[0] - ); -} - -#[tokio::test] -async fn cheap_conjunct_in_a_count_filter_runs_first() { - let mock = Arc::new(Mock::default()); - let result = run( - mock.clone(), - "SELECT count(*) AS n FROM (VALUES ('a','x'),('b','y'),('c','z')) t(name, body) \ - WHERE name <> 'c' AND prompt_jev(body, 'Urgent?') > 0.5", - ) - .await - .unwrap(); - assert!(display(&result).contains('2'), "{}", display(&result)); - let calls = mock.calls.lock().unwrap(); - assert_eq!( - calls - .iter() - .map(|c| c["questions"].as_object().unwrap().len()) - .sum::(), - 2 - ); -} - -#[tokio::test] -async fn a_folded_limit_survives_the_filter_split() { - // With one partition, LimitPushdown folds LIMIT into the FilterExec's fetch. - // The split must keep it, or the query returns every matching row. - let mock = Arc::new(Mock::default()); - let ctx = SessionContext::new_with_config(SessionConfig::new().with_target_partitions(1)); - datafusion_jev::register(&ctx, mock.clone()); - let df = datafusion_jev::sql( - &ctx, - "SELECT id FROM (VALUES (1,'a'),(2,'b'),(3,'c'),(4,'d')) t(id, body) \ - WHERE id > 1 AND prompt_jev(body, 'Urgent?') > 0.5 LIMIT 1", - ) - .await - .unwrap(); - let result = df.collect().await.unwrap(); - assert_eq!(result.iter().map(|b| b.num_rows()).sum::(), 1); - let asked: usize = mock - .calls - .lock() - .unwrap() - .iter() - .map(|c| c["questions"].as_object().unwrap().len()) - .sum(); - assert!( - asked <= 3, - "rows failing `id > 1` must not be asked about: {asked}" - ); -} - -#[tokio::test] -async fn quoted_and_uppercase_names_are_the_same_function() { - let mock = Arc::new(Mock::default()); - let result = run( - mock.clone(), - "SELECT \"prompt_jev\"('a', 'q') AS a, PROMPT_JEV('a', 'q') AS b, \"PROMPT_JEV\"('a', 'q') AS c", - ) - .await - .unwrap(); - assert_eq!(result.iter().map(|b| b.num_rows()).sum::(), 1); - assert!(display(&result).contains("0.9")); -} - -#[tokio::test] -async fn oversized_instructions_and_labels_fail_at_plan_time() { - let long_q = "q".repeat(5_000); - let long_label = "L".repeat(300); - let long_desc = "d".repeat(2_000); - for (sql, needle) in [ - ( - format!("SELECT prompt_jev('a', '{long_q}')"), - "4000 characters", - ), - ( - format!("SELECT prompt_jev('a', 'q', choice := ['{long_label}', 'b'])"), - "labels are limited", - ), - ( - format!( - "SELECT prompt_jev('a', 'q', choice := [{{label: 'a', description: '{long_desc}'}}, 'b'])" - ), - "descriptions to 1024", - ), - ] { - let mock = Arc::new(Mock::default()); - let err = run(mock.clone(), &sql).await.err().unwrap().to_string(); - assert!(err.contains(needle), "{needle}: {err}"); - assert!(mock.calls.lock().unwrap().is_empty()); - } -} diff --git a/tests/sql/common.rs b/tests/sql/common.rs new file mode 100644 index 0000000..e43d5aa --- /dev/null +++ b/tests/sql/common.rs @@ -0,0 +1,112 @@ +//! Shared harness for the `prompt_jev` integration tests: a mock provider and +//! helpers that register the crate, run a query, and format the result. +use async_trait::async_trait; +use datafusion::{ + arrow::{ + array::{ArrayRef, RecordBatch}, + datatypes::{Field, Schema}, + }, + common::Result, + datasource::MemTable, + prelude::*, +}; +use datafusion_jev::{JevClient, RequestError}; +use serde_json::{Value, json}; +use std::sync::{Arc, Mutex}; + +#[derive(Debug, Default)] +pub struct Mock { + pub calls: Mutex>, + pub unavailable: bool, + pub fatal: bool, + pub malformed: bool, + /// When set, every choice probability takes this value instead of a + /// one-hot distribution, so a test can control the probability sum. + pub probability: Option, +} +#[async_trait] +impl JevClient for Mock { + async fn request(&self, body: Value) -> std::result::Result { + self.calls.lock().unwrap().push(body.clone()); + if self.unavailable { + return Err(RequestError::Unavailable); + } + if self.fatal { + return Err(RequestError::Fatal("HTTP 422".into())); + } + if self.malformed { + return Ok(json!({"answers":{}})); + } + let mut answers = serde_json::Map::new(); + for (key, q) in body["questions"].as_object().unwrap() { + let value = match q["type"].as_str().unwrap() { + "noul" => json!({"type":"noul","noul":0.9}), + "choice" => { + let labels: Vec<_> = q["criteria"].as_object().unwrap().keys().collect(); + let probabilities: serde_json::Map = labels + .iter() + .enumerate() + .map(|(i, k)| { + let p = self.probability.unwrap_or(if i == 0 { 1.0 } else { 0.0 }); + ((*k).clone(), json!(p)) + }) + .collect(); + json!({"type":"choice","choice":labels[0],"confidence":0.95,"probabilities":probabilities}) + } + "score" => { + let n = q["criteria"].as_array().unwrap().len(); + let probabilities: serde_json::Map = (0..n) + .map(|i| { + ( + i.to_string(), + json!(if i == 0 { + 0.25 + } else if i == n - 1 { + 0.75 + } else { + 0.0 + }), + ) + }) + .collect(); + json!({"type":"score","score":0.75*(n-1) as f64,"confidence":0.7,"probabilities":probabilities}) + } + _ => unreachable!(), + }; + answers.insert(key.clone(), value); + } + Ok(json!({"answers":answers})) + } +} +pub async fn run(client: Arc, sql: &str) -> Result> { + plan_and_run(SessionContext::new(), client, sql).await +} +/// Run `sql` against a single-column table `t(body)` built from `body`. +pub async fn run_table( + client: Arc, + sql: &str, + body: ArrayRef, +) -> Result> { + let ctx = SessionContext::new(); + let schema = Arc::new(Schema::new(vec![Field::new( + "body", + body.data_type().clone(), + true, + )])); + let batch = RecordBatch::try_new(schema.clone(), vec![body])?; + ctx.register_table("t", Arc::new(MemTable::try_new(schema, vec![vec![batch]])?))?; + plan_and_run(ctx, client, sql).await +} +pub async fn plan_and_run( + ctx: SessionContext, + client: Arc, + sql: &str, +) -> Result> { + datafusion_jev::register(&ctx, client); + datafusion_jev::sql(&ctx, sql).await?.collect().await +} +pub fn display(batches: &[RecordBatch]) -> String { + datafusion::arrow::util::pretty::pretty_format_batches(batches) + .unwrap() + .to_string() +} diff --git a/tests/sql/execution.rs b/tests/sql/execution.rs new file mode 100644 index 0000000..4edf85a --- /dev/null +++ b/tests/sql/execution.rs @@ -0,0 +1,448 @@ +//! Running queries against the mock provider: result shapes, batching, +//! deduplication, NULL handling, retries, cancellation, and concurrency. +use crate::common::{Mock, display, run, run_table}; +use async_trait::async_trait; +use datafusion::arrow::{ + array::{ + Array, ArrayRef, DictionaryArray, Float64Array, ListArray, RecordBatch, StringArray, + StructArray, + }, + datatypes::{Int32Type, UInt32Type}, +}; +use datafusion_jev::{JevClient, RequestError}; +use serde_json::{Value, json}; +use std::sync::Arc; + +#[tokio::test] +async fn exact_user_syntax_and_typed_fields() { + let mock = Arc::new(Mock::default()); + let result = run( + mock.clone(), + r#" + WITH classified AS ( + SELECT conversation_id, prompt_jev(transcript, 'Identify the customer''s main complaint', choice := [ + {label: 'billing', description: 'Payments, invoices, and refunds'}, + {label: 'technical', description: 'Errors, outages, and integrations'}, + {label: 'sales', description: 'Pricing and upgrades'}, + {label: 'account', description: 'Cancellations and account administration'} + ]) AS classification + FROM (VALUES (1, 'Please refund this payment')) t(conversation_id, transcript) + ) SELECT conversation_id, classification.choice, classification.confidence, classification.probabilities FROM classified + "#, + ) + .await + .unwrap(); + let text = display(&result); + assert!(text.contains("0.95"), "{text}"); + assert!(text.contains("account"), "{text}"); + assert_eq!( + mock.calls.lock().unwrap().len(), + 1, + "reading fields must not repeat inference" + ); +} +#[tokio::test] +async fn batching_and_null_input() { + let mock = Arc::new(Mock::default()); + let result = run( + mock.clone(), + "SELECT id, prompt_jev(body, 'Urgent?', batch_size := 2) AS p \ + FROM (VALUES (1,'a'), (2,NULL), (3,'b'), (4,'c')) t(id,body)", + ) + .await + .unwrap(); + assert_eq!(result.iter().map(|b| b.num_rows()).sum::(), 4); + let calls = mock.calls.lock().unwrap(); + assert_eq!(calls.len(), 2); + assert_eq!( + calls + .iter() + .map(|v| v["questions"].as_object().unwrap().len()) + .sum::(), + 3 + ); + assert!(display(&result).contains("0.9")); +} +#[tokio::test] +async fn score_and_noul_types() { + let result = run( + Arc::new(Mock::default()), + "SELECT prompt_jev('bad', 'Severity?', score := ['low','medium','high']) AS s, \ + prompt_jev('bad','Urgent?') AS p", + ) + .await + .unwrap(); + let text = display(&result); + assert!(text.contains("1.5"), "{text}"); + assert!(text.contains("0.9"), "{text}"); + assert!(text.contains("medium"), "{text}"); +} +#[tokio::test] +async fn unavailable_is_retried_then_null() { + let mock = Arc::new(Mock { + unavailable: true, + ..Default::default() + }); + let result = run(mock.clone(), "SELECT prompt_jev('a','q') AS p") + .await + .unwrap(); + assert!(result[0].column(0).is_null(0)); + assert_eq!(mock.calls.lock().unwrap().len(), 3); +} +#[tokio::test] +async fn provider_validation_and_malformed_answers_fail() { + for mock in [ + Mock { + fatal: true, + ..Default::default() + }, + Mock { + malformed: true, + ..Default::default() + }, + ] { + let mock = Arc::new(mock); + assert!( + run(mock.clone(), "SELECT prompt_jev('a','q') AS p") + .await + .is_err() + ); + assert_eq!(mock.calls.lock().unwrap().len(), 1); + } +} +#[tokio::test] +async fn nulls_and_empty_results_make_no_requests() { + let mock = Arc::new(Mock::default()); + run(mock.clone(), "SELECT prompt_jev(NULL,'q') AS p") + .await + .unwrap(); + run( + mock.clone(), + "SELECT prompt_jev(body,'q') AS p FROM (VALUES ('a')) t(body) WHERE false", + ) + .await + .unwrap(); + assert!(mock.calls.lock().unwrap().is_empty()); +} +#[tokio::test] +async fn semantic_filter_and_aggregate() { + let mock = Arc::new(Mock::default()); + let result = run( + mock, + "SELECT count(*) AS n FROM (VALUES ('a'),('b')) t(body) \ + WHERE prompt_jev(body,'Urgent?') > 0.8", + ) + .await + .unwrap(); + assert!(display(&result).contains("2")); +} +#[derive(Debug)] +struct Hanging { + started: Arc, + dropped: Arc, +} +struct DropSignal(Arc); +impl Drop for DropSignal { + fn drop(&mut self) { + self.0.notify_one(); + } +} +#[async_trait] +impl JevClient for Hanging { + async fn request(&self, _: Value) -> std::result::Result { + let _signal = DropSignal(self.dropped.clone()); + self.started.notify_one(); + std::future::pending().await + } +} +#[tokio::test] +async fn cancelling_query_drops_inflight_request() { + let started = Arc::new(tokio::sync::Notify::new()); + let dropped = Arc::new(tokio::sync::Notify::new()); + let client = Arc::new(Hanging { + started: started.clone(), + dropped: dropped.clone(), + }); + let task = tokio::spawn(async move { run(client, "SELECT prompt_jev('a','q') AS p").await }); + tokio::time::timeout(std::time::Duration::from_secs(5), started.notified()) + .await + .unwrap(); + task.abort(); + let _ = task.await; + tokio::time::timeout(std::time::Duration::from_secs(5), dropped.notified()) + .await + .unwrap(); +} +#[tokio::test] +async fn isolated_batches_keep_state_as_single_text() { + let mock = Arc::new(Mock::default()); + run( + mock.clone(), + "SELECT prompt_jev(body,'q',batch_size := 1) FROM (VALUES ('alpha'),('beta')) t(body)", + ) + .await + .unwrap(); + let calls = mock.calls.lock().unwrap(); + assert_eq!(calls.len(), 2); + assert!(calls.iter().all(|c| c["state"].is_string())); +} +#[derive(Debug, Default)] +struct Concurrent { + mock: Mock, + active: std::sync::atomic::AtomicUsize, + peak: std::sync::atomic::AtomicUsize, +} +#[async_trait] +impl JevClient for Concurrent { + async fn request(&self, body: Value) -> std::result::Result { + use std::sync::atomic::Ordering::SeqCst; + let active = self.active.fetch_add(1, SeqCst) + 1; + self.peak.fetch_max(active, SeqCst); + tokio::time::sleep(std::time::Duration::from_millis(10)).await; + let result = self.mock.request(body).await; + self.active.fetch_sub(1, SeqCst); + result + } +} +#[tokio::test] +async fn inference_concurrency_is_bounded() { + let mock = Arc::new(Concurrent::default()); + let values = (0..64) + .map(|i| format!("('row {i}')")) + .collect::>() + .join(","); + let sql = format!("SELECT prompt_jev(body,'q',batch_size := 1) FROM (VALUES {values}) t(body)"); + run(mock.clone(), &sql).await.unwrap(); + let peak = mock.peak.load(std::sync::atomic::Ordering::SeqCst); + assert!(peak > 1 && peak <= 8, "peak requests: {peak}"); + assert_eq!(mock.mock.calls.lock().unwrap().len(), 64); +} +#[tokio::test] +async fn identical_inputs_are_requested_once() { + let mock = Arc::new(Mock::default()); + let result = run( + mock.clone(), + "SELECT prompt_jev(body,'q') AS p FROM (VALUES ('a'),('a'),('b'),('a')) t(body)", + ) + .await + .unwrap(); + let calls = mock.calls.lock().unwrap(); + assert_eq!(calls.len(), 1, "repeated text must be asked about once"); + assert_eq!(calls[0]["questions"].as_object().unwrap().len(), 2); + let column = result[0].column_by_name("p").unwrap(); + assert_eq!(result.iter().map(|b| b.num_rows()).sum::(), 4); + assert_eq!(column.null_count(), 0, "every row keeps an answer"); +} +#[tokio::test] +async fn constant_input_is_requested_once() { + let mock = Arc::new(Mock::default()); + let result = run( + mock.clone(), + "SELECT prompt_jev('x','q') AS p FROM (VALUES (1),(2),(3)) t(id)", + ) + .await + .unwrap(); + assert_eq!(result.iter().map(|b| b.num_rows()).sum::(), 3); + let calls = mock.calls.lock().unwrap(); + assert_eq!(calls.len(), 1); + assert_eq!(calls[0]["questions"].as_object().unwrap().len(), 1); +} +fn many_labels(n: usize) -> String { + (0..n) + .map(|i| format!("'c{i}'")) + .collect::>() + .join(",") +} +#[tokio::test] +async fn probability_sum_tolerance_scales_with_criteria_count() { + // 200 criteria rounded to three decimals can drift far from one. + let mock = Arc::new(Mock { + probability: Some(0.004), + ..Default::default() + }); + let sql = format!( + "SELECT prompt_jev('a','q',choice := [{}]) AS p", + many_labels(200) + ); + run(mock, &sql).await.expect("rounding drift is tolerated"); + // Garbage is still rejected: 200 * 0.0075 = 1.5. + let mock = Arc::new(Mock { + probability: Some(0.0075), + ..Default::default() + }); + let sql = format!( + "SELECT prompt_jev('a','q',choice := [{}]) AS p", + many_labels(200) + ); + assert!(run(mock, &sql).await.is_err()); + // A small question keeps the tight tolerance. + let mock = Arc::new(Mock { + probability: Some(0.4), + ..Default::default() + }); + assert!( + run(mock, "SELECT prompt_jev('a','q',choice := ['x','y']) AS p") + .await + .is_err() + ); +} +#[tokio::test] +async fn oversized_encoded_input_is_rejected_without_batch_size_advice() { + let mock = Arc::new(Mock::default()); + // ~60k control characters: six bytes each once JSON-encoded. + let body = Arc::new(StringArray::from(vec!["\u{1}".repeat(60_000)])) as ArrayRef; + let error = run_table( + mock.clone(), + "SELECT prompt_jev(body,'q',batch_size := 1) AS p FROM t", + body, + ) + .await + .unwrap_err() + .to_string(); + assert!(error.contains("input exceeds"), "{error}"); + assert!(!error.contains("batch_size"), "{error}"); + assert!(mock.calls.lock().unwrap().is_empty()); +} +#[tokio::test] +async fn dictionary_encoded_input_is_accepted() { + let mock = Arc::new(Mock::default()); + let body = Arc::new( + vec!["alpha", "beta", "alpha"] + .into_iter() + .collect::>(), + ) as ArrayRef; + let result = run_table( + mock.clone(), + "SELECT prompt_jev(body,'q') AS p FROM t", + body, + ) + .await + .unwrap(); + assert_eq!(result.iter().map(|b| b.num_rows()).sum::(), 3); + assert_eq!(result[0].column_by_name("p").unwrap().null_count(), 0); + let calls = mock.calls.lock().unwrap(); + assert_eq!(calls.len(), 1); + assert_eq!(calls[0]["questions"].as_object().unwrap().len(), 2); +} +#[tokio::test] +async fn batches_split_on_encoded_size() { + let mock = Arc::new(Mock::default()); + // Three distinct ~20 KiB rows: the 32 KiB batch budget allows only one each. + let body = Arc::new(StringArray::from( + (0..3) + .map(|i| format!("{i}{}", "x".repeat(20_000))) + .collect::>(), + )) as ArrayRef; + run_table( + mock.clone(), + "SELECT prompt_jev(body,'q') AS p FROM t", + body, + ) + .await + .unwrap(); + let calls = mock.calls.lock().unwrap(); + assert_eq!( + calls.len(), + 3, + "default batch_size must still split by size" + ); + assert!(calls.iter().all(|c| c["state"].is_string())); +} +/// The struct field `name` of the `probabilities` list entries, as (value, probability). +fn probabilities(batch: &RecordBatch, name: &str) -> (Vec, Vec, Option>) { + let column = batch.column_by_name(name).unwrap(); + let outer = column.as_any().downcast_ref::().unwrap(); + let list = outer + .column_by_name("probabilities") + .unwrap() + .as_any() + .downcast_ref::() + .unwrap(); + let entries = list.value(0); + let entries = entries.as_any().downcast_ref::().unwrap(); + let values = entries + .column_by_name("value") + .unwrap() + .as_any() + .downcast_ref::() + .unwrap(); + let probs = entries + .column_by_name("probability") + .unwrap() + .as_any() + .downcast_ref::() + .unwrap(); + let index = entries.column_by_name("index").map(|c| { + c.as_any() + .downcast_ref::>() + .unwrap() + .values() + .to_vec() + }); + ( + values.iter().map(|v| v.unwrap().to_owned()).collect(), + probs.values().to_vec(), + index, + ) +} +#[tokio::test] +async fn probabilities_follow_caller_order() { + let result = run( + Arc::new(Mock::default()), + "SELECT prompt_jev('a','q',choice := ['billing','technical','sales']) AS c, \ + prompt_jev('a','q',score := ['low','medium','high']) AS s", + ) + .await + .unwrap(); + let (labels, probs, index) = probabilities(&result[0], "c"); + assert_eq!(labels, ["billing", "technical", "sales"]); + assert_eq!(probs, [1.0, 0.0, 0.0]); + assert_eq!(index, None, "choice carries no index"); + let (labels, probs, index) = probabilities(&result[0], "s"); + assert_eq!(labels, ["low", "medium", "high"]); + assert_eq!(probs, [0.25, 0.0, 0.75]); + assert_eq!(index, Some(vec![0, 1, 2])); +} +#[tokio::test] +async fn score_criteria_carry_their_labels() { + let mock = Arc::new(Mock::default()); + run( + mock.clone(), + "SELECT prompt_jev('a','q',score := [\ + {label: 'low', description: 'minor annoyance'}, \ + {label: 'high', description: 'outage'}]) AS s", + ) + .await + .unwrap(); + let calls = mock.calls.lock().unwrap(); + let criteria = calls[0]["questions"]["row_0"]["criteria"] + .as_array() + .unwrap(); + assert_eq!( + criteria, + &vec![json!("low: minor annoyance"), json!("high: outage")] + ); +} +#[derive(Debug, Default)] +struct HangsOnce { + mock: Mock, + calls: std::sync::atomic::AtomicUsize, +} +#[async_trait] +impl JevClient for HangsOnce { + async fn request(&self, body: Value) -> std::result::Result { + if self.calls.fetch_add(1, std::sync::atomic::Ordering::SeqCst) == 0 { + std::future::pending::<()>().await; + } + self.mock.request(body).await + } +} +#[tokio::test(start_paused = true)] +async fn timed_out_request_is_retried() { + let client = Arc::new(HangsOnce::default()); + let result = run(client.clone(), "SELECT prompt_jev('a','q') AS p") + .await + .unwrap(); + assert!(display(&result).contains("0.9")); + assert_eq!(client.calls.load(std::sync::atomic::Ordering::SeqCst), 2); +} diff --git a/tests/sql/main.rs b/tests/sql/main.rs new file mode 100644 index 0000000..65f0a56 --- /dev/null +++ b/tests/sql/main.rs @@ -0,0 +1,5 @@ +//! Integration tests for `prompt_jev`. +mod common; +mod execution; +mod planning; +mod syntax; diff --git a/tests/sql/planning.rs b/tests/sql/planning.rs new file mode 100644 index 0000000..d1c39d3 --- /dev/null +++ b/tests/sql/planning.rs @@ -0,0 +1,217 @@ +//! Plan-level behaviour: hoisting calls out of nodes DataFusion cannot plan +//! them in, keeping inference below cheap filters, and evaluating a repeated +//! call once. +use crate::common::{Mock, display, run}; +use datafusion::prelude::*; +use std::sync::Arc; + +#[tokio::test] +async fn repeated_expressions_share_one_request_and_keep_columns() { + let mock = Arc::new(Mock::default()); + let result = run( + mock.clone(), + "WITH c AS (SELECT id, \ + prompt_jev(body,'a',choice := ['alpha','beta']) AS x, \ + prompt_jev(body,'b') AS y, \ + prompt_jev(body,'a',choice := ['alpha','beta']) AS z, \ + body FROM (VALUES (1,'m'),(2,'n')) t(id,body)) \ + SELECT id, x.choice AS xc, y, z.choice AS zc, body FROM c", + ) + .await + .unwrap(); + assert_eq!( + mock.calls.lock().unwrap().len(), + 2, + "the repeated expression must not be evaluated twice" + ); + let names: Vec<_> = result[0] + .schema() + .fields() + .iter() + .map(|f| f.name().clone()) + .collect(); + assert_eq!(names, ["id", "xc", "y", "zc", "body"]); + let text = display(&result); + for expected in ["alpha", "0.9", "| m ", "| n "] { + assert!(text.contains(expected), "{expected} missing from {text}"); + } + let xc = result[0].column_by_name("xc").unwrap(); + let zc = result[0].column_by_name("zc").unwrap(); + assert_eq!(xc.as_ref(), zc.as_ref(), "both slots hold the same answer"); +} +#[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()); +} +#[tokio::test] +async fn cheap_conjuncts_filter_rows_before_inference() { + // `WHERE id = 1 AND prompt_jev(...) > 0.5` must not send rows 2 and 3. + let mock = Arc::new(Mock::default()); + let result = run( + mock.clone(), + "SELECT id FROM (VALUES (1,'keep'),(2,'drop'),(3,'drop')) t(id, body) \ + WHERE id = 1 AND prompt_jev(body, 'Urgent?') > 0.5", + ) + .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 passing the cheap predicate is asked about: {}", + calls[0] + ); +} +#[tokio::test] +async fn cheap_conjunct_in_a_count_filter_runs_first() { + let mock = Arc::new(Mock::default()); + let result = run( + mock.clone(), + "SELECT count(*) AS n FROM (VALUES ('a','x'),('b','y'),('c','z')) t(name, body) \ + WHERE name <> 'c' AND prompt_jev(body, 'Urgent?') > 0.5", + ) + .await + .unwrap(); + assert!(display(&result).contains('2'), "{}", display(&result)); + let calls = mock.calls.lock().unwrap(); + assert_eq!( + calls + .iter() + .map(|c| c["questions"].as_object().unwrap().len()) + .sum::(), + 2 + ); +} +#[tokio::test] +async fn a_folded_limit_survives_the_filter_split() { + // With one partition, LimitPushdown folds LIMIT into the FilterExec's fetch. + // The split must keep it, or the query returns every matching row. + let mock = Arc::new(Mock::default()); + let ctx = SessionContext::new_with_config(SessionConfig::new().with_target_partitions(1)); + datafusion_jev::register(&ctx, mock.clone()); + let df = datafusion_jev::sql( + &ctx, + "SELECT id FROM (VALUES (1,'a'),(2,'b'),(3,'c'),(4,'d')) t(id, body) \ + WHERE id > 1 AND prompt_jev(body, 'Urgent?') > 0.5 LIMIT 1", + ) + .await + .unwrap(); + let result = df.collect().await.unwrap(); + assert_eq!(result.iter().map(|b| b.num_rows()).sum::(), 1); + let asked: usize = mock + .calls + .lock() + .unwrap() + .iter() + .map(|c| c["questions"].as_object().unwrap().len()) + .sum(); + assert!( + asked <= 3, + "rows failing `id > 1` must not be asked about: {asked}" + ); +} diff --git a/tests/sql/syntax.rs b/tests/sql/syntax.rs new file mode 100644 index 0000000..722373d --- /dev/null +++ b/tests/sql/syntax.rs @@ -0,0 +1,126 @@ +//! Argument validation, the SQL rewrite itself, dialect handling, name +//! matching, and the plan-time size caps. +use crate::common::{Mock, display, run}; +use datafusion::prelude::*; +use std::sync::Arc; + +#[tokio::test] +async fn invalid_arguments_never_call_provider() { + for sql in [ + "SELECT prompt_jev('a','q',choice := ['a'])", + "SELECT prompt_jev('a','q',choice := ['a','a'])", + "SELECT prompt_jev('a','q',choice := ['a','b'], score := ['x','y'])", + "SELECT prompt_jev('a','q',batch_size := 0)", + "SELECT prompt_jev('a','q',batch_size := 65)", + "SELECT prompt_jev('a',body) FROM (VALUES ('q')) t(body)", + "SELECT prompt_jev('a','q',choice := body) FROM (VALUES ('q')) t(body)", + "SELECT prompt_jev('a','q',noul := ['yes','no'])", + "SELECT prompt_jev('a','q',model := 'foo')", + "SELECT prompt_jev(42,'q')", + "SELECT prompt_jev('a','q',noul := [])", + ] { + let mock = Arc::new(Mock::default()); + assert!(run(mock.clone(), sql).await.is_err(), "{sql}"); + assert!(mock.calls.lock().unwrap().is_empty()); + } +} +#[tokio::test] +async fn explain_does_not_run_inference() { + let mock = Arc::new(Mock::default()); + run(mock.clone(), "EXPLAIN SELECT prompt_jev('a','q') AS p") + .await + .unwrap(); + assert!(mock.calls.lock().unwrap().is_empty()); +} +#[tokio::test] +async fn drop_in_sql_handles_ddl_set_and_options() { + let mock = Arc::new(Mock::default()); + let ctx = SessionContext::new(); + datafusion_jev::register(&ctx, mock.clone()); + datafusion_jev::sql(&ctx, "CREATE TABLE t AS VALUES ('a'), ('b')") + .await + .unwrap() + .collect() + .await + .unwrap(); + datafusion_jev::sql(&ctx, "SET datafusion.execution.batch_size = 1024") + .await + .unwrap() + .collect() + .await + .unwrap(); + let result = datafusion_jev::sql(&ctx, "SELECT prompt_jev(column1, 'q') AS p FROM t") + .await + .unwrap() + .collect() + .await + .unwrap(); + assert_eq!(result.iter().map(|b| b.num_rows()).sum::(), 2); + assert_eq!(mock.calls.lock().unwrap().len(), 1); + let forbidden = datafusion_jev::sql_with_options( + &ctx, + "CREATE TABLE u AS VALUES (1)", + datafusion::execution::context::SQLOptions::new().with_allow_ddl(false), + ) + .await; + assert!(forbidden.is_err()); +} +#[tokio::test] +async fn strict_dialect_gets_a_clear_error() { + let ctx = SessionContext::new_with_config( + SessionConfig::new().set_str("datafusion.sql_parser.dialect", "postgresql"), + ); + datafusion_jev::register(&ctx, Arc::new(Mock::default())); + let err = datafusion_jev::sql(&ctx, "SELECT prompt_jev('a','q', choice := ['x','y'])") + .await + .err() + .unwrap() + .to_string(); + assert!(err.contains("generic (default) or duckdb"), "{err}"); + assert!(err.contains("postgresql"), "{err}"); + let plain = datafusion_jev::sql(&ctx, "SELECT 1 +") + .await + .err() + .unwrap() + .to_string(); + assert!(!plain.contains("prompt_jev"), "{plain}"); +} +#[tokio::test] +async fn quoted_and_uppercase_names_are_the_same_function() { + let mock = Arc::new(Mock::default()); + let result = run( + mock.clone(), + "SELECT \"prompt_jev\"('a', 'q') AS a, PROMPT_JEV('a', 'q') AS b, \"PROMPT_JEV\"('a', 'q') AS c", + ) + .await + .unwrap(); + assert_eq!(result.iter().map(|b| b.num_rows()).sum::(), 1); + assert!(display(&result).contains("0.9")); +} +#[tokio::test] +async fn oversized_instructions_and_labels_fail_at_plan_time() { + let long_q = "q".repeat(5_000); + let long_label = "L".repeat(300); + let long_desc = "d".repeat(2_000); + for (sql, needle) in [ + ( + format!("SELECT prompt_jev('a', '{long_q}')"), + "4000 characters", + ), + ( + format!("SELECT prompt_jev('a', 'q', choice := ['{long_label}', 'b'])"), + "labels are limited", + ), + ( + format!( + "SELECT prompt_jev('a', 'q', choice := [{{label: 'a', description: '{long_desc}'}}, 'b'])" + ), + "descriptions to 1024", + ), + ] { + let mock = Arc::new(Mock::default()); + let err = run(mock.clone(), &sql).await.err().unwrap().to_string(); + assert!(err.contains(needle), "{needle}: {err}"); + assert!(mock.calls.lock().unwrap().is_empty()); + } +} From 2e74b9752c93b07da815a3d09cfb119d89deb3f8 Mon Sep 17 00:00:00 2001 From: Eddie A Tejeda <669988+eddietejeda@users.noreply.github.com> Date: Thu, 24 Sep 2026 09:05:08 -0700 Subject: [PATCH 2/2] refactor: build the rewritten argument list directly --- src/sql.rs | 24 ++++++++++++------------ 1 file changed, 12 insertions(+), 12 deletions(-) diff --git a/src/sql.rs b/src/sql.rs index 0fdf3d9..a5a80c2 100644 --- a/src/sql.rs +++ b/src/sql.rs @@ -3,7 +3,8 @@ use crate::names::INTERNAL_FUNCTION; use datafusion::{ common::{Result, plan_datafusion_err}, sql::sqlparser::ast::{ - self, Expr, FunctionArg, FunctionArgExpr, FunctionArguments, Ident, ObjectName, Value, + self, Expr, FunctionArg, FunctionArgExpr, FunctionArgumentList, FunctionArguments, Ident, + ObjectName, Value, }, }; use serde::{Deserialize, Serialize}; @@ -267,18 +268,17 @@ pub fn rewrite_statement(statement: &mut ast::Statement) -> Result<()> { let (input, q) = parse_call(f)?; q.validate()?; let config = serde_json::to_string(&q).map_err(|e| plan_datafusion_err!("{e}"))?; - let FunctionArguments::List(args) = &mut f.args else { - return Err(plan_datafusion_err!( - "prompt_jev requires input and instructions" - )); - }; f.name = ObjectName::from(vec![Ident::new(INTERNAL_FUNCTION)]); - args.args = vec![ - FunctionArg::Unnamed(FunctionArgExpr::Expr(input)), - FunctionArg::Unnamed(FunctionArgExpr::Expr(Expr::Value( - Value::SingleQuotedString(config).into(), - ))), - ]; + f.args = FunctionArguments::List(FunctionArgumentList { + duplicate_treatment: None, + args: vec![ + FunctionArg::Unnamed(FunctionArgExpr::Expr(input)), + FunctionArg::Unnamed(FunctionArgExpr::Expr(Expr::Value( + Value::SingleQuotedString(config).into(), + ))), + ], + clauses: vec![], + }); Ok(()) })(); match result {