use std::cell::RefCell;
pub mod inputs;
mod statements;
use super::transaction::DirectRoutineCommandGuard;
use super::{CreateFunction, RoutineContext, RoutineOutcome, SQLError, SQLParam, SQLResult, Value};
use uqa_sql::{
assignment::routines::coerce_routine_value_from,
ast::FunctionBody,
binding::statements::AnalyzedResult,
plan::{CommandPlan, UnifiedPlan},
plpgsql::runtime_diagnostics::result_row_values,
routines::{
body_parameters::{is_sql_body_parameter, sql_body_parameter_scope},
declaration::RoutineTypeCatalog,
resolution::RoutineOverloadContext,
result_check::{
check_sql_function_result, sql_function_result_layout, validate_sql_function_record,
SQLFunctionResultKind, SQLFunctionResultLayout,
},
routine_returns_anonymous_record, SQLUserFunction,
},
};
pub struct SQLBody<'a> {
pub function: &'a SQLUserFunction,
pub definition: &'a CreateFunction,
pub plans: &'a [UnifiedPlan],
}
pub fn execute_sql_language(
context: RoutineContext<'_>,
types: &dyn RoutineTypeCatalog,
overloads: &RoutineOverloadContext<'_>,
body: SQLBody<'_>,
bound: &[Value],
record_target: Option<&[uqa_sql::routines::result_check::SQLFunctionResultColumn]>,
) -> Result<RoutineOutcome, SQLError> {
let SQLBody {
function,
definition: def,
plans,
} = body;
let call_params = def.call_params();
if call_params.len() != bound.len() {
return Err(SQLError::Internal(format!(
"routine `{}` received {} values for {} concrete call parameters",
def.name,
bound.len(),
call_params.len()
)));
}
let params = bound
.iter()
.cloned()
.zip(call_params)
.filter(|(_, parameter)| is_sql_body_parameter(parameter))
.map(|(value, parameter)| {
let ty = uqa_sql::ast::ColumnType::from_sql_name(¶meter.type_name)
.ok()
.or_else(|| {
context
.expressions
.catalog_column_type(¶meter.type_name)
})
.ok_or_else(|| {
SQLError::TypeMismatch(format!("unknown type `{}`", parameter.type_name))
})?;
Ok(SQLParam::typed_scalar(value, ty))
})
.collect::<Result<Vec<_>, SQLError>>()?;
if plans.is_empty() {
check_sql_function_result(types, def, None)?;
}
let layout = RefCell::new(None);
let check_result = |result: &AnalyzedResult| {
let result_layout = sql_function_result_layout(types, def, Some(result))?;
if result_layout.kind == SQLFunctionResultKind::Tuple
&& routine_returns_anonymous_record(def)
{
if let Some(target) = record_target {
uqa_sql::routines::result_check::validate_anonymous_record_result(
types,
result.column_types().unwrap_or_default(),
target,
Some(SQLFunctionResultKind::Tuple),
)?;
}
}
*layout.borrow_mut() = Some(result_layout);
Ok(())
};
let parameters = matches!(def.body, FunctionBody::Source(_))
.then(|| sql_body_parameter_scope(def, ¶ms))
.transpose()?;
let identity = inputs::SQLBodyIdentity::new(function)?;
let parameter_types = params
.iter()
.map(|param| {
param
.declared_scalar_type()
.cloned()
.ok_or_else(|| SQLError::Internal("SQL body parameter has no concrete type".into()))
})
.collect::<Result<Vec<_>, _>>()?;
let statements = statements::SQLStatements {
context,
types,
overloads,
params: ¶ms,
parameters,
identity,
parameter_types,
inputs: context.statements.body_input_context(),
};
let mut last = SQLResult::empty();
for (position, plan) in plans.iter().enumerate() {
let check = (position + 1 == plans.len())
.then_some(&check_result as super::context::StatementResultCheck<'_>);
let statement = statements.prepare(plan, position, check)?;
let _direct_routine_command = matches!(
&statement,
UnifiedPlan::Command(command)
if matches!(
command.as_ref(),
CommandPlan::Call { .. } | CommandPlan::DoBlock { .. }
)
)
.then(|| DirectRoutineCommandGuard::enter(context.session));
last = context
.statements
.execute_body_statement(statement, ¶ms, check)?;
}
let layout = layout
.into_inner()
.map_or_else(|| sql_function_result_layout(types, def, None), Ok)?;
shape_result(context.expressions, types, def, &last, &layout)
}
fn shape_result(
expressions: &dyn uqa_sql::assignment::routines::RoutineValueContext,
types: &dyn RoutineTypeCatalog,
def: &CreateFunction,
last: &SQLResult,
layout: &SQLFunctionResultLayout,
) -> Result<RoutineOutcome, SQLError> {
let outputs = def.output_params();
let anonymous = routine_returns_anonymous_record(def);
let record_types = anonymous
.then(|| {
if layout.kind == SQLFunctionResultKind::Value {
layout.source_record.clone()
} else {
Some(last.column_types.clone())
}
})
.flatten();
let shape = |values| result_value(expressions, types, layout, last, values);
if def.returns_set() {
if layout.kind == SQLFunctionResultKind::Value {
uqa_sql::routines::result_check::validate_sql_function_record_rows(last)?;
}
let mut rows = Vec::with_capacity(last.rows.len());
for index in 0..last.rows.len() {
let value = shape(result_row_values(last, index).unwrap_or_default())?;
rows.push(output_values(def, value, outputs.len())?);
}
return Ok(RoutineOutcome {
value: Value::Null,
out_values: vec![Value::Null; outputs.len()],
set_rows: rows,
anonymous_record_column_types: record_types,
sql_result_kind: Some(layout.kind),
});
}
let value = result_row_values(last, 0)
.map(shape)
.transpose()?
.unwrap_or(Value::Null);
let (value, out_values) = if outputs.is_empty() {
(value, Vec::new())
} else {
(Value::Null, output_values(def, value, outputs.len())?)
};
Ok(RoutineOutcome {
value,
out_values,
set_rows: Vec::new(),
anonymous_record_column_types: record_types,
sql_result_kind: Some(layout.kind),
})
}
fn result_value(
expressions: &dyn uqa_sql::assignment::routines::RoutineValueContext,
types: &dyn RoutineTypeCatalog,
layout: &SQLFunctionResultLayout,
last: &SQLResult,
mut values: Vec<Value>,
) -> Result<Value, SQLError> {
match layout.kind {
SQLFunctionResultKind::Void => Ok(Value::Null),
SQLFunctionResultKind::Value => {
let value = values.pop().unwrap_or(Value::Null);
if layout.declared_type == uqa_sql::ColumnType::Record {
if let Some(expected) = &layout.columns {
if !matches!(value, Value::Null) {
let target = expected
.iter()
.map(|column| column.ty.clone())
.collect::<Vec<_>>();
if let Value::Row(row) = &value {
if let Some(source) = row.field_types() {
uqa_sql::routines::result_check::validate_sql_function_record_identity(types, source, &target)?;
} else if let Some(source) = &layout.source_record {
validate_sql_function_record(types, source, &target)?;
}
} else if let Some(source) = &layout.source_record {
validate_sql_function_record(types, source, &target)?;
}
}
let fields = row_fields(value, expected.len())?;
return Ok(Value::Record(
expected
.iter()
.map(|column| column.name.clone())
.zip(fields)
.collect(),
));
}
}
coerce_routine_value_from(
expressions,
&value,
&layout.declared_type.catalog_name(),
last.column_types.first().and_then(Option::as_ref),
)
}
SQLFunctionResultKind::Tuple => {
let columns = if let Some(expected) = &layout.columns {
for (index, (value, column)) in values.iter_mut().zip(expected).enumerate() {
*value = coerce_routine_value_from(
expressions,
value,
&column.ty.catalog_name(),
last.column_types.get(index).and_then(Option::as_ref),
)?;
}
expected
.iter()
.map(|column| column.name.clone())
.collect::<Vec<_>>()
} else {
last.columns.clone()
};
Ok(Value::Record(columns.into_iter().zip(values).collect()))
}
}
}
fn output_values(
def: &CreateFunction,
value: Value,
output_count: usize,
) -> Result<Vec<Value>, SQLError> {
if output_count == 0 || (output_count == 1 && !def.is_procedure) {
Ok(vec![value])
} else {
row_fields(value, output_count)
}
}
pub(super) fn row_fields(value: Value, expected: usize) -> Result<Vec<Value>, SQLError> {
let fields = match value {
Value::Record(fields) => fields
.into_iter()
.map(|(_, value)| value)
.collect::<Vec<_>>(),
Value::Row(fields) => fields.into_values(),
Value::Null => return Ok(vec![Value::Null; expected]),
_ => {
return Err(SQLError::Internal(
"a SQL routine row result contained a scalar value".into(),
))
}
};
if fields.len() == expected {
Ok(fields)
} else {
Err(SQLError::Diagnostic {
sqlstate: "42804".into(),
message: "function return row and query-specified return row do not match".into(),
detail: Some(format!(
"Returned row contains {} attribute{}, but query expects {expected}.",
fields.len(),
if fields.len() == 1 { "" } else { "s" }
)),
hint: None,
})
}
}
#[cfg(test)]
mod tests;