Skip to main content

uqa_sql/plpgsql/
lowering_expression.rs

1//
2// Unified Query Algebra
3//
4// Copyright (c) 2023-2026 Cognica, Inc.
5//
6
7//! Embedded SQL expression and assignment-target lowering.
8
9use super::{
10    expect_tag, json_i64_or_zero, require_nonempty_str, Expr, JSONValue, PLpgSQLCompileMode,
11    PLpgSQLCursorArgument, PLpgSQLCursorArguments, PLpgSQLExpression, PLpgSQLFragment,
12    PLpgSQLParseMode, PLpgSQLSource, PLpgSQLStatement, Result, SQLError, Statement,
13};
14
15pub(super) fn lower_expr_list(
16    raw: Option<&JSONValue>,
17    mode: PLpgSQLCompileMode,
18) -> Result<Vec<PLpgSQLExpression>> {
19    let Some(raw) = raw else {
20        return Ok(Vec::new());
21    };
22    let list = raw
23        .as_array()
24        .ok_or_else(|| SQLError::Internal("PL/pgSQL expression list is not an array".into()))?;
25    let mut out = Vec::with_capacity(list.len());
26    for item in list {
27        out.push(lower_expr(item, mode)?);
28    }
29    Ok(out)
30}
31
32/// Lower a `PLpgSQL_expr` node whose text is a scalar expression
33/// (parse modes 2 = expression, 3/4/5 = assignment source).
34fn source(raw: &JSONValue) -> Result<std::sync::Arc<PLpgSQLSource>> {
35    let (query, mode) = expr_text(raw)?;
36    Ok(std::sync::Arc::new(PLpgSQLSource {
37        query: query.into(),
38        mode: mode.try_into()?,
39    }))
40}
41
42pub(super) fn lower_expr(raw: &JSONValue, mode: PLpgSQLCompileMode) -> Result<PLpgSQLExpression> {
43    let source = source(raw)?;
44    if source.mode == PLpgSQLParseMode::Statement {
45        return Err(SQLError::Internal(
46            "PL/pgSQL scalar expression has statement parse mode".into(),
47        ));
48    }
49    let validation = mode
50        .validates()
51        .then(|| compile_expression_source(&source))
52        .transpose()?;
53    Ok(PLpgSQLFragment::new(source, validation))
54}
55
56pub(super) fn compile_expression_source(source: &PLpgSQLSource) -> Result<Expr> {
57    let query = &*source.query;
58    let mode = source.mode;
59    let parse_mode = mode.parser_mode();
60    let node = parse_one_raw_node(query, parse_mode)?;
61    match (mode, node.node.as_ref()) {
62        (PLpgSQLParseMode::Expression, Some(pg_query::NodeEnum::SelectStmt(select))) => {
63            compile_single_select_expression(select, query)
64        }
65        (
66            PLpgSQLParseMode::Assignment1
67            | PLpgSQLParseMode::Assignment2
68            | PLpgSQLParseMode::Assignment3,
69            Some(pg_query::NodeEnum::PlassignStmt(assign)),
70        ) => {
71            let expected_names = match mode {
72                PLpgSQLParseMode::Assignment1 => 1,
73                PLpgSQLParseMode::Assignment2 => 2,
74                PLpgSQLParseMode::Assignment3 => 3,
75                _ => unreachable!(),
76            };
77            if assign.nnames != expected_names {
78                return Err(SQLError::Internal(format!(
79                    "PL/pgSQL assignment parser returned {} target names for parse mode {mode:?}",
80                    assign.nnames
81                )));
82            }
83            let value = assign
84                .val
85                .as_deref()
86                .ok_or_else(|| SQLError::Internal("PL/pgSQL assignment has no value".into()))?;
87            compile_single_select_expression(value, query)
88        }
89        (_, Some(other)) => Err(SQLError::Internal(format!(
90            "PL/pgSQL parse mode {mode:?} returned unexpected node {other:?}"
91        ))),
92        (_, None) => Err(SQLError::Internal(
93            "PL/pgSQL expression parser returned an empty node".into(),
94        )),
95    }
96}
97
98/// Lower a `PLpgSQL_expr` node holding a complete SQL statement
99/// (parse mode 0: queries, PERFORM bodies, CALL statements).
100pub(super) fn lower_full_statement(
101    raw: &JSONValue,
102    mode: PLpgSQLCompileMode,
103) -> Result<PLpgSQLStatement> {
104    lower_sourced_statement(raw, mode).map(|(statement, _)| statement)
105}
106
107pub(super) fn lower_sourced_statement(
108    raw: &JSONValue,
109    mode: PLpgSQLCompileMode,
110) -> Result<(PLpgSQLStatement, String)> {
111    let source = source(raw)?;
112    if source.mode != PLpgSQLParseMode::Statement {
113        return Err(SQLError::Internal(
114            "embedded PL/pgSQL statement has expression parse mode".into(),
115        ));
116    }
117    let validation = mode
118        .validates()
119        .then(|| compile_statement_source(&source))
120        .transpose()?;
121    let query = source.query.to_string();
122    Ok((PLpgSQLFragment::new(source, validation), query))
123}
124
125pub(super) fn compile_statement_source(source: &PLpgSQLSource) -> Result<Statement> {
126    let mut statements = crate::compile(&source.query)?;
127    match statements.len() {
128        1 => Ok(statements.remove(0)),
129        n => Err(SQLError::Internal(format!(
130            "embedded PL/pgSQL query compiled to {n} statements"
131        ))),
132    }
133}
134
135pub(super) fn lower_cursor_arguments(
136    raw: Option<&JSONValue>,
137    mode: PLpgSQLCompileMode,
138) -> Result<PLpgSQLCursorArguments> {
139    let Some(raw) = raw else {
140        return Ok(PLpgSQLCursorArguments::empty());
141    };
142    let source = source(raw)?;
143    if source.mode != PLpgSQLParseMode::Expression {
144        return Err(SQLError::Internal(
145            "PL/pgSQL cursor arguments have non-expression parse mode".into(),
146        ));
147    }
148    let validation = mode
149        .validates()
150        .then(|| compile_cursor_arguments_source(&source))
151        .transpose()?;
152    Ok(PLpgSQLFragment::new(source, validation))
153}
154
155pub(super) fn compile_cursor_arguments_source(
156    source: &PLpgSQLSource,
157) -> Result<Vec<PLpgSQLCursorArgument>> {
158    let query = &*source.query;
159    if query.is_empty() {
160        return Ok(Vec::new());
161    }
162    let node = parse_one_raw_node(query, pg_query::ParseMode::PlPgSqlExpr)?;
163    let Some(pg_query::NodeEnum::SelectStmt(select)) = node.node.as_ref() else {
164        return Err(SQLError::Internal(format!(
165            "PL/pgSQL cursor arguments did not parse as a SELECT target list: {query}"
166        )));
167    };
168    validate_select_expression_envelope(select, query)?;
169    if select.target_list.is_empty() {
170        return Err(SQLError::Parse(format!(
171            "PL/pgSQL cursor argument list is empty: {query}"
172        )));
173    }
174    Ok(
175        crate::compiler::compile_pg_projections(&select.target_list)?
176            .into_iter()
177            .map(|projection| PLpgSQLCursorArgument {
178                name: projection.alias.map(|name| name.to_ascii_lowercase()),
179                expr: projection.expr,
180            })
181            .collect(),
182    )
183}
184
185pub(super) fn expr_text(raw: &JSONValue) -> Result<(String, i64)> {
186    let expr = expect_tag(raw, "PLpgSQL_expr", "expression")?;
187    let query = require_nonempty_str(expr, "query", "PLpgSQL expression")?;
188    // RAW_PARSE_DEFAULT is encoded as zero and therefore omitted by
189    // libpg_query's JSON serializer.
190    let mode = json_i64_or_zero(expr, "parseMode")?;
191    Ok((query, mode))
192}
193
194/// Compile a bare expression through `PostgreSQL`'s PL/pgSQL expression parser.
195pub fn compile_expression_text(text: &str) -> Result<Expr> {
196    let node = parse_one_raw_node(text, pg_query::ParseMode::PlPgSqlExpr)?;
197    let Some(pg_query::NodeEnum::SelectStmt(select)) = node.node.as_ref() else {
198        return Err(SQLError::Parse(format!("not an expression: {text}")));
199    };
200    compile_single_select_expression(select, text)
201}
202
203fn parse_one_raw_node(text: &str, mode: pg_query::ParseMode) -> Result<pg_query::protobuf::Node> {
204    let parsed = crate::parser::parse_with_mode(text, mode)?;
205    let mut statements = parsed.protobuf.stmts;
206    if statements.len() != 1 {
207        return Err(SQLError::Parse(format!(
208            "PL/pgSQL fragment parsed to {} statements: {text}",
209            statements.len()
210        )));
211    }
212    statements
213        .remove(0)
214        .stmt
215        .map(|node| *node)
216        .ok_or_else(|| SQLError::Internal("PL/pgSQL parser returned an empty statement".into()))
217}
218
219fn compile_single_select_expression(
220    select: &pg_query::protobuf::SelectStmt,
221    text: &str,
222) -> Result<Expr> {
223    if select.target_list.len() != 1 {
224        return Err(SQLError::Parse(format!("not a single expression: {text}")));
225    }
226    if validate_select_expression_envelope(select, text).is_err() {
227        return crate::compiler::compile_pg_select(select)
228            .map(Box::new)
229            .map(Expr::ScalarSubquery);
230    }
231    let Some(pg_query::NodeEnum::ResTarget(target)) = select.target_list[0].node.as_ref() else {
232        return Err(SQLError::Internal(
233            "PL/pgSQL expression target is not a ResTarget".into(),
234        ));
235    };
236    let value = target
237        .val
238        .as_deref()
239        .ok_or_else(|| SQLError::Internal("PL/pgSQL expression target has no value".into()))?;
240    crate::compiler::compile_pg_expression(value)
241}
242
243fn validate_select_expression_envelope(
244    select: &pg_query::protobuf::SelectStmt,
245    text: &str,
246) -> Result<()> {
247    if !select.distinct_clause.is_empty()
248        || select.into_clause.is_some()
249        || !select.from_clause.is_empty()
250        || select.where_clause.is_some()
251        || !select.group_clause.is_empty()
252        || select.group_distinct
253        || select.having_clause.is_some()
254        || !select.window_clause.is_empty()
255        || !select.values_lists.is_empty()
256        || !select.sort_clause.is_empty()
257        || select.limit_offset.is_some()
258        || select.limit_count.is_some()
259        || select.limit_option != pg_query::protobuf::LimitOption::Default as i32
260        || !select.locking_clause.is_empty()
261        || select.with_clause.is_some()
262        || select.op != pg_query::protobuf::SetOperation::SetopNone as i32
263        || select.all
264        || select.larg.is_some()
265        || select.rarg.is_some()
266    {
267        return Err(SQLError::Parse(format!(
268            "PL/pgSQL fragment contains non-expression SELECT state: {text}"
269        )));
270    }
271    Ok(())
272}
273
274#[cfg(test)]
275mod tests {
276    use super::*;
277
278    fn parsed_expression_select() -> pg_query::protobuf::SelectStmt {
279        let node = parse_one_raw_node("value + 1", pg_query::ParseMode::PlPgSqlExpr).unwrap();
280        let Some(pg_query::NodeEnum::SelectStmt(select)) = node.node else {
281            panic!("expression mode did not return SelectStmt");
282        };
283        *select
284    }
285
286    #[test]
287    fn expression_envelope_rejects_every_non_expression_select_field() {
288        let base = parsed_expression_select();
289        validate_select_expression_envelope(&base, "value + 1").unwrap();
290
291        let mut malformed = Vec::new();
292        let mut select = base.clone();
293        select
294            .distinct_clause
295            .push(pg_query::protobuf::Node::default());
296        malformed.push(select);
297        let mut select = base.clone();
298        select.into_clause = Some(Box::default());
299        malformed.push(select);
300        let mut select = base.clone();
301        select.group_distinct = true;
302        malformed.push(select);
303        let mut select = base.clone();
304        select.limit_option = pg_query::protobuf::LimitOption::WithTies as i32;
305        malformed.push(select);
306        let mut select = base.clone();
307        select.op = pg_query::protobuf::SetOperation::SetopUnion as i32;
308        malformed.push(select);
309        let mut select = base;
310        select.all = true;
311        malformed.push(select);
312
313        for select in malformed {
314            assert!(matches!(
315                validate_select_expression_envelope(&select, "malformed"),
316                Err(SQLError::Parse(message))
317                    if message.contains("non-expression SELECT state")
318            ));
319        }
320    }
321
322    #[test]
323    fn expression_with_from_lowers_to_a_scalar_subquery() {
324        let expression = compile_expression_text("max(a) FROM xacttest").unwrap();
325        let Expr::ScalarSubquery(query) = expression else {
326            panic!("PL/pgSQL query-shaped expression was not preserved as a scalar subquery");
327        };
328        assert_eq!(query.projections.len(), 1);
329        assert!(query.from.is_some());
330    }
331}