use datafusion::arrow::array::{Array, Float64Array};
use datafusion::arrow::datatypes::DataType;
use datafusion::error::{DataFusionError, Result};
use datafusion::prelude::SessionContext;
use ddx_core::test_utils::{
close_at_scale, eval, gen_adversarial_sql, gen_expr, max_intermediate_mag, seeded,
Conditioning, Failures, Rng, Var,
};
use ddx_core::{ColRef, Ddx};
const ROWS: usize = 16;
const SEEDS: u64 = 120;
const RTOL: f64 = 1e-9;
const ATOL: f64 = 1e-12;
struct Sim {
ddx: SessionContext,
plain: SessionContext,
names: SessionContext,
pts: Vec<(f64, f64)>,
ipts: Vec<(f64, f64)>,
}
impl Sim {
async fn new() -> Result<Sim> {
let mut rng = Rng::new(0x5EED_D47A);
let pts: Vec<(f64, f64)> = (0..ROWS)
.map(|_| (rng.range(0.2, 1.8), rng.range(0.2, 1.8)))
.collect();
let ipts: Vec<(f64, f64)> = (0..ROWS)
.map(|k| (1.0 + (k % 4) as f64, 1.0 + (k / 4) as f64))
.collect();
let ddx = SessionContext::new();
ddx_datafusion::install(&ddx);
let plain = SessionContext::new();
let names = SessionContext::new();
names.register_udf(ddx_datafusion::grad_udf());
names.register_udf(ddx_datafusion::jvp_udf());
for ctx in [&ddx, &plain, &names] {
seed_tables(ctx, &pts, &ipts).await?;
}
Ok(Sim {
ddx,
plain,
names,
pts,
ipts,
})
}
}
async fn seed_tables(ctx: &SessionContext, pts: &[(f64, f64)], ipts: &[(f64, f64)]) -> Result<()> {
let values = |ps: &[(f64, f64)]| {
ps.iter()
.enumerate()
.map(|(i, (x, y))| format!("({i}, {x:?}, {y:?})"))
.collect::<Vec<_>>()
.join(", ")
};
for ddl in [
format!(
"CREATE TABLE t AS SELECT * FROM (VALUES {}) AS v(i, x, y)",
values(pts)
),
"CREATE TABLE t_upper AS SELECT i, x AS \"X\", y AS \"Y\" FROM t".into(),
"CREATE TABLE t_kw AS SELECT i, x AS \"order\", y AS \"select\" FROM t".into(),
"CREATE TABLE t_space AS SELECT i, x AS \"my x\", y AS \"my y\" FROM t".into(),
format!(
"CREATE TABLE ti_double AS SELECT * FROM (VALUES {}) AS v(i, x, y)",
values(ipts)
),
"CREATE TABLE ti_bigint AS SELECT i, CAST(x AS BIGINT) AS x, CAST(y AS BIGINT) AS y \
FROM ti_double"
.into(),
"CREATE TABLE ti_decimal AS SELECT i, CAST(x AS DECIMAL(20, 6)) AS x, \
CAST(y AS DECIMAL(20, 6)) AS y FROM ti_double"
.into(),
] {
ctx.sql(&ddl).await?.collect().await?;
}
Ok(())
}
async fn d_column(ctx: &SessionContext, sql: &str) -> Result<(Vec<Option<f64>>, DataType)> {
let batches = ctx.sql(sql).await?.collect().await?;
let Some(first) = batches.first() else {
return Err(DataFusionError::Execution(format!(
"query returned no batches: {sql}"
)));
};
let idx = first.schema().index_of("d")?;
let ty = first.schema().field(idx).data_type().clone();
if ty != DataType::Float64 {
return Ok((Vec::new(), ty));
}
let mut out = Vec::new();
for b in &batches {
let a = b
.column(idx)
.as_any()
.downcast_ref::<Float64Array>()
.ok_or_else(|| {
DataFusionError::Execution(format!("the `d` column changed type mid-stream: {sql}"))
})?;
out.extend((0..a.len()).map(|i| (!a.is_null(i)).then(|| a.value(i))));
}
Ok((out, ty))
}
fn compare_rows(
pts: &[(f64, f64)],
gate: &[&ddx_core::sqlparser::ast::Expr],
lhs: &[Option<f64>],
rhs: &[Option<f64>],
) -> (u32, Option<String>) {
let cond = Conditioning::default();
let mut compared = 0u32;
if lhs.len() != rhs.len() {
return (
0,
Some(format!("row counts differ: {} vs {}", lhs.len(), rhs.len())),
);
}
for (row, &(x, y)) in pts.iter().enumerate().take(lhs.len()) {
if !cond.admits(gate, x, y) {
continue;
}
let scale = gate
.iter()
.filter_map(|e| max_intermediate_mag(e, x, y))
.fold(1.0f64, f64::max);
match (lhs[row], rhs[row]) {
(None, None) => compared += 1,
(Some(a), Some(b)) if a.is_nan() && b.is_nan() => compared += 1,
(Some(a), Some(b)) => {
compared += 1;
if !close_at_scale(a, b, scale, RTOL, ATOL) {
return (
compared,
Some(format!(
"row {row} (x={x:.6}, y={y:.6}): {a} vs {b} \
(scale {scale:.3e}, allowed {:.3e})",
ATOL + RTOL * scale.max(1.0)
)),
);
}
}
(a, b) => {
return (
compared,
Some(format!(
"row {row} (x={x:.6}, y={y:.6}): NULL mismatch — {a:?} vs {b:?}"
)),
)
}
}
}
(compared, None)
}
struct Case {
primal: String,
wrt: Var,
d: ddx_core::sqlparser::ast::Expr,
f: ddx_core::sqlparser::ast::Expr,
}
impl Case {
fn marker(&self) -> String {
format!("grad({}, {})", self.primal, self.wrt.name())
}
}
fn gen_case(rng: &mut Rng, ddx: &Ddx, depth: u32) -> Option<Case> {
let primal = gen_expr(rng, depth);
let wrt = if rng.below(2) == 0 { Var::X } else { Var::Y };
let f = ddx_core::test_utils::try_parse(&primal).ok()?;
let d = ddx.differentiate(&f, &ColRef::bare(wrt.name())).ok()?;
Some(Case { primal, wrt, d, f })
}
fn rename_vars(text: &str, xname: &str, yname: &str) -> String {
let b = text.as_bytes();
let mut out = String::with_capacity(text.len());
let mut i = 0;
while i < b.len() {
if b[i].is_ascii_alphabetic() || b[i] == b'_' {
let start = i;
while i < b.len() && (b[i].is_ascii_alphanumeric() || b[i] == b'_') {
i += 1;
}
match &text[start..i] {
"x" => out.push_str(xname),
"y" => out.push_str(yname),
word => out.push_str(word),
}
} else {
out.push(b[i] as char);
i += 1;
}
}
out
}
#[tokio::test]
async fn path_a_and_path_b_agree_on_random_expressions() -> Result<()> {
let sim = Sim::new().await?;
let ddx = Ddx::for_datafusion();
let mut fail = Failures::new();
for seed in 0..SEEDS {
let mut rng = seeded(seed, 0x9A7B_0001);
let Some(case) = gen_case(&mut rng, &ddx, 2 + (seed % 3) as u32) else {
continue;
};
let sql = format!("SELECT i, {} AS d FROM t ORDER BY i", case.marker());
let via_b = d_column(&sim.ddx, &sql).await;
let via_a = match ddx_datafusion::ddx_sql(&sim.plain, &sql).await {
Ok(df) => df.collect().await.and_then(|batches| {
let idx = batches[0].schema().index_of("d")?;
let mut out = Vec::new();
for b in &batches {
let a = b
.column(idx)
.as_any()
.downcast_ref::<Float64Array>()
.ok_or_else(|| {
DataFusionError::Execution("path A `d` is not Float64".into())
})?;
out.extend((0..a.len()).map(|i| (!a.is_null(i)).then(|| a.value(i))));
}
Ok(out)
}),
Err(e) => Err(e),
};
match (via_a, via_b) {
(Ok(a), Ok((b, ty))) => {
if ty != DataType::Float64 {
fail.push(seed, format!("Path B returned {ty:?}, not Float64: {sql}"));
continue;
}
let (compared, bad) = compare_rows(&sim.pts, &[&case.f, &case.d], &a, &b);
if let Some(bad) = bad {
fail.push(seed, format!("paths disagree on `{sql}`\n {bad}"));
} else if compared > 0 {
fail.tested();
}
}
(Err(a), Err(b)) => {
let _ = (a, b);
}
(Ok(_), Err(e)) => fail.push(
seed,
format!("Path A ran but Path B failed on `{sql}`: {e}"),
),
(Err(e), Ok(_)) => fail.push(
seed,
format!("Path B ran but Path A failed on `{sql}`: {e}"),
),
}
}
fail.assert_clean("path A vs path B", 40);
Ok(())
}
#[tokio::test]
async fn the_engine_agrees_with_the_reference_interpreter() -> Result<()> {
let sim = Sim::new().await?;
let ddx = Ddx::for_datafusion();
let mut fail = Failures::new();
for seed in 0..SEEDS {
let mut rng = seeded(seed, 0x9A7B_0002);
let Some(case) = gen_case(&mut rng, &ddx, 2 + (seed % 3) as u32) else {
continue;
};
let sql = format!("SELECT i, {} AS d FROM t ORDER BY i", case.marker());
let Ok((engine, _)) = d_column(&sim.ddx, &sql).await else {
continue;
};
let reference: Vec<Option<f64>> =
sim.pts.iter().map(|&(x, y)| eval(&case.d, x, y)).collect();
let (compared, bad) = compare_rows(&sim.pts, &[&case.f, &case.d], &reference, &engine);
if let Some(bad) = bad {
fail.push(
seed,
format!(
"engine disagrees with the reference interpreter\n primal = {}\n \
d/d{} = {}\n {bad}",
case.primal,
case.wrt.name(),
case.d
),
);
} else if compared > 0 {
fail.tested();
}
}
fail.assert_clean("engine vs reference interpreter", 40);
Ok(())
}
#[tokio::test]
async fn a_derivative_does_not_depend_on_the_column_name() -> Result<()> {
let sim = Sim::new().await?;
let ddx = Ddx::for_datafusion();
let mut fail = Failures::new();
let renamings: &[(&str, &str, &str)] = &[
("t_upper", "\"X\"", "\"Y\""),
("t_kw", "\"order\"", "\"select\""),
("t_space", "\"my x\"", "\"my y\""),
];
for seed in 0..SEEDS {
let mut rng = seeded(seed, 0x9A7B_0003);
let Some(case) = gen_case(&mut rng, &ddx, 2 + (seed % 3) as u32) else {
continue;
};
let base_sql = format!("SELECT i, {} AS d FROM t ORDER BY i", case.marker());
let Ok((base, _)) = d_column(&sim.ddx, &base_sql).await else {
continue;
};
let mut counted = false;
for &(table, xn, yn) in renamings {
let renamed = format!(
"SELECT i, grad({}, {}) AS d FROM {table} ORDER BY i",
rename_vars(&case.primal, xn, yn),
match case.wrt {
Var::X => xn,
Var::Y => yn,
}
);
match d_column(&sim.ddx, &renamed).await {
Err(e) => fail.push(
seed,
format!("renaming to {table} broke a working query: {renamed}\n {e}"),
),
Ok((got, ty)) => {
if ty != DataType::Float64 {
fail.push(seed, format!("{table} returned {ty:?}, not Float64"));
continue;
}
let (compared, bad) = compare_rows(&sim.pts, &[&case.f, &case.d], &base, &got);
if let Some(bad) = bad {
fail.push(
seed,
format!(
"renaming the columns changed the derivative\n \
base = {base_sql}\n renamed = {renamed}\n {bad}"
),
);
} else if compared > 0 && !counted {
counted = true;
fail.tested();
}
}
}
}
}
fail.assert_clean("column-rename invariance", 40);
Ok(())
}
fn agreeing_rows(
pts: &[(f64, f64)],
gate: &[&ddx_core::sqlparser::ast::Expr],
a: &[Option<f64>],
b: &[Option<f64>],
) -> Vec<usize> {
let cond = Conditioning::default();
let mut keep = Vec::new();
for (row, &(x, y)) in pts.iter().enumerate() {
if row >= a.len() || row >= b.len() || !cond.admits(gate, x, y) {
continue;
}
let scale = gate
.iter()
.filter_map(|e| max_intermediate_mag(e, x, y))
.fold(1.0f64, f64::max);
let same = match (a[row], b[row]) {
(None, None) => true,
(Some(p), Some(q)) if p.is_nan() && q.is_nan() => true,
(Some(p), Some(q)) => close_at_scale(p, q, scale, RTOL, ATOL),
_ => false,
};
if same {
keep.push(row);
}
}
keep
}
#[tokio::test]
async fn a_derivative_does_not_depend_on_the_column_storage_type() -> Result<()> {
let sim = Sim::new().await?;
let ddx = Ddx::for_datafusion();
let mut fail = Failures::new();
for seed in 0..SEEDS {
let mut rng = seeded(seed, 0x9A7B_0004);
let Some(case) = gen_case(&mut rng, &ddx, 2 + (seed % 3) as u32) else {
continue;
};
let grad_query =
|table: &str| format!("SELECT i, {} AS d FROM {table} ORDER BY i", case.marker());
let primal_query = |table: &str| {
format!(
"SELECT i, CAST(({}) AS DOUBLE) AS d FROM {table} ORDER BY i",
case.primal
)
};
let (Ok((base, base_ty)), Ok((base_primal, _))) = (
d_column(&sim.ddx, &grad_query("ti_double")).await,
d_column(&sim.ddx, &primal_query("ti_double")).await,
) else {
continue;
};
if base_ty != DataType::Float64 {
fail.push(seed, format!("ti_double returned {base_ty:?}"));
continue;
}
let mut counted = false;
for (table, compare_values) in [("ti_bigint", true), ("ti_decimal", false)] {
let Ok((primal, _)) = d_column(&sim.ddx, &primal_query(table)).await else {
continue;
};
let keep = agreeing_rows(&sim.ipts, &[&case.f], &base_primal, &primal);
if keep.is_empty() {
continue;
}
match d_column(&sim.ddx, &grad_query(table)).await {
Err(e) => fail.push(
seed,
format!(
"{table} failed a query that works on ti_double: {}\n {e}",
grad_query(table)
),
),
Ok((got, ty)) => {
if ty != DataType::Float64 {
fail.push(
seed,
format!(
"{table} returned {ty:?}, not Float64 — every derivative is \
DOUBLE: {}",
grad_query(table)
),
);
continue;
}
if !compare_values {
if !counted {
counted = true;
fail.tested();
}
continue;
}
let pts: Vec<(f64, f64)> = keep.iter().map(|&r| sim.ipts[r]).collect();
let pick = |v: &[Option<f64>]| -> Vec<Option<f64>> {
keep.iter().map(|&r| v.get(r).copied().flatten()).collect()
};
let (compared, bad) =
compare_rows(&pts, &[&case.f, &case.d], &pick(&base), &pick(&got));
if let Some(bad) = bad {
fail.push(
seed,
format!(
"storage type changed the derivative at a row where the \
primal did NOT change (ti_double vs {table})\n {}\n {bad}",
grad_query(table)
),
);
} else if compared > 0 && !counted {
counted = true;
fail.tested();
}
}
}
}
}
fail.assert_clean("storage-type invariance", 20);
Ok(())
}
fn value_placements(marker: &str) -> Vec<(&'static str, String)> {
vec![
(
"nested subquery",
format!("SELECT i, d FROM (SELECT i, {marker} AS d FROM t) ORDER BY i"),
),
(
"case arm",
format!(
"SELECT i, CASE WHEN i >= 0 THEN {marker} ELSE NULL END AS d FROM t ORDER BY i"
),
),
(
"group-by key",
format!("SELECT i, {marker} AS d FROM t GROUP BY i, {marker} ORDER BY i"),
),
(
"window over a unique partition",
format!("SELECT i, max({marker}) OVER (PARTITION BY i) AS d FROM t ORDER BY i"),
),
(
"aliased DISTINCT ON",
format!("SELECT i, d FROM (SELECT DISTINCT ON (i) i, {marker} AS d FROM t) ORDER BY i"),
),
(
"join key passthrough",
format!(
"SELECT a.i AS i, a.d AS d FROM (SELECT i, {marker} AS d FROM t) a \
JOIN t b ON a.i = b.i ORDER BY a.i"
),
),
]
}
#[tokio::test]
async fn marker_values_do_not_depend_on_plan_placement() -> Result<()> {
let sim = Sim::new().await?;
let ddx = Ddx::for_datafusion();
let mut fail = Failures::new();
for seed in 0..SEEDS {
let mut rng = seeded(seed, 0x9A7B_0005);
let Some(case) = gen_case(&mut rng, &ddx, 2 + (seed % 2) as u32) else {
continue;
};
let marker = case.marker();
let base_sql = format!("SELECT i, {marker} AS d FROM t ORDER BY i");
let Ok((base, _)) = d_column(&sim.ddx, &base_sql).await else {
continue;
};
let mut counted = false;
for (label, sql) in value_placements(&marker) {
match d_column(&sim.ddx, &sql).await {
Err(e) => fail.push(
seed,
format!("placement `{label}` failed a query the bare projection runs:\n {sql}\n {e}"),
),
Ok((got, ty)) => {
if ty != DataType::Float64 {
fail.push(seed, format!("placement `{label}` returned {ty:?}: {sql}"));
continue;
}
let (compared, bad) =
compare_rows(&sim.pts, &[&case.f, &case.d], &base, &got);
if let Some(bad) = bad {
fail.push(
seed,
format!("placement `{label}` changed the value:\n {sql}\n {bad}"),
);
} else if compared > 0 && !counted {
counted = true;
fail.tested();
}
}
}
}
}
fail.assert_clean("plan-placement value invariance", 40);
Ok(())
}
fn all_placements(marker: &str) -> Vec<(&'static str, String)> {
let mut v: Vec<(&'static str, String)> = vec![
("projection", format!("SELECT i, {marker} AS d FROM t")),
("where", format!("SELECT i FROM t WHERE {marker} > 0")),
(
"having",
format!("SELECT i FROM t GROUP BY i HAVING sum({marker}) > 0"),
),
("order by", format!("SELECT i FROM t ORDER BY {marker}")),
("distinct", format!("SELECT DISTINCT {marker} AS d FROM t")),
(
"distinct on",
format!("SELECT DISTINCT ON (i) {marker} AS d FROM t"),
),
(
"union branch",
format!("SELECT {marker} AS d FROM t UNION ALL SELECT y AS d FROM t"),
),
(
"scalar subquery",
format!("SELECT i FROM t WHERE x > (SELECT avg({marker}) FROM t)"),
),
(
"IN subquery",
format!("SELECT i FROM t WHERE x IN (SELECT {marker} FROM t)"),
),
(
"EXISTS subquery",
format!("SELECT i FROM t WHERE EXISTS (SELECT 1 FROM t WHERE {marker} > 0)"),
),
];
for op in ["ALL", "ANY", "SOME"] {
v.push((
match op {
"ALL" => "ALL subquery",
"ANY" => "ANY subquery",
_ => "SOME subquery",
},
format!("SELECT i FROM t WHERE x > {op} (SELECT {marker} FROM t)"),
));
}
v
}
#[tokio::test]
async fn an_installed_marker_never_reaches_execution() -> Result<()> {
let sim = Sim::new().await?;
let ddx = Ddx::for_datafusion();
let mut fail = Failures::new();
for seed in 0..SEEDS {
let mut rng = seeded(seed, 0x9A7B_0006);
let Some(case) = gen_case(&mut rng, &ddx, 1 + (seed % 2) as u32) else {
continue;
};
let mut counted = false;
for (label, sql) in all_placements(&case.marker()) {
let err = match sim.ddx.sql(&sql).await {
Err(e) => Some(e),
Ok(df) => df.collect().await.err(),
};
if !counted {
counted = true;
fail.tested();
}
if let Some(e) = err {
let msg = e.to_string();
if msg.contains("reached execution") {
fail.push(
seed,
format!(
"a marker survived the rewrite in placement `{label}`:\n {sql}\n \
the rewrite never reached it, so it ran as a row function and \
errored"
),
);
}
}
}
}
fail.assert_clean("marker reachability", 40);
Ok(())
}
fn naming_shapes(marker: &str) -> Vec<(&'static str, String)> {
vec![
("projection", format!("SELECT {marker} FROM t")),
("aggregate", format!("SELECT avg({marker}) FROM t")),
(
"window",
format!("SELECT sum({marker}) OVER (ORDER BY i) FROM t"),
),
(
"distinct on",
format!("SELECT DISTINCT ON (i) {marker} FROM t"),
),
("distinct", format!("SELECT DISTINCT {marker} FROM t")),
]
}
#[tokio::test]
async fn marker_columns_keep_their_name_across_plan_shapes() -> Result<()> {
let sim = Sim::new().await?;
let ddx = Ddx::for_datafusion();
let mut fail = Failures::new();
for seed in 0..SEEDS {
let mut rng = seeded(seed, 0x9A7B_0007);
let Some(case) = gen_case(&mut rng, &ddx, 1 + (seed % 2) as u32) else {
continue;
};
let mut counted = false;
for (label, sql) in naming_shapes(&case.marker()) {
let Ok(want) = sim.names.sql(&sql).await else {
continue;
};
let want = want.schema().field(0).name().clone();
let Ok(df) = sim.ddx.sql(&sql).await else {
continue;
};
let Ok(batches) = df.collect().await else {
continue;
};
let Some(first) = batches.first() else {
continue;
};
let got = first.schema().field(0).name().clone();
if !counted {
counted = true;
fail.tested();
}
if got != want {
fail.push(
seed,
format!(
"the rewrite renamed the output field in placement `{label}`:\n \
{sql}\n expected `{want}`\n got `{got}`"
),
);
}
}
}
fail.assert_clean("field-name preservation", 40);
Ok(())
}
#[tokio::test]
async fn the_rewrite_leaves_non_marker_columns_untouched() -> Result<()> {
let sim = Sim::new().await?;
let ddx = Ddx::for_datafusion();
let mut fail = Failures::new();
for seed in 0..SEEDS {
let mut rng = seeded(seed, 0x9A7B_0008);
let Some(case) = gen_case(&mut rng, &ddx, 2 + (seed % 3) as u32) else {
continue;
};
let bystander = "(x * 3.0 - y / 7.0)";
let with_marker = format!(
"SELECT i, {bystander} AS d, {} AS g FROM t WHERE y > 0.0 ORDER BY i",
case.marker()
);
let without = format!("SELECT i, {bystander} AS d FROM t WHERE y > 0.0 ORDER BY i");
let Ok((got, got_ty)) = d_column(&sim.ddx, &with_marker).await else {
continue;
};
let (want, want_ty) = d_column(&sim.plain, &without).await?;
fail.tested();
if got_ty != want_ty {
fail.push(
seed,
format!("the bystander column's type changed: {want_ty:?} -> {got_ty:?}"),
);
}
if got != want {
fail.push(
seed,
format!(
"a column that contains no marker changed when a marker was added \
elsewhere in the same query:\n with = {with_marker}\n \
without = {without}\n {want:?}\n {got:?}"
),
);
}
}
fail.assert_clean("bystander-column invariance", 40);
Ok(())
}
#[tokio::test]
async fn jvp_equals_the_tangent_times_grad_on_the_engine() -> Result<()> {
let sim = Sim::new().await?;
let ddx = Ddx::for_datafusion();
let mut fail = Failures::new();
for seed in 0..SEEDS {
let mut rng = seeded(seed, 0x9A7B_0009);
let Some(case) = gen_case(&mut rng, &ddx, 2 + (seed % 2) as u32) else {
continue;
};
let tangent_depth = 1 + rng.below(2) as u32;
let tangent = gen_expr(&mut rng, tangent_depth);
let Ok(tangent_expr) = ddx_core::test_utils::try_parse(&tangent) else {
continue;
};
let jvp_sql = format!(
"SELECT i, jvp({}, {}, {tangent}) AS d FROM t ORDER BY i",
case.primal,
case.wrt.name()
);
let scaled_sql = format!(
"SELECT i, ({tangent}) * {} AS d FROM t ORDER BY i",
case.marker()
);
let (Ok((jvp, jvp_ty)), Ok((scaled, _))) = (
d_column(&sim.ddx, &jvp_sql).await,
d_column(&sim.ddx, &scaled_sql).await,
) else {
continue;
};
if jvp_ty != DataType::Float64 {
fail.push(seed, format!("jvp returned {jvp_ty:?}: {jvp_sql}"));
continue;
}
let (compared, bad) =
compare_rows(&sim.pts, &[&case.f, &case.d, &tangent_expr], &jvp, &scaled);
if let Some(bad) = bad {
fail.push(
seed,
format!("jvp ≠ tangent · grad:\n {jvp_sql}\n {scaled_sql}\n {bad}"),
);
} else if compared > 0 {
fail.tested();
}
}
fail.assert_clean("jvp = tangent · grad", 40);
Ok(())
}
#[tokio::test]
async fn nesting_a_marker_is_repeated_differentiation() -> Result<()> {
let sim = Sim::new().await?;
let ddx = Ddx::for_datafusion();
let mut fail = Failures::new();
for seed in 0..SEEDS {
let mut rng = seeded(seed, 0x9A7B_000A);
let Some(case) = gen_case(&mut rng, &ddx, 2 + (seed % 2) as u32) else {
continue;
};
let Ok(dd) = ddx.differentiate(&case.d, &ColRef::bare(case.wrt.name())) else {
continue;
};
let nested_sql = format!(
"SELECT i, grad({}, {w}) AS d FROM t ORDER BY i",
case.marker(),
w = case.wrt.name()
);
let direct_sql = format!("SELECT i, ({dd}) AS d FROM t ORDER BY i");
let (Ok((nested, ty)), Ok((direct, _))) = (
d_column(&sim.ddx, &nested_sql).await,
d_column(&sim.plain, &direct_sql).await,
) else {
continue;
};
if ty != DataType::Float64 {
fail.push(seed, format!("nested grad returned {ty:?}: {nested_sql}"));
continue;
}
let (compared, bad) = compare_rows(&sim.pts, &[&case.f, &case.d, &dd], &nested, &direct);
if let Some(bad) = bad {
fail.push(
seed,
format!(
"grad(grad(f)) ≠ the twice-differentiated expression:\n {nested_sql}\n \
{direct_sql}\n {bad}"
),
);
} else if compared > 0 {
fail.tested();
}
}
fail.assert_clean("higher-order nesting", 30);
Ok(())
}
#[test]
fn the_analyzer_never_panics_on_adversarial_sql() -> Result<()> {
let rt = tokio::runtime::Builder::new_current_thread()
.enable_all()
.build()
.expect("a current-thread runtime");
let sim = rt.block_on(Sim::new())?;
let mut fail = Failures::new();
for seed in 0..SEEDS * 4 {
let mut rng = seeded(seed, 0x9A7B_000B);
let sql = gen_adversarial_sql(&mut rng);
fail.tested();
let outcome = std::panic::catch_unwind(std::panic::AssertUnwindSafe(|| {
rt.block_on(async {
match sim.ddx.sql(&sql).await {
Ok(df) => df.collect().await.is_ok(),
Err(_) => false,
}
})
}));
if let Err(payload) = outcome {
let msg = payload
.downcast_ref::<&str>()
.map(|s| s.to_string())
.or_else(|| payload.downcast_ref::<String>().cloned())
.unwrap_or_else(|| "<non-string panic>".to_string());
fail.push(
seed,
format!("PANICKED (must be a typed DataFusionError) on {sql:?}\n panic = {msg}"),
);
}
}
fail.assert_clean("analyzer never-panic", 100);
Ok(())
}