1use super::walk_schema_expr_mut;
10use crate::ast::{ColumnType, Expr};
11use crate::expr::EngineHook;
12use crate::plan::{QueryPlan, UnifiedPlan};
13use crate::{SQLError, ScalarExpr};
14use uqa_core::{ArrayValue, Value};
15
16pub trait OidAliasInput {
18 fn resolve_oid_alias_input(&self, ty: &ColumnType, name: &str)
19 -> Result<Option<i64>, SQLError>;
20}
21
22impl<T: EngineHook + ?Sized> OidAliasInput for T {
23 fn resolve_oid_alias_input(
24 &self,
25 ty: &ColumnType,
26 name: &str,
27 ) -> Result<Option<i64>, SQLError> {
28 match ty {
29 ColumnType::Regclass => EngineHook::resolve_regclass_input(self, name),
30 ColumnType::Regtype => EngineHook::resolve_regtype_input(self, name),
31 ColumnType::Regproc => EngineHook::resolve_regproc(self, name),
32 ColumnType::Regprocedure => EngineHook::resolve_regprocedure_input(self, name),
33 ColumnType::Regnamespace => EngineHook::resolve_regnamespace(self, name),
34 ColumnType::Regrole => EngineHook::resolve_regrole(self, name),
35 other => Err(SQLError::Internal(format!(
36 "{} is not an OID alias type read at analysis",
37 other.sql_name()
38 ))),
39 }
40 }
41}
42
43fn is_alias(ty: &ColumnType) -> bool {
45 matches!(
46 ty,
47 ColumnType::Regclass
48 | ColumnType::Regtype
49 | ColumnType::Regproc
50 | ColumnType::Regprocedure
51 | ColumnType::Regnamespace
52 )
53}
54
55fn alias_type(ty: &str) -> Option<(ColumnType, bool)> {
57 if !ty
58 .as_bytes()
59 .windows(3)
60 .any(|window| window.eq_ignore_ascii_case(b"reg"))
61 {
62 return None;
63 }
64 match ColumnType::from_sql_name(ty).ok()? {
65 ColumnType::Array(element) if is_alias(&element) => Some((*element, true)),
66 element if is_alias(&element) => Some((element, false)),
67 _ => None,
68 }
69}
70
71fn missing_object(ty: &ColumnType, name: &str) -> SQLError {
73 let (sqlstate, object) = match ty {
74 ColumnType::Regclass => ("42P01", "relation"),
75 ColumnType::Regtype => ("42704", "type"),
76 ColumnType::Regrole => ("42704", "role"),
77 ColumnType::Regproc | ColumnType::Regprocedure => ("42883", "function"),
78 _ => ("3F000", "schema"),
79 };
80 SQLError::Routine {
81 sqlstate: sqlstate.into(),
82 message: format!("{object} \"{name}\" does not exist"),
83 }
84}
85
86fn read_name<C: OidAliasInput + ?Sized>(
88 catalog: &C,
89 ty: &ColumnType,
90 name: &str,
91) -> Result<i64, SQLError> {
92 catalog
93 .resolve_oid_alias_input(ty, name)?
94 .ok_or_else(|| missing_object(ty, name))
95}
96
97fn read_constant<C: OidAliasInput + ?Sized>(
99 catalog: &C,
100 ty: &ColumnType,
101 text: &str,
102 array: bool,
103) -> Result<Value, SQLError> {
104 if !array {
105 return read_name(catalog, ty, text).map(Value::Int);
106 }
107 let array = crate::expr::parse_pg_array_literal(text)?;
108 let lower_bounds = array.lower_bounds().to_vec();
109 let mut elements = array.into_elements();
110 read_array_elements(catalog, ty, &mut elements)?;
111 ArrayValue::with_lower_bounds(elements, lower_bounds)
112 .map(Value::Array)
113 .ok_or_else(|| {
114 SQLError::Internal(format!("{} array literal lost its shape", ty.sql_name()))
115 })
116}
117
118fn read_array_elements<C: OidAliasInput + ?Sized>(
119 catalog: &C,
120 ty: &ColumnType,
121 elements: &mut [Value],
122) -> Result<(), SQLError> {
123 for element in elements {
124 match element {
125 Value::Null => {}
126 Value::Str(name) => *element = Value::Int(read_name(catalog, ty, name)?),
127 Value::List(nested) => read_array_elements(catalog, ty, nested)?,
128 other => {
129 return Err(SQLError::TypeMismatch(format!(
130 "cannot read {other:?} as {}",
131 ty.sql_name(),
132 )))
133 }
134 }
135 }
136 Ok(())
137}
138
139pub(crate) fn read_unknown_constant(
141 catalog: &dyn OidAliasInput,
142 ty: &ColumnType,
143 text: &str,
144) -> Result<Option<Value>, SQLError> {
145 match ty {
146 ColumnType::Array(element)
147 if is_alias(element) || matches!(element.as_ref(), ColumnType::Regrole) =>
148 {
149 read_constant(catalog, element, text, true).map(Some)
150 }
151 scalar if is_alias(scalar) || matches!(scalar, ColumnType::Regrole) => {
152 read_constant(catalog, scalar, text, false).map(Some)
153 }
154 _ => Ok(None),
155 }
156}
157
158fn constant_type(ty: ColumnType, array: bool) -> ColumnType {
160 if array {
161 ColumnType::Array(Box::new(ty))
162 } else {
163 ty
164 }
165}
166
167pub fn read_oid_alias_constants<C: OidAliasInput + ?Sized>(
169 catalog: &C,
170 expression: &mut Expr,
171) -> Result<(), SQLError> {
172 let mut failure = None;
173 let outcome = walk_schema_expr_mut(expression, &mut |node| {
174 let Expr::Cast { expr, ty, .. } = node else {
175 return Ok(());
176 };
177 let Some((alias, array)) = alias_type(ty) else {
178 return Ok(());
179 };
180 let Expr::Literal(Value::Str(text)) = expr.as_ref() else {
181 return Ok(());
182 };
183 match read_constant(catalog, &alias, text, array) {
184 Ok(value) => {
185 **expr = Expr::TypedLiteral {
186 value,
187 ty: constant_type(alias, array).catalog_name(),
188 };
189 Ok(())
190 }
191 Err(error) => {
192 failure = Some(error);
193 Err(String::new())
194 }
195 }
196 });
197 match (outcome, failure) {
198 (Ok(()), _) => Ok(()),
199 (Err(_), Some(error)) => Err(error),
200 (Err(message), None) => Err(SQLError::Internal(message)),
201 }
202}
203
204fn sequence_argument_mut(expression: &mut ScalarExpr) -> Option<&mut ScalarExpr> {
206 let ScalarExpr::Func { name, args, .. } = expression else {
207 return None;
208 };
209 if !is_sequence_function(name) {
210 return None;
211 }
212 args.first_mut()
213 .filter(|argument| matches!(argument, ScalarExpr::Literal(Value::Str(_))))
214}
215
216fn sequence_argument(expression: &ScalarExpr) -> Option<&str> {
217 let ScalarExpr::Func { name, args, .. } = expression else {
218 return None;
219 };
220 if !is_sequence_function(name) {
221 return None;
222 }
223 match args.first() {
224 Some(ScalarExpr::Literal(Value::Str(text))) => Some(text),
225 _ => None,
226 }
227}
228
229fn is_sequence_function(name: &str) -> bool {
230 let lower = name.to_ascii_lowercase();
231 let local = lower.strip_prefix("pg_catalog.").unwrap_or(&lower);
232 matches!(local, "nextval" | "currval" | "setval")
233 && (!lower.contains('.') || lower.starts_with("pg_catalog."))
234}
235
236fn read_scalar_constant<C: OidAliasInput + ?Sized>(
238 catalog: &C,
239 expression: &mut ScalarExpr,
240 keep_relations: bool,
241 failure: &mut Option<SQLError>,
242) {
243 if failure.is_some() {
244 return;
245 }
246 if let Some(argument) = sequence_argument_mut(expression) {
247 let ScalarExpr::Literal(Value::Str(text)) = &*argument else {
248 unreachable!("sequence argument is an unknown literal");
249 };
250 match read_constant(catalog, &ColumnType::Regclass, text, false) {
251 Ok(value) => {
252 if !keep_relations {
253 *argument = ScalarExpr::TypedLiteral {
254 value,
255 ty: ColumnType::Regclass.catalog_name(),
256 bound_type: Some(ColumnType::Regclass),
257 parameter_index: None,
258 };
259 }
260 }
261 Err(error) => *failure = Some(error),
262 }
263 return;
264 }
265 let ScalarExpr::Cast { expr, ty, .. } = expression else {
266 return;
267 };
268 let Some((alias, array)) = alias_type(ty) else {
269 return;
270 };
271 let ScalarExpr::Literal(Value::Str(text)) = expr.as_ref() else {
272 return;
273 };
274 match read_constant(catalog, &alias, text, array) {
275 Ok(value) => {
276 if keep_relations && matches!(alias, ColumnType::Regclass) {
277 return;
278 }
279 let bound_type = constant_type(alias, array);
280 **expr = ScalarExpr::TypedLiteral {
281 value,
282 ty: bound_type.catalog_name(),
283 bound_type: Some(bound_type),
284 parameter_index: None,
285 };
286 }
287 Err(error) => *failure = Some(error),
288 }
289}
290
291pub fn read_oid_alias_constants_in_plan<C: OidAliasInput + ?Sized>(
293 catalog: &C,
294 plan: &mut QueryPlan,
295) -> Result<(), SQLError> {
296 let mut failure = None;
297 plan.rewrite_scalar_expressions(&mut |root| {
298 root.visit_mut(&mut |expression| {
299 read_scalar_constant(catalog, expression, false, &mut failure);
300 });
301 });
302 failure.map_or(Ok(()), Err)
303}
304
305pub fn read_prepared_oid_alias_constants<C: OidAliasInput + ?Sized>(
307 catalog: &C,
308 plan: &mut UnifiedPlan,
309) -> Result<(), SQLError> {
310 let mut failure = None;
311 plan.rewrite_scalar_expressions(&mut |root| {
312 root.visit_mut(&mut |expression| {
313 read_scalar_constant(catalog, expression, false, &mut failure);
314 });
315 });
316 failure.map_or(Ok(()), Err)
317}
318
319pub fn check_statement_oid_alias_constants<C: OidAliasInput + ?Sized>(
321 catalog: &C,
322 plan: &UnifiedPlan,
323) -> Result<(), SQLError> {
324 let mut failure = None;
325 plan.visit_scalar_expressions(&mut |root| {
326 root.visit(&mut |expression| {
327 if failure.is_some() {
328 return;
329 }
330 if let Some(text) = sequence_argument(expression) {
331 if let Err(error) = read_constant(catalog, &ColumnType::Regclass, text, false) {
332 failure = Some(error);
333 }
334 return;
335 }
336 let ScalarExpr::Cast { expr, ty, .. } = expression else {
337 return;
338 };
339 let Some((alias, array)) = alias_type(ty) else {
340 return;
341 };
342 let ScalarExpr::Literal(Value::Str(text)) = expr.as_ref() else {
343 return;
344 };
345 if let Err(error) = read_constant(catalog, &alias, text, array) {
346 failure = Some(error);
347 }
348 });
349 });
350 failure.map_or(Ok(()), Err)
351}