use sqlparser::ast::Expr;
use sqlparser::dialect::Dialect;
use sqlparser::parser::Parser;
use crate::colref::{ColRef, IdentCasing};
use crate::engine::{differentiate, jvp, Rule, RuleRegistry};
use crate::error::{DiffError, Result};
use crate::rewrite::{self, Explanation};
#[derive(Clone)]
pub struct Ddx {
rules: RuleRegistry,
casing: IdentCasing,
}
impl Default for Ddx {
fn default() -> Self {
Self::new()
}
}
impl Ddx {
pub fn new() -> Self {
Ddx {
rules: RuleRegistry::new(),
casing: IdentCasing::FoldUnquoted,
}
}
pub fn for_datafusion() -> Self {
Self::with_casing(IdentCasing::FoldUnquoted)
}
pub fn for_duckdb() -> Self {
Self::with_casing(IdentCasing::FoldAll)
}
pub fn with_casing(casing: IdentCasing) -> Self {
Ddx {
rules: RuleRegistry::new(),
casing,
}
}
pub fn casing(&self) -> IdentCasing {
self.casing
}
pub fn register(&mut self, name: &str, rule: Rule) {
self.rules.register(name, rule);
}
pub fn rewrite_sql(&self, sql: &str, dialect: &dyn Dialect) -> Result<String> {
rewrite::rewrite_sql(sql, dialect, self.casing, &self.rules)
}
pub fn explain(&self, sql: &str, dialect: &dyn Dialect) -> Result<Explanation> {
rewrite::explain_sql(sql, dialect, self.casing, &self.rules)
}
pub fn differentiate(&self, e: &Expr, wrt: &ColRef) -> Result<Expr> {
differentiate(e, wrt, self.casing, &self.rules)
}
pub fn jvp(&self, e: &Expr, seeds: &[(ColRef, Expr)]) -> Result<Expr> {
jvp(e, seeds, self.casing, &self.rules)
}
pub fn differentiate_sql(
&self,
expr: &str,
wrt: &str,
dialect: &dyn Dialect,
) -> Result<String> {
let parsed = parse_expr(expr, dialect)?;
let wrt_expr = parse_expr(wrt, dialect)?;
let wrt_col = ColRef::from_wrt_arg("differentiate_sql", &wrt_expr)?;
let derivative = self.differentiate(&parsed, &wrt_col)?;
Ok(derivative.to_string())
}
}
fn parse_expr(text: &str, dialect: &dyn Dialect) -> Result<Expr> {
Parser::new(dialect)
.try_with_sql(text)
.and_then(|mut p| p.parse_expr())
.map_err(|e| DiffError::Parse(format!("failed to parse expression `{text}`: {e}")))
}