1use anyhow::{Result, bail};
2use itertools::Itertools;
3use rowan::Direction;
4use squawk_line_index::{LineEnding, UniversalNewlines, find_newline};
5use squawk_syntax::ast::{self, AstNode, LitKind};
6use squawk_syntax::quote::quote_column_alias;
7use squawk_syntax::{SyntaxKind, SyntaxNode, SyntaxToken};
8use tiny_pretty::Doc;
9use tiny_pretty::{LineBreak, PrintOptions, print};
10
11fn build_source_file(source_file: &ast::SourceFile) -> Doc<'_> {
15 let mut doc = Doc::nil();
16 for el in source_file.syntax().children_with_tokens() {
17 match el {
18 rowan::NodeOrToken::Node(node) => {
19 if let Some(stmt) = ast::Stmt::cast(node) {
20 match stmt {
21 ast::Stmt::Select(select) => {
22 doc = doc.append(build_select_doc(&select));
23 }
24 ast::Stmt::CreateTable(create_table) => {
25 doc = doc.append(build_create_table(&create_table));
26 }
27 _ => (),
28 }
29 }
30 }
31 rowan::NodeOrToken::Token(token) => {
32 if token.kind() == SyntaxKind::COMMENT {
33 doc = doc.append(Doc::text(token.text().to_string()));
34 } else if token.kind() == SyntaxKind::WHITESPACE {
35 let lines = token.text().universal_newlines().count();
37 if lines >= 2 {
38 doc = doc.append(Doc::empty_line()).append(Doc::empty_line());
39 } else {
40 doc = doc.append(Doc::empty_line());
41 }
42 }
43 }
44 }
45 }
46 doc
47}
48
49fn build_create_table<'a>(create_table: &ast::CreateTable) -> Doc<'a> {
50 let mut doc = Doc::text("create")
51 .append(Doc::space())
52 .append(Doc::text("table"))
53 .append(Doc::space())
54 .append(Doc::text(
55 create_table.table_name().unwrap().syntax().to_string(),
56 ))
57 .append(Doc::text("("))
58 .append(
59 Doc::line_or_nil()
60 .append(Doc::list(
61 Itertools::intersperse(
62 create_table
63 .table_arg_list()
64 .unwrap()
65 .args()
66 .map(build_table_arg),
67 Doc::text(",").append(Doc::hard_line()),
68 )
69 .collect(),
70 ))
71 .nest(2)
72 .append(Doc::line_or_nil())
73 .group(),
74 )
75 .append(Doc::text(")"));
76
77 doc = doc.append(build_semicolon(create_table.semicolon_token()));
78
79 doc
80}
81
82fn build_table_arg<'a>(create_table: ast::TableArg) -> Doc<'a> {
83 match create_table {
84 ast::TableArg::Column(column) => build_column_name(column.name().unwrap())
85 .append(Doc::space())
86 .append(Doc::text(column.ty().unwrap().syntax().to_string())),
87 ast::TableArg::LikeClause(_like_clause) => todo!(),
88 ast::TableArg::TableConstraint(_table_constraint) => todo!(),
89 }
90}
91
92fn build_select_doc<'a>(select: &ast::Select) -> Doc<'a> {
93 let mut doc = Doc::text("select").append(Doc::line_or_space());
94
95 if let Some(select_clause) = select.select_clause() {
96 match select_clause.select_quantifier() {
97 Some(ast::SelectQuantifier::DistinctClause(distinct_clause)) => {
98 doc = doc.append(leading_comments(distinct_clause.syntax()));
99 doc = doc.append(Doc::text("distinct")).append(Doc::space());
100 }
101 Some(ast::SelectQuantifier::All(all)) => {
102 doc = doc.append(leading_comments(all.syntax()));
103 doc = doc.append(Doc::text("all")).append(Doc::space());
104 }
105 None => (),
106 }
107 if let Some(target_list) = select_clause.target_list() {
108 doc = doc.append(leading_comments(target_list.syntax()));
109 doc = doc
110 .append(Doc::list(
111 Itertools::intersperse(
112 target_list.targets().flat_map(build_target),
113 Doc::text(",").append(Doc::line_or_space()),
114 )
115 .collect(),
116 ))
117 .nest(2);
118 }
119 }
120
121 if let Some(from) = &select.from_clause() {
122 doc = doc.append(
123 Doc::line_or_space()
124 .append(Doc::text("from"))
125 .append(Doc::space())
126 .append(Doc::text(
127 from.from_items().next().unwrap().syntax().to_string(),
128 )),
129 );
130 }
131
132 if let Some(group) = &select.group_by_clause() {
133 let mut group_doc = Doc::line_or_space().append(leading_comments(group.syntax()));
134 group_doc = group_doc.append(Doc::text("group")).append(Doc::space());
135 if let Some(by_token) = group.by_token() {
136 group_doc = group_doc.append(leading_comments_token(&by_token));
137 }
138 group_doc = group_doc.append(Doc::text("by")).append(Doc::space());
139 if let Some(list) = group.group_by_list() {
140 group_doc = group_doc.append(build_group_by_list(list));
141 }
142 doc = doc.append(group_doc);
143 }
144
145 doc = doc.append(build_semicolon(select.semicolon_token()));
146
147 doc.group()
148}
149
150fn build_group_by_list<'a>(list: ast::GroupByList) -> Doc<'a> {
151 leading_comments(list.syntax()).append(Doc::text(list.syntax().to_string()))
152}
153
154fn build_semicolon<'a>(semi: Option<SyntaxToken>) -> Doc<'a> {
155 let Some(semi) = semi else {
156 return Doc::nil();
157 };
158 let mut doc = Doc::nil();
159 let mut comments: Vec<SyntaxToken> = vec![];
160 for next in semi.siblings_with_tokens(Direction::Prev).skip(1) {
161 match next {
162 rowan::NodeOrToken::Node(_) => break,
163 rowan::NodeOrToken::Token(token) => {
164 if token.kind() == SyntaxKind::COMMENT {
165 comments.push(token);
166 } else if token.kind() == SyntaxKind::WHITESPACE {
167 continue;
168 } else {
169 break;
170 }
171 }
172 }
173 }
174 for comment in comments.iter().rev() {
175 doc = doc.append(Doc::text(comment.text().to_string()));
176 }
177 doc.append(Doc::text(";"))
178}
179
180fn build_expr<'a>(expr: ast::Expr) -> Doc<'a> {
181 match expr {
182 ast::Expr::ArrayExpr(array_expr) => {
183 let mut doc = Doc::nil();
184
185 if array_expr.array_token().is_some() {
187 doc = doc.append(Doc::text("array"));
188 };
189
190 if let Some(select) = array_expr.select() {
191 doc = doc
192 .append(Doc::text("("))
193 .append(build_select_doc(&select))
194 .append(Doc::text(")"))
195 } else {
196 doc = doc
197 .append(Doc::text("["))
198 .append(Doc::list(
199 Itertools::intersperse(
200 array_expr.exprs().map(build_expr),
201 Doc::text(",").append(Doc::space()),
202 )
203 .collect(),
204 ))
205 .append(Doc::text("]"));
206 }
207
208 doc
209 }
210 ast::Expr::BetweenExpr(between_expr) => {
211 let mut doc = build_expr(between_expr.target().unwrap());
212 if between_expr.not_token().is_some() {
213 doc = doc.append(Doc::space()).append(Doc::text("not"));
214 }
215 doc = doc.append(Doc::space()).append(Doc::text("between"));
216 match between_expr.between_symmetry() {
217 Some(ast::BetweenSymmetry::Asymmetric(_)) => {
218 doc = doc.append(Doc::space()).append(Doc::text("asymmetric"));
219 }
220 Some(ast::BetweenSymmetry::Symmetric(_)) => {
221 doc = doc.append(Doc::space()).append(Doc::text("symmetric"));
222 }
223 None => (),
224 }
225 doc.append(Doc::space())
226 .append(build_expr(between_expr.start().unwrap()))
227 .append(Doc::space())
228 .append(Doc::text("and"))
229 .append(Doc::space())
230 .append(build_expr(between_expr.end().unwrap()))
231 }
232 ast::Expr::BinExpr(bin_expr) => build_expr(bin_expr.lhs().unwrap())
233 .append(Doc::space())
234 .append(build_op(bin_expr.op().unwrap()))
235 .append(Doc::space())
236 .append(build_expr(bin_expr.rhs().unwrap())),
237 ast::Expr::CastExpr(cast_expr) => {
240 let mut doc = Doc::nil();
241 if cast_expr.colon_colon().is_some() {
242 doc = doc
243 .append(build_expr(cast_expr.expr().unwrap()))
244 .append(Doc::text("::"))
245 .append(build_type(cast_expr.ty().unwrap()))
246 } else if cast_expr.as_token().is_some() {
247 if cast_expr.cast_token().is_some() {
248 doc = doc.append(Doc::text("cast"))
249 } else if cast_expr.treat_token().is_some() {
250 doc = doc.append(Doc::text("treat"))
251 }
252 doc = doc
253 .append(Doc::text("("))
254 .append(build_expr(cast_expr.expr().unwrap()))
255 .append(Doc::space())
256 .append(Doc::text("as"))
257 .append(Doc::space())
258 .append(build_type(cast_expr.ty().unwrap()))
259 .append(Doc::text(")"))
260 } else {
261 doc = doc
262 .append(build_type(cast_expr.ty().unwrap()))
263 .append(Doc::space())
264 .append(build_literal(cast_expr.literal().unwrap()))
265 }
266 doc
267 }
268 ast::Expr::Collate(collate) => build_expr(collate.expr().unwrap())
269 .append(Doc::space())
270 .append(Doc::text("collate"))
271 .append(Doc::space())
272 .append(Doc::text(
273 collate.collation_ref().unwrap().syntax().to_string(),
274 )),
275 ast::Expr::Literal(literal) => build_literal(literal),
278 ast::Expr::PostfixExpr(postfix_expr) => {
281 let expr = build_expr(postfix_expr.expr().unwrap());
282 let op = match postfix_expr.op().unwrap() {
283 ast::PostfixOp::AtLocal(_) => Doc::text("at local"),
284 ast::PostfixOp::IsNull(_) => Doc::text("isnull"),
285 ast::PostfixOp::NotNull(_) => Doc::text("notnull"),
286 ast::PostfixOp::IsJson(n) => {
287 let mut doc = Doc::text("is json");
288 if let Some(clause) = n.json_keys_unique_clause() {
289 doc = doc
290 .append(Doc::space())
291 .append(build_json_keys_unique_clause(clause));
292 }
293 doc
294 }
295 ast::PostfixOp::IsJsonArray(n) => {
296 let mut doc = Doc::text("is json array");
297 if let Some(clause) = n.json_keys_unique_clause() {
298 doc = doc
299 .append(Doc::space())
300 .append(build_json_keys_unique_clause(clause));
301 }
302 doc
303 }
304 ast::PostfixOp::IsJsonObject(n) => {
305 let mut doc = Doc::text("is json object");
306 if let Some(clause) = n.json_keys_unique_clause() {
307 doc = doc
308 .append(Doc::space())
309 .append(build_json_keys_unique_clause(clause));
310 }
311 doc
312 }
313 ast::PostfixOp::IsJsonScalar(n) => {
314 let mut doc = Doc::text("is json scalar");
315 if let Some(clause) = n.json_keys_unique_clause() {
316 doc = doc
317 .append(Doc::space())
318 .append(build_json_keys_unique_clause(clause));
319 }
320 doc
321 }
322 ast::PostfixOp::IsJsonValue(n) => {
323 let mut doc = Doc::text("is json value");
324 if let Some(clause) = n.json_keys_unique_clause() {
325 doc = doc
326 .append(Doc::space())
327 .append(build_json_keys_unique_clause(clause));
328 }
329 doc
330 }
331 ast::PostfixOp::IsNormalized(n) => {
332 let mut doc = Doc::text("is");
333 if let Some(form) = n.unicode_normal_form() {
334 doc = doc
335 .append(Doc::space())
336 .append(build_unicode_normal_form(form));
337 }
338 doc.append(Doc::space()).append(Doc::text("normalized"))
339 }
340 ast::PostfixOp::IsNotJson(n) => {
341 let mut doc = Doc::text("is not json");
342 if let Some(clause) = n.json_keys_unique_clause() {
343 doc = doc
344 .append(Doc::space())
345 .append(build_json_keys_unique_clause(clause));
346 }
347 doc
348 }
349 ast::PostfixOp::IsNotJsonArray(n) => {
350 let mut doc = Doc::text("is not json array");
351 if let Some(clause) = n.json_keys_unique_clause() {
352 doc = doc
353 .append(Doc::space())
354 .append(build_json_keys_unique_clause(clause));
355 }
356 doc
357 }
358 ast::PostfixOp::IsNotJsonObject(n) => {
359 let mut doc = Doc::text("is not json object");
360 if let Some(clause) = n.json_keys_unique_clause() {
361 doc = doc
362 .append(Doc::space())
363 .append(build_json_keys_unique_clause(clause));
364 }
365 doc
366 }
367 ast::PostfixOp::IsNotJsonScalar(n) => {
368 let mut doc = Doc::text("is not json scalar");
369 if let Some(clause) = n.json_keys_unique_clause() {
370 doc = doc
371 .append(Doc::space())
372 .append(build_json_keys_unique_clause(clause));
373 }
374 doc
375 }
376 ast::PostfixOp::IsNotJsonValue(n) => {
377 let mut doc = Doc::text("is not json value");
378 if let Some(clause) = n.json_keys_unique_clause() {
379 doc = doc
380 .append(Doc::space())
381 .append(build_json_keys_unique_clause(clause));
382 }
383 doc
384 }
385 ast::PostfixOp::IsNotNormalized(n) => {
386 let mut doc = Doc::text("is not");
387 if let Some(form) = n.unicode_normal_form() {
388 doc = doc
389 .append(Doc::space())
390 .append(build_unicode_normal_form(form));
391 }
392 doc.append(Doc::space()).append(Doc::text("normalized"))
393 }
394 };
395 expr.append(Doc::space()).append(op)
396 }
397 _ => Doc::text(expr.syntax().to_string()),
401 }
402}
403
404fn build_json_keys_unique_clause<'a>(clause: ast::JsonKeysUniqueClause) -> Doc<'a> {
405 let prefix = match clause {
406 ast::JsonKeysUniqueClause::JsonWithoutUniqueKeys(_) => "without",
407 ast::JsonKeysUniqueClause::JsonWithUniqueKeys(_) => "with",
408 };
409 Doc::text(prefix)
410 .append(Doc::space())
411 .append(Doc::text("unique"))
412 .append(Doc::space())
413 .append(Doc::text("keys"))
414}
415
416fn build_unicode_normal_form<'a>(form: ast::UnicodeNormalForm) -> Doc<'a> {
417 if form.nfc_token().is_some() {
418 Doc::text("nfc")
419 } else if form.nfd_token().is_some() {
420 Doc::text("nfd")
421 } else if form.nfkc_token().is_some() {
422 Doc::text("nfkc")
423 } else {
424 Doc::text("nfkd")
425 }
426}
427
428fn build_keyword_node<'a>(node: &SyntaxNode) -> Doc<'a> {
429 let mut docs: Vec<Doc<'a>> = vec![];
430 for el in node.children_with_tokens() {
431 match el {
432 rowan::NodeOrToken::Token(token) => match token.kind() {
433 SyntaxKind::WHITESPACE => continue,
434 SyntaxKind::COMMENT => {
435 if !docs.is_empty() {
436 docs.push(Doc::space());
437 }
438 docs.push(Doc::text(token.text().to_string()));
439 }
440 _ => {
441 if !docs.is_empty() {
442 docs.push(Doc::space());
443 }
444 docs.push(Doc::text(token.text().to_ascii_lowercase()));
445 }
446 },
447 rowan::NodeOrToken::Node(_) => (),
448 }
449 }
450 Doc::list(docs)
451}
452
453fn build_op<'a>(op: ast::BinOp) -> Doc<'a> {
454 match op {
455 ast::BinOp::And(_) => Doc::text("and"),
456 ast::BinOp::AtTimeZone(n) => build_keyword_node(n.syntax()),
457 ast::BinOp::Caret(_) => Doc::text("^"),
458 ast::BinOp::ColonColon(_) => Doc::text("::"),
459 ast::BinOp::ColonEq(_) => Doc::text(":="),
460 ast::BinOp::CustomOp(custom_op) => Doc::text(custom_op.syntax().to_string()),
461 ast::BinOp::Eq(_) => Doc::text("="),
462 ast::BinOp::Escape(_) => Doc::text("escape"),
463 ast::BinOp::FatArrow(_) => Doc::text("=>"),
464 ast::BinOp::Gteq(_) => Doc::text(">="),
465 ast::BinOp::Ilike(_) => Doc::text("ilike"),
466 ast::BinOp::In(_) => Doc::text("in"),
467 ast::BinOp::Is(_) => Doc::text("is"),
468 ast::BinOp::IsDistinctFrom(n) => build_keyword_node(n.syntax()),
469 ast::BinOp::IsNot(n) => build_keyword_node(n.syntax()),
470 ast::BinOp::IsNotDistinctFrom(n) => build_keyword_node(n.syntax()),
471 ast::BinOp::LAngle(_) => Doc::text("<"),
472 ast::BinOp::Like(_) => Doc::text("like"),
473 ast::BinOp::Lteq(_) => Doc::text("<="),
474 ast::BinOp::Minus(_) => Doc::text("-"),
475 ast::BinOp::Neq(_) => Doc::text("!="),
476 ast::BinOp::Neqb(_) => Doc::text("<>"),
477 ast::BinOp::NotIlike(n) => build_keyword_node(n.syntax()),
478 ast::BinOp::NotIn(n) => build_keyword_node(n.syntax()),
479 ast::BinOp::NotLike(n) => build_keyword_node(n.syntax()),
480 ast::BinOp::NotSimilarTo(n) => build_keyword_node(n.syntax()),
481 ast::BinOp::OperatorCall(op) => Doc::text(op.syntax().to_string()),
482 ast::BinOp::Or(_) => Doc::text("or"),
483 ast::BinOp::Overlaps(_) => Doc::text("overlaps"),
484 ast::BinOp::Percent(_) => Doc::text("%"),
485 ast::BinOp::Plus(_) => Doc::text("+"),
486 ast::BinOp::RAngle(_) => Doc::text(">"),
487 ast::BinOp::SimilarTo(n) => build_keyword_node(n.syntax()),
488 ast::BinOp::Slash(_) => Doc::text("/"),
489 ast::BinOp::Star(_) => Doc::text("*"),
490 }
491}
492
493fn build_literal<'a>(lit: ast::Literal) -> Doc<'a> {
494 let Some(kind) = lit.kind() else {
495 return Doc::nil();
496 };
497 match kind {
498 LitKind::Default(_) => Doc::text("default"),
499 LitKind::False(_) => Doc::text("false"),
500 LitKind::IntNumber(t) => Doc::text(t.text().to_string()),
501 LitKind::Null(_) => Doc::text("null"),
502 LitKind::NumericNumber(t) => Doc::text(t.text().to_string()),
503 LitKind::PositionalParam(t) => Doc::text(t.text().to_string()),
504 LitKind::True(_) => Doc::text("true"),
505 LitKind::BitString(_)
506 | LitKind::ByteString(_)
507 | LitKind::DollarQuotedString(_)
508 | LitKind::EscString(_)
509 | LitKind::NationalString(_)
510 | LitKind::String(_)
511 | LitKind::UnicodeEscString(_) => build_string_literal(&lit),
512 }
513}
514
515fn build_string_literal<'a>(lit: &ast::Literal) -> Doc<'a> {
516 let parts: Vec<Doc<'a>> = lit
517 .syntax()
518 .children_with_tokens()
519 .filter_map(|el| match el {
520 rowan::NodeOrToken::Token(t) if t.kind() != SyntaxKind::WHITESPACE => {
521 Some(Doc::text(format_string_token(&t)))
522 }
523 _ => None,
524 })
525 .collect();
526 Doc::list(Itertools::intersperse(parts.into_iter(), Doc::hard_line()).collect())
527}
528
529fn format_string_token(t: &SyntaxToken) -> String {
530 let text = t.text();
531 if matches!(
532 t.kind(),
533 SyntaxKind::STRING | SyntaxKind::DOLLAR_QUOTED_STRING
534 ) {
535 return text.to_string();
536 }
537 match text.find('\'') {
538 Some(idx) => {
539 let (prefix, rest) = text.split_at(idx);
540 let mut s = String::with_capacity(text.len());
541 s.push_str(&prefix.to_ascii_lowercase());
542 s.push_str(rest);
543 s
544 }
545 None => text.to_string(),
546 }
547}
548
549fn build_column_name<'a>(name: ast::ColumnName) -> Doc<'a> {
550 Doc::text(quote_column_alias(&name.text()))
551}
552
553fn build_type<'a>(ty: ast::Type) -> Doc<'a> {
554 Doc::text(ty.syntax().to_string())
555}
556
557fn leading_comments_token<'a>(node: &SyntaxToken) -> Doc<'a> {
558 let mut doc = Doc::nil();
559 for next in node.siblings_with_tokens(Direction::Prev).skip(1) {
560 match next {
561 rowan::NodeOrToken::Node(_node) => {
562 break;
563 }
564 rowan::NodeOrToken::Token(token) => {
565 if token.kind() == SyntaxKind::COMMENT {
566 doc = doc
567 .append(Doc::text(token.text().to_string()))
568 .append(Doc::space());
569 } else if token.kind() == SyntaxKind::WHITESPACE {
570 continue;
571 } else {
572 break;
573 }
574 }
575 }
576 }
577 doc
578}
579
580fn leading_comments<'a>(node: &SyntaxNode) -> Doc<'a> {
581 let mut doc = Doc::nil();
582 for next in node.siblings_with_tokens(Direction::Prev).skip(1) {
583 match next {
584 rowan::NodeOrToken::Node(_node) => {
585 break;
586 }
587 rowan::NodeOrToken::Token(token) => {
588 if token.kind() == SyntaxKind::COMMENT {
589 let is_block = token.text().starts_with("--");
590 doc = doc
591 .append(Doc::text(token.text().to_string()))
592 .append(if is_block {
593 Doc::hard_line()
594 } else {
595 Doc::space()
596 });
597 } else if token.kind() == SyntaxKind::WHITESPACE {
598 continue;
599 } else {
600 break;
601 }
602 }
603 }
604 }
605 doc
606}
607
608fn trailing_comments<'a>(node: &SyntaxNode) -> Doc<'a> {
609 let mut doc = Doc::nil();
610 for next in node.siblings_with_tokens(Direction::Next).skip(1) {
611 match next {
612 rowan::NodeOrToken::Node(_node) => {
613 break;
614 }
615 rowan::NodeOrToken::Token(token) => {
616 if token.kind() == SyntaxKind::COMMENT {
617 doc = doc
618 .append(Doc::space())
619 .append(Doc::text(token.text().to_string()));
620 } else if token.kind() == SyntaxKind::WHITESPACE {
621 continue;
622 } else {
623 break;
624 }
625 }
626 }
627 }
628 doc
629}
630
631fn build_target<'a>(target: ast::Target) -> Option<Doc<'a>> {
632 let mut doc = leading_comments(target.syntax());
633
634 if target.star_token().is_some() {
635 return Some(doc.append(Doc::text("*")));
636 }
637 let expr = target.expr()?;
638 doc = doc.append(build_expr(expr));
639
640 if let Some(as_name) = target.as_name() {
641 if as_name.as_token().is_some() {
642 doc = doc.append(Doc::space()).append(Doc::text("as"))
643 }
644
645 if let Some(column_name) = as_name.name() {
646 doc = doc
647 .append(Doc::space())
648 .append(build_column_name(column_name));
649 }
650 }
651
652 doc = doc.append(trailing_comments(target.syntax()));
653
654 Some(doc)
655}
656
657pub fn fmt(text: &str) -> Result<String> {
658 let line_ending = find_newline(text)
659 .map(|(_, ending)| ending)
660 .unwrap_or_default();
661
662 let line_break = match line_ending {
663 LineEnding::Cr => bail!("CR line endings aren't supported"),
665 LineEnding::CrLf => LineBreak::Crlf,
666 LineEnding::Lf => LineBreak::Lf,
667 };
668
669 let parse = ast::SourceFile::parse(text);
670 let file = parse.tree();
671 debug_assert_eq!(
672 parse.errors(),
673 vec![],
674 "should bail out when there's parse errors"
675 );
676 let doc = build_source_file(&file);
677
678 Ok(print(
679 &doc,
680 &PrintOptions {
681 line_break,
682 ..Default::default()
683 },
684 ))
685}