1use super::binding::{bind_expr, bind_statement, ResolvedVariable, VariableResolver};
12use super::options::VariableConflict;
13use super::{Expr, Projection, Result, SelectStmt, Statement};
14use crate::ast::InternalColumnRef;
15use crate::binding::VariableSiteResolution;
16use crate::{SQLError, SQLParam, ScalarExpr};
17use uqa_core::Value;
18
19struct VariableSite {
21 reference: Expr,
22 binding: Expr,
23}
24
25pub type VariableSiteResolver<'a> =
27 &'a mut dyn FnMut(
28 Statement,
29 Vec<SQLParam>,
30 Vec<ScalarExpr>,
31 ) -> std::result::Result<Vec<VariableSiteResolution>, SQLError>;
32
33struct VariableSiteMarker<'a> {
35 inner: &'a mut dyn VariableResolver,
36 sites: Vec<VariableSite>,
37}
38
39impl VariableSiteMarker<'_> {
40 fn mark(&mut self, reference: Expr, binding: Option<Expr>) -> Option<Expr> {
41 let binding = binding?;
42 self.sites.push(VariableSite { reference, binding });
43 Some(Expr::Param(self.sites.len()))
44 }
45}
46
47impl VariableResolver for VariableSiteMarker<'_> {
48 fn resolve_name(&mut self, name: &str) -> Result<Option<ResolvedVariable>> {
49 self.inner.resolve_name(name)
50 }
51
52 fn resolve_qualified(
53 &mut self,
54 qualifier: &str,
55 column: &str,
56 ) -> Result<Option<ResolvedVariable>> {
57 self.inner.resolve_qualified(qualifier, column)
58 }
59
60 fn resolve_param(&mut self, index: usize) -> Result<Option<ResolvedVariable>> {
61 self.inner.resolve_param(index)
62 }
63
64 fn rewrite_name(&mut self, name: &str) -> Result<Option<Expr>> {
65 let binding = self.inner.rewrite_name(name)?;
66 Ok(self.mark(Expr::Column(name.to_string()), binding))
67 }
68
69 fn rewrite_qualified(&mut self, qualifier: &str, column: &str) -> Result<Option<Expr>> {
70 let binding = self.inner.rewrite_qualified(qualifier, column)?;
71 Ok(self.mark(
72 Expr::QualifiedColumn {
73 qualifier: qualifier.to_string(),
74 column: column.to_string(),
75 },
76 binding,
77 ))
78 }
79
80 fn rewrite_qualified_star(&mut self, qualifier: &str) -> Result<Option<Vec<Expr>>> {
81 self.inner.rewrite_qualified_star(qualifier)
82 }
83
84 fn rewrite_qualified_whole_row(&mut self, qualifier: &str) -> Result<Option<Expr>> {
85 self.inner.rewrite_qualified_whole_row(qualifier)
86 }
87
88 fn rewrite_param(&mut self, index: usize) -> Result<Option<Expr>> {
89 self.inner.rewrite_param(index)
90 }
91
92 fn rewrite_internal(&mut self, column: InternalColumnRef) -> Result<Option<Expr>> {
93 self.inner.rewrite_internal(column)
94 }
95}
96
97struct VariableSiteFilter<'a> {
99 inner: &'a mut dyn VariableResolver,
100 keep: &'a [bool],
101 next: usize,
102}
103
104impl VariableSiteFilter<'_> {
105 fn filter(&mut self, binding: Option<Expr>) -> Option<Expr> {
106 let binding = binding?;
107 let site = self.next;
108 self.next += 1;
109 (!self.keep.get(site).copied().unwrap_or(false)).then_some(binding)
110 }
111}
112
113impl VariableResolver for VariableSiteFilter<'_> {
114 fn resolve_name(&mut self, name: &str) -> Result<Option<ResolvedVariable>> {
115 self.inner.resolve_name(name)
116 }
117
118 fn resolve_qualified(
119 &mut self,
120 qualifier: &str,
121 column: &str,
122 ) -> Result<Option<ResolvedVariable>> {
123 self.inner.resolve_qualified(qualifier, column)
124 }
125
126 fn resolve_param(&mut self, index: usize) -> Result<Option<ResolvedVariable>> {
127 self.inner.resolve_param(index)
128 }
129
130 fn rewrite_name(&mut self, name: &str) -> Result<Option<Expr>> {
131 let binding = self.inner.rewrite_name(name)?;
132 Ok(self.filter(binding))
133 }
134
135 fn rewrite_qualified(&mut self, qualifier: &str, column: &str) -> Result<Option<Expr>> {
136 let binding = self.inner.rewrite_qualified(qualifier, column)?;
137 Ok(self.filter(binding))
138 }
139
140 fn rewrite_qualified_star(&mut self, qualifier: &str) -> Result<Option<Vec<Expr>>> {
141 self.inner.rewrite_qualified_star(qualifier)
142 }
143
144 fn rewrite_qualified_whole_row(&mut self, qualifier: &str) -> Result<Option<Expr>> {
145 self.inner.rewrite_qualified_whole_row(qualifier)
146 }
147
148 fn rewrite_param(&mut self, index: usize) -> Result<Option<Expr>> {
149 self.inner.rewrite_param(index)
150 }
151
152 fn rewrite_internal(&mut self, column: InternalColumnRef) -> Result<Option<Expr>> {
153 self.inner.rewrite_internal(column)
154 }
155}
156
157fn ambiguous_variable_error(reference: &Expr) -> SQLError {
159 let name = match reference {
160 Expr::QualifiedColumn { qualifier, column } => format!("{qualifier}.{column}"),
161 Expr::Column(name) => name.clone(),
162 other => format!("{other:?}"),
163 };
164 SQLError::Diagnostic {
165 sqlstate: "42702".into(),
166 message: format!("column reference \"{name}\" is ambiguous"),
167 detail: Some("It could refer to either a PL/pgSQL variable or a table column.".into()),
168 hint: None,
169 }
170}
171
172pub(super) fn site_parameter(binding: &Expr, resolver: &dyn VariableResolver) -> SQLParam {
174 match binding {
175 Expr::TypedLiteral { value, ty } => resolver.parameter_type(ty).map_or_else(
176 || SQLParam::scalar(value.clone()),
177 |ty| SQLParam::typed_scalar(value.clone(), ty),
178 ),
179 Expr::Literal(value) => SQLParam::scalar(value.clone()),
180 _ => SQLParam::scalar(Value::Null),
181 }
182}
183
184fn site_name(reference: &Expr) -> ScalarExpr {
186 match reference {
187 Expr::QualifiedColumn { qualifier, column } => ScalarExpr::QualifiedColumn {
188 qualifier: qualifier.clone(),
189 column: column.clone(),
190 },
191 Expr::Column(name) => ScalarExpr::Column(name.clone()),
192 _ => ScalarExpr::Literal(Value::Null),
193 }
194}
195
196fn kept_names(
198 sites: &[VariableSite],
199 resolutions: &[VariableSiteResolution],
200 conflict: VariableConflict,
201) -> std::result::Result<Vec<bool>, SQLError> {
202 sites
203 .iter()
204 .zip(resolutions)
205 .map(|(site, resolution)| match (resolution, conflict) {
206 (VariableSiteResolution::Output, _)
207 | (VariableSiteResolution::Column, VariableConflict::UseColumn) => Ok(true),
208 (VariableSiteResolution::Column, VariableConflict::Error) => {
209 Err(ambiguous_variable_error(&site.reference))
210 }
211 (VariableSiteResolution::Column, VariableConflict::UseVariable)
212 | (VariableSiteResolution::Variable, _) => Ok(false),
213 })
214 .collect()
215}
216
217pub fn bind_statement_variables(
219 statement: &Statement,
220 resolver: &mut dyn VariableResolver,
221 conflict: VariableConflict,
222 resolve_sites: VariableSiteResolver<'_>,
223) -> Result<Statement> {
224 let keep = statement_variable_names(statement, resolver, conflict, resolve_sites)?;
225 bind_statement(
226 statement,
227 &mut VariableSiteFilter {
228 inner: resolver,
229 keep: &keep,
230 next: 0,
231 },
232 )
233}
234
235pub(super) fn statement_variable_names(
237 statement: &Statement,
238 resolver: &mut dyn VariableResolver,
239 conflict: VariableConflict,
240 resolve_sites: VariableSiteResolver<'_>,
241) -> Result<Vec<bool>> {
242 let mut numbering = VariableSiteMarker {
243 inner: resolver,
244 sites: Vec::new(),
245 };
246 let numbered = bind_statement(statement, &mut numbering)?;
247 let sites = numbering.sites;
248 if sites.is_empty() {
249 return Ok(Vec::new());
250 }
251 let resolutions = resolve_sites(
252 numbered,
253 sites
254 .iter()
255 .map(|site| site_parameter(&site.binding, resolver))
256 .collect(),
257 sites
258 .iter()
259 .map(|site| site_name(&site.reference))
260 .collect(),
261 )?;
262 kept_names(&sites, &resolutions, conflict)
263}
264
265pub fn bind_expression_variables(
267 expression: &Expr,
268 resolver: &mut dyn VariableResolver,
269 conflict: VariableConflict,
270 resolve_sites: VariableSiteResolver<'_>,
271) -> Result<Expr> {
272 let queries = expression.any_node(&|node| {
273 matches!(
274 node,
275 Expr::ScalarSubquery(_) | Expr::Exists { .. } | Expr::InSubquery { .. }
276 )
277 });
278 if !queries {
279 return bind_expr(expression, resolver);
280 }
281 let mut numbering = VariableSiteMarker {
282 inner: resolver,
283 sites: Vec::new(),
284 };
285 let numbered = bind_expr(expression, &mut numbering)?;
286 let sites = numbering.sites;
287 if sites.is_empty() {
288 return Ok(numbered);
289 }
290 let resolutions = resolve_sites(
291 expression_query(numbered),
292 sites
293 .iter()
294 .map(|site| site_parameter(&site.binding, resolver))
295 .collect(),
296 sites
297 .iter()
298 .map(|site| site_name(&site.reference))
299 .collect(),
300 )?;
301 let keep = kept_names(&sites, &resolutions, conflict)?;
302 bind_expr(
303 expression,
304 &mut VariableSiteFilter {
305 inner: resolver,
306 keep: &keep,
307 next: 0,
308 },
309 )
310}
311
312pub(crate) fn expression_query(expression: Expr) -> Statement {
314 Statement::Select(Box::new(SelectStmt {
315 windows: Vec::new(),
316 projections: vec![Projection {
317 expr: expression,
318 alias: None,
319 }],
320 values: Vec::new(),
321 from: None,
322 r#where: None,
323 group_by: Vec::new(),
324 grouping_sets: Vec::new(),
325 group_distinct: false,
326 having: None,
327 order_by: Vec::new(),
328 limit: None,
329 with_ties: false,
330 offset: None,
331 with: Vec::new(),
332 set_op: None,
333 distinct: false,
334 distinct_on: Vec::new(),
335 locking: Vec::new(),
336 }))
337}