use std::collections::HashMap;
use crate::ast::generated::NodeIdWalk;
use crate::ast::generated::visit::{self, Visit};
use crate::ast::{FunctionCall, NoExt, NodeId, SourceStore, Span};
use crate::parser::Parsed;
use crate::tokenizer::{TriviaKind, TriviaRange};
#[derive(Clone, Debug, Default)]
pub struct NodeComments {
pub leading: Vec<TriviaRange>,
pub trailing: Vec<TriviaRange>,
pub dangling: Vec<TriviaRange>,
}
#[derive(Clone, Debug, Default)]
pub struct CommentAttachments {
by_anchor: HashMap<NodeId, NodeComments>,
span_to_id: HashMap<Span, NodeId>,
comments: Vec<TriviaRange>,
attached: usize,
}
impl CommentAttachments {
pub fn leading(&self, id: NodeId) -> &[TriviaRange] {
self.by_anchor.get(&id).map_or(&[], |c| &c.leading)
}
pub fn trailing(&self, id: NodeId) -> &[TriviaRange] {
self.by_anchor.get(&id).map_or(&[], |c| &c.trailing)
}
pub fn dangling(&self, id: NodeId) -> &[TriviaRange] {
self.by_anchor.get(&id).map_or(&[], |c| &c.dangling)
}
pub fn is_empty(&self) -> bool {
self.attached == 0
}
pub fn len(&self) -> usize {
self.attached
}
pub fn all_comments(&self) -> &[TriviaRange] {
&self.comments
}
pub fn leading_for(&self, span: Span) -> &[TriviaRange] {
self.span_to_id
.get(&span)
.map_or(&[], |id| self.leading(*id))
}
pub fn trailing_for(&self, span: Span) -> &[TriviaRange] {
self.span_to_id
.get(&span)
.map_or(&[], |id| self.trailing(*id))
}
pub fn dangling_for(&self, span: Span) -> &[TriviaRange] {
self.span_to_id
.get(&span)
.map_or(&[], |id| self.dangling(*id))
}
fn entry(&mut self, id: NodeId) -> &mut NodeComments {
self.by_anchor.entry(id).or_default()
}
pub fn compute<S: SourceStore>(parsed: &Parsed<S, NoExt>) -> CommentAttachments {
let comments: Vec<TriviaRange> = parsed
.trivia()
.iter()
.copied()
.filter(|t| matches!(t.kind(), TriviaKind::LineComment | TriviaKind::BlockComment))
.collect();
if comments.is_empty() {
return CommentAttachments::default();
}
let mut walk = NodeIdWalk::default();
for statement in parsed.statements() {
walk.visit_statement(statement);
}
let nodes: Vec<NodeInfo> = walk
.metas
.iter()
.filter(|m| !m.span.is_synthetic())
.map(|m| NodeInfo {
id: m.node_id,
span: m.span,
})
.collect();
let mut children: Vec<Vec<usize>> = vec![Vec::new(); nodes.len()];
let mut roots: Vec<usize> = Vec::new();
let mut stack: Vec<usize> = Vec::new();
for (i, node) in nodes.iter().enumerate() {
while let Some(&top) = stack.last() {
if contains(nodes[top].span, node.span) {
break;
}
stack.pop();
}
match stack.last() {
Some(&top) => children[top].push(i),
None => roots.push(i),
}
stack.push(i);
}
let mut brackets = EmptyBracketWalk::default();
for statement in parsed.statements() {
brackets.visit_statement(statement);
}
let source = parsed.source();
let mut span_to_id: HashMap<Span, NodeId> = HashMap::new();
for node in &nodes {
span_to_id.insert(node.span, node.id);
}
let mut out = CommentAttachments {
comments: comments.clone(),
..CommentAttachments::default()
};
for comment in &comments {
out.attached += 1;
classify(
&mut out,
*comment,
&nodes,
&children,
&roots,
&brackets.ids,
&span_to_id,
source,
parsed,
);
}
out.span_to_id = span_to_id;
out
}
}
struct NodeInfo {
id: NodeId,
span: Span,
}
fn contains(outer: Span, inner: Span) -> bool {
!outer.is_synthetic()
&& !inner.is_synthetic()
&& outer.start() <= inner.start()
&& inner.end() <= outer.end()
}
fn newline_between(source: &str, a: u32, b: u32) -> bool {
if a >= b {
return false;
}
source
.get(a as usize..b as usize)
.is_some_and(|s| s.contains('\n'))
}
#[allow(clippy::too_many_arguments)]
fn classify<S: SourceStore>(
out: &mut CommentAttachments,
comment: TriviaRange,
nodes: &[NodeInfo],
children: &[Vec<usize>],
roots: &[usize],
brackets: &std::collections::HashSet<NodeId>,
span_to_id: &HashMap<Span, NodeId>,
source: &str,
parsed: &Parsed<S, NoExt>,
) {
let anchor = |idx: usize| -> NodeId {
span_to_id
.get(&nodes[idx].span)
.copied()
.unwrap_or(nodes[idx].id)
};
let cspan = comment.span();
let mut enclosing: Option<usize> = None;
for (i, node) in nodes.iter().enumerate() {
if !contains(node.span, cspan) {
continue;
}
enclosing = Some(match enclosing {
None => i,
Some(best) => {
let (bl, nl) = (nodes[best].span.len(), node.span.len());
if nl < bl || (nl == bl && i > best) {
i
} else {
best
}
}
});
}
let Some(enc) = enclosing else {
attach_top_level(out, comment, nodes, roots, span_to_id);
return;
};
let kids = &children[enc];
let mut preceding: Option<usize> = None;
let mut following: Option<usize> = None;
for &k in kids {
let ks = nodes[k].span;
if ks.end() <= cspan.start() && preceding.is_none_or(|p| ks.end() > nodes[p].span.end()) {
preceding = Some(k);
}
if ks.start() >= cspan.end() && following.is_none_or(|f| ks.start() < nodes[f].span.start())
{
following = Some(k);
}
}
let enc_id = anchor(enc);
let preceding_end = preceding.map_or(nodes[enc].span.start(), |p| nodes[p].span.end());
let clause_kw_before = parsed
.clause_marks_in(nodes[enc].span)
.iter()
.any(|mark| mark.offset() >= preceding_end && mark.offset() < cspan.start());
match (preceding, following) {
(Some(p), Some(f)) => {
let own_line = newline_between(source, nodes[p].span.end(), cspan.start());
if own_line || clause_kw_before {
out.entry(anchor(f)).leading.push(comment);
} else {
out.entry(anchor(p)).trailing.push(comment);
}
}
(Some(p), None) => {
if brackets.contains(&enc_id) {
out.entry(enc_id).dangling.push(comment);
} else {
out.entry(anchor(p)).trailing.push(comment);
}
}
(None, Some(f)) => {
out.entry(anchor(f)).leading.push(comment);
}
(None, None) => {
out.entry(enc_id).dangling.push(comment);
}
}
}
fn attach_top_level(
out: &mut CommentAttachments,
comment: TriviaRange,
nodes: &[NodeInfo],
roots: &[usize],
span_to_id: &HashMap<Span, NodeId>,
) {
let anchor = |idx: usize| -> NodeId {
span_to_id
.get(&nodes[idx].span)
.copied()
.unwrap_or(nodes[idx].id)
};
let cspan = comment.span();
let mut preceding: Option<usize> = None;
let mut following: Option<usize> = None;
for &r in roots {
let rs = nodes[r].span;
if rs.end() <= cspan.start() && preceding.is_none_or(|p| rs.end() > nodes[p].span.end()) {
preceding = Some(r);
}
if rs.start() >= cspan.end() && following.is_none_or(|f| rs.start() < nodes[f].span.start())
{
following = Some(r);
}
}
match (following, preceding) {
(Some(f), _) => out.entry(anchor(f)).leading.push(comment),
(None, Some(p)) => out.entry(anchor(p)).trailing.push(comment),
(None, None) => {}
}
}
#[derive(Default)]
struct EmptyBracketWalk {
ids: std::collections::HashSet<NodeId>,
}
impl<'ast> Visit<'ast, NoExt> for EmptyBracketWalk {
fn visit_function_call(&mut self, node: &'ast FunctionCall<NoExt>) {
if node.args.is_empty() {
self.ids.insert(node.meta.node_id);
}
visit::walk_function_call(self, node);
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::dialect::Ansi;
use crate::parse_with;
fn attach(sql: &str) -> (CommentAttachments, crate::Parsed) {
let parsed =
parse_with(sql, crate::ParseConfig::new(Ansi).capture_trivia(true)).expect("parses");
let attachments = CommentAttachments::compute(&parsed);
(attachments, parsed)
}
fn texts(ranges: &[TriviaRange], source: &str) -> Vec<String> {
ranges
.iter()
.map(|r| source[r.span().start() as usize..r.span().end() as usize].to_owned())
.collect()
}
#[test]
fn empty_when_no_trivia_captured() {
let parsed = crate::parse_with("SELECT 1", crate::ParseConfig::new(Ansi)).expect("parses");
assert!(CommentAttachments::compute(&parsed).is_empty());
}
#[test]
fn leading_comment_before_statement() {
let (a, p) = attach("-- hi\nSELECT 1");
assert_eq!(a.len(), 1);
let all: Vec<String> = a
.by_anchor
.values()
.flat_map(|c| texts(&c.leading, p.source()))
.collect();
assert_eq!(all, vec!["-- hi".to_string()]);
}
#[test]
fn comment_before_clause_keyword_anchors_leading_not_trailing() {
let sql = "SELECT a FROM t WHERE a = 1\n-- note\nGROUP BY a";
let (a, p) = attach(sql);
let leading_all: Vec<String> = a
.by_anchor
.values()
.flat_map(|c| texts(&c.leading, p.source()))
.collect();
assert!(
leading_all.contains(&"-- note".to_string()),
"clause-keyword comment should anchor leading; got {leading_all:?}"
);
let trailing_all: Vec<String> = a
.by_anchor
.values()
.flat_map(|c| texts(&c.trailing, p.source()))
.collect();
assert!(!trailing_all.contains(&"-- note".to_string()));
}
#[test]
fn empty_call_dangles_comment_inside() {
let (a, p) = attach("SELECT count(/* c */)");
let dangling_all: Vec<String> = a
.by_anchor
.values()
.flat_map(|c| texts(&c.dangling, p.source()))
.collect();
assert_eq!(dangling_all, vec!["/* c */".to_string()]);
}
#[test]
fn trailing_eol_comment_stays_trailing() {
let (a, p) = attach("SELECT a -- tag\nFROM t");
let trailing_all: Vec<String> = a
.by_anchor
.values()
.flat_map(|c| texts(&c.trailing, p.source()))
.collect();
assert_eq!(trailing_all, vec!["-- tag".to_string()]);
}
}