Skip to main content

uqa_sql/ir/
call_arguments.rs

1//
2// Unified Query Algebra
3//
4// Copyright (c) 2023-2026 Cognica, Inc.
5//
6
7//! Named and explicit variadic SQL argument validation.
8
9use super::{ScalarExpr, ScalarOrder};
10use crate::ast::{FunctionBinding, FunctionDispatch};
11use crate::SQLError;
12use uqa_core::{
13    memory::{Produced, ProductionControl, ProductionVec},
14    Value,
15};
16
17/// A SQL call argument after removing the compiler's named and explicit `VARIADIC` syntax markers.
18#[doc(hidden)]
19#[derive(Debug, Clone, Copy, PartialEq)]
20pub struct ScalarCallArgument<'a> {
21    pub name: Option<&'a str>,
22    pub value: &'a ScalarExpr,
23    pub explicit_variadic: bool,
24}
25
26/// Decode and validate all compiler-owned call-argument markers. `PostgreSQL` permits one explicit `VARIADIC` argument and requires it to be the final argument.
27#[doc(hidden)]
28pub fn scalar_call_arguments(
29    arguments: &[ScalarExpr],
30) -> Result<Vec<ScalarCallArgument<'_>>, SQLError> {
31    scalar_call_arguments_with_control(arguments, &ProductionControl::uncontrolled()).map(
32        |decoded| {
33            decoded
34                .into_uncontrolled()
35                .expect("ordinary call argument decoding has no reservation")
36        },
37    )
38}
39
40/// Decode the same borrowed markers into an admitted temporary container. Names and expression nodes remain borrowed from the input IR owner.
41pub fn scalar_call_arguments_with_control<'a>(
42    arguments: &'a [ScalarExpr],
43    control: &ProductionControl<'_>,
44) -> Result<Produced<Vec<ScalarCallArgument<'a>>>, SQLError> {
45    let mut decoded = ProductionVec::new(*control);
46    decoded.reserve(arguments.len())?;
47    for argument in arguments {
48        control.check()?;
49        decoded.push_copy(scalar_call_argument(argument)?)?;
50    }
51    validate_scalar_call_arguments(&decoded)?;
52    decoded.finish().map_err(Into::into)
53}
54
55/// Validate cross-argument invariants after individual syntax markers have been decoded, returning whether the call used explicit `VARIADIC` syntax.
56#[doc(hidden)]
57pub fn validate_scalar_call_arguments(
58    arguments: &[ScalarCallArgument<'_>],
59) -> Result<bool, SQLError> {
60    let mut count = 0;
61    let mut last_position = None;
62    for (position, argument) in arguments.iter().enumerate() {
63        if argument.explicit_variadic {
64            count += 1;
65            last_position = Some(position);
66        }
67    }
68    if count > 1 {
69        return Err(malformed_call_argument(
70            "call contains more than one explicit VARIADIC argument",
71        ));
72    }
73    if last_position.is_some_and(|position| position + 1 != arguments.len()) {
74        return Err(malformed_call_argument(
75            "explicit VARIADIC argument must be the final call argument",
76        ));
77    }
78    Ok(count != 0)
79}
80
81/// Decode one compiler-owned call-argument marker. Use [`scalar_call_arguments`] for a complete call so duplicate and ordering invariants are also checked.
82#[doc(hidden)]
83pub fn scalar_call_argument(expression: &ScalarExpr) -> Result<ScalarCallArgument<'_>, SQLError> {
84    let ScalarExpr::Func {
85        order_syntax,
86        name,
87        args,
88        binding,
89        distinct,
90        order_by,
91        filter,
92    } = expression
93    else {
94        return Ok(ScalarCallArgument {
95            name: None,
96            value: expression,
97            explicit_variadic: false,
98        });
99    };
100    if binding.as_ref().and_then(|binding| binding.dispatch)
101        == Some(FunctionDispatch::NamedArgument)
102    {
103        validate_marker_shape(
104            binding.as_ref(),
105            FunctionDispatch::NamedArgument,
106            *distinct,
107            order_by,
108            *order_syntax,
109            filter.as_deref(),
110            name,
111        )?;
112        let [ScalarExpr::Literal(Value::Str(argument_name)), value] = args.as_slice() else {
113            return Err(malformed_call_argument(
114                "named argument marker must contain a string name and one value",
115            ));
116        };
117        let (value, explicit_variadic) = direct_variadic_argument(value)?;
118        if !explicit_variadic
119            && matches!(
120                value,
121                ScalarExpr::Func { binding, .. }
122                    if binding.as_ref().and_then(|binding| binding.dispatch)
123                        == Some(FunctionDispatch::NamedArgument)
124            )
125        {
126            return Err(malformed_call_argument(
127                "call argument contains nested syntax markers",
128            ));
129        }
130        return Ok(ScalarCallArgument {
131            name: Some(argument_name),
132            value,
133            explicit_variadic,
134        });
135    }
136    let (value, explicit_variadic) = direct_variadic_argument(expression)?;
137    Ok(ScalarCallArgument {
138        name: None,
139        value,
140        explicit_variadic,
141    })
142}
143
144fn direct_variadic_argument(expression: &ScalarExpr) -> Result<(&ScalarExpr, bool), SQLError> {
145    let ScalarExpr::Func {
146        order_syntax,
147        name,
148        args,
149        binding,
150        distinct,
151        order_by,
152        filter,
153    } = expression
154    else {
155        return Ok((expression, false));
156    };
157    if binding.as_ref().and_then(|binding| binding.dispatch)
158        != Some(FunctionDispatch::VariadicArgument)
159    {
160        return Ok((expression, false));
161    }
162    validate_marker_shape(
163        binding.as_ref(),
164        FunctionDispatch::VariadicArgument,
165        *distinct,
166        order_by,
167        *order_syntax,
168        filter.as_deref(),
169        name,
170    )?;
171    let [value] = args.as_slice() else {
172        return Err(malformed_call_argument(
173            "VARIADIC argument marker must contain exactly one value",
174        ));
175    };
176    if matches!(
177        value,
178        ScalarExpr::Func { binding, .. }
179            if matches!(
180                binding.as_ref().and_then(|binding| binding.dispatch),
181                Some(FunctionDispatch::VariadicArgument | FunctionDispatch::NamedArgument)
182            )
183    ) {
184        return Err(malformed_call_argument(
185            "call argument contains nested syntax markers",
186        ));
187    }
188    Ok((value, true))
189}
190
191fn validate_marker_shape(
192    binding: Option<&FunctionBinding>,
193    expected_dispatch: FunctionDispatch,
194    distinct: bool,
195    order_by: &[ScalarOrder],
196    order_syntax: crate::ast::FunctionOrderSyntax,
197    filter: Option<&ScalarExpr>,
198    name: &str,
199) -> Result<(), SQLError> {
200    if binding.is_none_or(|binding| {
201        !binding.builtin
202            || binding.dispatch != Some(expected_dispatch)
203            || !binding.argument_types.is_empty()
204            || binding.invocation.is_some()
205            || binding.resolution_error.is_some()
206    }) || distinct
207        || !order_by.is_empty()
208        || order_syntax == crate::ast::FunctionOrderSyntax::WithinGroup
209        || filter.is_some()
210    {
211        return Err(malformed_call_argument(&format!(
212            "{name} syntax marker contains function-call metadata"
213        )));
214    }
215    Ok(())
216}
217
218fn malformed_call_argument(message: &str) -> SQLError {
219    SQLError::Internal(format!("malformed call argument: {message}"))
220}
221
222/// Decode SQL call markers carried by expression plans without evaluating arguments.
223pub fn analyze_expression_call_arguments(
224    arguments: &[crate::plan::ExpressionPlan],
225) -> Result<(Vec<ScalarCallArgument<'_>>, bool), SQLError> {
226    let decoded = arguments
227        .iter()
228        .map(|argument| scalar_call_argument(&argument.scalar))
229        .collect::<Result<Vec<_>, _>>()?;
230    let explicit_variadic = validate_scalar_call_arguments(&decoded)?;
231    Ok((decoded, explicit_variadic))
232}
233
234#[cfg(test)]
235mod production_tests;