1use 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
32fn 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
98pub(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 let mode = json_i64_or_zero(expr, "parseMode")?;
191 Ok((query, mode))
192}
193
194pub 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}