use crate::syntax::{SyntaxKind, SyntaxNode, SyntaxToken};
use rowan::TextSize;
use std::collections::BTreeMap;
const INDENT_WIDTH: usize = 4;
#[derive(Debug, Clone, Copy, Default, PartialEq, Eq, PartialOrd, Ord)]
pub(crate) enum Sep {
#[default]
None,
Space,
Newline,
}
#[derive(Debug, Clone, Copy, Default, PartialEq, Eq)]
enum Width {
#[default]
Open,
Blank,
Settled,
}
#[derive(Debug, Clone, Copy, Default, PartialEq, Eq)]
struct Gap {
sep: Sep,
width: Width,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub(crate) enum RowFamily {
Instantiation,
ParameterDefinition,
EnumEntry,
Other,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, PartialOrd, Ord)]
pub(crate) enum AlignPoint {
InstType,
InstName,
InstReset,
InstAddress,
InstStride,
InstAlign,
ParamName,
ParamDefault,
EnumValue,
TrailingComment,
}
#[derive(Debug, Clone, Copy)]
struct Marker {
point: AlignPoint,
pos: usize,
base_space: bool,
}
#[derive(Debug)]
struct Row {
family: RowFamily,
start: Option<usize>,
end: Option<usize>,
markers: Vec<Marker>,
}
#[derive(Debug, Default)]
struct Scope {
rows: Vec<usize>,
}
#[derive(Debug, Clone, Copy)]
struct PendingMarker {
row: usize,
point: AlignPoint,
}
pub(crate) struct Formatter<'a> {
src: &'a str,
out: String,
indent: usize,
gap: Gap,
blank_lines: bool,
saw_newline: bool,
after_comment: bool,
eol: &'static str,
rows: Vec<Row>,
row_stack: Vec<usize>,
scopes: Vec<Scope>,
scope_stack: Vec<usize>,
pending_markers: Vec<PendingMarker>,
}
pub(crate) fn line_ending(src: &str) -> &'static str {
match src.find('\n') {
Some(i) if src.as_bytes()[..i].last() == Some(&b'\r') => "\r\n",
_ => "\n",
}
}
impl<'a> Formatter<'a> {
pub(crate) fn new(src: &'a str) -> Self {
Formatter {
src,
out: String::with_capacity(src.len()),
indent: 0,
gap: Gap::default(),
blank_lines: true,
saw_newline: false,
after_comment: false,
eol: line_ending(src),
rows: Vec::new(),
row_stack: Vec::new(),
scopes: vec![Scope::default()],
scope_stack: vec![0],
pending_markers: Vec::new(),
}
}
pub(crate) fn finish(mut self) -> String {
self.align();
let trimmed = self.out.trim_end().len();
self.out.truncate(trimmed);
if !self.out.is_empty() {
self.out.push_str(self.eol);
}
self.out
}
pub(crate) fn request(&mut self, sep: Sep) {
self.gap.sep = self.gap.sep.max(sep);
}
pub(crate) fn blank_line(&mut self) {
if self.blank_lines && self.gap.width == Width::Open {
self.gap.width = Width::Blank;
}
}
pub(crate) fn pin(&mut self, sep: Sep) {
self.gap = Gap {
sep,
width: Width::Settled,
};
}
pub(crate) fn settle_width(&mut self) {
self.gap.width = Width::Settled;
}
pub(crate) fn allow_blank_lines(&mut self, allow: bool) -> bool {
std::mem::replace(&mut self.blank_lines, allow)
}
pub(crate) fn open_alignment_scope(&mut self) {
let id = self.scopes.len();
self.scopes.push(Scope::default());
self.scope_stack.push(id);
}
pub(crate) fn close_alignment_scope(&mut self) {
debug_assert!(self.scope_stack.len() > 1);
self.scope_stack.pop();
}
pub(crate) fn begin_row(&mut self, family: RowFamily) {
let id = self.rows.len();
self.rows.push(Row {
family,
start: None,
end: None,
markers: Vec::new(),
});
let scope = *self.scope_stack.last().expect("root alignment scope");
self.scopes[scope].rows.push(id);
self.row_stack.push(id);
}
pub(crate) fn end_row(&mut self) {
let row = self.row_stack.pop().expect("end_row without begin_row");
self.pending_markers.retain(|marker| marker.row != row);
}
pub(crate) fn align_before(&mut self, point: AlignPoint) {
if let Some(&row) = self.row_stack.last() {
self.pending_markers.push(PendingMarker { row, point });
}
}
pub(crate) fn indent(&mut self) {
self.indent += 1;
}
pub(crate) fn dedent(&mut self) {
self.indent = self.indent.saturating_sub(1);
}
fn materialize(&mut self) {
let gap = std::mem::take(&mut self.gap);
if self.out.is_empty() {
return;
}
match gap.sep {
Sep::None => self.materialize_markers(false),
Sep::Space => {
self.materialize_markers(true);
self.out.push(' ');
}
Sep::Newline => {
self.newline(if gap.width == Width::Blank { 2 } else { 1 });
self.materialize_markers(false);
}
}
}
fn materialize_markers(&mut self, base_space: bool) {
for pending in self.pending_markers.drain(..) {
self.rows[pending.row].markers.push(Marker {
point: pending.point,
pos: self.out.len(),
base_space,
});
}
}
fn newline(&mut self, count: usize) {
for _ in 0..count {
self.out.push_str(self.eol);
}
for _ in 0..self.indent * INDENT_WIDTH {
self.out.push(' ');
}
}
fn write_raw(&mut self, text: &str) {
self.materialize();
self.out.push_str(text);
}
pub(crate) fn token(&mut self, tok: &SyntaxToken) {
debug_assert!(!tok.kind().is_trivia(), "trivia must go through trivia()");
self.materialize();
let start = self.out.len();
for &row in &self.row_stack {
self.rows[row].start.get_or_insert(start);
}
self.out.push_str(tok.text());
let end = self.out.len();
for &row in &self.row_stack {
self.rows[row].end = Some(end);
}
self.saw_newline = false;
self.after_comment = false;
}
pub(crate) fn trivia(&mut self, tok: &SyntaxToken) {
match tok.kind() {
SyntaxKind::WHITESPACE => {
let newlines = tok.text().bytes().filter(|&b| b == b'\n').count();
if newlines >= 2 {
self.blank_line();
} else if newlines == 1 && self.after_comment {
self.request(Sep::Newline);
}
self.saw_newline |= newlines >= 1;
}
kind if kind.is_directive() => {
self.request(Sep::Newline);
self.write_raw(tok.text().trim_end());
self.request(Sep::Newline);
self.saw_newline = false;
self.after_comment = false;
}
kind if kind.is_comment() => {
let inline = !self.saw_newline;
if self.saw_newline {
self.request(Sep::Newline);
} else if kind == SyntaxKind::LINE_COMMENT {
self.pin(Sep::Space);
} else {
self.request(Sep::Space);
}
if inline {
self.attach_trailing_comment();
}
self.write_raw(tok.text());
if kind == SyntaxKind::LINE_COMMENT {
self.request(Sep::Newline);
} else {
self.request(Sep::Space);
}
self.saw_newline = false;
self.after_comment = true;
}
kind => unreachable!("not trivia: {kind:?}"),
}
}
pub(crate) fn verbatim(&mut self, node: &SyntaxNode) {
let src = self.src;
let mut start = node.text_range().start();
let mut end = node.text_range().end();
for tok in leading_trivia(node) {
self.trivia(&tok);
start = tok.text_range().end();
}
let trailing = trailing_trivia(node, start);
if let Some(first) = trailing.first() {
end = first.text_range().start();
}
if start < end {
self.write_raw(&src[usize::from(start)..usize::from(end)]);
self.saw_newline = false;
self.after_comment = false;
}
for tok in &trailing {
self.trivia(tok);
}
}
fn attach_trailing_comment(&mut self) {
let line_start = self.out.rfind('\n').map_or(0, |pos| pos + 1);
let Some((_, row)) = self
.rows
.iter_mut()
.enumerate()
.rev()
.find(|(_, row)| row.end.is_some_and(|end| end >= line_start))
else {
return;
};
row.markers.push(Marker {
point: AlignPoint::TrailingComment,
pos: self.out.len(),
base_space: true,
});
}
fn align(&mut self) {
let mut insertions: BTreeMap<usize, usize> = BTreeMap::new();
for scope in &self.scopes {
let mut run: Vec<usize> = Vec::new();
let mut previous: Option<usize> = None;
for &row_id in &scope.rows {
let row = &self.rows[row_id];
let eligible = row.family != RowFamily::Other
&& row.start.zip(row.end).is_some_and(|(start, end)| {
!self.out[start..end].contains('\n') && !row.markers.is_empty()
});
let continues = eligible
&& previous.is_some_and(|prev_id| {
let prev = &self.rows[prev_id];
prev.family == row.family && !self.breaks_run(prev, row)
});
if !continues {
self.align_run(&run, &mut insertions);
run.clear();
}
if eligible {
run.push(row_id);
previous = Some(row_id);
} else {
previous = None;
}
}
self.align_run(&run, &mut insertions);
}
if insertions.is_empty() {
return;
}
let mut aligned =
String::with_capacity(self.out.len() + insertions.values().copied().sum::<usize>());
let mut cursor = 0;
for (pos, count) in insertions {
aligned.push_str(&self.out[cursor..pos]);
aligned.extend(std::iter::repeat_n(' ', count));
cursor = pos;
}
aligned.push_str(&self.out[cursor..]);
self.out = aligned;
}
fn breaks_run(&self, previous: &Row, current: &Row) -> bool {
let (Some(end), Some(start)) = (previous.end, current.start) else {
return true;
};
let between = &self.out[end..start];
let mut physical = between.split('\n');
physical.next();
let mut interior: Vec<&str> = physical.collect();
interior.pop();
interior.iter().any(|line| line.trim().is_empty())
|| crate::syntax::lex(between)
.iter()
.any(|(kind, _)| kind.is_directive())
}
fn align_run(&self, run: &[usize], insertions: &mut BTreeMap<usize, usize>) {
if run.len() < 2 {
return;
}
for point in [
AlignPoint::InstType,
AlignPoint::InstName,
AlignPoint::InstReset,
AlignPoint::InstAddress,
AlignPoint::InstStride,
AlignPoint::InstAlign,
AlignPoint::ParamName,
AlignPoint::ParamDefault,
AlignPoint::EnumValue,
AlignPoint::TrailingComment,
] {
let mut group: Vec<(usize, Marker)> = Vec::new();
for &row_id in run {
let marker = self.rows[row_id]
.markers
.iter()
.find(|marker| marker.point == point)
.copied();
if let Some(marker) = marker {
group.push((row_id, marker));
} else {
self.align_column_group(&group, insertions);
group.clear();
}
}
self.align_column_group(&group, insertions);
}
}
fn align_column_group(
&self,
group: &[(usize, Marker)],
insertions: &mut BTreeMap<usize, usize>,
) {
if group.len() < 2 {
return;
}
let widths: Vec<usize> = group
.iter()
.map(|&(row_id, marker)| self.cell_width(&self.rows[row_id], marker))
.collect();
let maximum = widths.iter().copied().max().unwrap_or(0);
let separator = usize::from(maximum > 0);
for ((_, marker), width) in group.iter().zip(widths) {
let base = usize::from(marker.base_space);
let padding = maximum - width + separator.saturating_sub(base);
if padding > 0 {
insertions
.entry(marker.pos)
.and_modify(|old| *old = (*old).max(padding))
.or_insert(padding);
}
}
}
fn cell_width(&self, row: &Row, marker: Marker) -> usize {
let start = row
.markers
.iter()
.filter(|candidate| candidate.pos < marker.pos)
.map(|candidate| candidate.pos)
.max()
.or(row.start)
.unwrap_or(marker.pos);
self.out[start..marker.pos]
.trim_start_matches([' ', '\t'])
.chars()
.count()
}
}
fn leading_trivia(node: &SyntaxNode) -> impl Iterator<Item = SyntaxToken> {
let end = node.text_range().end();
std::iter::successors(node.first_token(), |tok: &SyntaxToken| tok.next_token())
.take_while(move |tok| tok.text_range().end() <= end && tok.kind().is_trivia())
}
fn trailing_trivia(node: &SyntaxNode, floor: TextSize) -> Vec<SyntaxToken> {
let mut out: Vec<SyntaxToken> =
std::iter::successors(node.last_token(), |tok: &SyntaxToken| tok.prev_token())
.take_while(|tok| tok.text_range().start() >= floor && tok.kind().is_trivia())
.collect();
out.reverse();
out
}