Skip to main content

uqa_sql/plpgsql/
parsing.rs

1//
2// Unified Query Algebra
3//
4// Copyright (c) 2023-2026 Cognica, Inc.
5//
6
7//! PL/pgSQL parser invocation, datum lowering, and condition normalization.
8
9use super::lowering_expression::lower_sourced_statement;
10use super::options::{compile_options, CompileOptions, VariableConflict};
11use super::{
12    condition_sqlstate, ensure_single_tag, expect_tag, json_bool_or_false, json_kind,
13    json_optional_i64, json_usize_or_zero, lower_block, lower_cursor_scroll_options, lower_expr,
14    normalize_plpgsql_type, optional_array, require, require_nonempty_str,
15    validate_assignable_datum, CreateFunction, FunctionBody, FunctionParamMode, FunctionReturns,
16    JSONValue, PLpgSQLCompilationIdentity, PLpgSQLCompileMode, PLpgSQLCursor, PLpgSQLDatum,
17    PLpgSQLFunction, PLpgSQLRowField, PLpgSQLVar, Result, RoutineColumnTypeReference, SQLError,
18};
19
20pub fn parse_function(def: &CreateFunction) -> Result<PLpgSQLFunction> {
21    let FunctionBody::Source(body) = &def.body else {
22        return Err(SQLError::Internal(
23            "PL/pgSQL parser invoked on a SQL-standard body".into(),
24        ));
25    };
26    let text = synthesize_create_text(def, body, &|type_name| Ok(type_name.to_string()))?;
27    Ok(with_compile_options(parse_plpgsql_text(&text)?, body))
28}
29
30/// Parse a stored routine using the engine's catalog type snapshot.
31pub fn parse_function_with_catalog(
32    def: &CreateFunction,
33    catalog: &pg_query::PlpgsqlCatalog,
34) -> Result<PLpgSQLFunction> {
35    parse_function_with_catalog_mode(def, catalog, PLpgSQLCompileMode::Validate)
36}
37
38pub fn parse_function_with_catalog_mode(
39    def: &CreateFunction,
40    catalog: &pg_query::PlpgsqlCatalog,
41    mode: PLpgSQLCompileMode,
42) -> Result<PLpgSQLFunction> {
43    let FunctionBody::Source(body) = &def.body else {
44        return Err(SQLError::Internal(
45            "PL/pgSQL parser invoked on a SQL-standard body".into(),
46        ));
47    };
48    let text = synthesize_create_text(def, body, &|type_name| {
49        catalog_type_spelling(catalog, type_name)
50    })?;
51    Ok(with_compile_options(
52        lower_plpgsql_json(
53            &crate::parser::parse_plpgsql_mode(&text, Some(catalog), mode)?,
54            mode,
55        )?,
56        body,
57    ))
58}
59
60/// Signatures name user-defined types by identity; the synthesized declaration spells them by their current qualified name, which the snapshot resolves to the same OID.
61fn catalog_type_spelling(catalog: &pg_query::PlpgsqlCatalog, type_name: &str) -> Result<String> {
62    let Some(identity) = crate::ast::UserTypeIdentity::parse(type_name) else {
63        return Ok(type_name.to_string());
64    };
65    let missing = || SQLError::Internal(format!("cache lookup failed for type {}", identity.oid));
66    let ty = catalog
67        .types
68        .iter()
69        .find(|ty| ty.oid == identity.oid)
70        .ok_or_else(missing)?;
71    let schema = catalog
72        .namespaces
73        .iter()
74        .find_map(|(name, oid)| (*oid == ty.namespace_oid).then_some(name))
75        .ok_or_else(missing)?;
76    Ok(format!(
77        "{}.{}{}",
78        quote_ident(schema),
79        quote_ident(&ty.name),
80        "[]".repeat(identity.dimensions)
81    ))
82}
83
84/// Parse an anonymous block using the engine's catalog type snapshot.
85pub fn parse_do_block_with_catalog(
86    body: &str,
87    catalog: &pg_query::PlpgsqlCatalog,
88) -> Result<PLpgSQLFunction> {
89    let tag = fresh_dollar_tag(body);
90    Ok(with_compile_options(
91        lower_plpgsql_json(
92            &crate::parser::parse_plpgsql(
93                &format!("DO {tag}{body}{tag} LANGUAGE plpgsql;"),
94                Some(catalog),
95            )?,
96            PLpgSQLCompileMode::Validate,
97        )?,
98        body,
99    ))
100}
101
102/// Parse a `DO $$ ... $$` body through `PostgreSQL`'s native inline-code path.
103pub fn parse_do_block(body: &str) -> Result<PLpgSQLFunction> {
104    let tag = fresh_dollar_tag(body);
105    let text = format!("DO {tag}{body}{tag} LANGUAGE plpgsql;");
106    Ok(with_compile_options(parse_plpgsql_text(&text)?, body))
107}
108
109/// Canonical `CREATE FUNCTION` / `CREATE PROCEDURE` text used solely
110/// to feed the `PL/pgSQL` parser (parameter DEFAULTs are resolved at
111/// call time and intentionally omitted).
112pub(super) fn synthesize_create_text(
113    def: &CreateFunction,
114    body: &str,
115    spell_type: &dyn Fn(&str) -> Result<String>,
116) -> Result<String> {
117    let mut sql = String::new();
118    sql.push_str(if def.is_procedure {
119        "CREATE PROCEDURE "
120    } else {
121        "CREATE FUNCTION "
122    });
123    sql.push_str(&quote_ident(&def.name));
124    sql.push('(');
125    let mut first = true;
126    for p in &def.params {
127        if matches!(p.mode, FunctionParamMode::Table) {
128            continue;
129        }
130        if !first {
131            sql.push_str(", ");
132        }
133        first = false;
134        match p.mode {
135            FunctionParamMode::Out => sql.push_str("OUT "),
136            FunctionParamMode::InOut => sql.push_str("INOUT "),
137            FunctionParamMode::Variadic => sql.push_str("VARIADIC "),
138            FunctionParamMode::In | FunctionParamMode::Table => {}
139        }
140        if !p.name.is_empty() {
141            sql.push_str(&quote_ident(&p.name));
142            sql.push(' ');
143        }
144        sql.push_str(&spell_type(&p.type_name)?);
145    }
146    sql.push(')');
147    match &def.returns {
148        FunctionReturns::None => {}
149        FunctionReturns::Scalar { type_name } => {
150            sql.push_str(" RETURNS ");
151            sql.push_str(&spell_type(type_name)?);
152        }
153        FunctionReturns::SetOf { type_name } => {
154            sql.push_str(" RETURNS SETOF ");
155            sql.push_str(&spell_type(type_name)?);
156        }
157        FunctionReturns::Table => {
158            sql.push_str(" RETURNS TABLE(");
159            let mut first_col = true;
160            for p in &def.params {
161                if !matches!(p.mode, FunctionParamMode::Table) {
162                    continue;
163                }
164                if !first_col {
165                    sql.push_str(", ");
166                }
167                first_col = false;
168                sql.push_str(&quote_ident(&p.name));
169                sql.push(' ');
170                sql.push_str(&spell_type(&p.type_name)?);
171            }
172            sql.push(')');
173        }
174    }
175    let tag = fresh_dollar_tag(body);
176    sql.push_str(" AS ");
177    sql.push_str(&tag);
178    sql.push_str(body);
179    sql.push_str(&tag);
180    sql.push_str(" LANGUAGE plpgsql;");
181    Ok(sql)
182}
183
184pub(super) fn quote_ident(name: &str) -> String {
185    format!("\"{}\"", name.replace('"', "\"\""))
186}
187
188/// Dollar-quote tag guaranteed not to collide with the body text.
189pub(super) fn fresh_dollar_tag(body: &str) -> String {
190    let mut n = 0usize;
191    loop {
192        let tag = if n == 0 {
193            "$$".to_string()
194        } else {
195            format!("$plpgsql{n}$")
196        };
197        if !body.contains(&tag) {
198            return tag;
199        }
200        n += 1;
201    }
202}
203
204pub(super) fn parse_plpgsql_text(text: &str) -> Result<PLpgSQLFunction> {
205    lower_plpgsql_json(
206        &crate::parser::parse_plpgsql(text, None)?,
207        PLpgSQLCompileMode::Validate,
208    )
209}
210
211fn lower_plpgsql_json(json: &JSONValue, mode: PLpgSQLCompileMode) -> Result<PLpgSQLFunction> {
212    crate::parser::without_notices(|| lower_parsed_plpgsql(json, mode))
213}
214
215fn lower_parsed_plpgsql(json: &JSONValue, mode: PLpgSQLCompileMode) -> Result<PLpgSQLFunction> {
216    let functions = json
217        .as_array()
218        .ok_or_else(|| SQLError::Internal("PL/pgSQL parse returned no function list".into()))?;
219    if functions.len() != 1 {
220        return Err(SQLError::Internal(format!(
221            "PL/pgSQL parse returned {} functions; expected exactly one",
222            functions.len()
223        )));
224    }
225    let function = expect_tag(&functions[0], "PLpgSQL_function", "parsed function")?;
226    lower_function(function, mode)
227}
228
229// ---------------------------------------------------------------------
230// JSON lowering
231// ---------------------------------------------------------------------
232
233pub(super) fn lower_function(
234    function: &JSONValue,
235    mode: PLpgSQLCompileMode,
236) -> Result<PLpgSQLFunction> {
237    let raw_datums = function
238        .get("datums")
239        .and_then(JSONValue::as_array)
240        .ok_or_else(|| SQLError::Internal("PL/pgSQL function without datums".into()))?;
241    let mut datums = Vec::with_capacity(raw_datums.len());
242    for raw in raw_datums {
243        datums.push(lower_datum(raw, mode)?);
244    }
245    validate_datums(&datums)?;
246    let trigger_datum = |field: &str, name: &str| -> Result<Option<usize>> {
247        let explicit = match json_optional_i64(function, field)? {
248            Some(index) if index >= 0 => {
249                let index = usize::try_from(index).map_err(|_| {
250                    SQLError::Internal(format!(
251                        "PL/pgSQL {field} {index} does not fit this platform"
252                    ))
253                })?;
254                if index >= datums.len() {
255                    return Err(SQLError::Internal(format!(
256                        "PL/pgSQL {field} has out-of-range datum index {index}"
257                    )));
258                }
259                Some(index)
260            }
261            Some(index) => {
262                return Err(SQLError::Internal(format!(
263                    "PL/pgSQL {field} has invalid datum index {index}"
264                )))
265            }
266            None => None,
267        };
268        Ok(explicit.or_else(|| {
269            datums.iter().position(|datum| {
270                datum
271                    .name()
272                    .is_some_and(|datum_name| datum_name.eq_ignore_ascii_case(name))
273            })
274        }))
275    };
276    let new_datum = trigger_datum("new_varno", "new")?;
277    let old_datum = trigger_datum("old_varno", "old")?;
278    let found_datum = datums
279        .iter()
280        .position(|d| matches!(d, PLpgSQLDatum::Var(v) if v.name.eq_ignore_ascii_case("found")));
281    let raw_action = require(function, "action")?;
282    let action = expect_tag(raw_action, "PLpgSQL_stmt_block", "function body")?;
283    let action = lower_block(action, &datums, mode)?;
284    Ok(PLpgSQLFunction {
285        compilation: PLpgSQLCompilationIdentity::default(),
286        datums,
287        action,
288        new_datum,
289        old_datum,
290        found_datum,
291        options: CompileOptions::default(),
292        variable_conflict: VariableConflict::default(),
293    })
294}
295
296/// The function a body compiles to, with the options the body declares.
297fn with_compile_options(mut function: PLpgSQLFunction, body: &str) -> PLpgSQLFunction {
298    function.options = compile_options(body);
299    function.variable_conflict = function.options.variable_conflict.unwrap_or_default();
300    function
301}
302
303fn has_percent_type_suffix(type_name: &str) -> bool {
304    type_name
305        .get(type_name.len().saturating_sub("%type".len())..)
306        .is_some_and(|suffix| suffix.eq_ignore_ascii_case("%type"))
307}
308
309fn lower_percent_type_reference(
310    datatype: &JSONValue,
311    variable_name: &str,
312) -> Result<RoutineColumnTypeReference> {
313    let identifiers = require(datatype, "typname_identifiers")?
314        .as_array()
315        .ok_or_else(|| {
316            SQLError::Internal(format!(
317                "PL/pgSQL variable `{variable_name}` type metadata `typname_identifiers` must be an array"
318            ))
319        })?;
320    let identifiers = identifiers
321        .iter()
322        .enumerate()
323        .map(|(index, identifier)| match identifier.as_str() {
324            Some(identifier) if !identifier.is_empty() => Ok(identifier.to_string()),
325            _ => Err(SQLError::Internal(format!(
326                "PL/pgSQL variable `{variable_name}` type metadata identifier {index} must be a non-empty string"
327            ))),
328        })
329        .collect::<Result<Vec<_>>>()?;
330    match identifiers.as_slice() {
331        [relation, column] => Ok(RoutineColumnTypeReference::new(
332            None,
333            relation.clone(),
334            column.clone(),
335        )),
336        [schema, relation, column] => Ok(RoutineColumnTypeReference::new(
337            Some(schema.clone()),
338            relation.clone(),
339            column.clone(),
340        )),
341        _ => Err(SQLError::TypeMismatch(format!(
342            "PL/pgSQL variable `{variable_name}` %TYPE must identify a relation column"
343        ))),
344    }
345}
346
347pub(super) fn lower_datum(raw: &JSONValue, mode: PLpgSQLCompileMode) -> Result<PLpgSQLDatum> {
348    ensure_single_tag(raw, "datum")?;
349    if let Some(var) = raw.get("PLpgSQL_var") {
350        let name = require_nonempty_str(var, "refname", "variable datum")?;
351        let datatype = require(var, "datatype")?;
352        let datatype = expect_tag(datatype, "PLpgSQL_type", "variable datatype")?;
353        let type_name = normalize_plpgsql_type(&require_nonempty_str(
354            datatype,
355            "typname",
356            "variable datatype",
357        )?);
358        if type_name.is_empty() {
359            return Err(SQLError::Internal(format!(
360                "PL/pgSQL variable `{name}` has an empty normalized type"
361            )));
362        }
363        let type_reference = has_percent_type_suffix(&type_name)
364            .then(|| lower_percent_type_reference(datatype, &name))
365            .transpose()?;
366        let default = match var.get("default_val") {
367            Some(node) => Some(lower_expr(node, mode)?),
368            None => None,
369        };
370        let cursor = if let Some(query) = var.get("cursor_explicit_expr") {
371            let (query, source_sql) = lower_sourced_statement(query, mode)?;
372            Some(PLpgSQLCursor {
373                query,
374                source_sql: source_sql.into(),
375                argument_row: match json_optional_i64(var, "cursor_explicit_argrow")? {
376                    None | Some(-1) => None,
377                    Some(index) if index >= 0 => Some(usize::try_from(index).map_err(|_| {
378                        SQLError::Internal(format!(
379                            "PL/pgSQL cursor `{name}` argument row {index} does not fit this platform"
380                        ))
381                    })?),
382                    Some(index) => {
383                        return Err(SQLError::Internal(format!(
384                            "PL/pgSQL cursor `{name}` has invalid argument row {index}"
385                        )));
386                    }
387                },
388                scroll: lower_cursor_scroll_options(var, "cursor declaration")?,
389            })
390        } else {
391            if var.get("cursor_explicit_argrow").is_some() {
392                return Err(SQLError::Internal(format!(
393                    "PL/pgSQL cursor variable `{name}` has arguments but no query"
394                )));
395            }
396            None
397        };
398        return Ok(PLpgSQLDatum::Var(Box::new(PLpgSQLVar {
399            name,
400            type_oid: json_optional_i64(datatype, "typoid")?
401                .map(|oid| {
402                    u32::try_from(oid).map_err(|_| {
403                        SQLError::Internal("PL/pgSQL variable has an invalid type OID".into())
404                    })
405                })
406                .transpose()?,
407            type_name,
408            type_reference,
409            default,
410            constant: json_bool_or_false(var, "isconst")?,
411            not_null: json_bool_or_false(var, "notnull")?,
412            cursor,
413            lineno: json_optional_i64(var, "lineno")?,
414        })));
415    }
416    if let Some(rec) = raw.get("PLpgSQL_rec") {
417        return Ok(PLpgSQLDatum::Rec {
418            name: require_nonempty_str(rec, "refname", "record datum")?,
419        });
420    }
421    if let Some(field) = raw.get("PLpgSQL_recfield") {
422        return Ok(PLpgSQLDatum::RecField {
423            field: require_nonempty_str(field, "fieldname", "record-field datum")?,
424            // libpg_query omits a zero-valued recparentno.
425            parent: json_usize_or_zero(field, "recparentno")?,
426        });
427    }
428    if let Some(row) = raw.get("PLpgSQL_row") {
429        return Ok(PLpgSQLDatum::Row {
430            fields: lower_row_fields(row)?,
431        });
432    }
433    Err(SQLError::Unsupported(format!(
434        "PL/pgSQL datum {}",
435        json_kind(raw)
436    )))
437}
438
439pub(super) fn lower_row_fields(row: &JSONValue) -> Result<Vec<PLpgSQLRowField>> {
440    let mut out = Vec::new();
441    if let Some(fields) = optional_array(row, "fields")? {
442        for f in fields {
443            // libpg_query's JSON dump omits zero-valued fields, so a
444            // missing varno means datum 0.
445            out.push(PLpgSQLRowField {
446                name: require_nonempty_str(f, "name", "row target field")?,
447                varno: json_usize_or_zero(f, "varno")?,
448            });
449        }
450    }
451    Ok(out)
452}
453
454pub(super) fn validate_datums(datums: &[PLpgSQLDatum]) -> Result<()> {
455    for (idx, datum) in datums.iter().enumerate() {
456        match datum {
457            PLpgSQLDatum::RecField { parent, .. } => {
458                let Some(parent_datum) = datums.get(*parent) else {
459                    return Err(SQLError::Internal(format!(
460                        "PL/pgSQL record-field datum {idx} references missing parent datum {parent}"
461                    )));
462                };
463                if !matches!(parent_datum, PLpgSQLDatum::Rec { .. }) {
464                    return Err(SQLError::Internal(format!(
465                        "PL/pgSQL record-field datum {idx} parent {parent} is not a record"
466                    )));
467                }
468            }
469            PLpgSQLDatum::Row { fields } => {
470                if fields.is_empty() {
471                    return Err(SQLError::Internal(format!(
472                        "PL/pgSQL row datum {idx} has no fields"
473                    )));
474                }
475                for field in fields {
476                    validate_assignable_datum(datums, field.varno, "row target field")?;
477                }
478            }
479            PLpgSQLDatum::Var(var) => {
480                if let Some(cursor) = &var.cursor {
481                    if var.type_name != "refcursor" {
482                        return Err(SQLError::Internal(format!(
483                            "PL/pgSQL bound cursor `{}` is not a refcursor datum",
484                            var.name
485                        )));
486                    }
487                    if let Some(argument_row) = cursor.argument_row {
488                        if !matches!(datums.get(argument_row), Some(PLpgSQLDatum::Row { .. })) {
489                            return Err(SQLError::Internal(format!(
490                                "PL/pgSQL cursor `{}` references invalid argument row {argument_row}",
491                                var.name
492                            )));
493                        }
494                    }
495                }
496            }
497            PLpgSQLDatum::Rec { .. } => {}
498        }
499    }
500    Ok(())
501}
502
503pub(super) fn normalize_condition(value: String, allow_others: bool) -> Result<String> {
504    let lower = value.to_ascii_lowercase();
505    if allow_others && lower == "others" {
506        return Ok(lower);
507    }
508    if condition_sqlstate(&lower).is_some() {
509        return Ok(lower);
510    }
511    let upper = value.to_ascii_uppercase();
512    if upper.len() == 5
513        && upper
514            .bytes()
515            .all(|byte| byte.is_ascii_uppercase() || byte.is_ascii_digit())
516    {
517        return Ok(upper);
518    }
519    Err(SQLError::Internal(format!(
520        "unrecognized PL/pgSQL exception condition `{value}`"
521    )))
522}