uqa-sql 0.4.0

PostgreSQL-compatible SQL compiler built on libpg_query
Documentation
//
// Unified Query Algebra
//
// Copyright (c) 2023-2026 Cognica, Inc.
//

//! Named and explicit variadic SQL argument validation.

use super::{ScalarExpr, ScalarOrder};
use crate::ast::{FunctionBinding, FunctionDispatch};
use crate::SQLError;
use uqa_core::{
    memory::{Produced, ProductionControl, ProductionVec},
    Value,
};

/// A SQL call argument after removing the compiler's named and explicit `VARIADIC` syntax markers.
#[doc(hidden)]
#[derive(Debug, Clone, Copy, PartialEq)]
pub struct ScalarCallArgument<'a> {
    pub name: Option<&'a str>,
    pub value: &'a ScalarExpr,
    pub explicit_variadic: bool,
}

/// Decode and validate all compiler-owned call-argument markers. `PostgreSQL` permits one explicit `VARIADIC` argument and requires it to be the final argument.
#[doc(hidden)]
pub fn scalar_call_arguments(
    arguments: &[ScalarExpr],
) -> Result<Vec<ScalarCallArgument<'_>>, SQLError> {
    scalar_call_arguments_with_control(arguments, &ProductionControl::uncontrolled()).map(
        |decoded| {
            decoded
                .into_uncontrolled()
                .expect("ordinary call argument decoding has no reservation")
        },
    )
}

/// Decode the same borrowed markers into an admitted temporary container. Names and expression nodes remain borrowed from the input IR owner.
pub fn scalar_call_arguments_with_control<'a>(
    arguments: &'a [ScalarExpr],
    control: &ProductionControl<'_>,
) -> Result<Produced<Vec<ScalarCallArgument<'a>>>, SQLError> {
    let mut decoded = ProductionVec::new(*control);
    decoded.reserve(arguments.len())?;
    for argument in arguments {
        control.check()?;
        decoded.push_copy(scalar_call_argument(argument)?)?;
    }
    validate_scalar_call_arguments(&decoded)?;
    decoded.finish().map_err(Into::into)
}

/// Validate cross-argument invariants after individual syntax markers have been decoded, returning whether the call used explicit `VARIADIC` syntax.
#[doc(hidden)]
pub fn validate_scalar_call_arguments(
    arguments: &[ScalarCallArgument<'_>],
) -> Result<bool, SQLError> {
    let mut count = 0;
    let mut last_position = None;
    for (position, argument) in arguments.iter().enumerate() {
        if argument.explicit_variadic {
            count += 1;
            last_position = Some(position);
        }
    }
    if count > 1 {
        return Err(malformed_call_argument(
            "call contains more than one explicit VARIADIC argument",
        ));
    }
    if last_position.is_some_and(|position| position + 1 != arguments.len()) {
        return Err(malformed_call_argument(
            "explicit VARIADIC argument must be the final call argument",
        ));
    }
    Ok(count != 0)
}

/// Decode one compiler-owned call-argument marker. Use [`scalar_call_arguments`] for a complete call so duplicate and ordering invariants are also checked.
#[doc(hidden)]
pub fn scalar_call_argument(expression: &ScalarExpr) -> Result<ScalarCallArgument<'_>, SQLError> {
    let ScalarExpr::Func {
        name,
        args,
        binding,
        distinct,
        order_by,
        filter,
    } = expression
    else {
        return Ok(ScalarCallArgument {
            name: None,
            value: expression,
            explicit_variadic: false,
        });
    };
    if binding.as_ref().and_then(|binding| binding.dispatch)
        == Some(FunctionDispatch::NamedArgument)
    {
        validate_marker_shape(
            binding.as_ref(),
            FunctionDispatch::NamedArgument,
            *distinct,
            order_by,
            filter.as_deref(),
            name,
        )?;
        let [ScalarExpr::Literal(Value::Str(argument_name)), value] = args.as_slice() else {
            return Err(malformed_call_argument(
                "named argument marker must contain a string name and one value",
            ));
        };
        let (value, explicit_variadic) = direct_variadic_argument(value)?;
        if !explicit_variadic
            && matches!(
                value,
                ScalarExpr::Func { binding, .. }
                    if binding.as_ref().and_then(|binding| binding.dispatch)
                        == Some(FunctionDispatch::NamedArgument)
            )
        {
            return Err(malformed_call_argument(
                "call argument contains nested syntax markers",
            ));
        }
        return Ok(ScalarCallArgument {
            name: Some(argument_name),
            value,
            explicit_variadic,
        });
    }
    let (value, explicit_variadic) = direct_variadic_argument(expression)?;
    Ok(ScalarCallArgument {
        name: None,
        value,
        explicit_variadic,
    })
}

fn direct_variadic_argument(expression: &ScalarExpr) -> Result<(&ScalarExpr, bool), SQLError> {
    let ScalarExpr::Func {
        name,
        args,
        binding,
        distinct,
        order_by,
        filter,
    } = expression
    else {
        return Ok((expression, false));
    };
    if binding.as_ref().and_then(|binding| binding.dispatch)
        != Some(FunctionDispatch::VariadicArgument)
    {
        return Ok((expression, false));
    }
    validate_marker_shape(
        binding.as_ref(),
        FunctionDispatch::VariadicArgument,
        *distinct,
        order_by,
        filter.as_deref(),
        name,
    )?;
    let [value] = args.as_slice() else {
        return Err(malformed_call_argument(
            "VARIADIC argument marker must contain exactly one value",
        ));
    };
    if matches!(
        value,
        ScalarExpr::Func { binding, .. }
            if matches!(
                binding.as_ref().and_then(|binding| binding.dispatch),
                Some(FunctionDispatch::VariadicArgument | FunctionDispatch::NamedArgument)
            )
    ) {
        return Err(malformed_call_argument(
            "call argument contains nested syntax markers",
        ));
    }
    Ok((value, true))
}

fn validate_marker_shape(
    binding: Option<&FunctionBinding>,
    expected_dispatch: FunctionDispatch,
    distinct: bool,
    order_by: &[ScalarOrder],
    filter: Option<&ScalarExpr>,
    name: &str,
) -> Result<(), SQLError> {
    if binding.is_none_or(|binding| {
        !binding.builtin
            || binding.dispatch != Some(expected_dispatch)
            || !binding.argument_types.is_empty()
            || binding.invocation.is_some()
            || binding.resolution_error.is_some()
    }) || distinct
        || !order_by.is_empty()
        || filter.is_some()
    {
        return Err(malformed_call_argument(&format!(
            "{name} syntax marker contains function-call metadata"
        )));
    }
    Ok(())
}

fn malformed_call_argument(message: &str) -> SQLError {
    SQLError::Internal(format!("malformed call argument: {message}"))
}

/// Decode SQL call markers carried by expression plans without evaluating arguments.
pub fn analyze_expression_call_arguments(
    arguments: &[crate::plan::ExpressionPlan],
) -> Result<(Vec<ScalarCallArgument<'_>>, bool), SQLError> {
    let decoded = arguments
        .iter()
        .map(|argument| scalar_call_argument(&argument.scalar))
        .collect::<Result<Vec<_>, _>>()?;
    let explicit_variadic = validate_scalar_call_arguments(&decoded)?;
    Ok((decoded, explicit_variadic))
}

#[cfg(test)]
mod production_tests;