1use pg_query::protobuf::Node;
10use pg_query::NodeEnum;
11
12use crate::ast::{ColumnType, IntervalFields, RangeSubtype};
13use crate::error::{Result, SQLError};
14
15use super::tree::extract_string;
16
17mod declarations;
18pub(crate) use declarations::compile_retained_type_declaration;
19pub(crate) use declarations::compile_retained_type_reference;
20pub(super) use declarations::preserve_alter_type_declaration;
21
22#[derive(Debug, Clone, PartialEq, Eq)]
24pub struct ParsedRegtypeName {
25 pub names: Vec<String>,
26 pub array_dimensions: usize,
27 pub has_type_modifiers: bool,
28}
29
30#[derive(Debug, Clone, PartialEq, Eq)]
32pub struct ParsedRegprocedureName {
33 pub names: Vec<String>,
34 pub argument_types: Option<Vec<ParsedRegtypeName>>,
35}
36
37const POSTGRES_IDENTIFIER_MAX_BYTES: usize = 63;
38const POSTGRES_FUNCTION_MAX_ARGUMENTS: usize = 100;
39
40fn scanner_isspace(byte: u8) -> bool {
41 matches!(byte, b' ' | b'\t' | b'\n' | b'\r' | 0x0b | 0x0c)
42}
43
44fn truncate_postgres_identifier(mut identifier: String) -> String {
45 if identifier.len() <= POSTGRES_IDENTIFIER_MAX_BYTES {
46 return identifier;
47 }
48 let mut end = POSTGRES_IDENTIFIER_MAX_BYTES;
49 while !identifier.is_char_boundary(end) {
50 end -= 1;
51 }
52 identifier.truncate(end);
53 identifier
54}
55
56#[must_use]
58pub fn parse_regobject_name(input: &str) -> Option<Vec<String>> {
59 let bytes = input.as_bytes();
60 let mut offset = 0usize;
61 while bytes.get(offset).is_some_and(|byte| scanner_isspace(*byte)) {
62 offset += 1;
63 }
64 if offset == bytes.len() {
65 return None;
66 }
67
68 let mut names = Vec::new();
69 loop {
70 let component = if bytes[offset] == b'"' {
71 offset += 1;
72 let mut quoted = String::new();
73 loop {
74 let relative = bytes[offset..].iter().position(|byte| *byte == b'"')?;
75 let quote = offset + relative;
76 quoted.push_str(&input[offset..quote]);
77 offset = quote + 1;
78 if bytes.get(offset) == Some(&b'"') {
79 quoted.push('"');
80 offset += 1;
81 continue;
82 }
83 break;
84 }
85 quoted
86 } else {
87 let start = offset;
88 while bytes
89 .get(offset)
90 .is_some_and(|byte| *byte != b'.' && !scanner_isspace(*byte))
91 {
92 offset += 1;
93 }
94 if offset == start {
95 return None;
96 }
97 input[start..offset].to_ascii_lowercase()
98 };
99 names.push(truncate_postgres_identifier(component));
100
101 while bytes.get(offset).is_some_and(|byte| scanner_isspace(*byte)) {
102 offset += 1;
103 }
104 match bytes.get(offset) {
105 None => return Some(names),
106 Some(b'.') => {
107 offset += 1;
108 while bytes.get(offset).is_some_and(|byte| scanner_isspace(*byte)) {
109 offset += 1;
110 }
111 if offset == bytes.len() {
112 return None;
113 }
114 }
115 Some(_) => return None,
116 }
117 }
118}
119
120pub fn parse_regtype_name(input: &str) -> Result<Option<ParsedRegtypeName>> {
122 if input.bytes().all(scanner_isspace) {
123 return Ok(None);
124 }
125 let parsed = crate::parser::parse_with_mode(input, pg_query::ParseMode::TypeName)?;
126 let [raw] = parsed.protobuf.stmts.as_slice() else {
127 return Ok(None);
128 };
129 let Some(NodeEnum::List(names)) = raw.stmt.as_ref().and_then(|node| node.node.as_ref()) else {
130 return Ok(None);
131 };
132 let names = names
133 .items
134 .iter()
135 .map(extract_string)
136 .collect::<Result<Vec<_>>>()?;
137 if names.is_empty() {
138 return Ok(None);
139 }
140 let scanned = crate::parser::scan(input)?;
141 let tokens = scanned
142 .tokens
143 .iter()
144 .filter_map(|token| pg_query::protobuf::Token::try_from(token.token).ok())
145 .collect::<Vec<_>>();
146 if tokens.contains(&pg_query::protobuf::Token::Setof) {
147 return Ok(None);
148 }
149 let bracket_dimensions = tokens
150 .iter()
151 .filter(|token| **token == pg_query::protobuf::Token::Ascii91)
152 .count();
153 let array_dimensions = bracket_dimensions.max(usize::from(
154 tokens.contains(&pg_query::protobuf::Token::Array),
155 ));
156 Ok(Some(ParsedRegtypeName {
157 names,
158 array_dimensions,
159 has_type_modifiers: tokens.contains(&pg_query::protobuf::Token::Ascii40),
160 }))
161}
162
163pub fn parse_regprocedure_name(input: &str) -> Result<Option<ParsedRegprocedureName>> {
165 let mut in_quote = false;
166 let left_parenthesis = input.bytes().enumerate().find_map(|(offset, byte)| {
167 if byte == b'"' {
168 in_quote = !in_quote;
169 None
170 } else if byte == b'(' && !in_quote {
171 Some(offset)
172 } else {
173 None
174 }
175 });
176 let Some(left_parenthesis) = left_parenthesis else {
177 return parse_regobject_name(input)
178 .map(|names| {
179 Some(ParsedRegprocedureName {
180 names,
181 argument_types: None,
182 })
183 })
184 .ok_or_else(|| regprocedure_syntax_error("expected a left parenthesis".into()));
185 };
186 let Some(names) = parse_regobject_name(&input[..left_parenthesis]) else {
187 return Ok(None);
188 };
189
190 let bytes = input.as_bytes();
191 let mut end = bytes.len();
192 while end > left_parenthesis + 1 && scanner_isspace(bytes[end - 1]) {
193 end -= 1;
194 }
195 if end <= left_parenthesis + 1 || bytes[end - 1] != b')' {
196 return Err(regprocedure_syntax_error(format!(
197 "expected a right parenthesis in routine identity \"{input}\""
198 )));
199 }
200 let arguments = &input[left_parenthesis + 1..end - 1];
201 let argument_types = parse_regprocedure_argument_types(input, arguments)?;
202
203 Ok(Some(ParsedRegprocedureName {
204 names,
205 argument_types: Some(argument_types),
206 }))
207}
208
209fn parse_regprocedure_argument_types(
210 input: &str,
211 arguments: &str,
212) -> Result<Vec<ParsedRegtypeName>> {
213 let argument_bytes = arguments.as_bytes();
214 let mut argument_types = Vec::new();
215 let mut offset = 0usize;
216 let mut had_comma = false;
217 loop {
218 while argument_bytes
219 .get(offset)
220 .is_some_and(|byte| scanner_isspace(*byte))
221 {
222 offset += 1;
223 }
224 if offset == argument_bytes.len() {
225 if had_comma {
226 return Err(regprocedure_syntax_error(format!(
227 "expected a type name in routine identity \"{input}\""
228 )));
229 }
230 break;
231 }
232
233 let start = offset;
234 let mut quoted = false;
235 let mut nesting = 0i32;
236 while let Some(byte) = argument_bytes.get(offset).copied() {
237 if byte == b'"' {
238 quoted = !quoted;
239 } else if byte == b',' && !quoted && nesting == 0 {
240 break;
241 } else if !quoted {
242 match byte {
243 b'(' | b'[' => nesting += 1,
244 b')' | b']' => nesting -= 1,
245 _ => {}
246 }
247 }
248 offset += 1;
249 }
250 if quoted || nesting != 0 {
251 return Err(regprocedure_syntax_error(format!(
252 "improper type name in routine identity \"{input}\""
253 )));
254 }
255 let mut type_end = offset;
256 while type_end > start && scanner_isspace(argument_bytes[type_end - 1]) {
257 type_end -= 1;
258 }
259 let Some(type_name) = parse_regtype_name(&arguments[start..type_end])? else {
260 return Err(SQLError::Parse(format!(
261 "invalid type name in routine identity \"{input}\""
262 )));
263 };
264 if argument_types.len() == POSTGRES_FUNCTION_MAX_ARGUMENTS {
265 return Err(SQLError::Routine {
266 sqlstate: "54023".into(),
267 message: format!("too many arguments in routine identity \"{input}\""),
268 });
269 }
270 argument_types.push(type_name);
271 had_comma = argument_bytes.get(offset) == Some(&b',');
272 if had_comma {
273 offset += 1;
274 }
275 }
276
277 Ok(argument_types)
278}
279
280fn regprocedure_syntax_error(message: String) -> SQLError {
281 SQLError::Routine {
282 sqlstate: "22P02".into(),
283 message,
284 }
285}
286
287pub(super) fn compile_foreign_key_action(raw: &str) -> Result<crate::ast::ForeignKeyAction> {
288 use crate::ast::ForeignKeyAction;
289 match raw.as_bytes().first().copied() {
290 None | Some(0) | Some(b'a') => Ok(ForeignKeyAction::NoAction),
291 Some(b'r') => Ok(ForeignKeyAction::Restrict),
292 Some(b'c') => Ok(ForeignKeyAction::Cascade),
293 Some(b'n') => Ok(ForeignKeyAction::SetNull),
294 Some(b'd') => Ok(ForeignKeyAction::SetDefault),
295 Some(other) => Err(SQLError::Unsupported(format!(
296 "unsupported FOREIGN KEY action byte {other:?}"
297 ))),
298 }
299}
300
301pub(super) fn compile_foreign_key_match(raw: &str) -> Result<crate::ast::ForeignKeyMatch> {
302 use crate::ast::ForeignKeyMatch;
303 match raw.as_bytes().first().copied() {
304 None | Some(0) | Some(b's') => Ok(ForeignKeyMatch::Simple),
305 Some(b'f') => Ok(ForeignKeyMatch::Full),
306 Some(b'p') => Err(SQLError::Unsupported(
307 "FOREIGN KEY MATCH PARTIAL is not implemented by PostgreSQL".into(),
308 )),
309 Some(other) => Err(SQLError::Unsupported(format!(
310 "unsupported FOREIGN KEY match byte {other:?}"
311 ))),
312 }
313}
314
315pub(super) fn validate_foreign_key_set_columns(
316 local_columns: &[String],
317 set_columns: &[String],
318 raw_delete_action: &str,
319) -> Result<()> {
320 if set_columns.is_empty() {
321 return Ok(());
322 }
323 let action = compile_foreign_key_action(raw_delete_action)?;
324 if !matches!(
325 action,
326 crate::ast::ForeignKeyAction::SetNull | crate::ast::ForeignKeyAction::SetDefault
327 ) {
328 return Err(SQLError::Unsupported(
329 "FOREIGN KEY column lists are only valid for ON DELETE SET NULL/DEFAULT".into(),
330 ));
331 }
332 for col in set_columns {
333 if !local_columns.iter().any(|local| local == col) {
334 return Err(SQLError::Unsupported(format!(
335 "FOREIGN KEY SET column `{col}` is not part of the local key"
336 )));
337 }
338 }
339 Ok(())
340}
341
342pub(super) fn raw_type_name(col: &pg_query::protobuf::ColumnDef) -> Result<Option<String>> {
343 let Some(type_name) = col.type_name.as_ref() else {
344 return Ok(None);
345 };
346 let names = type_name
347 .names
348 .iter()
349 .map(extract_string)
350 .collect::<Result<Vec<_>>>()?;
351 Ok(names.last().map(|name| name.to_lowercase()))
352}
353
354pub(super) fn compile_type_name(col: &pg_query::protobuf::ColumnDef) -> Result<ColumnType> {
355 let Some(type_name) = col.type_name.as_ref() else {
356 return Err(SQLError::Internal(format!(
357 "column `{}` has no type",
358 col.colname
359 )));
360 };
361 compile_pg_type_name(type_name, &col.colname)
362}
363
364#[expect(
365 clippy::too_many_lines,
366 reason = "ordered PostgreSQL lowering preserves syntax and error precedence"
367)]
368pub(super) fn compile_pg_type_name(
369 type_name: &pg_query::protobuf::TypeName,
370 column_name: &str,
371) -> Result<ColumnType> {
372 let names = type_name
373 .names
374 .iter()
375 .map(extract_string)
376 .collect::<Result<Vec<_>>>()?;
377 let raw = names
378 .last()
379 .ok_or_else(|| {
380 SQLError::Internal(format!(
381 "type name for `{column_name}` has no name components"
382 ))
383 })?
384 .to_lowercase();
385 let bind_named = names
386 .first()
387 .is_some_and(|schema| names.len() > 1 && schema != "pg_catalog")
388 || names.last().is_some_and(|name| *name != raw)
389 || matches!(raw.as_str(), "integer" | "smallint" | "bigint" | "boolean");
390 let base = if bind_named {
391 compile_named_type(&names, type_name)
392 } else {
393 match raw.as_str() {
394 "smallint" | "int2" | "smallserial" | "serial2" => Ok(ColumnType::SmallInteger),
395 "int" | "int4" | "integer" | "serial" | "serial4" => Ok(ColumnType::Integer),
396 "bigint" | "int8" | "bigserial" | "serial8" => Ok(ColumnType::BigInteger),
397 "oid" => Ok(ColumnType::Oid),
398 "xid" => Ok(ColumnType::Xid),
399 "void" => Ok(ColumnType::Void),
400 "text" => Ok(ColumnType::Text),
401 "name" => Ok(ColumnType::Name),
402 "uuid" => Ok(ColumnType::Uuid),
403 "varchar" | "character varying" => {
404 if type_name.typmods.len() > 1 {
405 return Err(SQLError::TypeMismatch(format!(
406 "CHARACTER VARYING accepts at most one length modifier, got {}",
407 type_name.typmods.len()
408 )));
409 }
410 let length = type_name
411 .typmods
412 .first()
413 .map(|node| expect_positive_character_length(node, "varchar"))
414 .transpose()?;
415 Ok(ColumnType::Varchar(length))
416 }
417 "char" => Ok(ColumnType::InternalChar),
418 "character" | "bpchar" => {
419 if type_name.typmods.len() > 1 {
420 return Err(SQLError::TypeMismatch(format!(
421 "CHARACTER accepts at most one length modifier, got {}",
422 type_name.typmods.len()
423 )));
424 }
425 let length = type_name
426 .typmods
427 .first()
428 .map(|node| expect_positive_character_length(node, "bpchar"))
429 .transpose()?
430 .unwrap_or(1);
431 Ok(ColumnType::Character(length))
432 }
433 "bool" | "boolean" => Ok(ColumnType::Boolean),
434 "real" | "float4" => Ok(ColumnType::Real),
435 "float8" | "double" | "double precision" => Ok(ColumnType::DoublePrecision),
436 "numeric" | "decimal" => {
437 if type_name.typmods.len() > 2 {
438 return Err(SQLError::Routine {
439 sqlstate: "22023".into(),
440 message: "invalid NUMERIC type modifier".into(),
441 });
442 }
443 let mut typmods_iter = type_name.typmods.iter();
444 let precision = typmods_iter
445 .next()
446 .map(|n| {
447 let value = expect_integer_const(n)?;
448 if !(1..=1000).contains(&value) {
449 return Err(SQLError::Routine {
450 sqlstate: "22023".into(),
451 message: format!(
452 "NUMERIC precision {value} must be between 1 and 1000"
453 ),
454 });
455 }
456 Ok(value as u32)
457 })
458 .transpose()?;
459 let scale = typmods_iter
460 .next()
461 .map(|n| {
462 let value = expect_integer_const(n)?;
463 if !(-1000..=1000).contains(&value) {
464 return Err(SQLError::Routine {
465 sqlstate: "22023".into(),
466 message: format!(
467 "NUMERIC scale {value} must be between -1000 and 1000"
468 ),
469 });
470 }
471 Ok(value as i32)
472 })
473 .transpose()?;
474 let scale = scale.or(precision.map(|_| 0));
477 Ok(ColumnType::Numeric { precision, scale })
478 }
479 "date" => Ok(ColumnType::Date),
480 "time" | "time without time zone" => Ok(ColumnType::Time),
481 "timetz" | "time with time zone" => Ok(ColumnType::TimeTz),
482 "timestamp" | "datetime" | "timestamp without time zone" => Ok(ColumnType::Timestamp),
483 "timestamptz" | "timestamp with time zone" => Ok(ColumnType::TimestampTz),
484 "interval" => {
485 let fields = type_name
486 .typmods
487 .first()
488 .map(expect_integer_const)
489 .transpose()?
490 .unwrap_or(32767);
491 let fields = IntervalFields::from_modifier_mask(fields)
492 .ok_or_else(|| SQLError::TypeMismatch("invalid interval fields".into()))?;
493 let precision = type_name
494 .typmods
495 .get(1)
496 .map(expect_integer_const)
497 .transpose()?;
498 ColumnType::with_interval_modifiers(fields, precision)
499 }
500 "int4range" => Ok(ColumnType::Range(RangeSubtype::Integer)),
501 "int8range" => Ok(ColumnType::Range(RangeSubtype::BigInteger)),
502 "numrange" => Ok(ColumnType::Range(RangeSubtype::Numeric)),
503 "daterange" => Ok(ColumnType::Range(RangeSubtype::Date)),
504 "tsrange" => Ok(ColumnType::Range(RangeSubtype::Timestamp)),
505 "tstzrange" => Ok(ColumnType::Range(RangeSubtype::TimestampTz)),
506 "int4multirange" => Ok(ColumnType::Multirange(RangeSubtype::Integer)),
507 "int8multirange" => Ok(ColumnType::Multirange(RangeSubtype::BigInteger)),
508 "nummultirange" => Ok(ColumnType::Multirange(RangeSubtype::Numeric)),
509 "datemultirange" => Ok(ColumnType::Multirange(RangeSubtype::Date)),
510 "tsmultirange" => Ok(ColumnType::Multirange(RangeSubtype::Timestamp)),
511 "tstzmultirange" => Ok(ColumnType::Multirange(RangeSubtype::TimestampTz)),
512 "json" => Ok(ColumnType::Json),
513 "jsonb" => Ok(ColumnType::JsonB),
514 "bytea" => Ok(ColumnType::Bytea),
515 "regproc" => Ok(ColumnType::Regproc),
516 "regprocedure" => Ok(ColumnType::Regprocedure),
517 "regclass" => Ok(ColumnType::Regclass),
518 "regnamespace" => Ok(ColumnType::Regnamespace),
519 "regrole" => Ok(ColumnType::Regrole),
520 "regtype" => Ok(ColumnType::Regtype),
521 "pg_node_tree" => Ok(ColumnType::PgNodeTree),
522 "aclitem" => Ok(ColumnType::AclItem),
523 "int2vector" => Ok(ColumnType::Int2Vector),
524 "oidvector" => Ok(ColumnType::OidVector),
525 "anyarray" => Ok(ColumnType::AnyArray),
526 "record" => Ok(ColumnType::Record),
527 "vector" => {
528 let [arg] = type_name.typmods.as_slice() else {
530 return Err(SQLError::Unsupported(
531 "VECTOR requires exactly one dimension".into(),
532 ));
533 };
534 let raw_dim = expect_integer_const(arg)?;
535 let dim = u32::try_from(raw_dim).map_err(|_| {
536 SQLError::TypeMismatch(format!(
537 "VECTOR dimension must be between 1 and {}, got {raw_dim}",
538 u32::MAX
539 ))
540 })?;
541 if dim == 0 {
542 return Err(SQLError::TypeMismatch(
543 "VECTOR dimension must be greater than zero".into(),
544 ));
545 }
546 Ok(ColumnType::Vector(dim))
547 }
548 "tensor" => {
549 let [arg] = type_name.typmods.as_slice() else {
551 return Err(SQLError::Unsupported(
552 "TENSOR requires exactly one dimension".into(),
553 ));
554 };
555 let raw_dim = expect_integer_const(arg)?;
556 let dim = u32::try_from(raw_dim).map_err(|_| {
557 SQLError::TypeMismatch(format!(
558 "TENSOR dimension must be between 1 and {}, got {raw_dim}",
559 u32::MAX
560 ))
561 })?;
562 if dim == 0 {
563 return Err(SQLError::TypeMismatch(
564 "TENSOR dimension must be greater than zero".into(),
565 ));
566 }
567 Ok(ColumnType::Tensor(dim))
568 }
569 _ => compile_named_type(&names, type_name),
570 }
571 }?;
572 let base = if matches!(
573 base,
574 ColumnType::Time | ColumnType::TimeTz | ColumnType::Timestamp | ColumnType::TimestampTz
575 ) {
576 if type_name.typmods.len() > 1 {
577 return Err(SQLError::TypeMismatch(
578 "invalid temporal type modifier".into(),
579 ));
580 }
581 base.with_temporal_precision(
582 type_name
583 .typmods
584 .first()
585 .map(expect_integer_const)
586 .transpose()?,
587 )?
588 } else {
589 base
590 };
591 if matches!(base, ColumnType::Void) && !type_name.array_bounds.is_empty() {
592 return Err(SQLError::Routine {
593 sqlstate: "42704".into(),
594 message: "type \"void[]\" does not exist".into(),
595 });
596 }
597 Ok(type_name
598 .array_bounds
599 .iter()
600 .fold(base, |element, _| ColumnType::Array(Box::new(element))))
601}
602
603pub(super) fn compile_pg_type_reference(
605 type_name: &pg_query::protobuf::TypeName,
606 context: &str,
607) -> Result<ColumnType> {
608 let names = type_name
609 .names
610 .iter()
611 .map(extract_string)
612 .collect::<Result<Vec<_>>>()?;
613 if names.last().is_some_and(|name| {
614 matches!(
615 name.as_str(),
616 "serial" | "serial2" | "serial4" | "serial8" | "smallserial" | "bigserial"
617 )
618 }) {
619 let base = compile_named_type(&names, type_name)?;
620 return Ok(type_name
621 .array_bounds
622 .iter()
623 .fold(base, |element, _| ColumnType::Array(Box::new(element))));
624 }
625 compile_pg_type_name(type_name, context)
626}
627
628fn compile_named_type(
629 names: &[String],
630 type_name: &pg_query::protobuf::TypeName,
631) -> Result<ColumnType> {
632 let mut name = names
633 .iter()
634 .map(|name| format!("\"{}\"", name.replace('"', "\"\"")))
635 .collect::<Vec<_>>()
636 .join(".");
637 if !type_name.typmods.is_empty() {
638 let modifiers = type_name
639 .typmods
640 .iter()
641 .map(expect_integer_const)
642 .collect::<Result<Vec<_>>>()?;
643 name.push('(');
644 name.push_str(
645 &modifiers
646 .iter()
647 .map(ToString::to_string)
648 .collect::<Vec<_>>()
649 .join(","),
650 );
651 name.push(')');
652 }
653 Ok(ColumnType::Named(name))
654}
655
656fn expect_positive_character_length(node: &Node, type_name: &str) -> Result<u32> {
657 let length = expect_integer_const(node)?;
658 u32::try_from(length)
659 .ok()
660 .filter(|length| *length > 0)
661 .ok_or_else(|| SQLError::Routine {
662 sqlstate: "22023".into(),
663 message: format!("length for type {type_name} must be at least 1"),
664 })
665}
666
667fn expect_integer_const(node: &Node) -> Result<i64> {
668 let Some(inner) = node.node.as_ref() else {
669 return Err(SQLError::Internal("missing const node".into()));
670 };
671 match inner {
672 NodeEnum::AConst(c) => match &c.val {
673 Some(pg_query::protobuf::a_const::Val::Ival(i)) => Ok(i64::from(i.ival)),
674 Some(pg_query::protobuf::a_const::Val::Fval(f)) => {
675 f.fval.parse::<i64>().map_err(|_| {
676 SQLError::TypeMismatch(format!(
677 "type modifier must be an integer, got `{}`",
678 f.fval
679 ))
680 })
681 }
682 other => Err(SQLError::Internal(format!(
683 "expected integer constant, got {other:?}"
684 ))),
685 },
686 _ => Err(SQLError::Internal(format!(
687 "expected A_Const, got {inner:?}"
688 ))),
689 }
690}