use super::{
body_validation::routine_parameter_values,
compilation::{lower_sql_routine_statement, RoutineCompilationContext, SQLRoutineLowering},
routine_local_name,
};
use crate::{
ast::{CreateFunction, Expr, FunctionBody, FunctionParam, FunctionParamMode, Statement},
binding::{bind_routine_parameter_references, RoutineParameterScope},
catalog::stored_ast::{visit_stored_statement_expressions, visit_stored_statement_projections},
plan::UnifiedPlan,
SQLError, SQLParam, ScalarExpr,
};
use std::cell::Cell;
use std::collections::BTreeSet;
#[must_use]
pub fn sql_body_parameters(def: &CreateFunction) -> Vec<&FunctionParam> {
def.params
.iter()
.filter(|parameter| is_sql_body_parameter(parameter))
.collect()
}
#[must_use]
pub const fn is_sql_body_parameter(parameter: &FunctionParam) -> bool {
matches!(
parameter.mode,
FunctionParamMode::In | FunctionParamMode::InOut | FunctionParamMode::Variadic
)
}
#[must_use]
pub fn sql_body_parameter_names(def: &CreateFunction) -> Vec<String> {
match def.body {
FunctionBody::Statements(_) => def
.sql_body_parameters()
.into_iter()
.map(|parameter| parameter.name.to_string())
.collect(),
FunctionBody::Source(_) => sql_body_parameters(def)
.iter()
.map(|parameter| parameter.name.clone())
.collect(),
}
}
pub fn sql_body_parameter_scope(
def: &CreateFunction,
params: &[SQLParam],
) -> Result<RoutineParameterScope, SQLError> {
Ok(RoutineParameterScope::new(
&routine_local_name(&def.name)?,
sql_body_parameter_names(def),
params
.iter()
.map(|param| param.declared_scalar_type().cloned())
.collect(),
))
}
pub fn record_sql_standard_body_parameters(
context: &RoutineCompilationContext<'_>,
def: &mut CreateFunction,
) -> Result<bool, SQLError> {
let FunctionBody::Statements(statements) = &def.body else {
return Ok(false);
};
let params = routine_parameter_values(context.types, def);
let recording = ParameterRecording {
context,
scope: sql_body_parameter_scope(def, ¶ms)?,
params,
sites: ParameterSites {
function: routine_local_name(&def.name)?,
names: sql_body_parameter_names(def),
},
};
let mut recorded = statements.clone();
let mut changed = false;
for statement in &mut recorded {
changed |= recording.record(statement)?;
}
if changed {
def.body = FunctionBody::Statements(recorded);
}
Ok(changed)
}
struct ParameterSites {
function: String,
names: Vec<String>,
}
impl ParameterSites {
fn position(&self, expression: &Expr) -> Option<usize> {
let name = match expression {
Expr::Column(name) => name,
Expr::QualifiedColumn { qualifier, column } if *qualifier == self.function => column,
_ => return None,
};
self.names
.iter()
.position(|candidate| !candidate.is_empty() && candidate == name)
}
fn collect(&self, statement: &mut Statement) -> Result<Vec<usize>, SQLError> {
let mut positions = Vec::new();
visit_stored_statement_expressions(statement, &mut |expression| {
positions.extend(self.position(expression));
Ok(())
})?;
Ok(positions)
}
fn replace(&self, statement: &mut Statement, chosen: &BTreeSet<usize>) -> Result<(), SQLError> {
let ordinal = Cell::new(0);
visit_stored_statement_projections(
statement,
&mut |projection| {
let name = match &projection.expr {
Expr::Column(name) | Expr::QualifiedColumn { column: name, .. } => name,
_ => return Ok(()),
};
if projection.alias.is_none()
&& self.position(&projection.expr).is_some()
&& chosen.contains(&ordinal.get())
{
projection.alias = Some(name.clone());
}
Ok(())
},
&mut |expression| {
if let Some(position) = self.position(expression) {
if chosen.contains(&ordinal.get()) {
*expression = Expr::Param(position + 1);
}
ordinal.set(ordinal.get() + 1);
}
Ok(())
},
)
}
}
struct ParameterRecording<'a> {
context: &'a RoutineCompilationContext<'a>,
scope: RoutineParameterScope,
params: Vec<SQLParam>,
sites: ParameterSites,
}
impl ParameterRecording<'_> {
fn record(&self, statement: &mut Statement) -> Result<bool, SQLError> {
let candidates = self.sites.collect(statement)?;
if candidates.is_empty() {
return Ok(false);
}
let original = self.resolved_counts(statement.clone())?;
let mut chosen = BTreeSet::new();
for (ordinal, position) in candidates.iter().enumerate() {
if original[*position] == 0 {
continue;
}
let mut probe = statement.clone();
self.sites.replace(&mut probe, &BTreeSet::from([ordinal]))?;
if self
.resolved_counts(probe)
.is_ok_and(|counts| counts[*position] < original[*position])
{
chosen.insert(ordinal);
}
}
if chosen.is_empty() {
return Ok(false);
}
self.sites.replace(statement, &chosen)?;
Ok(true)
}
fn resolved_counts(&self, statement: Statement) -> Result<Vec<usize>, SQLError> {
let mut plan = lower_sql_routine_statement(
self.context,
statement,
SQLRoutineLowering {
bind_catalog_dependencies: true,
persisted_definition: false,
preserve_target_expressions: false,
},
)?;
let before = parameter_counts(&mut plan, self.params.len());
let snapshot = self.context.catalog.binding_snapshot()?;
bind_routine_parameter_references(
self.context.routines,
&mut plan,
&self.params,
&snapshot.context(),
&self.scope,
)?;
Ok(parameter_counts(&mut plan, self.params.len())
.into_iter()
.zip(before)
.map(|(after, before)| after.saturating_sub(before))
.collect())
}
}
fn parameter_counts(plan: &mut UnifiedPlan, count: usize) -> Vec<usize> {
let mut counts = vec![0; count];
plan.rewrite_scalar_expressions(&mut |expression| {
if let ScalarExpr::Param(index) = expression {
if let Some(slot) = index.checked_sub(1).and_then(|slot| counts.get_mut(slot)) {
*slot += 1;
}
}
});
counts
}