use crate::{logical_plan_to_proof_plan, PlannerResult, PoSqlContextProvider};
use alloc::{sync::Arc, vec::Vec};
use datafusion::{
config::ConfigOptions,
logical_expr::LogicalPlan,
optimizer::{Analyzer, Optimizer, OptimizerContext, OptimizerRule},
sql::planner::{ParserOptions, SqlToRel},
};
use indexmap::IndexSet;
use proof_of_sql::{
base::database::{ParseError, SchemaAccessor, TableRef},
sql::proof_plans::DynProofPlan,
};
use sqlparser::ast::{visit_relations, Statement};
use std::ops::ControlFlow;
pub fn optimizer() -> Optimizer {
let recommended_rules: Vec<Arc<dyn OptimizerRule + Send + Sync>> = Optimizer::new().rules;
let filtered_rules = recommended_rules
.into_iter()
.filter(|rule| rule.name() != "common_sub_expression_eliminate")
.collect::<Vec<_>>();
Optimizer::with_rules(filtered_rules)
}
fn sql_to_posql_plans<T, F, A>(
statements: &[Statement],
schemas: &A,
config: &ConfigOptions,
planner_converter: F,
) -> PlannerResult<Vec<T>>
where
F: Fn(&LogicalPlan, &A) -> PlannerResult<T>,
A: SchemaAccessor + Clone,
{
let context_provider = PoSqlContextProvider::new(schemas.clone());
statements
.iter()
.map(|ast| -> PlannerResult<T> {
let raw_logical_plan = SqlToRel::new_with_options(
&context_provider,
ParserOptions {
parse_float_as_decimal: config.sql_parser.parse_float_as_decimal,
enable_ident_normalization: config.sql_parser.enable_ident_normalization,
},
)
.sql_statement_to_plan(ast.clone())?;
let analyzer = Analyzer::new();
let analyzed_logical_plan =
analyzer.execute_and_check(raw_logical_plan, config, |_, _| {})?;
let optimizer = optimizer();
let optimizer_context = OptimizerContext::default();
let optimized_logical_plan =
optimizer.optimize(analyzed_logical_plan, &optimizer_context, |_, _| {})?;
planner_converter(&optimized_logical_plan, schemas)
})
.collect::<PlannerResult<Vec<_>>>()
}
pub fn sql_to_proof_plans<A: SchemaAccessor + Clone>(
statements: &[Statement],
schemas: &A,
config: &ConfigOptions,
) -> PlannerResult<Vec<DynProofPlan>> {
sql_to_posql_plans(statements, schemas, config, logical_plan_to_proof_plan)
}
pub fn get_table_refs_from_statement(
statement: &Statement,
) -> Result<IndexSet<TableRef>, ParseError> {
let mut table_refs: IndexSet<TableRef> = IndexSet::<TableRef>::new();
visit_relations(statement, |object_name| {
match object_name.to_string().as_str().try_into() {
Ok(table_ref) => {
table_refs.insert(table_ref);
ControlFlow::Continue(())
}
e => ControlFlow::Break(e),
}
})
.break_value()
.transpose()?;
Ok(table_refs)
}
#[cfg(test)]
mod tests {
use super::get_table_refs_from_statement;
use crate::{
conversion::sql_to_posql_plans, sql_to_proof_plans, AggregatePlanError,
LogicalPlanNodeKind, PlannerError, PlannerResult,
};
use ahash::AHasher;
use datafusion::{config::ConfigOptions, logical_expr::LogicalPlan};
use indexmap::{indexmap_with_default, IndexSet};
use proof_of_sql::{
base::database::{
ColumnType, SchemaAccessor, SchemaAccessorImpl, TableRef, TableTestAccessor,
},
proof_primitive::dory::DynamicDoryEvaluationProof,
sql::proof_plans::{AggregateExecError, DynProofPlan},
};
use sqlparser::{dialect::GenericDialect, parser::Parser};
#[expect(non_snake_case)]
fn SQL_SCHEMAS() -> impl SchemaAccessor + Clone {
SchemaAccessorImpl::new(indexmap_with_default! {AHasher;
TableRef::new("", "test_table") => vec![
("id".into(), ColumnType::BigInt),
("name".into(), ColumnType::VarChar),
("payload".into(), ColumnType::VarBinary),
],
})
}
#[test]
fn we_can_get_table_references() {
let statement = Parser::parse_sql(
&GenericDialect {},
"SELECT e.employee_id, e.employee_name, d.department_name, p.project_name, s.salary
FROM employees e
JOIN departments d ON e.department_id = d.department_id
JOIN management.projects p ON e.employee_id = p.employee_id
JOIN internal.salaries s ON e.employee_id = s.employee_id
WHERE e.department_id IN (
SELECT department_id
FROM departments
WHERE department_name = 'Sales'
)
AND p.project_id IN (
SELECT project_id
FROM project_assignments
WHERE employee_id = e.employee_id
)
AND s.salary > (
SELECT AVG(salary)
FROM internal.salaries
WHERE department_id = e.department_id
);
",
)
.unwrap()[0]
.clone();
let table_refs = get_table_refs_from_statement(&statement).unwrap();
let expected_table_refs: IndexSet<TableRef> = [
("", "departments"),
("", "employees"),
("management", "projects"),
("", "project_assignments"),
("internal", "salaries"),
]
.map(|(s, t)| TableRef::new(s, t))
.into_iter()
.collect();
assert_eq!(table_refs, expected_table_refs);
}
#[test]
fn we_can_use_abs() {
let statements = Parser::parse_sql(&GenericDialect {}, "SELECT ABS(-1-1);").unwrap();
sql_to_posql_plans(
&statements,
&TableTestAccessor::<DynamicDoryEvaluationProof>::default(),
&ConfigOptions::default(),
|a, _| -> PlannerResult<LogicalPlan> { Ok(a.clone()) },
)
.unwrap();
}
#[test]
fn sql_distinct_reports_aggregate_construction_error() {
let statements =
Parser::parse_sql(&GenericDialect {}, "SELECT DISTINCT name FROM test_table;").unwrap();
let err =
sql_to_proof_plans(&statements, &SQL_SCHEMAS(), &ConfigOptions::default()).unwrap_err();
assert!(matches!(
err,
PlannerError::UnsupportedAggregatePlan {
source: AggregatePlanError::AggregateExec {
source: AggregateExecError::UnsupportedGroupByExpressionType {
data_type: ColumnType::VarChar,
},
},
..
}
));
}
#[test]
fn sql_grouping_reports_aggregate_construction_error() {
let statements = Parser::parse_sql(
&GenericDialect {},
"SELECT payload FROM test_table GROUP BY payload;",
)
.unwrap();
let err =
sql_to_proof_plans(&statements, &SQL_SCHEMAS(), &ConfigOptions::default()).unwrap_err();
assert!(matches!(
err,
PlannerError::UnsupportedAggregatePlan {
source: AggregatePlanError::AggregateExec {
source: AggregateExecError::UnsupportedGroupByExpressionType {
data_type: ColumnType::VarBinary,
},
},
..
}
));
}
#[test]
fn sql_window_reports_unsupported_plan_node() {
let statements = Parser::parse_sql(
&GenericDialect {},
"SELECT ROW_NUMBER() OVER () FROM test_table;",
)
.unwrap();
let err =
sql_to_proof_plans(&statements, &SQL_SCHEMAS(), &ConfigOptions::default()).unwrap_err();
assert!(matches!(
err,
PlannerError::UnsupportedLogicalPlan {
node: LogicalPlanNodeKind::Window,
}
));
}
#[test]
fn sql_distinct_on_supported_type_still_converts() {
let statements =
Parser::parse_sql(&GenericDialect {}, "SELECT DISTINCT id FROM test_table;").unwrap();
let plans =
sql_to_proof_plans(&statements, &SQL_SCHEMAS(), &ConfigOptions::default()).unwrap();
assert!(matches!(plans.as_slice(), [DynProofPlan::Projection(_)]));
}
}