1use std::borrow::Cow;
28
29#[cfg(test)]
30use insta::assert_snapshot;
31use rowan::{GreenNodeData, GreenTokenData, NodeOrToken};
32
33#[cfg(test)]
34use crate::SourceFile;
35use rowan::Direction;
36
37use crate::ast;
38use crate::ast::AstNode;
39use crate::unescape::{escape_unicode_esc_str, uescape_char};
40use crate::{SyntaxKind, SyntaxNode, SyntaxToken, TokenText};
41
42use super::support;
43
44#[derive(Debug, Clone, PartialEq, Eq)]
45pub enum LitKind {
46 BitString(SyntaxToken),
47 ByteString(SyntaxToken),
48 Default(SyntaxToken),
49 DollarQuotedString(SyntaxToken),
50 EscString(SyntaxToken),
51 False(SyntaxToken),
52 IntNumber(SyntaxToken),
53 NationalString(SyntaxToken),
54 Null(SyntaxToken),
55 NumericNumber(SyntaxToken),
56 PositionalParam(SyntaxToken),
57 String(SyntaxToken),
58 True(SyntaxToken),
59 UnicodeEscString(SyntaxToken),
60}
61
62impl ast::Literal {
63 pub fn kind(&self) -> Option<LitKind> {
64 let token = self.syntax().first_child_or_token()?.into_token()?;
65 let kind = match token.kind() {
66 SyntaxKind::BIT_STRING => LitKind::BitString(token),
67 SyntaxKind::BYTE_STRING => LitKind::ByteString(token),
68 SyntaxKind::DEFAULT_KW => LitKind::Default(token),
69 SyntaxKind::DOLLAR_QUOTED_STRING => LitKind::DollarQuotedString(token),
70 SyntaxKind::ESC_STRING => LitKind::EscString(token),
71 SyntaxKind::FALSE_KW => LitKind::False(token),
72 SyntaxKind::INT_NUMBER => LitKind::IntNumber(token),
73 SyntaxKind::NATIONAL_STRING => LitKind::NationalString(token),
74 SyntaxKind::NULL_KW => LitKind::Null(token),
75 SyntaxKind::NUMERIC_NUMBER => LitKind::NumericNumber(token),
76 SyntaxKind::POSITIONAL_PARAM => LitKind::PositionalParam(token),
77 SyntaxKind::STRING => LitKind::String(token),
78 SyntaxKind::TRUE_KW => LitKind::True(token),
79 SyntaxKind::UNICODE_ESC_STRING => LitKind::UnicodeEscString(token),
80 _ => return None,
81 };
82 Some(kind)
83 }
84}
85
86impl ast::Constraint {
87 #[inline]
88 pub fn constraint_name(&self) -> Option<ast::ConstraintName> {
89 support::child::<ast::ConstraintNameClause>(self.syntax())
90 .and_then(|clause| clause.constraint_name())
91 }
92}
93
94impl ast::CreateSchema {
95 pub fn schema_name(&self) -> Option<ast::Name> {
96 match self.create_schema_target()? {
97 ast::CreateSchemaTarget::AuthorizationSchema(auth) => auth.role()?.name(),
98 ast::CreateSchemaTarget::NamedSchema(named) => named.schema()?.name(),
99 }
100 }
101}
102
103impl ast::FromItem {
104 pub fn alias(&self) -> Option<ast::Alias> {
105 match self {
106 ast::FromItem::ExprFromItem(it) => it.alias(),
107 ast::FromItem::FunctionFromItem(it) => it.alias(),
108 ast::FromItem::GraphTableFromItem(it) => it.alias(),
109 ast::FromItem::JsonTableFromItem(it) => it.alias(),
110 ast::FromItem::ParenFromItem(it) => it.alias(),
111 ast::FromItem::RelationFromItem(it) => it.alias(),
112 ast::FromItem::RowsFromItem(it) => it.alias(),
113 ast::FromItem::XmlTableFromItem(it) => it.alias(),
114 }
115 }
116
117 pub fn ordinality_token(&self) -> Option<SyntaxToken> {
118 match self {
119 ast::FromItem::FunctionFromItem(it) => it.ordinality_token(),
120 ast::FromItem::RowsFromItem(it) => it.ordinality_token(),
121 _ => None,
122 }
123 }
124}
125
126#[derive(Debug, Clone, PartialEq, Eq)]
127pub enum BinOp {
128 And(SyntaxToken),
129 AtTimeZone(ast::AtTimeZone),
130 Caret(SyntaxToken),
131 ColonColon(ast::ColonColon),
132 ColonEq(SyntaxToken),
133 CustomOp(ast::CustomOp),
134 Eq(SyntaxToken),
135 Escape(SyntaxToken),
136 FatArrow(SyntaxToken),
137 Gteq(SyntaxToken),
138 Ilike(SyntaxToken),
139 In(SyntaxToken),
140 Is(SyntaxToken),
141 IsDistinctFrom(ast::IsDistinctFrom),
142 IsNot(ast::IsNot),
143 IsNotDistinctFrom(ast::IsNotDistinctFrom),
144 LAngle(SyntaxToken),
145 Like(SyntaxToken),
146 Lteq(SyntaxToken),
147 Minus(SyntaxToken),
148 Neq(SyntaxToken),
149 Neqb(SyntaxToken),
150 NotIlike(ast::NotIlike),
151 NotIn(ast::NotIn),
152 NotLike(ast::NotLike),
153 NotSimilarTo(ast::NotSimilarTo),
154 OperatorCall(ast::OperatorCall),
155 Or(SyntaxToken),
156 Overlaps(SyntaxToken),
157 Percent(SyntaxToken),
158 Plus(SyntaxToken),
159 RAngle(SyntaxToken),
160 SimilarTo(ast::SimilarTo),
161 Slash(SyntaxToken),
162 Star(SyntaxToken),
163}
164
165#[derive(Debug, Clone, PartialEq, Eq)]
166pub enum PostfixOp {
167 AtLocal(ast::AtLocal),
168 IsJson(ast::IsJson),
169 IsJsonArray(ast::IsJsonArray),
170 IsJsonObject(ast::IsJsonObject),
171 IsJsonScalar(ast::IsJsonScalar),
172 IsJsonValue(ast::IsJsonValue),
173 IsNormalized(ast::IsNormalized),
174 IsNotJson(ast::IsNotJson),
175 IsNotJsonArray(ast::IsNotJsonArray),
176 IsNotJsonObject(ast::IsNotJsonObject),
177 IsNotJsonScalar(ast::IsNotJsonScalar),
178 IsNotJsonValue(ast::IsNotJsonValue),
179 IsNotNormalized(ast::IsNotNormalized),
180 IsNull(SyntaxToken),
181 NotNull(SyntaxToken),
182}
183
184impl ast::BinExpr {
185 #[inline]
186 pub fn lhs(&self) -> Option<ast::Expr> {
187 support::children(self.syntax()).next()
188 }
189
190 #[inline]
191 pub fn rhs(&self) -> Option<ast::Expr> {
192 support::children(self.syntax()).nth(1)
193 }
194
195 pub fn op(&self) -> Option<BinOp> {
196 let lhs = self.lhs()?;
197 for child in lhs.syntax().siblings_with_tokens(Direction::Next).skip(1) {
198 match child {
199 NodeOrToken::Token(token) => {
200 let op = match token.kind() {
201 SyntaxKind::AND_KW => BinOp::And(token),
202 SyntaxKind::CARET => BinOp::Caret(token),
203 SyntaxKind::COLON_EQ => BinOp::ColonEq(token),
204 SyntaxKind::EQ => BinOp::Eq(token),
205 SyntaxKind::ESCAPE_KW => BinOp::Escape(token),
206 SyntaxKind::FAT_ARROW => BinOp::FatArrow(token),
207 SyntaxKind::GTEQ => BinOp::Gteq(token),
208 SyntaxKind::ILIKE_KW => BinOp::Ilike(token),
209 SyntaxKind::IN_KW => BinOp::In(token),
210 SyntaxKind::IS_KW => BinOp::Is(token),
211 SyntaxKind::L_ANGLE => BinOp::LAngle(token),
212 SyntaxKind::LIKE_KW => BinOp::Like(token),
213 SyntaxKind::LTEQ => BinOp::Lteq(token),
214 SyntaxKind::MINUS => BinOp::Minus(token),
215 SyntaxKind::NEQ => BinOp::Neq(token),
216 SyntaxKind::NEQB => BinOp::Neqb(token),
217 SyntaxKind::OR_KW => BinOp::Or(token),
218 SyntaxKind::OVERLAPS_KW => BinOp::Overlaps(token),
219 SyntaxKind::PERCENT => BinOp::Percent(token),
220 SyntaxKind::PLUS => BinOp::Plus(token),
221 SyntaxKind::R_ANGLE => BinOp::RAngle(token),
222 SyntaxKind::SLASH => BinOp::Slash(token),
223 SyntaxKind::STAR => BinOp::Star(token),
224 _ => continue,
225 };
226 return Some(op);
227 }
228 NodeOrToken::Node(node) => {
229 let op = match node.kind() {
230 SyntaxKind::AT_TIME_ZONE => {
231 BinOp::AtTimeZone(ast::AtTimeZone { syntax: node })
232 }
233 SyntaxKind::COLON_COLON => {
234 BinOp::ColonColon(ast::ColonColon { syntax: node })
235 }
236 SyntaxKind::CUSTOM_OP => BinOp::CustomOp(ast::CustomOp { syntax: node }),
237 SyntaxKind::IS_DISTINCT_FROM => {
238 BinOp::IsDistinctFrom(ast::IsDistinctFrom { syntax: node })
239 }
240 SyntaxKind::IS_NOT => BinOp::IsNot(ast::IsNot { syntax: node }),
241 SyntaxKind::IS_NOT_DISTINCT_FROM => {
242 BinOp::IsNotDistinctFrom(ast::IsNotDistinctFrom { syntax: node })
243 }
244 SyntaxKind::NOT_ILIKE => BinOp::NotIlike(ast::NotIlike { syntax: node }),
245 SyntaxKind::NOT_IN => BinOp::NotIn(ast::NotIn { syntax: node }),
246 SyntaxKind::NOT_LIKE => BinOp::NotLike(ast::NotLike { syntax: node }),
247 SyntaxKind::NOT_SIMILAR_TO => {
248 BinOp::NotSimilarTo(ast::NotSimilarTo { syntax: node })
249 }
250 SyntaxKind::OPERATOR_CALL => {
251 BinOp::OperatorCall(ast::OperatorCall { syntax: node })
252 }
253 SyntaxKind::SIMILAR_TO => BinOp::SimilarTo(ast::SimilarTo { syntax: node }),
254 _ => continue,
255 };
256 return Some(op);
257 }
258 }
259 }
260 None
261 }
262}
263
264impl ast::PostfixExpr {
265 pub fn op(&self) -> Option<PostfixOp> {
266 let lhs = self.expr()?;
267
268 let siblings = lhs.syntax().siblings_with_tokens(Direction::Next).skip(1);
269 for child in siblings {
270 match child {
271 NodeOrToken::Token(token) => {
272 let op = match token.kind() {
273 SyntaxKind::ISNULL_KW => PostfixOp::IsNull(token),
274 SyntaxKind::NOTNULL_KW => PostfixOp::NotNull(token),
275 _ => continue,
276 };
277 return Some(op);
278 }
279 NodeOrToken::Node(node) => {
280 let op = match node.kind() {
281 SyntaxKind::AT_LOCAL => PostfixOp::AtLocal(ast::AtLocal { syntax: node }),
282 SyntaxKind::IS_JSON => PostfixOp::IsJson(ast::IsJson { syntax: node }),
283 SyntaxKind::IS_JSON_ARRAY => {
284 PostfixOp::IsJsonArray(ast::IsJsonArray { syntax: node })
285 }
286 SyntaxKind::IS_JSON_OBJECT => {
287 PostfixOp::IsJsonObject(ast::IsJsonObject { syntax: node })
288 }
289 SyntaxKind::IS_JSON_SCALAR => {
290 PostfixOp::IsJsonScalar(ast::IsJsonScalar { syntax: node })
291 }
292 SyntaxKind::IS_JSON_VALUE => {
293 PostfixOp::IsJsonValue(ast::IsJsonValue { syntax: node })
294 }
295 SyntaxKind::IS_NORMALIZED => {
296 PostfixOp::IsNormalized(ast::IsNormalized { syntax: node })
297 }
298 SyntaxKind::IS_NOT_JSON => {
299 PostfixOp::IsNotJson(ast::IsNotJson { syntax: node })
300 }
301 SyntaxKind::IS_NOT_JSON_ARRAY => {
302 PostfixOp::IsNotJsonArray(ast::IsNotJsonArray { syntax: node })
303 }
304 SyntaxKind::IS_NOT_JSON_OBJECT => {
305 PostfixOp::IsNotJsonObject(ast::IsNotJsonObject { syntax: node })
306 }
307 SyntaxKind::IS_NOT_JSON_SCALAR => {
308 PostfixOp::IsNotJsonScalar(ast::IsNotJsonScalar { syntax: node })
309 }
310 SyntaxKind::IS_NOT_JSON_VALUE => {
311 PostfixOp::IsNotJsonValue(ast::IsNotJsonValue { syntax: node })
312 }
313 SyntaxKind::IS_NOT_NORMALIZED => {
314 PostfixOp::IsNotNormalized(ast::IsNotNormalized { syntax: node })
315 }
316 _ => continue,
317 };
318 return Some(op);
319 }
320 }
321 }
322
323 None
324 }
325}
326
327impl ast::FieldExpr {
328 #[inline]
332 pub fn base(&self) -> Option<ast::Expr> {
333 support::children(self.syntax()).next()
334 }
335 #[inline]
336 pub fn field(&self) -> Option<ast::NameRef> {
337 support::children(self.syntax()).last()
338 }
339}
340
341impl ast::IndexExpr {
342 #[inline]
343 pub fn base(&self) -> Option<ast::Expr> {
344 support::children(&self.syntax).next()
345 }
346 #[inline]
347 pub fn index(&self) -> Option<ast::Expr> {
348 support::children(&self.syntax).nth(1)
349 }
350}
351
352impl ast::SliceExpr {
353 #[inline]
354 pub fn base(&self) -> Option<ast::Expr> {
355 support::children(&self.syntax).next()
356 }
357
358 #[inline]
359 pub fn start(&self) -> Option<ast::Expr> {
360 let colon = self.colon_token()?;
365 support::children(&self.syntax)
366 .skip(1)
367 .find(|expr: &ast::Expr| expr.syntax().text_range().end() <= colon.text_range().start())
368 }
369
370 #[inline]
371 pub fn end(&self) -> Option<ast::Expr> {
372 let colon = self.colon_token()?;
375 support::children(&self.syntax)
376 .find(|expr: &ast::Expr| expr.syntax().text_range().start() >= colon.text_range().end())
377 }
378}
379
380impl ast::RenameColumn {
381 #[inline]
382 pub fn from(&self) -> Option<ast::NameRef> {
383 support::children(&self.syntax).nth(0)
384 }
385 #[inline]
386 pub fn to(&self) -> Option<ast::NameRef> {
387 support::children(&self.syntax).nth(1)
388 }
389}
390
391impl ast::RenameValue {
392 #[inline]
393 pub fn from(&self) -> Option<ast::Literal> {
394 support::children(&self.syntax).nth(0)
395 }
396 #[inline]
397 pub fn to(&self) -> Option<ast::Literal> {
398 support::children(&self.syntax).nth(1)
399 }
400}
401
402impl ast::ForeignKeyConstraint {
403 #[inline]
404 pub fn from_columns(&self) -> Option<ast::ColumnRefList> {
405 support::children(&self.syntax).nth(0)
406 }
407 #[inline]
408 pub fn to_columns(&self) -> Option<ast::ColumnRefList> {
409 support::children(&self.syntax).nth(1)
410 }
411}
412
413impl ast::BetweenExpr {
414 #[inline]
415 pub fn target(&self) -> Option<ast::Expr> {
416 support::children(&self.syntax).nth(0)
417 }
418 #[inline]
419 pub fn start(&self) -> Option<ast::Expr> {
420 support::children(&self.syntax).nth(1)
421 }
422 #[inline]
423 pub fn end(&self) -> Option<ast::Expr> {
424 support::children(&self.syntax).nth(2)
425 }
426}
427
428impl ast::FrameBetween {
429 #[inline]
430 pub fn start(&self) -> Option<ast::FrameBound> {
431 support::children(&self.syntax).nth(0)
432 }
433 #[inline]
434 pub fn end(&self) -> Option<ast::FrameBound> {
435 support::children(&self.syntax).nth(1)
436 }
437}
438
439impl ast::WhenClause {
440 #[inline]
441 pub fn condition(&self) -> Option<ast::Expr> {
442 support::children(&self.syntax).next()
443 }
444 #[inline]
445 pub fn then(&self) -> Option<ast::Expr> {
446 support::children(&self.syntax).nth(1)
447 }
448}
449
450impl ast::CompoundSelect {
451 #[inline]
452 pub fn lhs(&self) -> Option<ast::SelectVariant> {
453 support::children(&self.syntax).next()
454 }
455 #[inline]
456 pub fn rhs(&self) -> Option<ast::SelectVariant> {
457 support::children(&self.syntax).nth(1)
458 }
459}
460
461impl ast::NameRef {
462 #[inline]
463 pub fn text(&self) -> String {
464 normalize_name_node(self.syntax())
465 }
466
467 #[inline]
468 pub fn is_quoted(&self) -> bool {
469 is_quoted(self.syntax())
470 }
471}
472
473impl ast::Name {
474 #[inline]
475 pub fn text(&self) -> String {
476 normalize_name_node(self.syntax())
477 }
478
479 #[inline]
480 pub fn is_quoted(&self) -> bool {
481 is_quoted(self.syntax())
482 }
483}
484
485fn is_quoted(node: &SyntaxNode) -> bool {
486 let text = node.text();
487 let first = text.char_at(0.into());
488 let second = text.char_at(1.into());
489 matches!(
490 (first, second),
491 (Some('u' | 'U'), Some('"')) | (Some('"'), Some(_))
492 )
493}
494
495fn normalize_name_node(node: &SyntaxNode) -> String {
497 let mut tokens = node
498 .children_with_tokens()
499 .filter_map(|el| el.into_token())
500 .filter(|t| !t.kind().is_trivia());
501
502 let Some(ident_token) = tokens.next() else {
503 return String::new();
504 };
505 let raw = ident_token.text();
506
507 let unicode_inner = raw
508 .strip_prefix(['u', 'U'])
509 .and_then(|s| s.strip_prefix("&\""))
510 .and_then(|s| s.strip_suffix('"'));
511
512 if let Some(inner) = unicode_inner {
513 let mut escape_char = '\\';
514 if let Some(uesc) = tokens.next()
515 && uesc.kind() == SyntaxKind::UESCAPE_KW
516 && let Some(token) = tokens.next()
517 && let Some(ch) = uescape_char(token.text())
518 {
519 escape_char = ch;
520 }
521
522 let inner = inner.replace(r#""""#, "\"");
523 let mut result = String::with_capacity(inner.len());
524 escape_unicode_esc_str(&inner, escape_char, |_range, r| {
525 if let Ok(ch) = r {
526 result.push(ch);
527 }
528 });
529 return result;
530 }
531
532 raw.strip_prefix('"')
533 .and_then(|t| t.strip_suffix('"'))
534 .map(|x| x.replace(r#""""#, "\""))
535 .unwrap_or_else(|| raw.to_ascii_lowercase())
536}
537
538impl ast::CharType {
539 #[inline]
540 pub fn text(&self) -> TokenText<'_> {
541 text_of_first_token(self.syntax())
542 }
543}
544
545fn string_literal_contents(token: &SyntaxToken) -> Option<&str> {
546 match token.kind() {
547 SyntaxKind::STRING => token.text().strip_prefix('\'')?.strip_suffix('\''),
548 SyntaxKind::ESC_STRING | SyntaxKind::NATIONAL_STRING => {
549 token.text().get(2..)?.strip_suffix('\'')
550 }
551 SyntaxKind::UNICODE_ESC_STRING => token.text().get(3..)?.strip_suffix('\''),
552 SyntaxKind::DOLLAR_QUOTED_STRING => {
553 let text = token.text();
554 let rest = text.strip_prefix('$')?;
555 let tag_len = rest.find('$')?;
556 let delimiter = text.get(..=tag_len + 1)?;
557 text.get(delimiter.len()..)?.strip_suffix(delimiter)
558 }
559 _ => None,
560 }
561}
562
563fn is_falsey_token(token: &SyntaxToken) -> bool {
564 match token.kind() {
565 SyntaxKind::FALSE_KW | SyntaxKind::NO_KW | SyntaxKind::OFF_KW => true,
566 SyntaxKind::INT_NUMBER => token.text() == "0",
567 SyntaxKind::STRING
568 | SyntaxKind::ESC_STRING
569 | SyntaxKind::NATIONAL_STRING
570 | SyntaxKind::UNICODE_ESC_STRING
571 | SyntaxKind::DOLLAR_QUOTED_STRING => string_literal_contents(token)
572 .is_some_and(|text| matches!(text.to_ascii_lowercase().as_str(), "false" | "off")),
573 _ => false,
574 }
575}
576
577fn is_falsey_vacuum_option_value(value: &ast::VacuumOptionValue) -> bool {
578 value
579 .syntax()
580 .first_token()
581 .is_some_and(|token| is_falsey_token(&token))
582}
583
584impl ast::Reindex {
585 pub fn is_concurrently(&self) -> bool {
586 self.concurrently_token().is_some()
587 || self.reindex_option_list().is_some_and(|options| {
588 options.reindex_options().any(|option| {
589 option.concurrently_token().is_some()
590 && !option.literal().is_some_and(|literal| {
591 literal
592 .syntax()
593 .first_token()
594 .is_some_and(|token| is_falsey_token(&token))
595 })
596 })
597 })
598 }
599}
600
601impl ast::Vacuum {
602 pub fn is_full(&self) -> bool {
603 self.full_token().is_some()
604 || self.vacuum_option_list().is_some_and(|opt_list| {
606 opt_list.vacuum_options().any(|opt| {
607 opt.name().is_some_and(|name| {
608 name.syntax()
609 .first_token()
610 .is_some_and(|token| token.text().eq_ignore_ascii_case("full"))
611 }) && opt
612 .vacuum_option_value()
613 .is_none_or(|value| !is_falsey_vacuum_option_value(&value))
614 })
615 })
616 }
617}
618
619impl ast::OpSig {
620 #[inline]
621 pub fn lhs(&self) -> Option<ast::Type> {
622 support::children(self.syntax()).next()
623 }
624
625 #[inline]
626 pub fn rhs(&self) -> Option<ast::Type> {
627 support::children(self.syntax()).nth(1)
628 }
629}
630
631impl ast::CastSig {
632 #[inline]
633 pub fn lhs(&self) -> Option<ast::Type> {
634 support::children(self.syntax()).next()
635 }
636
637 #[inline]
638 pub fn rhs(&self) -> Option<ast::Type> {
639 support::children(self.syntax()).nth(1)
640 }
641}
642
643impl ast::ColumnConstraint {
644 #[inline]
645 pub fn constraint_name(&self) -> Option<ast::ConstraintName> {
646 match self {
647 ast::ColumnConstraint::CheckConstraint(check_constraint) => check_constraint
648 .constraint_name_clause()
649 .and_then(|clause| clause.constraint_name()),
650 ast::ColumnConstraint::DefaultConstraint(default_constraint) => default_constraint
651 .constraint_name_clause()
652 .and_then(|clause| clause.constraint_name()),
653 ast::ColumnConstraint::ExcludeConstraint(exclude_constraint) => exclude_constraint
654 .constraint_name_clause()
655 .and_then(|clause| clause.constraint_name()),
656 ast::ColumnConstraint::GeneratedConstraint(generated_constraint) => {
657 generated_constraint
658 .constraint_name_clause()
659 .and_then(|clause| clause.constraint_name())
660 }
661 ast::ColumnConstraint::NotNullConstraint(not_null_constraint) => not_null_constraint
662 .constraint_name_clause()
663 .and_then(|clause| clause.constraint_name()),
664 ast::ColumnConstraint::NullConstraint(null_constraint) => null_constraint
665 .constraint_name_clause()
666 .and_then(|clause| clause.constraint_name()),
667 ast::ColumnConstraint::PrimaryKeyConstraint(primary_key_constraint) => {
668 primary_key_constraint
669 .constraint_name_clause()
670 .and_then(|clause| clause.constraint_name())
671 }
672 ast::ColumnConstraint::ReferencesConstraint(references_constraint) => {
673 references_constraint
674 .constraint_name_clause()
675 .and_then(|clause| clause.constraint_name())
676 }
677 ast::ColumnConstraint::UniqueConstraint(unique_constraint) => unique_constraint
678 .constraint_name_clause()
679 .and_then(|clause| clause.constraint_name()),
680 }
681 }
682}
683
684impl ast::TableConstraint {
685 #[inline]
686 pub fn constraint_name(&self) -> Option<ast::ConstraintName> {
687 match self {
688 ast::TableConstraint::CheckConstraint(check_constraint) => check_constraint
689 .constraint_name_clause()
690 .and_then(|clause| clause.constraint_name()),
691 ast::TableConstraint::ExcludeConstraint(exclude_constraint) => exclude_constraint
692 .constraint_name_clause()
693 .and_then(|clause| clause.constraint_name()),
694 ast::TableConstraint::ForeignKeyConstraint(foreign_key_constraint) => {
695 foreign_key_constraint
696 .constraint_name_clause()
697 .and_then(|clause| clause.constraint_name())
698 }
699 ast::TableConstraint::PrimaryKeyConstraint(primary_key_constraint) => {
700 primary_key_constraint
701 .constraint_name_clause()
702 .and_then(|clause| clause.constraint_name())
703 }
704 ast::TableConstraint::UniqueConstraint(unique_constraint) => unique_constraint
705 .constraint_name_clause()
706 .and_then(|clause| clause.constraint_name()),
707 }
708 }
709}
710
711pub(crate) fn text_of_first_token(node: &SyntaxNode) -> TokenText<'_> {
712 fn first_token(green_ref: &GreenNodeData) -> &GreenTokenData {
713 green_ref
714 .children()
715 .next()
716 .and_then(NodeOrToken::into_token)
717 .unwrap()
718 }
719
720 match node.green() {
721 Cow::Borrowed(green_ref) => TokenText::borrowed(first_token(green_ref).text()),
722 Cow::Owned(green) => TokenText::owned(first_token(&green).to_owned()),
723 }
724}
725
726impl ast::WithQuery {
727 #[inline]
728 pub fn with_clause(&self) -> Option<ast::WithClause> {
729 support::child(self.syntax())
730 }
731}
732
733impl ast::CreateTableAsQuery {
734 #[inline]
735 pub fn select_variant(&self) -> Option<ast::SelectVariant> {
736 match self {
737 ast::CreateTableAsQuery::Execute(_) => None,
738 ast::CreateTableAsQuery::SelectVariant(select_variant) => Some(select_variant.clone()),
739 }
740 }
741}
742
743impl ast::SelectVariant {
744 #[inline]
745 pub fn target_list(&self) -> Option<ast::TargetList> {
746 match self {
747 ast::SelectVariant::Select(select) => {
748 return select.select_clause()?.target_list();
749 }
750 ast::SelectVariant::SelectInto(select_into) => {
751 return select_into.select_clause()?.target_list();
752 }
753 ast::SelectVariant::ParenSelect(paren_select) => {
754 return paren_select.select()?.target_list();
755 }
756 _ => return None,
757 }
758 }
759}
760
761impl ast::HasParamList {
762 #[inline]
763 pub fn param_list(&self) -> Option<ast::ParamList> {
764 support::child(self.syntax())
765 }
766 #[inline]
767 pub fn path(&self) -> Option<ast::Path> {
768 match self {
769 ast::HasParamList::CreateFunction(function) => function.name()?.path(),
770 ast::HasParamList::CreateProcedure(procedure) => procedure.name()?.path(),
771 _ => support::child(self.syntax()),
772 }
773 }
774 #[inline]
775 pub fn path_ref(&self) -> Option<ast::PathRef> {
776 match self {
777 ast::HasParamList::FunctionSig(signature) => signature.function_name_ref()?.path_ref(),
778 ast::HasParamList::ProcedureSig(signature) => {
779 signature.procedure_name_ref()?.path_ref()
780 }
781 ast::HasParamList::RoutineSig(signature) => signature.routine_name_ref()?.path_ref(),
782 _ => support::child(self.syntax()),
783 }
784 }
785}
786
787impl ast::NameLike for ast::Name {
788 #[inline]
789 fn text(&self) -> String {
790 self.text()
791 }
792}
793impl ast::NameLike for ast::NameRef {
794 #[inline]
795 fn text(&self) -> String {
796 self.text()
797 }
798}
799
800impl ast::HasWithClause for ast::Select {}
801impl ast::HasWithClause for ast::SelectInto {}
802impl ast::HasWithClause for ast::Insert {}
803impl ast::HasWithClause for ast::Update {}
804impl ast::HasWithClause for ast::Delete {}
805
806impl ast::HasCreateTable for ast::CreateTable {}
807impl ast::HasCreateTable for ast::CreateForeignTable {}
808impl ast::HasCreateTable for ast::CreateTableLike {}
809
810#[test]
811fn name() {
812 assert_snapshot!(extract_name("select 1 foo"), @"foo");
813 assert_snapshot!(extract_name("select 1 FOO"), @"foo");
814 assert_snapshot!(extract_name(r#"select 1 "foo""#), @"foo");
815 assert_snapshot!(extract_name(r#"select 1 "Foo""#), @"Foo");
816 assert_snapshot!(extract_name(r#"select 1 "FOO""#), @"FOO");
817 assert_snapshot!(extract_name(r#"select 1 U&"\0066\006f\006f""#), @"foo");
818 assert_snapshot!(extract_name(r#"select 1 U&"@0066@006f@006f" uescape '@'"#), @"foo");
819
820 fn extract_name(source_code: &str) -> String {
821 let parse = SourceFile::parse(source_code);
822 assert!(parse.errors().is_empty());
823 let stmt = parse.tree().stmts().next().unwrap();
824 let ast::Stmt::Select(select) = stmt else {
825 unreachable!()
826 };
827 let name = select
828 .select_clause()
829 .unwrap()
830 .target_list()
831 .unwrap()
832 .targets()
833 .next()
834 .unwrap()
835 .as_name()
836 .unwrap()
837 .name()
838 .unwrap();
839 name.text().to_string()
840 }
841}
842
843#[test]
844fn name_ref() {
845 assert_snapshot!(extract_name_ref("select foo"), @"foo");
846 assert_snapshot!(extract_name_ref("select FOO"), @"foo");
847 assert_snapshot!(extract_name_ref(r#"select "foo""#), @"foo");
848 assert_snapshot!(extract_name_ref(r#"select "Foo""#), @"Foo");
849 assert_snapshot!(extract_name_ref(r#"select "FOO""#), @"FOO");
850 assert_snapshot!(extract_name_ref(r#"select U&"\0066\006f\006f""#), @"foo");
851 assert_snapshot!(extract_name_ref(r#"select U&"@0066@006f@006f" uescape '@'"#), @"foo");
852
853 fn extract_name_ref(source_code: &str) -> String {
854 let parse = SourceFile::parse(source_code);
855 assert!(parse.errors().is_empty());
856 let stmt = parse.tree().stmts().next().unwrap();
857 let ast::Stmt::Select(select) = stmt else {
858 unreachable!()
859 };
860 let select_clause = select.select_clause().unwrap();
861 let target = select_clause
862 .target_list()
863 .unwrap()
864 .targets()
865 .next()
866 .unwrap();
867 let ast::Expr::NameRef(name_ref) = target.expr().unwrap() else {
868 unreachable!()
869 };
870 name_ref.text().to_string()
871 }
872}
873
874#[test]
875fn unicode_quoted_name_keeps_doubled_single_quotes() {
876 let parse = SourceFile::parse(r#"select 1 U&"a''b""#);
877 assert!(parse.errors().is_empty());
878 let stmt = parse.tree().stmts().next().unwrap();
879 let ast::Stmt::Select(select) = stmt else {
880 unreachable!()
881 };
882 let name = select
883 .select_clause()
884 .unwrap()
885 .target_list()
886 .unwrap()
887 .targets()
888 .next()
889 .unwrap()
890 .as_name()
891 .unwrap()
892 .name()
893 .unwrap();
894
895 assert_snapshot!(name.text().to_string(), @"a''b");
896}
897
898#[test]
899fn index_expr() {
900 let source_code = "
901 select foo[bar];
902 ";
903 let parse = SourceFile::parse(source_code);
904 assert!(parse.errors().is_empty());
905 let stmt = parse.tree().stmts().next().unwrap();
906 let ast::Stmt::Select(select) = stmt else {
907 unreachable!()
908 };
909 let select_clause = select.select_clause().unwrap();
910 let target = select_clause
911 .target_list()
912 .unwrap()
913 .targets()
914 .next()
915 .unwrap();
916 let ast::Expr::IndexExpr(index_expr) = target.expr().unwrap() else {
917 unreachable!()
918 };
919 let base = index_expr.base().unwrap();
920 let index = index_expr.index().unwrap();
921 assert_eq!(base.syntax().text(), "foo");
922 assert_eq!(index.syntax().text(), "bar");
923}
924
925#[test]
926fn slice_expr() {
927 use insta::assert_snapshot;
928 let source_code = "
929 select x[1:2], x[2:], x[:3], x[:];
930 ";
931 let parse = SourceFile::parse(source_code);
932 assert!(parse.errors().is_empty());
933 let stmt = parse.tree().stmts().next().unwrap();
934 let ast::Stmt::Select(select) = stmt else {
935 unreachable!()
936 };
937 let select_clause = select.select_clause().unwrap();
938 let mut targets = select_clause.target_list().unwrap().targets();
939
940 let ast::Expr::SliceExpr(slice) = targets.next().unwrap().expr().unwrap() else {
941 unreachable!()
942 };
943 assert_snapshot!(slice.syntax(), @"x[1:2]");
944 assert_eq!(slice.base().unwrap().syntax().text(), "x");
945 assert_eq!(slice.start().unwrap().syntax().text(), "1");
946 assert_eq!(slice.end().unwrap().syntax().text(), "2");
947
948 let ast::Expr::SliceExpr(slice) = targets.next().unwrap().expr().unwrap() else {
949 unreachable!()
950 };
951 assert_snapshot!(slice.syntax(), @"x[2:]");
952 assert_eq!(slice.base().unwrap().syntax().text(), "x");
953 assert_eq!(slice.start().unwrap().syntax().text(), "2");
954 assert!(slice.end().is_none());
955
956 let ast::Expr::SliceExpr(slice) = targets.next().unwrap().expr().unwrap() else {
957 unreachable!()
958 };
959 assert_snapshot!(slice.syntax(), @"x[:3]");
960 assert_eq!(slice.base().unwrap().syntax().text(), "x");
961 assert!(slice.start().is_none());
962 assert_eq!(slice.end().unwrap().syntax().text(), "3");
963
964 let ast::Expr::SliceExpr(slice) = targets.next().unwrap().expr().unwrap() else {
965 unreachable!()
966 };
967 assert_snapshot!(slice.syntax(), @"x[:]");
968 assert_eq!(slice.base().unwrap().syntax().text(), "x");
969 assert!(slice.start().is_none());
970 assert!(slice.end().is_none());
971}
972
973#[test]
974fn field_expr() {
975 let source_code = "
976 select foo.bar;
977 ";
978 let parse = SourceFile::parse(source_code);
979 assert!(parse.errors().is_empty());
980 let stmt = parse.tree().stmts().next().unwrap();
981 let ast::Stmt::Select(select) = stmt else {
982 unreachable!()
983 };
984 let select_clause = select.select_clause().unwrap();
985 let target = select_clause
986 .target_list()
987 .unwrap()
988 .targets()
989 .next()
990 .unwrap();
991 let ast::Expr::FieldExpr(field_expr) = target.expr().unwrap() else {
992 unreachable!()
993 };
994 let base = field_expr.base().unwrap();
995 let field = field_expr.field().unwrap();
996 assert_eq!(base.syntax().text(), "foo");
997 assert_eq!(field.syntax().text(), "bar");
998}
999
1000#[test]
1001fn between_expr() {
1002 let source_code = "
1003 select 2 between 1 and 3;
1004 ";
1005 let parse = SourceFile::parse(source_code);
1006 assert!(parse.errors().is_empty());
1007 let stmt = parse.tree().stmts().next().unwrap();
1008 let ast::Stmt::Select(select) = stmt else {
1009 unreachable!()
1010 };
1011 let select_clause = select.select_clause().unwrap();
1012 let target = select_clause
1013 .target_list()
1014 .unwrap()
1015 .targets()
1016 .next()
1017 .unwrap();
1018 let ast::Expr::BetweenExpr(between_expr) = target.expr().unwrap() else {
1019 unreachable!()
1020 };
1021 let target = between_expr.target().unwrap();
1022 let start = between_expr.start().unwrap();
1023 let end = between_expr.end().unwrap();
1024 assert_eq!(target.syntax().text(), "2");
1025 assert_eq!(start.syntax().text(), "1");
1026 assert_eq!(end.syntax().text(), "3");
1027}
1028
1029#[test]
1030fn cast_expr() {
1031 use insta::assert_snapshot;
1032
1033 let cast = extract_expr("select cast('123' as int)");
1034 assert!(cast.expr().is_some());
1035 assert_snapshot!(cast.expr().unwrap().syntax(), @"'123'");
1036 assert!(cast.ty().is_some());
1037 assert_snapshot!(cast.ty().unwrap().syntax(), @"int");
1038
1039 let cast = extract_expr("select cast('123' as pg_catalog.int4)");
1040 assert!(cast.expr().is_some());
1041 assert_snapshot!(cast.expr().unwrap().syntax(), @"'123'");
1042 assert!(cast.ty().is_some());
1043 assert_snapshot!(cast.ty().unwrap().syntax(), @"pg_catalog.int4");
1044
1045 let cast = extract_expr("select int '123'");
1046 assert!(cast.expr().is_some());
1047 assert_snapshot!(cast.expr().unwrap().syntax(), @"'123'");
1048 assert!(cast.ty().is_some());
1049 assert_snapshot!(cast.ty().unwrap().syntax(), @"int");
1050
1051 let cast = extract_expr("select pg_catalog.int4 '123'");
1052 assert!(cast.expr().is_some());
1053 assert_snapshot!(cast.expr().unwrap().syntax(), @"'123'");
1054 assert!(cast.ty().is_some());
1055 assert_snapshot!(cast.ty().unwrap().syntax(), @"pg_catalog.int4");
1056
1057 let cast = extract_expr("select '123'::int");
1058 assert!(cast.expr().is_some());
1059 assert_snapshot!(cast.expr().unwrap().syntax(), @"'123'");
1060 assert!(cast.ty().is_some());
1061 assert_snapshot!(cast.ty().unwrap().syntax(), @"int");
1062
1063 let cast = extract_expr("select '123'::int4");
1064 assert!(cast.expr().is_some());
1065 assert_snapshot!(cast.expr().unwrap().syntax(), @"'123'");
1066 assert!(cast.ty().is_some());
1067 assert_snapshot!(cast.ty().unwrap().syntax(), @"int4");
1068
1069 let cast = extract_expr("select '123'::pg_catalog.int4");
1070 assert!(cast.expr().is_some());
1071 assert_snapshot!(cast.expr().unwrap().syntax(), @"'123'");
1072 assert!(cast.ty().is_some());
1073 assert_snapshot!(cast.ty().unwrap().syntax(), @"pg_catalog.int4");
1074
1075 let cast = extract_expr("select '{123}'::pg_catalog.varchar(10)[]");
1076 assert!(cast.expr().is_some());
1077 assert_snapshot!(cast.expr().unwrap().syntax(), @"'{123}'");
1078 assert!(cast.ty().is_some());
1079 assert_snapshot!(cast.ty().unwrap().syntax(), @"pg_catalog.varchar(10)[]");
1080
1081 let cast = extract_expr("select cast('{123}' as pg_catalog.varchar(10)[])");
1082 assert!(cast.expr().is_some());
1083 assert_snapshot!(cast.expr().unwrap().syntax(), @"'{123}'");
1084 assert!(cast.ty().is_some());
1085 assert_snapshot!(cast.ty().unwrap().syntax(), @"pg_catalog.varchar(10)[]");
1086
1087 let cast = extract_expr("select pg_catalog.varchar(10) '{123}'");
1088 assert!(cast.expr().is_some());
1089 assert_snapshot!(cast.expr().unwrap().syntax(), @"'{123}'");
1090 assert!(cast.ty().is_some());
1091 assert_snapshot!(cast.ty().unwrap().syntax(), @"pg_catalog.varchar(10)");
1092
1093 let cast = extract_expr("select interval '1' month");
1094 assert!(cast.expr().is_some());
1095 assert_snapshot!(cast.expr().unwrap().syntax(), @"'1'");
1096 assert!(cast.ty().is_some());
1097 assert_snapshot!(cast.ty().unwrap().syntax(), @"interval");
1098
1099 fn extract_expr(sql: &str) -> ast::CastExpr {
1100 let parse = SourceFile::parse(sql);
1101 assert!(parse.errors().is_empty());
1102 let node = parse
1103 .tree()
1104 .stmts()
1105 .map(|x| match x {
1106 ast::Stmt::Select(select) => select
1107 .select_clause()
1108 .unwrap()
1109 .target_list()
1110 .unwrap()
1111 .targets()
1112 .next()
1113 .unwrap()
1114 .expr()
1115 .unwrap()
1116 .clone(),
1117 _ => unreachable!(),
1118 })
1119 .next()
1120 .unwrap();
1121 match node {
1122 ast::Expr::CastExpr(cast) => cast,
1123 _ => unreachable!(),
1124 }
1125 }
1126}
1127
1128#[test]
1129fn op_sig() {
1130 let source_code = "
1131 alter operator p.+ (int4, int8)
1132 owner to u;
1133 ";
1134 let parse = SourceFile::parse(source_code);
1135 assert!(parse.errors().is_empty());
1136 let stmt = parse.tree().stmts().next().unwrap();
1137 let ast::Stmt::AlterOperator(alter_op) = stmt else {
1138 unreachable!()
1139 };
1140 let op_sig = alter_op.op_sig().unwrap();
1141 let lhs = op_sig.lhs().unwrap();
1142 let rhs = op_sig.rhs().unwrap();
1143 assert_snapshot!(lhs.syntax().text(), @"int4");
1144 assert_snapshot!(rhs.syntax().text(), @"int8");
1145}
1146
1147#[test]
1148fn cast_sig() {
1149 let source_code = "
1150 drop cast (text as int);
1151 ";
1152 let parse = SourceFile::parse(source_code);
1153 assert!(parse.errors().is_empty());
1154 let stmt = parse.tree().stmts().next().unwrap();
1155 let ast::Stmt::DropCast(alter_op) = stmt else {
1156 unreachable!()
1157 };
1158 let cast_sig = alter_op.cast_sig().unwrap();
1159 let lhs = cast_sig.lhs().unwrap();
1160 let rhs = cast_sig.rhs().unwrap();
1161 assert_snapshot!(lhs.syntax().text(), @"text");
1162 assert_snapshot!(rhs.syntax().text(), @"int");
1163}
1164
1165#[cfg(test)]
1166fn extract_vacuum(sql: &str) -> ast::Vacuum {
1167 let parse = SourceFile::parse(sql);
1168 assert!(parse.errors().is_empty());
1169 let stmt = parse.tree().stmts().next().unwrap();
1170 let ast::Stmt::Vacuum(vacuum) = stmt else {
1171 unreachable!()
1172 };
1173 vacuum
1174}
1175
1176#[test]
1177fn vacuum_full_is_full() {
1178 assert!(extract_vacuum("VACUUM FULL foo;").is_full());
1179}
1180
1181#[test]
1182fn vacuum_option_list_full_is_full() {
1183 assert!(extract_vacuum("VACUUM (FULL) foo;").is_full());
1184}
1185
1186#[test]
1187fn vacuum_full_true_is_full() {
1188 assert!(extract_vacuum("VACUUM (FULL TRUE) foo;").is_full());
1189}
1190
1191#[test]
1192fn vacuum_full_on_is_full() {
1193 assert!(extract_vacuum("VACUUM (FULL ON) foo;").is_full());
1194}
1195
1196#[test]
1197fn vacuum_full_1_is_full() {
1198 assert!(extract_vacuum("VACUUM (FULL 1) foo;").is_full());
1199}
1200
1201#[test]
1202fn vacuum_no_full_is_not_full() {
1203 assert!(!extract_vacuum("VACUUM foo;").is_full());
1204}
1205
1206#[test]
1207fn vacuum_other_option_is_not_full() {
1208 assert!(!extract_vacuum("VACUUM (FREEZE) foo;").is_full());
1209}
1210
1211#[test]
1212fn vacuum_full_false_is_not_full() {
1213 assert!(!extract_vacuum("VACUUM (FULL FALSE) foo;").is_full());
1214}
1215
1216#[test]
1217fn vacuum_full_off_is_not_full() {
1218 assert!(!extract_vacuum("VACUUM (FULL OFF) foo;").is_full());
1219}
1220
1221#[test]
1222fn vacuum_full_no_is_not_full() {
1223 assert!(!extract_vacuum("VACUUM (FULL NO) foo;").is_full());
1224}
1225
1226#[test]
1227fn vacuum_full_quoted_off_is_not_full() {
1228 assert!(!extract_vacuum("VACUUM (FULL 'off') foo;").is_full());
1229}
1230
1231#[test]
1232fn vacuum_full_escaped_string_off_is_not_full() {
1233 assert!(!extract_vacuum("VACUUM (FULL E'off') foo;").is_full());
1234}
1235
1236#[test]
1237fn vacuum_full_unicode_escaped_string_off_is_not_full() {
1238 assert!(!extract_vacuum("VACUUM (FULL U&'off') foo;").is_full());
1239}
1240
1241#[test]
1242fn vacuum_full_dollar_quoted_off_is_not_full() {
1243 assert!(!extract_vacuum("VACUUM (FULL $$off$$) t;").is_full());
1244}
1245
1246#[test]
1247fn vacuum_full_0_is_not_full() {
1248 assert!(!extract_vacuum("VACUUM (FULL 0) foo;").is_full());
1249}