use datafusion::arrow::datatypes::DataType;
use datafusion::error::{DataFusionError, Result};
use datafusion::logical_expr::{
ColumnarValue, ScalarFunctionArgs, ScalarUDF, ScalarUDFImpl, Signature, Volatility,
};
pub const GRAD: &str = "grad";
pub const JVP: &str = "jvp";
#[derive(Debug, PartialEq, Eq, Hash)]
struct Marker {
name: &'static str,
signature: Signature,
}
impl Marker {
fn new(name: &'static str, arg_count: usize) -> Self {
Marker {
name,
signature: Signature::any(arg_count, Volatility::Immutable),
}
}
}
impl ScalarUDFImpl for Marker {
fn name(&self) -> &str {
self.name
}
fn signature(&self) -> &Signature {
&self.signature
}
fn return_type(&self, _arg_types: &[DataType]) -> Result<DataType> {
Ok(DataType::Float64)
}
fn invoke_with_args(&self, _args: ScalarFunctionArgs) -> Result<ColumnarValue> {
Err(DataFusionError::Execution(format!(
"ddx: `{name}` reached execution, which never happens in a correct \
rewrite — it is a compile-time marker, not a row function.\n\n\
The `{name}()` call was not rewritten away before planning finished. \
Either the ddx analyzer rule is not installed on this SessionContext \
(use `ddx_datafusion::install(&ctx)`), or the marker sits somewhere \
the rule does not reach — in which case rewrite the SQL text instead \
with `ddx_datafusion::ddx_sql(&ctx, sql)`.",
name = self.name
)))
}
}
pub fn grad_udf() -> ScalarUDF {
ScalarUDF::new_from_impl(Marker::new(GRAD, 2))
}
pub fn jvp_udf() -> ScalarUDF {
ScalarUDF::new_from_impl(Marker::new(JVP, 3))
}
pub(crate) fn marker_kind(name: &str) -> Option<&'static str> {
if name.eq_ignore_ascii_case(GRAD) {
Some(GRAD)
} else if name.eq_ignore_ascii_case(JVP) {
Some(JVP)
} else {
None
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn markers_are_named_and_arity_checked() {
assert_eq!(grad_udf().name(), "grad");
assert_eq!(jvp_udf().name(), "jvp");
}
#[test]
fn marker_names_are_matched_case_insensitively() {
assert_eq!(marker_kind("grad"), Some(GRAD));
assert_eq!(marker_kind("GRAD"), Some(GRAD));
assert_eq!(marker_kind("Jvp"), Some(JVP));
assert_eq!(marker_kind("gradient"), None);
assert_eq!(marker_kind("mygrad"), None);
}
#[test]
fn executing_a_marker_is_a_loud_error_not_a_number() {
let udf = grad_udf();
let err = udf
.invoke_with_args(ScalarFunctionArgs {
args: vec![],
arg_fields: vec![],
number_rows: 1,
return_field: std::sync::Arc::new(datafusion::arrow::datatypes::Field::new(
"d",
DataType::Float64,
true,
)),
config_options: std::sync::Arc::new(datafusion::config::ConfigOptions::default()),
})
.expect_err("a marker must never execute successfully");
let msg = err.to_string();
assert!(msg.contains("reached execution"), "unexpected: {msg}");
assert!(msg.contains("install"), "no remedy in message: {msg}");
assert!(msg.contains("ddx_sql"), "no fallback in message: {msg}");
}
}