uqa_sql/plpgsql/
lowering_expression.rs1use super::{
10 expect_tag, json_i64_or_zero, require_nonempty_str, Expr, JSONValue, PLpgSQLCursorArgument,
11 Result, SQLError, Statement,
12};
13
14pub(super) fn lower_expr_list(raw: Option<&JSONValue>) -> Result<Vec<Expr>> {
15 let Some(raw) = raw else {
16 return Ok(Vec::new());
17 };
18 let list = raw
19 .as_array()
20 .ok_or_else(|| SQLError::Internal("PL/pgSQL expression list is not an array".into()))?;
21 let mut out = Vec::with_capacity(list.len());
22 for item in list {
23 out.push(lower_expr(item)?);
24 }
25 Ok(out)
26}
27
28pub(super) fn lower_expr(raw: &JSONValue) -> Result<Expr> {
31 let (query, mode) = expr_text(raw)?;
32 let parse_mode = match mode {
33 2 => pg_query::ParseMode::PlPgSqlExpr,
34 3 => pg_query::ParseMode::PlPgSqlAssign1,
35 4 => pg_query::ParseMode::PlPgSqlAssign2,
36 5 => pg_query::ParseMode::PlPgSqlAssign3,
37 other => {
38 return Err(SQLError::Internal(format!(
39 "PL/pgSQL scalar expression has invalid parse mode {other}"
40 )));
41 }
42 };
43 let node = parse_one_raw_node(&query, parse_mode)?;
44 match (mode, node.node.as_ref()) {
45 (2, Some(pg_query::NodeEnum::SelectStmt(select))) => {
46 compile_single_select_expression(select, &query)
47 }
48 (3..=5, Some(pg_query::NodeEnum::PlassignStmt(assign))) => {
49 let expected_names = i32::try_from(mode - 2).map_err(|_| {
50 SQLError::Internal(format!("invalid PL/pgSQL assignment parse mode {mode}"))
51 })?;
52 if assign.nnames != expected_names {
53 return Err(SQLError::Internal(format!(
54 "PL/pgSQL assignment parser returned {} target names for parse mode {mode}",
55 assign.nnames
56 )));
57 }
58 let value = assign
59 .val
60 .as_deref()
61 .ok_or_else(|| SQLError::Internal("PL/pgSQL assignment has no value".into()))?;
62 compile_single_select_expression(value, &query)
63 }
64 (_, Some(other)) => Err(SQLError::Internal(format!(
65 "PL/pgSQL parse mode {mode} returned unexpected node {other:?}"
66 ))),
67 (_, None) => Err(SQLError::Internal(
68 "PL/pgSQL expression parser returned an empty node".into(),
69 )),
70 }
71}
72
73pub(super) fn lower_full_statement(raw: &JSONValue) -> Result<Statement> {
76 lower_sourced_statement(raw).map(|(statement, _)| statement)
77}
78
79pub(super) fn lower_sourced_statement(raw: &JSONValue) -> Result<(Statement, String)> {
80 let (query, mode) = expr_text(raw)?;
81 if mode != 0 {
82 return Err(SQLError::Internal(format!(
83 "embedded PL/pgSQL statement has invalid parse mode {mode}"
84 )));
85 }
86 let mut stmts = crate::compile(&query)?;
87 match stmts.len() {
88 1 => Ok((stmts.remove(0), query)),
89 n => Err(SQLError::Internal(format!(
90 "embedded PL/pgSQL query compiled to {n} statements"
91 ))),
92 }
93}
94
95pub(super) fn lower_cursor_arguments(
96 raw: Option<&JSONValue>,
97) -> Result<Vec<PLpgSQLCursorArgument>> {
98 let Some(raw) = raw else {
99 return Ok(Vec::new());
100 };
101 let (query, mode) = expr_text(raw)?;
102 if mode != 2 {
103 return Err(SQLError::Internal(format!(
104 "PL/pgSQL cursor arguments have invalid parse mode {mode}"
105 )));
106 }
107 let node = parse_one_raw_node(&query, pg_query::ParseMode::PlPgSqlExpr)?;
108 let Some(pg_query::NodeEnum::SelectStmt(select)) = node.node.as_ref() else {
109 return Err(SQLError::Internal(format!(
110 "PL/pgSQL cursor arguments did not parse as a SELECT target list: {query}"
111 )));
112 };
113 validate_select_expression_envelope(select, &query)?;
114 if select.target_list.is_empty() {
115 return Err(SQLError::Parse(format!(
116 "PL/pgSQL cursor argument list is empty: {query}"
117 )));
118 }
119 let projections = crate::compiler::compile_pg_projections(&select.target_list)?;
120 Ok(projections
121 .into_iter()
122 .map(|projection| PLpgSQLCursorArgument {
123 name: projection.alias.map(|name| name.to_ascii_lowercase()),
124 expr: projection.expr,
125 })
126 .collect())
127}
128
129pub(super) fn expr_text(raw: &JSONValue) -> Result<(String, i64)> {
130 let expr = expect_tag(raw, "PLpgSQL_expr", "expression")?;
131 let query = require_nonempty_str(expr, "query", "PLpgSQL expression")?;
132 let mode = json_i64_or_zero(expr, "parseMode")?;
135 Ok((query, mode))
136}
137
138pub fn compile_expression_text(text: &str) -> Result<Expr> {
140 let node = parse_one_raw_node(text, pg_query::ParseMode::PlPgSqlExpr)?;
141 let Some(pg_query::NodeEnum::SelectStmt(select)) = node.node.as_ref() else {
142 return Err(SQLError::Parse(format!("not an expression: {text}")));
143 };
144 compile_single_select_expression(select, text)
145}
146
147fn parse_one_raw_node(text: &str, mode: pg_query::ParseMode) -> Result<pg_query::protobuf::Node> {
148 let parsed = pg_query::parse_with_mode(text, mode)?;
149 let mut statements = parsed.protobuf.stmts;
150 if statements.len() != 1 {
151 return Err(SQLError::Parse(format!(
152 "PL/pgSQL fragment parsed to {} statements: {text}",
153 statements.len()
154 )));
155 }
156 statements
157 .remove(0)
158 .stmt
159 .map(|node| *node)
160 .ok_or_else(|| SQLError::Internal("PL/pgSQL parser returned an empty statement".into()))
161}
162
163fn compile_single_select_expression(
164 select: &pg_query::protobuf::SelectStmt,
165 text: &str,
166) -> Result<Expr> {
167 if select.target_list.len() != 1 {
168 return Err(SQLError::Parse(format!("not a single expression: {text}")));
169 }
170 if validate_select_expression_envelope(select, text).is_err() {
171 return crate::compiler::compile_pg_select(select)
172 .map(Box::new)
173 .map(Expr::ScalarSubquery);
174 }
175 let Some(pg_query::NodeEnum::ResTarget(target)) = select.target_list[0].node.as_ref() else {
176 return Err(SQLError::Internal(
177 "PL/pgSQL expression target is not a ResTarget".into(),
178 ));
179 };
180 let value = target
181 .val
182 .as_deref()
183 .ok_or_else(|| SQLError::Internal("PL/pgSQL expression target has no value".into()))?;
184 crate::compiler::compile_pg_expression(value)
185}
186
187fn validate_select_expression_envelope(
188 select: &pg_query::protobuf::SelectStmt,
189 text: &str,
190) -> Result<()> {
191 if !select.distinct_clause.is_empty()
192 || select.into_clause.is_some()
193 || !select.from_clause.is_empty()
194 || select.where_clause.is_some()
195 || !select.group_clause.is_empty()
196 || select.group_distinct
197 || select.having_clause.is_some()
198 || !select.window_clause.is_empty()
199 || !select.values_lists.is_empty()
200 || !select.sort_clause.is_empty()
201 || select.limit_offset.is_some()
202 || select.limit_count.is_some()
203 || select.limit_option != pg_query::protobuf::LimitOption::Default as i32
204 || !select.locking_clause.is_empty()
205 || select.with_clause.is_some()
206 || select.op != pg_query::protobuf::SetOperation::SetopNone as i32
207 || select.all
208 || select.larg.is_some()
209 || select.rarg.is_some()
210 {
211 return Err(SQLError::Parse(format!(
212 "PL/pgSQL fragment contains non-expression SELECT state: {text}"
213 )));
214 }
215 Ok(())
216}
217
218#[cfg(test)]
219mod tests {
220 use super::*;
221
222 fn parsed_expression_select() -> pg_query::protobuf::SelectStmt {
223 let node = parse_one_raw_node("value + 1", pg_query::ParseMode::PlPgSqlExpr).unwrap();
224 let Some(pg_query::NodeEnum::SelectStmt(select)) = node.node else {
225 panic!("expression mode did not return SelectStmt");
226 };
227 *select
228 }
229
230 #[test]
231 fn expression_envelope_rejects_every_non_expression_select_field() {
232 let base = parsed_expression_select();
233 validate_select_expression_envelope(&base, "value + 1").unwrap();
234
235 let mut malformed = Vec::new();
236 let mut select = base.clone();
237 select
238 .distinct_clause
239 .push(pg_query::protobuf::Node::default());
240 malformed.push(select);
241 let mut select = base.clone();
242 select.into_clause = Some(Box::default());
243 malformed.push(select);
244 let mut select = base.clone();
245 select.group_distinct = true;
246 malformed.push(select);
247 let mut select = base.clone();
248 select.limit_option = pg_query::protobuf::LimitOption::WithTies as i32;
249 malformed.push(select);
250 let mut select = base.clone();
251 select.op = pg_query::protobuf::SetOperation::SetopUnion as i32;
252 malformed.push(select);
253 let mut select = base;
254 select.all = true;
255 malformed.push(select);
256
257 for select in malformed {
258 assert!(matches!(
259 validate_select_expression_envelope(&select, "malformed"),
260 Err(SQLError::Parse(message))
261 if message.contains("non-expression SELECT state")
262 ));
263 }
264 }
265
266 #[test]
267 fn expression_with_from_lowers_to_a_scalar_subquery() {
268 let expression = compile_expression_text("max(a) FROM xacttest").unwrap();
269 let Expr::ScalarSubquery(query) = expression else {
270 panic!("PL/pgSQL query-shaped expression was not preserved as a scalar subquery");
271 };
272 assert_eq!(query.projections.len(), 1);
273 assert!(query.from.is_some());
274 }
275}