use std::{
fmt::{Debug, Display},
ops::Deref,
};
use crate::syntax::{
lexer::token::Token,
location::{OriginTable, Pos, Span}, parser::token::Tt,
};
use super::kind::TreeKind;
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub struct NodeHandle(pub u32);
impl NodeHandle {
pub fn with_cst(self, cst: &Cst) -> NodeRef {
NodeRef { node: self, cst }
}
}
#[derive(Debug, Clone)]
pub enum Node {
Token(Token),
Open { kind: TreeKind, close: NodeHandle },
Close { kind: TreeKind, open: NodeHandle },
}
impl Node {
pub fn is_open(&self, expected: TreeKind) -> bool {
matches!(self, Node::Open { kind, .. } if *kind == expected)
}
pub fn is_open_where(&self, pred: impl Fn(TreeKind) -> bool) -> bool {
matches!(self, Node::Open { kind, .. } if (pred(*kind)))
}
pub fn is_tt(&self, expected: impl Into<Tt>) -> bool {
matches!(self, Node::Token(token) if Tt::from_kind(&token.kind) == Some(expected.into()))
}
pub fn token(&self) -> Option<&Token> {
match self {
Node::Token(token) => Some(token),
_ => None,
}
}
pub fn open_kind(&self) -> Option<TreeKind> {
match self {
Node::Open { kind, .. } => Some(*kind),
_ => None,
}
}
}
#[derive(Debug)]
pub struct Cst {
pub nodes: Vec<Node>,
}
impl Cst {
pub fn new(nodes: Vec<Node>) -> Self {
Self { nodes }
}
pub fn root(&self) -> NodeRef {
NodeHandle(0).with_cst(self)
}
pub fn token_at(&self, pos: Pos, ot: &OriginTable) -> Option<NodeRef> {
let check_span = |span: &Span| {
if span.origin != pos.origin {
return false;
}
let start = span.start().text_pos;
let end = span.end(ot).text_pos;
let pos = pos.text_pos;
start <= pos && pos <= end
};
self.nodes
.iter()
.enumerate()
.filter_map(|(i, n)| n.token().map(|t| (i, t)))
.find(|(_, t)| check_span(&t.span))
.map(|(i, _)| NodeHandle(i as u32).with_cst(self))
}
pub fn node_at(&self, pos: Pos, ot: &OriginTable) -> NodeRef {
let check_span = |span: &Span| {
if span.origin != pos.origin {
return false;
}
let end = span.end(ot).text_pos;
let pos = pos.text_pos;
pos <= end
};
self.nodes
.iter()
.enumerate()
.filter_map(|(i, n)| n.token().map(|t| (i, t)))
.find(|(_, t)| check_span(&t.span))
.map(|(i, _)| NodeHandle(i as u32).with_cst(self))
.unwrap_or(self.root())
}
pub fn get(&self, node: NodeHandle) -> &Node {
&self.nodes[node.0 as usize]
}
pub fn parent(&self, child: NodeHandle) -> Option<NodeHandle> {
let mut node = child.0;
loop {
if node == 0 {
return None;
}
node -= 1;
match self.nodes[node as usize] {
Node::Token(_) => (),
Node::Open { .. } => return Some(NodeHandle(node)),
Node::Close { open, .. } => node = open.0,
}
}
}
pub fn prev_sibling(&self, node: NodeHandle) -> Option<NodeHandle> {
let node = node.0;
if node == 0 {
return None;
}
let node = node - 1;
match &self.nodes[node as usize] {
Node::Token(_) => Some(NodeHandle(node)),
Node::Open { .. } => None,
Node::Close { open, .. } => Some(*open),
}
}
pub fn next_sibling(&self, node: NodeHandle) -> Option<NodeHandle> {
let node = node.0;
if node as usize == self.nodes.len() - 1 {
return None;
}
let node = node + 1;
match &self.nodes[node as usize] {
Node::Token(_) | Node::Open { .. } => Some(NodeHandle(node)),
Node::Close { .. } => None,
}
}
pub fn children(&self, parent: NodeHandle) -> Children<'_> {
match self.get(parent) {
Node::Token(_) | Node::Close { .. } => Children {
cst: self,
next: NodeHandle(self.nodes.len() as u32 - 1),
},
Node::Open { .. } => Children {
cst: self,
next: NodeHandle(parent.0 + 1),
},
}
}
pub fn is_ancestor(&self, ancestor: NodeHandle, descendant: NodeHandle) -> bool {
let Node::Open { close, .. } = self.get(ancestor) else {
return false;
};
let start = ancestor.0;
let end = close.0;
let descendant = descendant.0;
start <= descendant && descendant <= end
}
pub fn span(&self, node: NodeHandle, ot: &OriginTable) -> Option<Span> {
let range = match self.get(node) {
Node::Token(token) => return Some(token.span),
Node::Open { close, .. } => node.0 as usize..close.0 as usize,
Node::Close { open, .. } => open.0 as usize..node.0 as usize,
};
let start = &self.nodes[range.clone()].iter().find_map(Node::token)?;
let end = &self.nodes[range].iter().rev().find_map(Node::token)?;
let start = start.span.start();
let end = end.span.end(ot);
let span = Span::cross_origin(start, end, ot);
Some(span)
}
pub fn iter(&self) -> impl Iterator<Item = NodeRef<'_>> {
self.nodes
.iter()
.enumerate()
.map(|(idx, _)| NodeHandle(idx as u32).with_cst(self))
}
pub fn tokens(&self) -> impl Iterator<Item = (NodeHandle, &Token)> {
self.iter()
.filter_map(|n| Some((n.handle(), self.get(n.handle()).token()?)))
}
}
impl Display for Cst {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
let mut depth = 0;
fn print_indent(f: &mut std::fmt::Formatter<'_>, depth: usize) -> std::fmt::Result {
for _ in 0..depth {
write!(f, " ")?
}
Ok(())
}
for node in &self.nodes {
match node {
Node::Token(token) => {
print_indent(f, depth)?;
writeln!(f, "{} {:?}", token.kind, token.span)?;
}
Node::Open { kind, .. } => {
print_indent(f, depth)?;
writeln!(f, "{kind:?}:")?;
depth += 1;
}
Node::Close { .. } => {
depth -= 1;
}
}
}
Ok(())
}
}
#[derive(Debug, Clone, Copy)]
pub struct Children<'cst> {
cst: &'cst Cst,
next: NodeHandle,
}
impl<'cst> Children<'cst> {
pub fn find_kind(mut self, kind: TreeKind) -> Option<NodeRef<'cst>> {
self.find(|n| n.is_open(kind))
}
pub fn find_kind_where(mut self, pred: impl Fn(TreeKind) -> bool) -> Option<NodeRef<'cst>> {
self.find(|n| n.is_open_where(&pred))
}
pub fn find_tt(mut self, tt: impl Into<Tt>) -> Option<NodeRef<'cst>> {
let tt = tt.into();
self.find(|n| n.is_tt(tt))
}
}
impl<'cst> Iterator for Children<'cst> {
type Item = NodeRef<'cst>;
fn next(&mut self) -> Option<Self::Item> {
let node = self.next;
match self.cst.get(node) {
Node::Token(_) => {
self.next = NodeHandle(node.0 + 1);
Some(node.with_cst(self.cst))
}
Node::Open { close, .. } => {
self.next = NodeHandle(close.0 + 1);
Some(node.with_cst(self.cst))
}
Node::Close { .. } => None,
}
}
}
#[derive(Clone, Copy)]
pub struct NodeRef<'cst> {
node: NodeHandle,
cst: &'cst Cst,
}
impl Debug for NodeRef<'_> {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("NodeRef").field("node", &self.node).finish()
}
}
impl PartialEq for NodeRef<'_> {
fn eq(&self, other: &Self) -> bool {
self.node == other.node
}
}
impl Eq for NodeRef<'_> {}
impl Deref for NodeRef<'_> {
type Target = Node;
fn deref(&self) -> &Self::Target {
self.cst.get(self.node)
}
}
impl<'cst> NodeRef<'cst> {
pub fn cst(&self) -> &Cst {
self.cst
}
pub fn handle(&self) -> NodeHandle {
self.node
}
pub fn project(self, f: impl FnOnce(&Cst, NodeHandle) -> NodeHandle) -> Self {
NodeRef {
node: (f)(self.cst, self.node),
cst: self.cst,
}
}
pub fn try_project(
self,
f: impl FnOnce(&Cst, NodeHandle) -> Option<NodeHandle>,
) -> Option<Self> {
Some(NodeRef {
node: (f)(self.cst, self.node)?,
cst: self.cst,
})
}
pub fn parent(self) -> Option<Self> {
self.try_project(Cst::parent)
}
pub fn next_sibling(self) -> Option<Self> {
self.try_project(Cst::next_sibling)
}
pub fn prev_sibling(self) -> Option<Self> {
self.try_project(Cst::prev_sibling)
}
pub fn children(self) -> Children<'cst> {
self.cst.children(self.node)
}
pub fn span(self, ot: &OriginTable) -> Option<Span> {
self.cst.span(self.node, ot)
}
pub fn contains(self, inner: NodeHandle) -> bool {
self.cst.is_ancestor(self.handle(), inner)
}
}