Skip to main content

squawk_syntax/ast/
node_ext.rs

1// via https://github.com/rust-lang/rust-analyzer/blob/d8887c0758bbd2d5f752d5bd405d4491e90e7ed6/crates/syntax/src/ast/node_ext.rs
2//
3// Permission is hereby granted, free of charge, to any
4// person obtaining a copy of this software and associated
5// documentation files (the "Software"), to deal in the
6// Software without restriction, including without
7// limitation the rights to use, copy, modify, merge,
8// publish, distribute, sublicense, and/or sell copies of
9// the Software, and to permit persons to whom the Software
10// is furnished to do so, subject to the following
11// conditions:
12//
13// The above copyright notice and this permission notice
14// shall be included in all copies or substantial portions
15// of the Software.
16//
17// THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF
18// ANY KIND, EXPRESS OR IMPLIED, INCLUDING BUT NOT LIMITED
19// TO THE WARRANTIES OF MERCHANTABILITY, FITNESS FOR A
20// PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT
21// SHALL THE AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY
22// CLAIM, DAMAGES OR OTHER LIABILITY, WHETHER IN AN ACTION
23// OF CONTRACT, TORT OR OTHERWISE, ARISING FROM, OUT OF OR
24// IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER
25// DEALINGS IN THE SOFTWARE.
26
27use 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    // We have NameRef as a variant of Expr which complicates things (and it
329    // might not be worth it).
330    // Rust analyzer doesn't do this so it doesn't have to special case this.
331    #[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        // With `select x[1:]`, we have two exprs, `x` and `1`.
361        // We skip over the first one, and then we want the second one, but we
362        // want to make sure we don't choose the end expr if instead we had:
363        // `select x[:1]`
364        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        // We want to make sure we get the last expr after the `:` which is the
373        // end of the slice, i.e., `2` in: `select x[:2]`
374        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
495// TODO: return a NewType wrapper around String?
496fn 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            // TODO: we need a better way of handling option lists
605            || 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}