use datafusion::error::Result;
use datafusion::prelude::SessionContext;
use ddx_datafusion::{ddx_sql, rewrite_sql};
mod common;
async fn ctx() -> Result<SessionContext> {
let ctx = SessionContext::new();
ctx.sql(
"CREATE TABLE t AS
SELECT * FROM (VALUES (1.0, 4.0), (2.0, 5.0), (3.0, 6.0)) AS v(x, y)",
)
.await?
.collect()
.await?;
Ok(ctx)
}
async fn col(ctx: &SessionContext, sql: &str) -> Result<Vec<f64>> {
let batches = ddx_sql(ctx, sql).await?.collect().await?;
Ok(common::f64_column(&batches))
}
#[tokio::test]
async fn grad_runs_on_a_context_with_no_ddx_setup() -> Result<()> {
let ctx = ctx().await?;
assert_eq!(
col(&ctx, "SELECT grad(x * x, x) AS d FROM t ORDER BY x").await?,
vec![2.0, 4.0, 6.0],
);
Ok(())
}
#[tokio::test]
async fn agrees_with_path_b_on_the_same_query() -> Result<()> {
let sql = "SELECT grad(sin(x * y), y) AS d FROM t ORDER BY x";
let a_ctx = ctx().await?;
let via_a = col(&a_ctx, sql).await?;
let b_ctx = ctx().await?;
ddx_datafusion::install(&b_ctx);
let batches = b_ctx.sql(sql).await?.collect().await?;
let via_b = common::f64_column(&batches);
assert_eq!(via_a.len(), via_b.len());
for (a, b) in via_a.iter().zip(&via_b) {
assert!((a - b).abs() < 1e-12, "Path A gave {a}, Path B gave {b}");
}
Ok(())
}
#[tokio::test]
async fn newton_iteration_in_a_recursive_cte() -> Result<()> {
let ctx = ctx().await?;
let sql = "
WITH RECURSIVE newton AS (
SELECT 0 AS i, 1.0 AS x
UNION ALL
SELECT i + 1 AS i, x - (x * x - 2.0) / grad(x * x - 2.0, x) AS x
FROM newton WHERE i < 6
)
SELECT x FROM newton WHERE i = 6";
let got = col(&ctx, sql).await?;
assert_eq!(got.len(), 1);
assert!(
(got[0] - std::f64::consts::SQRT_2).abs() < 1e-12,
"Newton did not converge to sqrt(2): got {}",
got[0]
);
Ok(())
}
#[tokio::test]
async fn a_marker_free_statement_is_returned_byte_identical() -> Result<()> {
let sql = "SELECT * FROM t /* odd spacing preserved */";
assert_eq!(rewrite_sql(sql)?, sql);
Ok(())
}
#[tokio::test]
async fn rewrite_is_inspectable_without_running_it() -> Result<()> {
let out = rewrite_sql("SELECT grad(sin(x), x) AS d FROM t")?;
assert_eq!(out, "SELECT (cos(x)) AS d FROM t");
Ok(())
}
#[tokio::test]
async fn path_a_fails_at_the_call_not_at_collect() -> Result<()> {
let ctx = ctx().await?;
let err = ddx_sql(&ctx, "SELECT grad(atan2(x, y), x) FROM t")
.await
.expect_err("atan2 has no rule yet");
let msg = err.to_string();
assert!(msg.contains("not implemented"), "unexpected: {msg}");
assert!(msg.contains("atan2"), "must name the culprit: {msg}");
Ok(())
}