use uqa_core::Value;
use uqa_sql::ast::{FunctionBody, FunctionParamMode, SQLBodyForm, Statement};
use uqa_sql::expr::quote_ident;
use uqa_sql::routines::SQLUserFunction;
use uqa_sql::SQLError;
use crate::catalog::context::CatalogContext;
use crate::catalog::{CatalogReadView, RelationNameResolution};
use super::builtin_routines::{BuiltinRoutineCatalogEntry, PG18_BUILTIN_ROUTINE_GROUPS};
mod builtin_body;
mod definition;
pub use definition::pg_get_functiondef_value;
enum Routine {
User(std::sync::Arc<SQLUserFunction>),
Builtin(BuiltinRoutineCatalogEntry),
}
pub(super) fn routine_parameter_default_text(
output: Option<&dyn uqa_sql::expr::EngineHook>,
catalog: &CatalogReadView,
resolution: &RelationNameResolution,
parameter: &uqa_sql::ast::FunctionParam,
) -> Result<Option<String>, SQLError> {
use uqa_sql::ast::{Expr, RoutineDefaultType};
let Some(default) = ¶meter.default else {
return Ok(None);
};
let typed;
let expression = if matches!(default, Expr::Literal(Value::Str(_) | Value::Null)) {
let ty = match ¶meter.default_type {
Some(RoutineDefaultType::Concrete(ty)) => ty.catalog_name(),
Some(RoutineDefaultType::Polymorphic(name)) => name.clone(),
None => "unknown".into(),
};
typed = Expr::Cast {
implicit: true,
expr: Box::new(default.clone()),
ty,
};
&typed
} else {
default
};
super::view_definition::stored_expression_text(output, catalog, resolution, expression)
.map(Some)
}
fn routine_oid_argument(name: &str, arguments: &[Value]) -> Result<Option<i64>, SQLError> {
match arguments {
[Value::Null] => Ok(None),
[Value::Int(oid)] => Ok(Some(*oid)),
[_] => Err(SQLError::TypeMismatch(format!("{name} requires an oid"))),
_ => Err(SQLError::BadArity {
name: name.into(),
expected: "1".into(),
actual: arguments.len(),
}),
}
}
fn find_routine(context: &CatalogContext<'_>, oid: i64) -> Result<Option<Routine>, SQLError> {
for function in context.catalog_read_view().all_sql_functions() {
if super::pg_proc::user_routine_catalog_oid(&function)? == oid {
return Ok(Some(Routine::User(function)));
}
}
Ok(PG18_BUILTIN_ROUTINE_GROUPS
.iter()
.flat_map(|group| group.iter())
.copied()
.chain(super::builtin_routines::native_foreign_handlers())
.find(|entry| entry.oid == oid)
.map(Routine::Builtin))
}
fn type_display(context: &CatalogContext<'_>, oid: i64) -> Result<String, SQLError> {
match super::format_type_value(context, &[Value::Int(oid), Value::Null])? {
Value::Str(name) => Ok(name),
other => Err(SQLError::Internal(format!(
"format_type returned {other:?} for type {oid}"
))),
}
}
fn declared_type_display(
context: &CatalogContext<'_>,
type_name: &str,
) -> Result<String, SQLError> {
let oid = super::regtypes::catalog_routine_type_oid(&context.catalog_read_view(), type_name);
type_display(context, oid)
}
fn user_arguments(
context: &CatalogContext<'_>,
function: &SQLUserFunction,
table_arguments: bool,
defaults: bool,
) -> Result<(String, usize), SQLError> {
let def = &function.def;
let catalog = context.catalog_read_view();
let resolution = context.session_execution_view().relation_name_resolution();
let mut printed = Vec::new();
for parameter in &def.params {
let (mode, input) = match parameter.mode {
FunctionParamMode::In if def.is_procedure => ("IN ", true),
FunctionParamMode::In => ("", true),
FunctionParamMode::InOut => ("INOUT ", true),
FunctionParamMode::Out => ("OUT ", false),
FunctionParamMode::Variadic => ("VARIADIC ", true),
FunctionParamMode::Table => ("", false),
};
if table_arguments != (parameter.mode == FunctionParamMode::Table) {
continue;
}
let mut argument = mode.to_string();
if !parameter.name.is_empty() {
argument.push_str("e_ident(¶meter.name));
argument.push(' ');
}
argument.push_str(&declared_type_display(context, ¶meter.type_name)?);
if defaults && input {
if let Some(default) = routine_parameter_default_text(
Some(&crate::catalog::projection::CatalogOutput(*context)),
&catalog,
&resolution,
parameter,
)? {
argument.push_str(" DEFAULT ");
argument.push_str(&default);
}
}
printed.push(argument);
}
Ok((printed.join(", "), printed.len()))
}
fn builtin_arguments(
context: &CatalogContext<'_>,
entry: &BuiltinRoutineCatalogEntry,
defaults: bool,
) -> Result<String, SQLError> {
let first_default = entry
.argument_types
.len()
.saturating_sub(entry.default_arguments);
let default_texts = if defaults {
builtin_body::argument_defaults(context, entry)?
} else {
Vec::new()
};
let types = entry.all_argument_types().unwrap_or(entry.argument_types);
let mut printed = Vec::with_capacity(types.len());
let mut input_index = 0;
for (index, oid) in types.iter().enumerate() {
let mode = entry
.argument_modes()
.and_then(|modes| modes.get(index))
.copied()
.unwrap_or(if entry.variadic_type() != 0 && index + 1 == types.len() {
"v"
} else {
"i"
});
if mode == "t" {
continue;
}
let input = mode != "o";
let mut argument = match mode {
"o" => "OUT ",
"b" => "INOUT ",
"v" => "VARIADIC ",
_ if entry.kind == "p" => "IN ",
_ => "",
}
.to_owned();
if let Some(name) = entry
.argument_names
.get(index)
.filter(|name| !name.is_empty())
{
argument.push_str("e_ident(name));
argument.push(' ');
}
argument.push_str(&type_display(context, *oid)?);
if defaults && input && input_index >= first_default {
if let Some(text) = default_texts.get(input_index - first_default) {
argument.push_str(" DEFAULT ");
argument.push_str(text);
}
}
input_index += usize::from(input);
printed.push(argument);
}
Ok(printed.join(", "))
}
pub fn pg_get_function_arguments_value(
context: &CatalogContext<'_>,
arguments: &[Value],
) -> Result<Value, SQLError> {
function_arguments(context, "pg_get_function_arguments", arguments, true)
}
pub fn pg_get_function_identity_arguments_value(
context: &CatalogContext<'_>,
arguments: &[Value],
) -> Result<Value, SQLError> {
function_arguments(
context,
"pg_get_function_identity_arguments",
arguments,
false,
)
}
fn function_arguments(
context: &CatalogContext<'_>,
name: &str,
arguments: &[Value],
defaults: bool,
) -> Result<Value, SQLError> {
let Some(oid) = routine_oid_argument(name, arguments)? else {
return Ok(Value::Null);
};
Ok(match find_routine(context, oid)? {
Some(Routine::User(function)) => {
Value::Str(user_arguments(context, &function, false, defaults)?.0)
}
Some(Routine::Builtin(entry)) => Value::Str(builtin_arguments(context, &entry, defaults)?),
None => Value::Null,
})
}
pub fn pg_get_function_result_value(
context: &CatalogContext<'_>,
arguments: &[Value],
) -> Result<Value, SQLError> {
let Some(oid) = routine_oid_argument("pg_get_function_result", arguments)? else {
return Ok(Value::Null);
};
let Some(routine) = find_routine(context, oid)? else {
return Ok(Value::Null);
};
routine_result(context, &routine)
}
fn routine_result(context: &CatalogContext<'_>, routine: &Routine) -> Result<Value, SQLError> {
let function = match routine {
Routine::Builtin(entry) => {
return if entry.kind == "p" {
Ok(Value::Null)
} else {
type_display(context, entry.return_type).map(|name| {
Value::Str(if entry.returns_set() {
format!("SETOF {name}")
} else {
name
})
})
};
}
Routine::User(function) => function,
};
let def = &function.def;
if def.is_procedure {
return Ok(Value::Null);
}
if def.returns_set() {
let (columns, count) = user_arguments(context, function, true, false)?;
if count > 0 {
return Ok(Value::Str(format!("TABLE({columns})")));
}
}
let result = declared_type_display(
context,
uqa_sql::routines::declaration::result_type_name(def),
)?;
Ok(Value::Str(if def.returns_set() {
format!("SETOF {result}")
} else {
result
}))
}
pub fn pg_get_function_sqlbody_value(
context: &CatalogContext<'_>,
arguments: &[Value],
) -> Result<Value, SQLError> {
let Some(oid) = routine_oid_argument("pg_get_function_sqlbody", arguments)? else {
return Ok(Value::Null);
};
let Some(routine) = find_routine(context, oid)? else {
return Ok(Value::Null);
};
routine_sqlbody(context, &routine)
}
fn routine_sqlbody(context: &CatalogContext<'_>, routine: &Routine) -> Result<Value, SQLError> {
let function = match routine {
Routine::User(function) => function,
Routine::Builtin(routine) => return builtin_body::definition(context, routine),
};
let FunctionBody::Statements(statements) = &function.def.body else {
return Ok(Value::Null);
};
let catalog = context.catalog_read_view();
let resolution = context.session_execution_view().relation_name_resolution();
let form = function
.def
.sql_body_form
.unwrap_or_else(|| legacy_body_form(statements));
super::view_definition::routine_body_definition(
Some(&crate::catalog::projection::CatalogOutput(*context)),
&catalog,
&resolution,
&function.def,
form,
statements,
)
.map(Value::Str)
}
fn legacy_body_form(statements: &[Statement]) -> SQLBodyForm {
match statements {
[Statement::Select(select)]
if select.from.is_none()
&& select.with.is_empty()
&& select.values.is_empty()
&& select.set_op.is_none()
&& select.r#where.is_none()
&& select.group_by.is_empty()
&& select.having.is_none()
&& select.order_by.is_empty()
&& select.limit.is_none()
&& select.offset.is_none()
&& !select.distinct
&& matches!(select.projections.as_slice(), [projection] if projection.alias.is_none()) =>
{
SQLBodyForm::Return
}
_ => SQLBodyForm::Atomic,
}
}