1use super::{
10 binding::{bind_statement, ResolvedVariable, VariableResolver},
11 variable_conflicts::{statement_variable_names, VariableSiteResolver},
12 VariableConflict,
13};
14use crate::{
15 ast::{Expr, Statement},
16 SQLError, SQLParam,
17};
18
19pub fn parameter_type_mismatch(
20 number: usize,
21 actual: &crate::ColumnType,
22 expected: &crate::ColumnType,
23) -> crate::ast::DeferredSQLError {
24 crate::ast::DeferredSQLError {
25 sqlstate: "42804".into(),
26 message: format!(
27 "type of parameter {number} ({}) does not match that when preparing the plan ({})",
28 actual.display_name(),
29 expected.display_name()
30 ),
31 }
32}
33
34#[derive(Debug, Clone)]
35pub enum PLpgSQLVariableReference {
36 Name(String),
37 Qualified { qualifier: String, column: String },
38 Parameter(usize),
39}
40impl PLpgSQLVariableReference {
41 pub fn read(&self, resolver: &mut dyn VariableResolver) -> Result<SQLParam, SQLError> {
42 let value = match self {
43 Self::Name(name) => resolver.resolve_name(name),
44 Self::Qualified { qualifier, column } => resolver.resolve_qualified(qualifier, column),
45 Self::Parameter(index) => resolver.resolve_param(*index),
46 }?
47 .ok_or_else(|| {
48 SQLError::Internal("prepared PL/pgSQL variable no longer resolves".into())
49 })?;
50 Ok(parameter(value, resolver))
51 }
52}
53fn parameter(value: ResolvedVariable, resolver: &dyn VariableResolver) -> SQLParam {
54 match value
55 .declared_type
56 .as_deref()
57 .and_then(|name| resolver.parameter_type(name))
58 {
59 Some(ty) => SQLParam::typed_scalar(value.value, ty),
60 None => SQLParam::scalar(value.value),
61 }
62}
63
64pub struct PLpgSQLVariableBindings {
65 pub statement: Statement,
66 pub references: Vec<PLpgSQLVariableReference>,
67 pub parameters: Vec<SQLParam>,
68}
69
70pub fn parameterize_statement_variables(
73 statement: &Statement,
74 resolver: &mut dyn VariableResolver,
75 conflict: VariableConflict,
76 resolve_sites: VariableSiteResolver<'_>,
77) -> Result<PLpgSQLVariableBindings, SQLError> {
78 let keep = statement_variable_names(statement, resolver, conflict, resolve_sites)?;
79 let mut binding = VariableBindings {
80 inner: resolver,
81 keep: &keep,
82 next_name: 0,
83 references: Vec::new(),
84 parameters: Vec::new(),
85 };
86 let statement = bind_statement(statement, &mut binding)?;
87 Ok(PLpgSQLVariableBindings {
88 statement,
89 references: binding.references,
90 parameters: binding.parameters,
91 })
92}
93struct VariableBindings<'a> {
94 inner: &'a mut dyn VariableResolver,
95 keep: &'a [bool],
96 next_name: usize,
97 references: Vec<PLpgSQLVariableReference>,
98 parameters: Vec<SQLParam>,
99}
100impl VariableBindings<'_> {
101 fn mark(
102 &mut self,
103 reference: PLpgSQLVariableReference,
104 value: Option<ResolvedVariable>,
105 named: bool,
106 ) -> Option<Expr> {
107 let value = value?;
108 if named {
109 let site = self.next_name;
110 self.next_name += 1;
111 if self.keep.get(site).copied().unwrap_or(false) {
112 return None;
113 }
114 }
115 self.references.push(reference);
116 self.parameters.push(parameter(value, self.inner));
117 Some(Expr::Param(self.references.len()))
118 }
119}
120impl VariableResolver for VariableBindings<'_> {
121 fn resolve_name(&mut self, name: &str) -> Result<Option<ResolvedVariable>, SQLError> {
122 self.inner.resolve_name(name)
123 }
124 fn resolve_qualified(
125 &mut self,
126 qualifier: &str,
127 column: &str,
128 ) -> Result<Option<ResolvedVariable>, SQLError> {
129 self.inner.resolve_qualified(qualifier, column)
130 }
131 fn resolve_param(&mut self, index: usize) -> Result<Option<ResolvedVariable>, SQLError> {
132 self.inner.resolve_param(index)
133 }
134 fn rewrite_name(&mut self, name: &str) -> Result<Option<Expr>, SQLError> {
135 let value = self.inner.resolve_name(name)?;
136 Ok(self.mark(
137 PLpgSQLVariableReference::Name(name.to_string()),
138 value,
139 true,
140 ))
141 }
142 fn rewrite_qualified(
143 &mut self,
144 qualifier: &str,
145 column: &str,
146 ) -> Result<Option<Expr>, SQLError> {
147 let value = self.inner.resolve_qualified(qualifier, column)?;
148 Ok(self.mark(
149 PLpgSQLVariableReference::Qualified {
150 qualifier: qualifier.to_string(),
151 column: column.to_string(),
152 },
153 value,
154 true,
155 ))
156 }
157 fn rewrite_param(&mut self, index: usize) -> Result<Option<Expr>, SQLError> {
158 let value = self.inner.resolve_param(index)?;
159 Ok(self.mark(PLpgSQLVariableReference::Parameter(index), value, false))
160 }
161}
162
163#[cfg(test)]
164mod tests;