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
53 changes: 33 additions & 20 deletions crates/squawk_fmt/src/comment.rs
Original file line number Diff line number Diff line change
Expand Up @@ -82,7 +82,7 @@ impl CommentRun {
Doc::list(docs)
}

fn separator_before<'a>(&self, separator: Doc<'a>) -> Doc<'a> {
pub(crate) fn separator_before<'a>(&self, separator: Doc<'a>) -> Doc<'a> {
if self
.tokens
.first()
Expand Down Expand Up @@ -148,6 +148,32 @@ impl CommentRun {
.append(self.separator_after(Doc::space()))
}

pub(crate) fn split_trailing(self) -> (CommentRun, CommentRun) {
let count = self
.tokens
.iter()
.take_while(|token| is_trailing_comment(token))
.position(is_line_comment)
.map_or(0, |index| index + 1);
let mut trailing = self.tokens;
let rest = trailing.split_off(count);
(Self { tokens: trailing }, Self { tokens: rest })
}

pub(crate) fn trailing<'a>(&self) -> Doc<'a> {
if self.is_empty() {
return Doc::nil();
}
Doc::space().append(self.doc())
}

pub(crate) fn leading<'a>(&self) -> Doc<'a> {
if self.is_empty() {
return Doc::nil();
}
self.doc().append(self.separator_after(Doc::space()))
}

pub(crate) fn before_closing_delimiter<'a>(&self, separator: Doc<'a>) -> (Doc<'a>, Doc<'a>) {
(self.doc(), self.separator_after(separator))
}
Expand All @@ -166,13 +192,7 @@ pub(crate) fn comment_run_after(el: &(impl Into<SyntaxElement> + Clone)) -> Comm
}

pub(crate) fn leading_comments<'a>(el: &(impl Into<SyntaxElement> + Clone)) -> Doc<'a> {
let comments = comment_run_before(el);
if comments.is_empty() {
return Doc::nil();
}
comments
.doc()
.append(comments.separator_after(Doc::space()))
comment_run_before(el).leading()
}

pub(crate) fn comments_before<'a>(el: &(impl Into<SyntaxElement> + Clone)) -> Doc<'a> {
Expand Down Expand Up @@ -203,18 +223,11 @@ pub(crate) fn hard_line_before<'a>(el: &(impl Into<SyntaxElement> + Clone)) -> D
comment_run_before(el).before_node(Doc::hard_line())
}

pub(crate) fn separator_before<'a>(
sep: Doc<'a>,
el: &(impl Into<SyntaxElement> + Clone),
) -> Doc<'a> {
comment_run_before(el).separator_before(sep)
}

pub(crate) fn separator_after<'a>(
sep: Doc<'a>,
el: &(impl Into<SyntaxElement> + Clone),
) -> Doc<'a> {
comment_run_after(el).separator_before(sep)
pub(crate) fn has_own_line_comments_before(el: &(impl Into<SyntaxElement> + Clone)) -> bool {
comment_run_before(el)
.tokens
.first()
.is_some_and(|token| !is_trailing_comment(token))
}

