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