Skip to main content

squawk_syntax/
plpgsql.rs

1use rowan::GreenNode;
2
3use crate::{
4    ast,
5    body::{Body, BodyLanguage},
6    parsing,
7    syntax_error::SyntaxError,
8};
9
10pub type Plpgsql = Body<ast::Plpgsql>;
11
12impl BodyLanguage for ast::Plpgsql {
13    const LANGUAGE: &'static str = "plpgsql";
14
15    fn parse_text(text: &str) -> (GreenNode, Vec<SyntaxError>) {
16        parsing::parse_plpgsql_text(text)
17    }
18}
19
20impl ast::CreateFunction {
21    pub fn plpgsql(&self) -> Option<Plpgsql> {
22        Plpgsql::from_options(self.option_list()?)
23    }
24}
25
26impl ast::CreateProcedure {
27    pub fn plpgsql(&self) -> Option<Plpgsql> {
28        Plpgsql::from_options(self.option_list()?)
29    }
30}
31
32impl ast::Do {
33    pub fn plpgsql(&self) -> Option<Plpgsql> {
34        if let Some(language) = self.do_language()
35            && !ast::Plpgsql::is_language(language.language_name())
36        {
37            return None;
38        }
39        Some(Plpgsql::parse(self.body()?.decoded_value()?))
40    }
41}
42
43#[cfg(test)]
44mod tests {
45    use super::*;
46    use crate::SourceFile;
47    use crate::ast::AstNode;
48    use crate::test::render_errors;
49    use insta::assert_snapshot;
50    use rowan::{TextRange, TextSize};
51
52    fn find(sql: &str) -> Option<Plpgsql> {
53        let parse = SourceFile::parse(sql);
54        assert!(parse.errors().is_empty(), "{:?}", parse.errors());
55
56        parse.tree().syntax().descendants().find_map(|node| {
57            ast::CreateFunction::cast(node.clone())
58                .and_then(|it| it.plpgsql())
59                .or_else(|| ast::CreateProcedure::cast(node.clone()).and_then(|it| it.plpgsql()))
60                .or_else(|| ast::Do::cast(node.clone()).and_then(|it| it.plpgsql()))
61        })
62    }
63
64    fn body(sql: &str) -> String {
65        let body = find(sql).expect("no plpgsql body");
66
67        let range = body.source_range(TextRange::up_to(TextSize::of(body.text())));
68        let start = usize::from(range.start());
69        let end = usize::from(range.end());
70
71        let mut out = format!("{:#?}", body.syntax());
72        out.push_str(&format!("---\nsource {range:?} {:?}\n", &sql[start..end]));
73        out.push_str(&render_errors(sql, &body.errors()));
74        out
75    }
76
77    #[test]
78    fn language_after_as() {
79        assert_snapshot!(
80            body("create function f() returns int as $$ begin null; end $$ language plpgsql;"),
81            @r#"
82        PLPGSQL@0..17
83          WHITESPACE@0..1 " "
84          PLPGSQL_BLOCK@1..16
85            BEGIN_KW@1..6 "begin"
86            WHITESPACE@6..7 " "
87            PLPGSQL_BODY@7..12
88              PLPGSQL_NULL_STMT@7..12
89                NULL_KW@7..11 "null"
90                SEMICOLON@11..12 ";"
91            WHITESPACE@12..13 " "
92            END_KW@13..16 "end"
93          WHITESPACE@16..17 " "
94        ---
95        source 37..54 " begin null; end "
96        "#
97        );
98    }
99
100    #[test]
101    fn language_before_as() {
102        assert_snapshot!(
103            body("create function f() returns int language plpgsql as $$ begin null; end $$;"),
104            @r#"
105        PLPGSQL@0..17
106          WHITESPACE@0..1 " "
107          PLPGSQL_BLOCK@1..16
108            BEGIN_KW@1..6 "begin"
109            WHITESPACE@6..7 " "
110            PLPGSQL_BODY@7..12
111              PLPGSQL_NULL_STMT@7..12
112                NULL_KW@7..11 "null"
113                SEMICOLON@11..12 ";"
114            WHITESPACE@12..13 " "
115            END_KW@13..16 "end"
116          WHITESPACE@16..17 " "
117        ---
118        source 54..71 " begin null; end "
119        "#
120        );
121    }
122
123    #[test]
124    fn other_language_is_not_a_body() {
125        assert!(find("create function f() returns int as $$ select 1 $$ language sql;").is_none());
126    }
127
128    #[test]
129    fn procedure() {
130        assert_snapshot!(
131            body("create procedure p() as $$ begin null; end $$ language plpgsql;"),
132            @r#"
133        PLPGSQL@0..17
134          WHITESPACE@0..1 " "
135          PLPGSQL_BLOCK@1..16
136            BEGIN_KW@1..6 "begin"
137            WHITESPACE@6..7 " "
138            PLPGSQL_BODY@7..12
139              PLPGSQL_NULL_STMT@7..12
140                NULL_KW@7..11 "null"
141                SEMICOLON@11..12 ";"
142            WHITESPACE@12..13 " "
143            END_KW@13..16 "end"
144          WHITESPACE@16..17 " "
145        ---
146        source 26..43 " begin null; end "
147        "#
148        );
149    }
150
151    #[test]
152    fn do_defaults_to_plpgsql() {
153        assert_snapshot!(body("do $$ begin null; end $$;"), @r#"
154        PLPGSQL@0..17
155          WHITESPACE@0..1 " "
156          PLPGSQL_BLOCK@1..16
157            BEGIN_KW@1..6 "begin"
158            WHITESPACE@6..7 " "
159            PLPGSQL_BODY@7..12
160              PLPGSQL_NULL_STMT@7..12
161                NULL_KW@7..11 "null"
162                SEMICOLON@11..12 ";"
163            WHITESPACE@12..13 " "
164            END_KW@13..16 "end"
165          WHITESPACE@16..17 " "
166        ---
167        source 5..22 " begin null; end "
168        "#
169        );
170    }
171
172    #[test]
173    fn do_with_other_language_is_not_a_body() {
174        assert!(find("do language plpython3u $$ return 1 $$;").is_none());
175    }
176
177    #[test]
178    fn escaped_body_maps_back_through_the_escapes() {
179        assert_snapshot!(
180            body(r"create function f() returns int as E'begin null; end\n' language plpgsql;"),
181            @r#"
182        PLPGSQL@0..16
183          PLPGSQL_BLOCK@0..15
184            BEGIN_KW@0..5 "begin"
185            WHITESPACE@5..6 " "
186            PLPGSQL_BODY@6..11
187              PLPGSQL_NULL_STMT@6..11
188                NULL_KW@6..10 "null"
189                SEMICOLON@10..11 ";"
190            WHITESPACE@11..12 " "
191            END_KW@12..15 "end"
192          WHITESPACE@15..16 "\n"
193        ---
194        source 37..54 "begin null; end\\n"
195        "#
196        );
197    }
198
199    #[test]
200    fn unparsed_tokens_are_errors() {
201        assert_snapshot!(body("do $$ begin null null; end $$;"), @r#"
202        PLPGSQL@0..22
203          WHITESPACE@0..1 " "
204          PLPGSQL_BLOCK@1..21
205            BEGIN_KW@1..6 "begin"
206            WHITESPACE@6..7 " "
207            PLPGSQL_BODY@7..17
208              ERROR@7..17
209                NULL_KW@7..11 "null"
210                WHITESPACE@11..12 " "
211                NULL_KW@12..16 "null"
212                SEMICOLON@16..17 ";"
213            WHITESPACE@17..18 " "
214            END_KW@18..21 "end"
215          WHITESPACE@21..22 " "
216        ---
217        source 5..27 " begin null null; end "
218        error[syntax-error]: expected a statement, got NULL_KW
219          ╭▸ 
220        1 │ do $$ begin null null; end $$;
221          ╰╴            ━
222        "#);
223    }
224
225    #[test]
226    fn errors_in_an_escaped_body_map_into_the_file() {
227        assert_snapshot!(body(r"do E'begin\n null null;\n end';"), @r#"
228        PLPGSQL@0..22
229          PLPGSQL_BLOCK@0..22
230            BEGIN_KW@0..5 "begin"
231            WHITESPACE@5..7 "\n "
232            PLPGSQL_BODY@7..17
233              ERROR@7..17
234                NULL_KW@7..11 "null"
235                WHITESPACE@11..12 " "
236                NULL_KW@12..16 "null"
237                SEMICOLON@16..17 ";"
238            WHITESPACE@17..19 "\n "
239            END_KW@19..22 "end"
240        ---
241        source 5..29 "begin\\n null null;\\n end"
242        error[syntax-error]: expected a statement, got NULL_KW
243          ╭▸ 
244        1 │ do E'begin\n null null;\n end';
245          ╰╴             ━
246        "#);
247    }
248}