pub(crate) fn has_comments_before(el: &(impl Into<SyntaxElement> + Clone)) -> bool {
Expand Down
88 changes: 40 additions & 48 deletions crates/squawk_fmt/src/fmt.rs
Original file line number Diff line number Diff line change
Expand Up @@ -12,8 +12,8 @@ use squawk_syntax::{SyntaxKind, SyntaxNode, SyntaxToken};

use crate::comment::{
CommentRun, build_comment, comment_run_after, comment_run_before, comments_before,
hard_line_before, has_comments_before, is_line_comment, leading_comments, line_before,
separator_after, separator_before, space_before, space_or_comments_before, trailing_comments,
hard_line_before, has_comments_before, has_own_line_comments_before, is_line_comment,
leading_comments, line_before, space_before, space_or_comments_before, trailing_comments,
};
use tiny_pretty::Doc;
use tiny_pretty::{LineBreak, PrintOptions, print};
Expand Down Expand Up @@ -67,7 +67,7 @@ fn build_source_file<'a>(ctx: &Ctx, source_file: &'a ast::SourceFile) -> Doc<'a>
} else {
Doc::empty_line()
};
doc = doc.append(separator_after(gap, &token));
doc = doc.append(comment_run_after(&token).separator_before(gap));
}
}
}
Expand Down Expand Up @@ -342,7 +342,7 @@ fn build_insert<'a>(ctx: &Ctx, insert: &ast::Insert) -> Doc<'a> {
.append(before_source)
.append(Doc::text("values"))
.group()
.append(build_values_rows(ctx, &values, true));
.append(build_values_rows(ctx, &values));
for clause in values.tail_clauses() {
doc = doc
.append(Doc::line_or_space())
Expand Down Expand Up @@ -8853,7 +8853,7 @@ fn build_values<'a>(ctx: &Ctx, values: &ast::Values) -> Doc<'a> {

doc = doc.append(
Doc::text("values")
.append(build_values_rows(ctx, values, false))
.append(build_values_rows(ctx, values))
.group(),
);

Expand All @@ -8867,7 +8867,7 @@ fn build_values<'a>(ctx: &Ctx, values: &ast::Values) -> Doc<'a> {
.group()
}

fn build_values_rows<'a>(ctx: &Ctx, values: &ast::Values, nest_rows: bool) -> Doc<'a> {
fn build_values_rows<'a>(ctx: &Ctx, values: &ast::Values) -> Doc<'a> {
let mut doc = Doc::nil();
if let Some(row_list) = values.row_list() {
let rows = row_list.rows().map(|row| {
Expand All @@ -8878,13 +8878,13 @@ fn build_values_rows<'a>(ctx: &Ctx, values: &ast::Values, nest_rows: bool) -> Do
});
if let Some(rows) = build_comma_separated_docs(rows) {
let multiple_rows = row_list.rows().count() > 1;
let rows = if nest_rows && multiple_rows {
let rows = if multiple_rows {
line_before(row_list.syntax())
} else {
space_before(row_list.syntax())
}
.append(rows);
doc = doc.append(if nest_rows && multiple_rows {
doc = doc.append(if multiple_rows {
rows.nest(ctx.indent).group()
} else {
rows
Expand Down Expand Up @@ -9078,8 +9078,12 @@ fn build_with_clause<'a>(ctx: &Ctx, with_clause: ast::WithClause) -> Doc<'a> {
table.syntax().clone(),
)
});
let separator = match with_clause.with_tables().next() {
Some(first) if has_own_line_comments_before(first.syntax()) => Doc::hard_line(),
_ => Doc::space(),
};
if let Some(tables) = build_comma_separated_docs(tables) {
doc = doc.append(Doc::space()).append(tables);
doc = doc.append(separator).append(tables);
}
doc
}
Expand Down Expand Up @@ -9116,10 +9120,16 @@ fn build_with_table<'a>(ctx: &Ctx, table: ast::WithTable) -> Doc<'a> {
if let Some(l_paren) = table.l_paren_token() {
doc = doc.append(space_before(&l_paren)).append(Doc::text("("));
}
let body = table
.query()
.map(|query| leading_comments(query.syntax()).append(build_with_query(ctx, query)))
.unwrap_or_else(Doc::nil);
let body = match table.query() {
Some(query) => {
let (trailing, leading) = comment_run_before(query.syntax()).split_trailing();
doc = doc.append(trailing.trailing());
leading
.before_node(Doc::nil())
.append(build_with_query(ctx, query))
}
None => Doc::nil(),
};
doc = doc
.append(wrap_hard_body(ctx, body, table.r_paren_token()))
.append(Doc::text(")"));
Expand Down Expand Up @@ -12527,7 +12537,7 @@ fn build_op_sig<'a>(ctx: &Ctx, sig: ast::OpSig) -> Doc<'a> {
.map(|op| leading_comments(op.syntax()).append(build_ddl_operator(&op)))
.unwrap_or_else(Doc::nil);
if let Some(l_paren) = sig.l_paren_token() {
doc = doc.append(comments_before(&l_paren));
doc = doc.append(space_before(&l_paren));
}
let has_none = sig.none_token().is_some();
let mut body = if let Some(none) = sig.none_token() {
Expand Down Expand Up @@ -15261,7 +15271,9 @@ fn ends_with_inline_bracketed_expr(expr: &ast::Expr) -> bool {

fn is_single_inline_target_list(target_list: &ast::TargetList) -> bool {
match target_list.targets().exactly_one() {
Ok(target) => target.expr().as_ref().is_some_and(is_inline_bracketed_expr),
Ok(target) => target.expr().as_ref().is_some_and(|expr| {
is_inline_bracketed_expr(expr) || matches!(expr, ast::Expr::CaseExpr(_))
}),
Err(_) => false,
}
}
Expand All @@ -15280,14 +15292,14 @@ fn has_single_inline_target(select: &ast::Select) -> bool {
}

fn build_select_doc_ungrouped<'a>(ctx: &Ctx, select: &ast::Select) -> Doc<'a> {
let mut doc = Doc::nil();
let mut prefix = Doc::nil();
if let Some(with_clause) = select.with_clause() {
doc = doc
prefix = prefix
.append(leading_comments(with_clause.syntax()))
.append(build_with_clause(ctx, with_clause));
doc = match select.select_clause() {
Some(select_clause) => doc.append(hard_line_before(select_clause.syntax())),
None => doc.append(Doc::hard_line()),
prefix = match select.select_clause() {
Some(select_clause) => prefix.append(hard_line_before(select_clause.syntax())),
None => prefix.append(Doc::hard_line()),
};
}
let mut select_doc = Doc::text("select");
Expand All @@ -15311,11 +15323,7 @@ fn build_select_doc_ungrouped<'a>(ctx: &Ctx, select: &ast::Select) -> Doc<'a> {
select_doc.append(Doc::line_or_space().append(select_body).nest(ctx.indent))
}
};
doc = if select.with_clause().is_some() {
doc.append(select_doc.group())
} else {
doc.append(select_doc)
};
let mut doc = select_doc;
if select.from_clause().is_some() {
doc = doc.group();
}
Expand Down Expand Up @@ -15357,7 +15365,11 @@ fn build_select_doc_ungrouped<'a>(ctx: &Ctx, select: &ast::Select) -> Doc<'a> {

doc = doc.append(build_semicolon(select.semicolon_token()));

doc
if select.with_clause().is_some() {
prefix.append(doc.group())
} else {
doc
}
}

fn build_from_clause<'a>(ctx: &Ctx, from: ast::FromClause) -> Doc<'a> {
Expand Down Expand Up @@ -18386,7 +18398,7 @@ fn join_comma_separated<'a>(
docs.push(
trailing_comments(&previous_syntax)
.append(Doc::text(","))
.append(separator_before(Doc::line_or_space(), &syntax))
.append(comment_run_before(&syntax).separator_before(Doc::line_or_space()))
.append(item),
);
previous_syntax = syntax;
Expand Down Expand Up @@ -19244,30 +19256,10 @@ fn build_postfix_expr<'a>(ctx: &Ctx, postfix_expr: ast::PostfixExpr) -> Doc<'a>
let expr = build_expr(ctx, postfix_expr.expr().unwrap());
let op = postfix_expr.op().unwrap();
expr.append(Doc::space())
.append(leading_comments_postfix_op(&op))
.append(leading_comments(&op.syntax_element()))
.append(build_postfix_op(op))
}

fn leading_comments_postfix_op<'a>(op: &ast::PostfixOp) -> Doc<'a> {
match op {
ast::PostfixOp::AtLocal(node) => leading_comments(node.syntax()),
ast::PostfixOp::IsJson(node) => leading_comments(node.syntax()),
ast::PostfixOp::IsJsonArray(node) => leading_comments(node.syntax()),
ast::PostfixOp::IsJsonObject(node) => leading_comments(node.syntax()),
ast::PostfixOp::IsJsonScalar(node) => leading_comments(node.syntax()),
ast::PostfixOp::IsJsonValue(node) => leading_comments(node.syntax()),
ast::PostfixOp::IsNormalized(node) => leading_comments(node.syntax()),
ast::PostfixOp::IsNotJson(node) => leading_comments(node.syntax()),
ast::PostfixOp::IsNotJsonArray(node) => leading_comments(node.syntax()),
ast::PostfixOp::IsNotJsonObject(node) => leading_comments(node.syntax()),
ast::PostfixOp::IsNotJsonScalar(node) => leading_comments(node.syntax()),
ast::PostfixOp::IsNotJsonValue(node) => leading_comments(node.syntax()),
ast::PostfixOp::IsNotNormalized(node) => leading_comments(node.syntax()),
ast::PostfixOp::IsNull(token) => leading_comments(token),
ast::PostfixOp::NotNull(token) => leading_comments(token),
}
}

fn build_postfix_op<'a>(op: ast::PostfixOp) -> Doc<'a> {
match op {
ast::PostfixOp::AtLocal(n) => {
Expand Down
11 changes: 8 additions & 3 deletions crates/squawk_fmt/tests/after/alter_operator.snap
Original file line number Diff line number Diff line change
Expand Up @@ -2,18 +2,23 @@
source: crates/squawk_fmt/tests/tests.rs
input_file: crates/squawk_fmt/tests/before/alter_operator.sql
---
alter /* operator */ operator /* signature */ public.+ /* left paren */(
alter /* operator */ operator /* signature */ public.+ /* left paren */ (
/* left type */ integer /* comma */,
/* right type */ integer /* right paren */
)
/* owner */ owner /* to */ to /* role */ application_owner /* end */;

alter operator public.##(none, integer) set (
alter operator public.## (none, integer) set (
restrict = schema_a.restrict_function_with_a_very_long_name,
join = schema_a.join_function_with_a_very_long_name,
hashes,
merges
);

alter operator public.+(integer, integer)
alter operator public.+ (integer, integer)
/* set */ set /* schema */ schema /* name */ archive;

alter operator <= (vector, vector) set (
restrict = scalarlesel,
join = scalarlejoinsel
);
6 changes: 2 additions & 4 deletions crates/squawk_fmt/tests/after/delete.snap
Original file line number Diff line number Diff line change
Expand Up @@ -41,8 +41,7 @@ where foo.id = doomed.id;
with deleted as (
delete from foo where id = 1 returning id
)
select *
from deleted;
select * from deleted;

delete from t3
using t1 join t2 using (a)
Expand Down Expand Up @@ -90,5 +89,4 @@ with deleted_job as (
and river_job.state != 'running'
returning river_job.*
)
select *
from deleted_job;
select * from deleted_job;
12 changes: 7 additions & 5 deletions crates/squawk_fmt/tests/after/drop_operator.snap
Original file line number Diff line number Diff line change
Expand Up @@ -2,24 +2,26 @@
source: crates/squawk_fmt/tests/tests.rs
input_file: crates/squawk_fmt/tests/before/drop_operator.sql
---
drop operator public.===(integer, integer);
drop operator public.=== (integer, integer);

drop operator if exists
extraordinarily_long_schema_name.===(
extraordinarily_long_schema_name.=== (
extraordinarily_long_schema_name.extraordinarily_long_left_operand_type,
extraordinarily_long_schema_name.extraordinarily_long_right_operand_type
),
public.<>(none, double precision)
public.<> (none, double precision)
cascade;

-- comments in every position
drop /* operator */ operator /* if */ if /* exists */ exists
/* first operator */ public /* path dot */.=== /* open */(
/* first operator */ public /* path dot */.=== /* open */ (
/* left type */ integer /* type comma */,
/* right type */ integer /* close */
) /* signature comma */,
/* second operator */ ~ /* second open */(
/* second operator */ ~ /* second open */ (
/* none */ none /* second comma */,
/* second type */ text /* second close */
)
/* behavior */ restrict /* end */;

drop operator - (vector, vector);
9 changes: 3 additions & 6 deletions crates/squawk_fmt/tests/after/insert.snap
Original file line number Diff line number Diff line change
Expand Up @@ -7,8 +7,7 @@ insert into foo (id, name) values (1, 'one') returning id;
with inserted as (
insert into foo default values returning id
)
select *
from inserted;
select * from inserted;

with inserted as (
insert into foo as f (id, name)
Expand All @@ -17,8 +16,7 @@ with inserted as (
on conflict on constraint foo_pkey do nothing
returning id
)
select *
from inserted;
select * from inserted;

insert into products (product_no, name, price) values
(1, 'Cheese', 9.99),
Expand Down Expand Up @@ -69,8 +67,7 @@ with /*a*/ inserted/*b*/ (
/*bc*/ new /*bd*/ as /*be*/ new_row /*bf*/
) /*bg*/ new_row.id /*bh*/
)
/*bi*/ select /*bj*/ result_id
/*bk*/ from /*bl*/ inserted /*bm*/;
/*bi*/ select /*bj*/ result_id /*bk*/ from /*bl*/ inserted /*bm*/;

-- top
insert into products (product_no, name, price) values -- after insert
Expand Down
Loading
Loading