1use 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
30pub 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
60fn 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
84pub 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
102pub 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
109pub(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("e_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("e_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("e_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
188pub(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
229pub(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
296fn 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 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 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}