use super::{
apply_attribute_change, type_contains_composite, AttributeChange, CompositeTypeCatalog,
};
use crate::{
ast::{
ColumnDef, ColumnType, Expr, PartitionBound, PartitionRangeDatum, PartitionSpec, Statement,
TableCheck,
},
plan::{QueryPlan, UnifiedPlan},
type_resolution::FunctionTypeResolver,
SQLError, ScalarExpr,
};
use uqa_core::Value;
pub struct CompositeConstantChange<'a> {
pub target: u32,
pub change: &'a AttributeChange,
pub catalog: &'a dyn CompositeTypeCatalog,
pub types: &'a dyn FunctionTypeResolver,
}
impl CompositeConstantChange<'_> {
pub fn value(&self, value: &mut Value, ty: &ColumnType) -> Result<bool, SQLError> {
if matches!(value, Value::Null) || !type_contains_composite(ty, self.target, self.catalog)?
{
return Ok(false);
}
*value = apply_attribute_change(
std::mem::replace(value, Value::Null),
ty,
self.target,
self.change,
self.catalog,
)?;
Ok(true)
}
fn typed_value(&self, value: &mut Value, name: &str) -> Result<bool, SQLError> {
let Some(ty) = self.types.resolve_type_name(name)? else {
return if crate::ast::UserTypeIdentity::parse(name).is_some() {
Err(SQLError::Internal(format!(
"stored constant type {name} disappeared"
)))
} else {
Ok(false)
};
};
self.value(value, &ty)
}
fn syntax_node(&self, node: &mut Expr) -> Result<bool, SQLError> {
match node {
Expr::TypedLiteral { value, ty } => self.typed_value(value, ty),
_ => Ok(false),
}
}
fn scalar_node(&self, node: &mut ScalarExpr) -> Result<bool, SQLError> {
match node {
ScalarExpr::TypedLiteral { value, ty, .. } => self.typed_value(value, ty),
_ => Ok(false),
}
}
pub fn expression(&self, expression: &mut Expr) -> Result<bool, SQLError> {
let mut changed = false;
crate::catalog::stored_ast::visit_stored_expression(expression, &mut |node| {
changed |= self.syntax_node(node)?;
Ok(())
})?;
Ok(changed)
}
pub fn statement(&self, statement: &mut Statement) -> Result<bool, SQLError> {
let mut changed = false;
crate::catalog::stored_ast::visit_stored_statement_expressions(statement, &mut |node| {
changed |= self.syntax_node(node)?;
Ok(())
})?;
Ok(changed)
}
pub fn query(&self, query: &mut QueryPlan) -> Result<bool, SQLError> {
self.plan_nodes(|visit| query.rewrite_scalar_expressions(visit))
}
pub fn plan(&self, plan: &mut UnifiedPlan) -> Result<bool, SQLError> {
self.plan_nodes(|visit| plan.rewrite_scalar_expressions(visit))
}
pub fn expression_plan(
&self,
plan: &mut crate::plan::ExpressionPlan,
) -> Result<bool, SQLError> {
let mut changed = self
.plan_nodes(|visit| crate::plan::rewrite_scalar_expression(&mut plan.scalar, visit))?;
for query in &mut plan.subqueries {
changed |= self.query(query)?;
}
Ok(changed)
}
fn plan_nodes(
&self,
visit: impl FnOnce(&mut dyn FnMut(&mut ScalarExpr)),
) -> Result<bool, SQLError> {
let mut changed = false;
let mut failure = None;
visit(&mut |node| {
if failure.is_some() {
return;
}
match self.scalar_node(node) {
Ok(value) => changed |= value,
Err(error) => failure = Some(error),
}
});
failure.map_or(Ok(changed), Err)
}
pub fn columns(&self, columns: &mut [ColumnDef]) -> Result<bool, SQLError> {
let mut changed = false;
for column in columns {
for expression in column.default.iter_mut().chain(column.check.iter_mut()) {
changed |= self.expression(expression)?;
}
if let Some(generated) = &mut column.generated {
changed |= self.expression(&mut generated.expression)?;
}
}
Ok(changed)
}
pub fn checks(
&self,
checks: &mut [TableCheck],
columns: &[ColumnDef],
) -> Result<bool, SQLError> {
let mut changed = false;
for check in checks {
changed |= self.expression(&mut check.expr)?;
if let Some(partition) = &mut check.partition_constraint {
changed |= self.bound(&mut partition.bound, &partition.spec, columns)?;
changed |= self.spec(&mut partition.spec)?;
}
}
Ok(changed)
}
pub fn spec(&self, spec: &mut PartitionSpec) -> Result<bool, SQLError> {
let mut changed = false;
for key in &mut spec.keys {
changed |= self.expression(key)?;
}
Ok(changed)
}
pub fn bound(
&self,
bound: &mut PartitionBound,
spec: &PartitionSpec,
columns: &[ColumnDef],
) -> Result<bool, SQLError> {
let mut changed = false;
for (position, key) in spec.keys.iter().enumerate() {
let ty = crate::semantics::partition::partition_key_type(self.types, key, columns)?;
match bound {
PartitionBound::List(values) => {
for expression in values {
changed |= self.bound_value(expression, &ty)?;
}
}
PartitionBound::Range { lower, upper } => {
for point in [lower.get_mut(position), upper.get_mut(position)]
.into_iter()
.flatten()
{
if let PartitionRangeDatum::Value(expression) = point {
changed |= self.bound_value(expression, &ty)?;
}
}
}
PartitionBound::Default | PartitionBound::Hash { .. } => {}
}
}
Ok(changed)
}
fn bound_value(&self, expression: &mut Expr, ty: &ColumnType) -> Result<bool, SQLError> {
match expression {
Expr::Literal(value) | Expr::TypedLiteral { value, .. } => self.value(value, ty),
_ => Err(SQLError::Internal(
"partition bound datum was not evaluated".into(),
)),
}
}
}
#[cfg(test)]
mod tests;