use crate::{Language, SyntaxNode, SyntaxSlot, SyntaxToken};
pub trait SyntaxRewriter {
type Language: Language;
fn transform(&mut self, node: SyntaxNode<Self::Language>) -> SyntaxNode<Self::Language>
where
Self: Sized,
{
match self.visit_node(node) {
VisitNodeSignal::Replace(updated) => updated,
VisitNodeSignal::Traverse(node) => traverse(node, self),
}
}
fn visit_node(&mut self, node: SyntaxNode<Self::Language>) -> VisitNodeSignal<Self::Language> {
VisitNodeSignal::Traverse(node)
}
fn visit_token(&mut self, token: SyntaxToken<Self::Language>) -> SyntaxToken<Self::Language> {
token
}
}
#[derive(Debug, Clone)]
pub enum VisitNodeSignal<L: Language> {
Replace(SyntaxNode<L>),
Traverse(SyntaxNode<L>),
}
fn traverse<R>(mut parent: SyntaxNode<R::Language>, rewriter: &mut R) -> SyntaxNode<R::Language>
where
R: SyntaxRewriter,
{
for slot in parent.slots() {
match slot {
SyntaxSlot::Node(node) => {
let original_key = node.key();
let index = node.index();
let updated = rewriter.transform(node);
if updated.key() != original_key {
parent = parent.splice_slots(index..=index, [Some(updated.into())]);
}
}
SyntaxSlot::Token(token) => {
let original_key = token.key();
let index = token.index();
let updated = rewriter.visit_token(token);
if updated.key() != original_key {
parent = parent.splice_slots(index..=index, [Some(updated.into())]);
}
}
SyntaxSlot::Empty { .. } => {
}
}
}
parent
}
#[cfg(test)]
mod tests {
use crate::raw_language::{RawLanguage, RawLanguageKind, RawSyntaxTreeBuilder};
use crate::{SyntaxNode, SyntaxRewriter, SyntaxToken, VisitNodeSignal};
#[test]
pub fn test_visits_each_node() {
let mut builder = RawSyntaxTreeBuilder::new();
builder.start_node(RawLanguageKind::ROOT);
builder.start_node(RawLanguageKind::LITERAL_EXPRESSION);
builder.token(RawLanguageKind::NUMBER_TOKEN, "5");
builder.finish_node();
builder.finish_node();
let root = builder.finish();
let mut recorder = RecordRewritter::default();
let transformed = recorder.transform(root.clone());
assert_eq!(
&root, &transformed,
"It should return the same node if the rewritter doesn't replace a node."
);
let literal_expression = root
.descendants()
.find(|node| node.kind() == RawLanguageKind::LITERAL_EXPRESSION)
.unwrap();
assert_eq!(&recorder.nodes, &[root.clone(), literal_expression]);
let number_literal = root.first_token().unwrap();
assert_eq!(&recorder.tokens, &[number_literal]);
}
#[derive(Default)]
struct RecordRewritter {
nodes: Vec<SyntaxNode<RawLanguage>>,
tokens: Vec<SyntaxToken<RawLanguage>>,
}
impl SyntaxRewriter for RecordRewritter {
type Language = RawLanguage;
fn visit_node(
&mut self,
node: SyntaxNode<Self::Language>,
) -> VisitNodeSignal<Self::Language> {
self.nodes.push(node.clone());
VisitNodeSignal::Traverse(node)
}
fn visit_token(
&mut self,
token: SyntaxToken<Self::Language>,
) -> SyntaxToken<Self::Language> {
self.tokens.push(token.clone());
token
}
}
}