diff --git a/crates/squawk_fmt/src/fmt.rs b/crates/squawk_fmt/src/fmt.rs index 49c89b55..6cd1a3d6 100644 --- a/crates/squawk_fmt/src/fmt.rs +++ b/crates/squawk_fmt/src/fmt.rs @@ -12042,14 +12042,13 @@ fn build_op_sig<'a>(sig: ast::OpSig) -> Doc<'a> { doc = doc.append(comments_before(&l_paren)); } let has_none = sig.none_token().is_some(); - let lhs = if let Some(none) = sig.none_token() { + let mut body = if let Some(none) = sig.none_token() { leading_comments(&none).append(Doc::text("none")) } else if let Some(lhs) = sig.lhs() { leading_comments(lhs.syntax()).append(build_type(lhs)) } else { Doc::nil() }; - let mut body = lhs; if let Some(comma) = sig.comma_token() { body = body .append(comments_before(&comma)) diff --git a/crates/squawk_ide/src/inlay_hints.rs b/crates/squawk_ide/src/inlay_hints.rs index 334c9614..b66cf633 100644 --- a/crates/squawk_ide/src/inlay_hints.rs +++ b/crates/squawk_ide/src/inlay_hints.rs @@ -292,7 +292,7 @@ mod test { #[test] fn single_param() { assert_snapshot!(check_inlay_hints(" -create function foo(a int) returns int as 'select $$1' language sql; +create function foo(a int) returns int as 'select $1' language sql; select foo(1); "), @" labels: @@ -303,7 +303,7 @@ select foo(1); targets: ╭▸ current.sql:2:21 │ - 2 │ create function foo(a int) returns int as 'select $$1' language sql; + 2 │ create function foo(a int) returns int as 'select $1' language sql; ╰╴ ─ 1. target "); } @@ -311,7 +311,7 @@ select foo(1); #[test] fn multiple_params() { assert_snapshot!(check_inlay_hints(" -create function add(a int, b int) returns int as 'select $$1 + $$2' language sql; +create function add(a int, b int) returns int as 'select $1 + $2' language sql; select add(1, 2); "), @" labels: @@ -324,7 +324,7 @@ select add(1, 2); targets: ╭▸ current.sql:2:21 │ - 2 │ create function add(a int, b int) returns int as 'select $$1 + $$2' language sql; + 2 │ create function add(a int, b int) returns int as 'select $1 + $2' language sql; │ ┬ ─ 2. target │ │ ╰╴ 1. target @@ -342,7 +342,7 @@ select foo(); #[test] fn with_schema() { assert_snapshot!(check_inlay_hints(" -create function public.foo(x int) returns int as 'select $$1' language sql; +create function public.foo(x int) returns int as 'select $1' language sql; select public.foo(42); "), @" labels: @@ -353,7 +353,7 @@ select public.foo(42); targets: ╭▸ current.sql:2:28 │ - 2 │ create function public.foo(x int) returns int as 'select $$1' language sql; + 2 │ create function public.foo(x int) returns int as 'select $1' language sql; ╰╴ ─ 1. target "); } @@ -362,7 +362,7 @@ select public.foo(42); fn with_search_path() { assert_snapshot!(check_inlay_hints(r#" set search_path to myschema; -create function foo(val int) returns int as 'select $$1' language sql; +create function foo(val int) returns int as 'select $1' language sql; select foo(100); "#), @" labels: @@ -373,7 +373,7 @@ select foo(100); targets: ╭▸ current.sql:3:21 │ - 3 │ create function foo(val int) returns int as 'select $$1' language sql; + 3 │ create function foo(val int) returns int as 'select $1' language sql; ╰╴ ─── 1. target "); } @@ -381,7 +381,7 @@ select foo(100); #[test] fn multiple_calls() { assert_snapshot!(check_inlay_hints(" -create function inc(n int) returns int as 'select $$1 + 1' language sql; +create function inc(n int) returns int as 'select $1 + 1' language sql; select inc(1), inc(2); "), @" labels: @@ -394,7 +394,7 @@ select inc(1), inc(2); targets: ╭▸ current.sql:2:21 │ - 2 │ create function inc(n int) returns int as 'select $$1 + 1' language sql; + 2 │ create function inc(n int) returns int as 'select $1 + 1' language sql; │ ┬ │ │ │ 1. target @@ -405,7 +405,7 @@ select inc(1), inc(2); #[test] fn more_args_than_params() { assert_snapshot!(check_inlay_hints(" -create function foo(a int) returns int as 'select $$1' language sql; +create function foo(a int) returns int as 'select $1' language sql; select foo(1, 2); "), @" labels: @@ -416,7 +416,7 @@ select foo(1, 2); targets: ╭▸ current.sql:2:21 │ - 2 │ create function foo(a int) returns int as 'select $$1' language sql; + 2 │ create function foo(a int) returns int as 'select $1' language sql; ╰╴ ─ 1. target "); } diff --git a/crates/squawk_parser/src/grammar.rs b/crates/squawk_parser/src/grammar.rs index 917bb2a6..e461da53 100644 --- a/crates/squawk_parser/src/grammar.rs +++ b/crates/squawk_parser/src/grammar.rs @@ -3470,11 +3470,10 @@ fn compound_select_operand(p: &mut Parser<'_>) -> Option { fn compound_select_bp( p: &mut Parser<'_>, - lhs: CompletedMarker, + mut lhs: CompletedMarker, min_bp: u8, r: &SelectRestrictions, ) -> CompletedMarker { - let mut lhs = lhs; while let Some(bp) = compound_op_bp(p) { if bp < min_bp { break; @@ -3580,9 +3579,8 @@ fn select_tail( p: &mut Parser, m: Marker, r: &SelectRestrictions, - out_kind: SyntaxKind, + mut out_kind: SyntaxKind, ) -> CompletedMarker { - let mut out_kind = out_kind; if opt_into_clause(p).is_some() { out_kind = SELECT_INTO; } @@ -4694,7 +4692,7 @@ enum ColumnDefKind { // select * from f() as t(a int, b text, c text collate foo.bar.buzz); // ^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^ // [ ( column_name [, ... ] ) ] -fn opt_column_list_with(p: &mut Parser<'_>, kind: ColumnDefKind) -> bool { +fn opt_column_list_with(p: &mut Parser<'_>, mut kind: ColumnDefKind) -> bool { if !p.at(L_PAREN) || // we're probably at (select) !p.nth_at_ts(1, COLUMN_FIRST) && !p.nth_at(1, R_PAREN) && !p.nth_at(1, COMMA) @@ -4717,7 +4715,6 @@ fn opt_column_list_with(p: &mut Parser<'_>, kind: ColumnDefKind) -> bool { if matches!(items, ListItems::Required) && p.at(R_PAREN) { p.error("Expected at least one item"); } - let mut kind = kind; let mut seen_period = false; while !p.at(EOF) && !p.at(R_PAREN) { if p.at(COMMA) { diff --git a/crates/squawk_parser/tests/data/ok/create_function.sql b/crates/squawk_parser/tests/data/ok/create_function.sql index 60197839..cbc2b99a 100644 --- a/crates/squawk_parser/tests/data/ok/create_function.sql +++ b/crates/squawk_parser/tests/data/ok/create_function.sql @@ -307,7 +307,7 @@ language sql; create function f(bitmask bit(8)) returns boolean -as '0' +as 'select 0' language sql; -- argmode diff --git a/crates/squawk_parser/tests/data/ok/create_procedure.sql b/crates/squawk_parser/tests/data/ok/create_procedure.sql index 1a2fc034..f9b378a4 100644 --- a/crates/squawk_parser/tests/data/ok/create_procedure.sql +++ b/crates/squawk_parser/tests/data/ok/create_procedure.sql @@ -17,9 +17,9 @@ set buzz from current return 10 + 1; -- security_variants -create procedure p() language sql security invoker as 'foo'; -create procedure p() language sql external security definer as 'foo'; -create procedure p() language sql security definer as 'foo'; +create procedure p() language sql security invoker as 'select 1'; +create procedure p() language sql external security definer as 'select 1'; +create procedure p() language sql security definer as 'select 1'; -- as_with_two_strings create procedure p() language c as 'foo', 'bar'; diff --git a/crates/squawk_parser/tests/snapshots/tests__create_function_ok.snap b/crates/squawk_parser/tests/snapshots/tests__create_function_ok.snap index 24e140ef..b5d26f69 100644 --- a/crates/squawk_parser/tests/snapshots/tests__create_function_ok.snap +++ b/crates/squawk_parser/tests/snapshots/tests__create_function_ok.snap @@ -2272,7 +2272,7 @@ SOURCE_FILE WHITESPACE " " AS_DEFINITION LITERAL - STRING "'0'" + STRING "'select 0'" WHITESPACE "\n" LANGUAGE_FUNC_OPTION LANGUAGE_KW "language" diff --git a/crates/squawk_parser/tests/snapshots/tests__create_procedure_ok.snap b/crates/squawk_parser/tests/snapshots/tests__create_procedure_ok.snap index 28696b64..14ade6a2 100644 --- a/crates/squawk_parser/tests/snapshots/tests__create_procedure_ok.snap +++ b/crates/squawk_parser/tests/snapshots/tests__create_procedure_ok.snap @@ -228,7 +228,7 @@ SOURCE_FILE WHITESPACE " " AS_DEFINITION LITERAL - STRING "'foo'" + STRING "'select 1'" SEMICOLON ";" WHITESPACE "\n" CREATE_PROCEDURE @@ -263,7 +263,7 @@ SOURCE_FILE WHITESPACE " " AS_DEFINITION LITERAL - STRING "'foo'" + STRING "'select 1'" SEMICOLON ";" WHITESPACE "\n" CREATE_PROCEDURE @@ -296,7 +296,7 @@ SOURCE_FILE WHITESPACE " " AS_DEFINITION LITERAL - STRING "'foo'" + STRING "'select 1'" SEMICOLON ";" WHITESPACE "\n\n" COMMENT "-- as_with_two_strings" diff --git a/crates/squawk_syntax/src/ast.rs b/crates/squawk_syntax/src/ast.rs index d330797f..9c4acc0f 100644 --- a/crates/squawk_syntax/src/ast.rs +++ b/crates/squawk_syntax/src/ast.rs @@ -36,7 +36,7 @@ use std::marker::PhantomData; use crate::syntax_node::{SyntaxNode, SyntaxNodeChildren, SyntaxToken}; use squawk_parser::SyntaxKind; -pub use self::node_ext::{is_quoted_name_node, normalize_name_node}; +pub use self::node_ext::{LanguageName, is_quoted_name_node, normalize_name_node}; pub use self::{ generated::tokens::*, node_ext::{BinOp, CastKind, LitKind, PostfixOp, PrefixOp}, diff --git a/crates/squawk_syntax/src/ast/node_ext.rs b/crates/squawk_syntax/src/ast/node_ext.rs index eba7454f..c19b7c22 100644 --- a/crates/squawk_syntax/src/ast/node_ext.rs +++ b/crates/squawk_syntax/src/ast/node_ext.rs @@ -29,7 +29,7 @@ use std::borrow::Cow; use either::Either; #[cfg(test)] use insta::assert_snapshot; -use rowan::{GreenNodeData, GreenTokenData, NodeOrToken, TextSize}; +use rowan::{GreenNodeData, GreenTokenData, NodeOrToken, TextRange, TextSize}; use squawk_line_index::{LineEnding, find_newline}; #[cfg(test)] @@ -997,6 +997,41 @@ pub fn is_quoted_name_node(node: &SyntaxNode) -> bool { ) } +pub struct LanguageName { + pub name: String, + pub range: TextRange, +} + +fn language_name( + language_ref: Option, + literal: Option, +) -> Option { + if let Some(language_ref) = language_ref { + return Some(LanguageName { + name: normalize_name_node(language_ref.syntax()), + range: language_ref.syntax().text_range(), + }); + } + + let literal = literal?; + Some(LanguageName { + name: literal.string_value()?, + range: literal.syntax().text_range(), + }) +} + +impl ast::LanguageFuncOption { + pub fn language_name(&self) -> Option { + language_name(self.language_ref(), self.literal()) + } +} + +impl ast::DoLanguage { + pub fn language_name(&self) -> Option { + language_name(self.language_ref(), self.literal()) + } +} + // TODO: return a NewType wrapper around String? pub fn normalize_name_node(node: &SyntaxNode) -> String { let mut tokens = node diff --git a/crates/squawk_syntax/src/body.rs b/crates/squawk_syntax/src/body.rs new file mode 100644 index 00000000..86b93cd6 --- /dev/null +++ b/crates/squawk_syntax/src/body.rs @@ -0,0 +1,99 @@ +use std::marker::PhantomData; + +use rowan::{GreenNode, TextRange}; + +use crate::{ + SyntaxNode, ast, ast::AstNode, decoded_text::DecodedText, syntax_error::SyntaxError, validation, +}; + +pub struct Body { + green: GreenNode, + errors: Vec, + decoded: DecodedText, + _ty: PhantomData T>, +} + +pub trait BodyLanguage: AstNode { + const LANGUAGE: &'static str; + + fn parse_text(text: &str) -> (GreenNode, Vec); + + fn is_language(language: Option) -> bool { + language.is_some_and(|language| language.name == Self::LANGUAGE) + } +} + +impl Body { + pub const LANGUAGE: &'static str = T::LANGUAGE; + + pub(crate) fn parse(decoded: DecodedText) -> Self { + let (green, errors) = T::parse_text(decoded.text()); + let errors = errors + .into_iter() + .map(|error| { + let range = decoded.source_range(error.range()); + error.with_range(range) + }) + .collect(); + + Self { + green, + errors, + decoded, + _ty: PhantomData, + } + } + + pub(crate) fn from_options(options: ast::FuncOptionList) -> Option { + let mut matches = false; + let mut body = None; + + for option in options.options() { + match option { + ast::FuncOption::LanguageFuncOption(option) => { + matches = T::is_language(option.language_name()); + } + ast::FuncOption::AsFuncOption(option) => { + if let Some(ast::AsFuncTarget::AsDefinition(definition)) = + option.as_func_target() + { + body = definition.literal(); + } + } + _ => (), + } + } + + matches.then_some(())?; + Some(Self::parse(body?.decoded_value()?)) + } + + pub fn syntax(&self) -> SyntaxNode { + SyntaxNode::new_root(self.green.clone()) + } + + pub fn tree(&self) -> T { + T::cast(self.syntax()).expect("root is always the body's node") + } + + pub fn text(&self) -> &str { + self.decoded.text() + } + + pub fn source_range(&self, range: TextRange) -> TextRange { + self.decoded.source_range(range) + } + + pub fn errors(&self) -> Vec { + let mut validation_errors = vec![]; + validation::validate(&self.syntax(), &mut validation_errors); + + let mut errors = self.errors.clone(); + errors.extend(validation_errors.into_iter().map(|error| { + let range = self.decoded.source_range(error.range()); + error.with_range(range) + })); + errors.sort_by_key(|error| error.range().start()); + errors + } +} diff --git a/crates/squawk_syntax/src/lib.rs b/crates/squawk_syntax/src/lib.rs index b26685da..8e74a126 100644 --- a/crates/squawk_syntax/src/lib.rs +++ b/crates/squawk_syntax/src/lib.rs @@ -25,6 +25,7 @@ // DEALINGS IN THE SOFTWARE. pub mod ast; +pub mod body; pub mod column_name; pub mod decoded_text; mod generated; @@ -32,6 +33,7 @@ mod parsing; pub mod plpgsql; mod ptr; pub mod quote; +pub mod sql_body; pub mod syntax_error; mod syntax_node; mod token_text; @@ -98,6 +100,8 @@ impl Parse { vec![] }; validation::validate(&self.syntax_node(), &mut errors); + let file = SourceFile::cast(self.syntax_node()).expect("parse root is always a SourceFile"); + errors.extend(file.sql_body_errors()); errors.sort_by_key(|error| error.range().start()); errors } diff --git a/crates/squawk_syntax/src/plpgsql.rs b/crates/squawk_syntax/src/plpgsql.rs index 01de146c..a993e3a8 100644 --- a/crates/squawk_syntax/src/plpgsql.rs +++ b/crates/squawk_syntax/src/plpgsql.rs @@ -1,111 +1,38 @@ -use rowan::{GreenNode, TextRange}; +use rowan::GreenNode; use crate::{ - SyntaxNode, ast, ast::AstNode, decoded_text::DecodedText, parsing, syntax_error::SyntaxError, - validation, + ast, + body::{Body, BodyLanguage}, + parsing, + syntax_error::SyntaxError, }; -pub struct Plpgsql { - green: GreenNode, - errors: Vec, - decoded: DecodedText, -} - -impl Plpgsql { - pub(crate) fn parse(decoded: DecodedText) -> Self { - let (green, errors) = parsing::parse_plpgsql_text(decoded.text()); - let errors = errors - .into_iter() - .map(|error| { - let range = decoded.source_range(error.range()); - error.with_range(range) - }) - .collect(); - - Self { - green, - errors, - decoded, - } - } +pub type Plpgsql = Body; - pub fn syntax(&self) -> SyntaxNode { - SyntaxNode::new_root(self.green.clone()) - } +impl BodyLanguage for ast::Plpgsql { + const LANGUAGE: &'static str = "plpgsql"; - pub fn tree(&self) -> ast::Plpgsql { - ast::Plpgsql::cast(self.syntax()).expect("root is always a Plpgsql") + fn parse_text(text: &str) -> (GreenNode, Vec) { + parsing::parse_plpgsql_text(text) } - - pub fn text(&self) -> &str { - self.decoded.text() - } - - pub fn source_range(&self, range: TextRange) -> TextRange { - self.decoded.source_range(range) - } - - pub fn errors(self) -> Vec { - let mut validation_errors = vec![]; - validation::validate(&self.syntax(), &mut validation_errors); - - let mut errors = self.errors; - errors.extend(validation_errors.into_iter().map(|error| { - let range = self.decoded.source_range(error.range()); - error.with_range(range) - })); - errors.sort_by_key(|error| error.range().start()); - errors - } -} - -fn is_plpgsql(language_ref: Option, literal: Option) -> bool { - let name = match (language_ref, literal) { - (Some(language_ref), _) => language_ref.syntax().text().to_string(), - (_, Some(literal)) => literal.string_value().unwrap_or_default(), - _ => return false, - }; - name.eq_ignore_ascii_case("plpgsql") -} - -fn from_options(options: ast::FuncOptionList) -> Option { - let mut plpgsql = false; - let mut body = None; - - for option in options.options() { - match option { - ast::FuncOption::LanguageFuncOption(option) => { - plpgsql = is_plpgsql(option.language_ref(), option.literal()); - } - ast::FuncOption::AsFuncOption(option) => { - if let Some(ast::AsFuncTarget::AsDefinition(definition)) = option.as_func_target() { - body = definition.literal(); - } - } - _ => (), - } - } - - plpgsql.then_some(())?; - Some(Plpgsql::parse(body?.decoded_value()?)) } impl ast::CreateFunction { pub fn plpgsql(&self) -> Option { - from_options(self.option_list()?) + Plpgsql::from_options(self.option_list()?) } } impl ast::CreateProcedure { pub fn plpgsql(&self) -> Option { - from_options(self.option_list()?) + Plpgsql::from_options(self.option_list()?) } } impl ast::Do { pub fn plpgsql(&self) -> Option { if let Some(language) = self.do_language() - && !is_plpgsql(language.language_ref(), language.literal()) + && !ast::Plpgsql::is_language(language.language_name()) { return None; } @@ -117,6 +44,7 @@ impl ast::Do { mod tests { use super::*; use crate::SourceFile; + use crate::ast::AstNode; use crate::test::render_errors; use insta::assert_snapshot; use rowan::{TextRange, TextSize}; @@ -243,7 +171,7 @@ mod tests { #[test] fn do_with_other_language_is_not_a_body() { - assert!(find("do language sql $$ select 1 $$;").is_none()); + assert!(find("do language plpython3u $$ return 1 $$;").is_none()); } #[test] diff --git a/crates/squawk_syntax/src/snapshots/squawk_syntax__sql_body__tests__errors_map_into_the_containing_file.snap b/crates/squawk_syntax/src/snapshots/squawk_syntax__sql_body__tests__errors_map_into_the_containing_file.snap new file mode 100644 index 00000000..1d159b1b --- /dev/null +++ b/crates/squawk_syntax/src/snapshots/squawk_syntax__sql_body__tests__errors_map_into_the_containing_file.snap @@ -0,0 +1,19 @@ +--- +source: crates/squawk_syntax/src/sql_body.rs +expression: "body(\"create function f() returns int\n language sql\n as $$ select from $$;\")" +--- +SOURCE_FILE@0..13 + WHITESPACE@0..1 " " + SELECT@1..12 + SELECT_CLAUSE@1..7 + SELECT_KW@1..7 "select" + WHITESPACE@7..8 " " + FROM_CLAUSE@8..12 + FROM_KW@8..12 "from" + WHITESPACE@12..13 " " +--- +source 54..67 " select from " +error[syntax-error]: expected from item, got EOF + ╭▸ +3 │ as $$ select from $$; + ╰╴ ━ diff --git a/crates/squawk_syntax/src/snapshots/squawk_syntax__sql_body__tests__function_language_after_as.snap b/crates/squawk_syntax/src/snapshots/squawk_syntax__sql_body__tests__function_language_after_as.snap new file mode 100644 index 00000000..1c4223c4 --- /dev/null +++ b/crates/squawk_syntax/src/snapshots/squawk_syntax__sql_body__tests__function_language_after_as.snap @@ -0,0 +1,15 @@ +--- +source: crates/squawk_syntax/src/sql_body.rs +expression: "body(\"create function f(int) returns int\n as 'select $1'\n language sql;\")" +--- +SOURCE_FILE@0..9 + SELECT@0..9 + SELECT_CLAUSE@0..9 + SELECT_KW@0..6 "select" + WHITESPACE@6..7 " " + TARGET_LIST@7..9 + TARGET@7..9 + LITERAL@7..9 + POSITIONAL_PARAM@7..9 "$1" +--- +source 41..50 "select $1" diff --git a/crates/squawk_syntax/src/snapshots/squawk_syntax__sql_body__tests__function_language_before_as.snap b/crates/squawk_syntax/src/snapshots/squawk_syntax__sql_body__tests__function_language_before_as.snap new file mode 100644 index 00000000..7f3c2fd6 --- /dev/null +++ b/crates/squawk_syntax/src/snapshots/squawk_syntax__sql_body__tests__function_language_before_as.snap @@ -0,0 +1,23 @@ +--- +source: crates/squawk_syntax/src/sql_body.rs +expression: "body(\"create function f(int) returns int\n language sql\n as $$ select $1 + 1 $$;\")" +--- +SOURCE_FILE@0..15 + WHITESPACE@0..1 " " + SELECT@1..14 + SELECT_CLAUSE@1..14 + SELECT_KW@1..7 "select" + WHITESPACE@7..8 " " + TARGET_LIST@8..14 + TARGET@8..14 + BIN_EXPR@8..14 + LITERAL@8..10 + POSITIONAL_PARAM@8..10 "$1" + WHITESPACE@10..11 " " + PLUS@11..12 "+" + WHITESPACE@12..13 " " + LITERAL@13..14 + INT_NUMBER@13..14 "1" + WHITESPACE@14..15 " " +--- +source 57..72 " select $1 + 1 " diff --git a/crates/squawk_syntax/src/snapshots/squawk_syntax__sql_body__tests__procedure.snap b/crates/squawk_syntax/src/snapshots/squawk_syntax__sql_body__tests__procedure.snap new file mode 100644 index 00000000..d0678664 --- /dev/null +++ b/crates/squawk_syntax/src/snapshots/squawk_syntax__sql_body__tests__procedure.snap @@ -0,0 +1,27 @@ +--- +source: crates/squawk_syntax/src/sql_body.rs +expression: "body(\"create procedure p()\n language sql\n as $$ select 1; select 2 $$;\")" +--- +SOURCE_FILE@0..20 + WHITESPACE@0..1 " " + SELECT@1..10 + SELECT_CLAUSE@1..9 + SELECT_KW@1..7 "select" + WHITESPACE@7..8 " " + TARGET_LIST@8..9 + TARGET@8..9 + LITERAL@8..9 + INT_NUMBER@8..9 "1" + SEMICOLON@9..10 ";" + WHITESPACE@10..11 " " + SELECT@11..19 + SELECT_CLAUSE@11..19 + SELECT_KW@11..17 "select" + WHITESPACE@17..18 " " + TARGET_LIST@18..19 + TARGET@18..19 + LITERAL@18..19 + INT_NUMBER@18..19 "2" + WHITESPACE@19..20 " " +--- +source 43..63 " select 1; select 2 " diff --git a/crates/squawk_syntax/src/snapshots/squawk_syntax__test__conflicting_options_validation.snap b/crates/squawk_syntax/src/snapshots/squawk_syntax__test__conflicting_options_validation.snap index 11845c5a..b5cc0d59 100644 --- a/crates/squawk_syntax/src/snapshots/squawk_syntax__test__conflicting_options_validation.snap +++ b/crates/squawk_syntax/src/snapshots/squawk_syntax__test__conflicting_options_validation.snap @@ -474,6 +474,10 @@ error[syntax-error]: Conflicting or redundant options. ╭▸ 2 │ do language plpgsql $$ x $$ language sql; ╰╴ ━━━━━━━━━━━━ +error[syntax-error]: SQL is not a valid language for DO statements. + ╭▸ +2 │ do language plpgsql $$ x $$ language sql; + ╰╴ ━━━ error[syntax-error]: Conflicting or redundant options. ╭▸ 7 │ language sql diff --git a/crates/squawk_syntax/src/snapshots/squawk_syntax__test__do_language_ok_validation.snap b/crates/squawk_syntax/src/snapshots/squawk_syntax__test__do_language_ok_validation.snap new file mode 100644 index 00000000..5e3106bc --- /dev/null +++ b/crates/squawk_syntax/src/snapshots/squawk_syntax__test__do_language_ok_validation.snap @@ -0,0 +1,64 @@ +--- +source: crates/squawk_syntax/src/test.rs +input_file: crates/squawk_syntax/test_data/validation/do_language_ok.sql +--- +SOURCE_FILE@0..181 + DO@0..25 + DO_KW@0..2 "do" + WHITESPACE@2..3 " " + LITERAL@3..24 + DOLLAR_QUOTED_STRING@3..24 "$$ begin null; end $$" + SEMICOLON@24..25 ";" + WHITESPACE@25..26 "\n" + DO@26..68 + DO_KW@26..28 "do" + WHITESPACE@28..29 " " + DO_LANGUAGE@29..45 + LANGUAGE_KW@29..37 "language" + WHITESPACE@37..38 " " + LANGUAGE_REF@38..45 + IDENT@38..45 "plpgsql" + WHITESPACE@45..46 " " + LITERAL@46..67 + DOLLAR_QUOTED_STRING@46..67 "$$ begin null; end $$" + SEMICOLON@67..68 ";" + WHITESPACE@68..69 "\n" + DO@69..112 + DO_KW@69..71 "do" + WHITESPACE@71..72 " " + LITERAL@72..86 + DOLLAR_QUOTED_STRING@72..86 "$$ anything $$" + WHITESPACE@86..87 " " + DO_LANGUAGE@87..111 + LANGUAGE_KW@87..95 "language" + WHITESPACE@95..96 " " + LANGUAGE_REF@96..111 + IDENT@96..111 "custom_language" + SEMICOLON@111..112 ";" + WHITESPACE@112..113 "\n" + DO@113..146 + DO_KW@113..115 "do" + WHITESPACE@115..116 " " + DO_LANGUAGE@116..130 + LANGUAGE_KW@116..124 "language" + WHITESPACE@124..125 " " + LANGUAGE_REF@125..130 + IDENT@125..130 "\"SQL\"" + WHITESPACE@130..131 " " + LITERAL@131..145 + DOLLAR_QUOTED_STRING@131..145 "$$ anything $$" + SEMICOLON@145..146 ";" + WHITESPACE@146..147 "\n" + DO@147..180 + DO_KW@147..149 "do" + WHITESPACE@149..150 " " + DO_LANGUAGE@150..164 + LANGUAGE_KW@150..158 "language" + WHITESPACE@158..159 " " + LITERAL@159..164 + STRING@159..164 "'SQL'" + WHITESPACE@164..165 " " + LITERAL@165..179 + DOLLAR_QUOTED_STRING@165..179 "$$ anything $$" + SEMICOLON@179..180 ";" + WHITESPACE@180..181 "\n" diff --git a/crates/squawk_syntax/src/snapshots/squawk_syntax__test__do_language_validation.snap b/crates/squawk_syntax/src/snapshots/squawk_syntax__test__do_language_validation.snap new file mode 100644 index 00000000..42bb276d --- /dev/null +++ b/crates/squawk_syntax/src/snapshots/squawk_syntax__test__do_language_validation.snap @@ -0,0 +1,74 @@ +--- +source: crates/squawk_syntax/src/test.rs +input_file: crates/squawk_syntax/test_data/validation/do_language.sql +--- +SOURCE_FILE@0..132 + DO@0..31 + DO_KW@0..2 "do" + WHITESPACE@2..3 " " + LITERAL@3..17 + DOLLAR_QUOTED_STRING@3..17 "$$ select 1 $$" + WHITESPACE@17..18 " " + DO_LANGUAGE@18..30 + LANGUAGE_KW@18..26 "language" + WHITESPACE@26..27 " " + LANGUAGE_REF@27..30 + SQL_KW@27..30 "sql" + SEMICOLON@30..31 ";" + WHITESPACE@31..32 "\n" + DO@32..63 + DO_KW@32..34 "do" + WHITESPACE@34..35 " " + DO_LANGUAGE@35..47 + LANGUAGE_KW@35..43 "language" + WHITESPACE@43..44 " " + LANGUAGE_REF@44..47 + SQL_KW@44..47 "SQL" + WHITESPACE@47..48 " " + LITERAL@48..62 + DOLLAR_QUOTED_STRING@48..62 "$$ select 1 $$" + SEMICOLON@62..63 ";" + WHITESPACE@63..64 "\n" + DO@64..97 + DO_KW@64..66 "do" + WHITESPACE@66..67 " " + DO_LANGUAGE@67..81 + LANGUAGE_KW@67..75 "language" + WHITESPACE@75..76 " " + LANGUAGE_REF@76..81 + IDENT@76..81 "\"sql\"" + WHITESPACE@81..82 " " + LITERAL@82..96 + DOLLAR_QUOTED_STRING@82..96 "$$ select 1 $$" + SEMICOLON@96..97 ";" + WHITESPACE@97..98 "\n" + DO@98..131 + DO_KW@98..100 "do" + WHITESPACE@100..101 " " + DO_LANGUAGE@101..115 + LANGUAGE_KW@101..109 "language" + WHITESPACE@109..110 " " + LITERAL@110..115 + STRING@110..115 "'sql'" + WHITESPACE@115..116 " " + LITERAL@116..130 + DOLLAR_QUOTED_STRING@116..130 "$$ select 1 $$" + SEMICOLON@130..131 ";" + WHITESPACE@131..132 "\n" + +error[syntax-error]: SQL is not a valid language for DO statements. + ╭▸ +1 │ do $$ select 1 $$ language sql; + ╰╴ ━━━ +error[syntax-error]: SQL is not a valid language for DO statements. + ╭▸ +2 │ do language SQL $$ select 1 $$; + ╰╴ ━━━ +error[syntax-error]: SQL is not a valid language for DO statements. + ╭▸ +3 │ do language "sql" $$ select 1 $$; + ╰╴ ━━━━━ +error[syntax-error]: SQL is not a valid language for DO statements. + ╭▸ +4 │ do language 'sql' $$ select 1 $$; + ╰╴ ━━━━━ diff --git a/crates/squawk_syntax/src/sql_body.rs b/crates/squawk_syntax/src/sql_body.rs new file mode 100644 index 00000000..803a81ea --- /dev/null +++ b/crates/squawk_syntax/src/sql_body.rs @@ -0,0 +1,158 @@ +use rowan::GreenNode; + +use crate::{ + ast, + ast::AstNode, + body::{Body, BodyLanguage}, + parsing, + syntax_error::SyntaxError, +}; + +pub type SqlBody = Body; + +impl BodyLanguage for ast::SourceFile { + const LANGUAGE: &'static str = "sql"; + + fn parse_text(text: &str) -> (GreenNode, Vec) { + parsing::parse_text(text) + } +} + +impl ast::SourceFile { + pub fn sql_body_errors(&self) -> Vec { + self.syntax() + .descendants() + .filter_map(|node| { + ast::CreateFunction::cast(node.clone()) + .and_then(|function| function.sql_body()) + .or_else(|| { + ast::CreateProcedure::cast(node).and_then(|procedure| procedure.sql_body()) + }) + }) + .flat_map(|body| body.errors()) + .collect() + } +} + +impl ast::CreateFunction { + pub fn sql_body(&self) -> Option { + SqlBody::from_options(self.option_list()?) + } +} + +impl ast::CreateProcedure { + pub fn sql_body(&self) -> Option { + SqlBody::from_options(self.option_list()?) + } +} + +#[cfg(test)] +mod tests { + use super::*; + use crate::SourceFile; + use crate::test::render_errors; + use insta::assert_snapshot; + use rowan::{TextRange, TextSize}; + + fn find(sql: &str) -> Option { + let parse = SourceFile::parse(sql); + + parse.tree().syntax().descendants().find_map(|node| { + ast::CreateFunction::cast(node.clone()) + .and_then(|it| it.sql_body()) + .or_else(|| ast::CreateProcedure::cast(node).and_then(|it| it.sql_body())) + }) + } + + fn body(sql: &str) -> String { + let body = find(sql).expect("no SQL body"); + + let range = body.source_range(TextRange::up_to(TextSize::of(body.text()))); + let start = usize::from(range.start()); + let end = usize::from(range.end()); + + let mut out = format!("{:#?}", body.syntax()); + out.push_str(&format!("---\nsource {range:?} {:?}\n", &sql[start..end])); + out.push_str(&render_errors(sql, &body.errors())); + out + } + + #[test] + fn function_language_before_as() { + assert_snapshot!(body( + "\ +create function f(int) returns int + language sql + as $$ select $1 + 1 $$;" + )); + } + + #[test] + fn function_language_after_as() { + assert_snapshot!(body( + "\ +create function f(int) returns int + as 'select $1' + language sql;" + )); + } + + #[test] + fn procedure() { + assert_snapshot!(body( + "\ +create procedure p() + language sql + as $$ select 1; select 2 $$;" + )); + } + + #[test] + fn other_language_is_not_sql() { + assert!( + find( + "\ +create function f() returns int + as $$ select 1 $$ + language plpgsql;" + ) + .is_none() + ); + } + + #[test] + fn string_language_is_case_sensitive() { + assert!( + find( + "\ +create function f() returns int + as 'not even sql' + language 'SQL';" + ) + .is_none() + ); + } + + #[test] + fn inline_body_is_not_a_string_body() { + assert!( + find( + "\ +create function f() returns int + language sql + return 1;" + ) + .is_none() + ); + } + + #[test] + fn errors_map_into_the_containing_file() { + assert_snapshot!(body( + "\ +create function f() returns int + language sql + as $$ select from $$;" + )); + } +} diff --git a/crates/squawk_syntax/src/test.rs b/crates/squawk_syntax/src/test.rs index 251cc850..26cbf783 100644 --- a/crates/squawk_syntax/src/test.rs +++ b/crates/squawk_syntax/src/test.rs @@ -78,9 +78,7 @@ fn regression_suite_validation(fixture: Fixture<&str>) { } let parse = SourceFile::parse(content); - let errors = parse.errors(); - - let mut errors = errors; + let mut errors = parse.errors(); if test_name == "errors" { assert!( diff --git a/crates/squawk_syntax/src/validation.rs b/crates/squawk_syntax/src/validation.rs index d6b002a9..1f9a471b 100644 --- a/crates/squawk_syntax/src/validation.rs +++ b/crates/squawk_syntax/src/validation.rs @@ -10,7 +10,9 @@ use either::Either; use crate::ast::{AstNode, LitKind, PrefixOp}; use crate::unescape::{escape_unicode_esc_str, uescape_char}; -use crate::{SyntaxNode, SyntaxToken, ast, match_ast, syntax_error::SyntaxError}; +use crate::{ + SyntaxNode, SyntaxToken, ast, match_ast, sql_body::SqlBody, syntax_error::SyntaxError, +}; use rowan::{TextRange, TextSize, WalkEvent}; use squawk_parser::{ SyntaxKind::*, is_col_name_keyword, is_reserved_keyword, is_type_func_name_keyword, @@ -1497,7 +1499,17 @@ fn validate_do(do_: ast::Do, acc: &mut Vec) { let mut seen_body = false; for part in do_.language_and_body() { let (seen, range) = match part { - Either::Left(language) => (&mut seen_language, language.syntax().text_range()), + Either::Left(language) => { + if let Some(language_name) = language.language_name() + && language_name.name == SqlBody::LANGUAGE + { + acc.push(SyntaxError::new( + "SQL is not a valid language for DO statements.", + language_name.range, + )); + } + (&mut seen_language, language.syntax().text_range()) + } Either::Right(body) => (&mut seen_body, body.syntax().text_range()), }; if *seen { diff --git a/crates/squawk_syntax/test_data/validation/do_language.sql b/crates/squawk_syntax/test_data/validation/do_language.sql new file mode 100644 index 00000000..b24aa31f --- /dev/null +++ b/crates/squawk_syntax/test_data/validation/do_language.sql @@ -0,0 +1,4 @@ +do $$ select 1 $$ language sql; +do language SQL $$ select 1 $$; +do language "sql" $$ select 1 $$; +do language 'sql' $$ select 1 $$; diff --git a/crates/squawk_syntax/test_data/validation/do_language_ok.sql b/crates/squawk_syntax/test_data/validation/do_language_ok.sql new file mode 100644 index 00000000..9c112e65 --- /dev/null +++ b/crates/squawk_syntax/test_data/validation/do_language_ok.sql @@ -0,0 +1,5 @@ +do $$ begin null; end $$; +do language plpgsql $$ begin null; end $$; +do $$ anything $$ language custom_language; +do language "SQL" $$ anything $$; +do language 'SQL' $$ anything $$; diff --git a/crates/xtask/src/sync_pg.rs b/crates/xtask/src/sync_pg.rs index 71298402..597dcbe3 100644 --- a/crates/xtask/src/sync_pg.rs +++ b/crates/xtask/src/sync_pg.rs @@ -12,25 +12,50 @@ const KWLIST_PATH: &str = "postgres/kwlist.h"; const PL_RESERVED_KWLIST_PATH: &str = "postgres/pl_reserved_kwlist.h"; const PL_UNRESERVED_KWLIST_PATH: &str = "postgres/pl_unreserved_kwlist.h"; -const START_END_MARKERS: &[(&str, &str)] = &[ +const START_END_MARKERS: &[(&str, &str, &str)] = &[ ( + "merge.sql", "MERGE INTO target t RANDOMWORD", "\tUPDATE SET balance = 0;", ), ( + "merge.sql", "-- incorrectly specifying INTO target", "\tINSERT INTO target DEFAULT VALUES;", ), - ("-- Multiple VALUES clause", "\tINSERT VALUES (1,1), (2,2);"), - ("-- SELECT query for INSERT", "\tINSERT SELECT (1, 1);"), - ("-- UPDATE tablename", "\tUPDATE target SET balance = 0;"), ( + "merge.sql", + "-- Multiple VALUES clause", + "\tINSERT VALUES (1,1), (2,2);", + ), + ( + "merge.sql", + "-- SELECT query for INSERT", + "\tINSERT SELECT (1, 1);", + ), + ( + "merge.sql", + "-- UPDATE tablename", + "\tUPDATE target SET balance = 0;", + ), + ( + "for_portion_of.sql", "-- TO is used for the bound but not the INTERVAL:", " WHERE id = '[1,2)';", ), - ("-- => is disallowed as an operator name now", ");"), + ( + "create_operator.sql", + "-- => is disallowed as an operator name now", + ");", + ), ]; +const AFTER_START_END_MARKERS: &[(&str, &str, &str)] = &[( + "create_function_sql.sql", + "-- Things that shouldn't work:", + " AS 'not even SQL';", +)]; + const IGNORED_LINES: &[&str] = &[ r#"SELECT JSON_TABLE('[]', '$');"#, r#"SELECT rank() OVER (PARTITION BY four, ORDER BY ten) FROM tenk1;"#, @@ -321,7 +346,7 @@ fn preprocess_files(files: &[Utf8PathBuf], output_dir: &Utf8Path) -> Result<()> let reader = std::io::BufReader::new(input_file); let mut processed_content = vec![]; - if let Err(e) = preprocess_sql(reader, &mut processed_content) { + if let Err(e) = preprocess_sql(reader, &mut processed_content, filename) { eprintln!("Error: Failed to process file: {e}"); continue; } @@ -359,7 +384,11 @@ fn sync_plpgsql_suite(clone_dir: &Utf8Path) -> Result<()> { // The regression suite from postgres has a mix of valid and invalid sql. We // don't have a good way to determine what is what, so we munge the data to // comment out any problematic code. -pub(crate) fn preprocess_sql(source: R, mut dest: W) -> Result<()> { +pub(crate) fn preprocess_sql( + source: R, + mut dest: W, + filename: &str, +) -> Result<()> { let template_vars_regex = Regex::new(r"^:'([^']+)'|^:([a-zA-Z_][a-zA-Z0-9_]*)").unwrap(); let mut in_copy_stdin = false; let mut in_bogus_cases = false; @@ -382,8 +411,8 @@ pub(crate) fn preprocess_sql(source: R, mut dest: W) -> Re in_copy_select_input = false; } - for &(start, end) in START_END_MARKERS { - if line.contains(start) { + for &(marker_filename, start, end) in START_END_MARKERS { + if filename == marker_filename && line.contains(start) { looking_for_end = Some(end); } } @@ -395,6 +424,12 @@ pub(crate) fn preprocess_sql(source: R, mut dest: W) -> Re } } + for &(marker_filename, start, end) in AFTER_START_END_MARKERS { + if filename == marker_filename && line.contains(start) { + looking_for_end = Some(end); + } + } + let line_lower = line.to_ascii_lowercase(); if (line_lower.starts_with("copy ") || line_lower.starts_with("\\copy")) && (line_lower.contains("from stdin") || line_lower.contains("from stdout")) @@ -553,7 +588,7 @@ mod tests { let input = sql.as_bytes(); let mut output = Vec::new(); let cursor = Cursor::new(input); - preprocess_sql(cursor, &mut output)?; + preprocess_sql(cursor, &mut output, "")?; String::from_utf8(output).map_err(Into::into) } diff --git a/postgres/regression_suite/create_function_sql.sql b/postgres/regression_suite/create_function_sql.sql index 8d9d85b6..cbeea167 100644 --- a/postgres/regression_suite/create_function_sql.sql +++ b/postgres/regression_suite/create_function_sql.sql @@ -461,12 +461,12 @@ INSERT INTO pt VALUES (1); INSERT INTO pt VALUES (1); -- Things that shouldn't work: - -CREATE FUNCTION test1 (int) RETURNS int LANGUAGE SQL - AS 'SELECT ''not an integer'';'; - -CREATE FUNCTION test1 (int) RETURNS int LANGUAGE SQL - AS 'not even SQL'; +-- +-- CREATE FUNCTION test1 (int) RETURNS int LANGUAGE SQL +-- AS 'SELECT ''not an integer'';'; +-- +-- CREATE FUNCTION test1 (int) RETURNS int LANGUAGE SQL +-- AS 'not even SQL'; CREATE FUNCTION test1 (int) RETURNS int LANGUAGE SQL AS 'SELECT 1, 2, 3;';