Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
3 changes: 1 addition & 2 deletions crates/squawk_fmt/src/fmt.rs
Original file line number Diff line number Diff line change
Expand Up @@ -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))
Expand Down
24 changes: 12 additions & 12 deletions crates/squawk_ide/src/inlay_hints.rs
Original file line number Diff line number Diff line change
Expand Up @@ -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:
Expand All @@ -303,15 +303,15 @@ 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
");
}

#[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:
Expand All @@ -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
Expand All @@ -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:
Expand All @@ -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
");
}
Expand All @@ -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:
Expand All @@ -373,15 +373,15 @@ 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
");
}

#[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:
Expand All @@ -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
Expand All @@ -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:
Expand All @@ -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
");
}
Expand Down
9 changes: 3 additions & 6 deletions crates/squawk_parser/src/grammar.rs
Original file line number Diff line number Diff line change
Expand Up @@ -3470,11 +3470,10 @@ fn compound_select_operand(p: &mut Parser<'_>) -> Option<CompletedMarker> {

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;
Expand Down Expand Up @@ -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;
}
Expand Down Expand Up @@ -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)
Expand All @@ -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) {
Expand Down
2 changes: 1 addition & 1 deletion crates/squawk_parser/tests/data/ok/create_function.sql
Original file line number Diff line number Diff line change
Expand Up @@ -307,7 +307,7 @@ language sql;

create function f(bitmask bit(8))
returns boolean
as '0'
as 'select 0'
language sql;

-- argmode
Expand Down
6 changes: 3 additions & 3 deletions crates/squawk_parser/tests/data/ok/create_procedure.sql
Original file line number Diff line number Diff line change
Expand Up @@ -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';
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -2272,7 +2272,7 @@ SOURCE_FILE
WHITESPACE " "
AS_DEFINITION
LITERAL
STRING "'0'"
STRING "'select 0'"
WHITESPACE "\n"
LANGUAGE_FUNC_OPTION
LANGUAGE_KW "language"
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -228,7 +228,7 @@ SOURCE_FILE
WHITESPACE " "
AS_DEFINITION
LITERAL
STRING "'foo'"
STRING "'select 1'"
SEMICOLON ";"
WHITESPACE "\n"
CREATE_PROCEDURE
Expand Down Expand Up @@ -263,7 +263,7 @@ SOURCE_FILE
WHITESPACE " "
AS_DEFINITION
LITERAL
STRING "'foo'"
STRING "'select 1'"
SEMICOLON ";"
WHITESPACE "\n"
CREATE_PROCEDURE
Expand Down Expand Up @@ -296,7 +296,7 @@ SOURCE_FILE
WHITESPACE " "
AS_DEFINITION
LITERAL
STRING "'foo'"
STRING "'select 1'"
SEMICOLON ";"
WHITESPACE "\n\n"
COMMENT "-- as_with_two_strings"
Expand Down
2 changes: 1 addition & 1 deletion crates/squawk_syntax/src/ast.rs
Original file line number Diff line number Diff line change
Expand Up @@ -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},
Expand Down
37 changes: 36 additions & 1 deletion crates/squawk_syntax/src/ast/node_ext.rs
Original file line number Diff line number Diff line change
Expand Up @@ -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)]
Expand Down Expand Up @@ -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<ast::LanguageRef>,
literal: Option<ast::Literal>,
) -> Option<LanguageName> {
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<LanguageName> {
language_name(self.language_ref(), self.literal())
}
}

impl ast::DoLanguage {
pub fn language_name(&self) -> Option<LanguageName> {
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
Expand Down
99 changes: 99 additions & 0 deletions crates/squawk_syntax/src/body.rs
Original file line number Diff line number Diff line change
@@ -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<T> {
green: GreenNode,
errors: Vec<SyntaxError>,
decoded: DecodedText,
_ty: PhantomData<fn() -> T>,
}

pub trait BodyLanguage: AstNode {
const LANGUAGE: &'static str;

fn parse_text(text: &str) -> (GreenNode, Vec<SyntaxError>);

fn is_language(language: Option<ast::LanguageName>) -> bool {
language.is_some_and(|language| language.name == Self::LANGUAGE)
}
}

impl<T: BodyLanguage> Body<T> {
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<Self> {
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<SyntaxError> {
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
}
}
Loading
Loading