1use std::{
2 fmt::{Debug, Display},
3 ops::Deref,
4};
5
6use crate::syntax::{
7 lexer::token::Token,
8 location::{OriginTable, Pos, Span},
9 parser::token::Tt,
10};
11
12use super::kind::TreeKind;
13
14#[derive(Debug, Clone, Copy, PartialEq, Eq)]
15pub struct NodeHandle(pub u32);
16
17impl NodeHandle {
18 pub fn with_cst(self, cst: &Cst) -> NodeRef<'_> {
19 NodeRef { node: self, cst }
20 }
21}
22
23#[derive(Debug, Clone)]
24pub enum Node {
25 Token(Token),
26 Open { kind: TreeKind, close: NodeHandle },
27 Close { kind: TreeKind, open: NodeHandle },
28}
29
30impl Node {
31 pub fn is_open(&self, expected: TreeKind) -> bool {
32 matches!(self, Node::Open { kind, .. } if *kind == expected)
33 }
34
35 pub fn is_open_where(&self, pred: impl Fn(TreeKind) -> bool) -> bool {
36 matches!(self, Node::Open { kind, .. } if (pred(*kind)))
37 }
38
39 pub fn is_tt(&self, expected: impl Into<Tt>) -> bool {
40 matches!(self, Node::Token(token) if Tt::from_kind(&token.kind) == Some(expected.into()))
41 }
42
43 pub fn token(&self) -> Option<&Token> {
44 match self {
45 Node::Token(token) => Some(token),
46 _ => None,
47 }
48 }
49
50 pub fn open_kind(&self) -> Option<TreeKind> {
51 match self {
52 Node::Open { kind, .. } => Some(*kind),
53 _ => None,
54 }
55 }
56}
57
58#[derive(Debug)]
59pub struct Cst {
60 pub nodes: Vec<Node>,
62}
63
64impl Cst {
65 pub fn new(nodes: Vec<Node>) -> Self {
66 Self { nodes }
67 }
68
69 pub fn root(&self) -> NodeRef<'_> {
70 NodeHandle(0).with_cst(self)
71 }
72
73 pub fn token_at(&self, pos: Pos, ot: &OriginTable) -> Option<NodeRef<'_>> {
74 let check_span = |span: &Span| {
76 if span.origin != pos.origin {
77 return false;
78 }
79
80 let start = span.start().text_pos;
81 let end = span.end(ot).text_pos;
82 let pos = pos.text_pos;
83
84 start <= pos && pos <= end
85 };
86
87 self.nodes
88 .iter()
89 .enumerate()
90 .filter_map(|(i, n)| n.token().map(|t| (i, t)))
91 .find(|(_, t)| check_span(&t.span))
92 .map(|(i, _)| NodeHandle(i as u32).with_cst(self))
93 }
94
95 pub fn node_at(&self, pos: Pos, ot: &OriginTable) -> NodeRef<'_> {
96 let check_span = |span: &Span| {
98 if span.origin != pos.origin {
99 return false;
100 }
101
102 let end = span.end(ot).text_pos;
103 let pos = pos.text_pos;
104
105 pos <= end
106 };
107
108 self.nodes
109 .iter()
110 .enumerate()
111 .filter_map(|(i, n)| n.token().map(|t| (i, t)))
112 .find(|(_, t)| check_span(&t.span))
113 .map(|(i, _)| NodeHandle(i as u32).with_cst(self))
114 .unwrap_or(self.root())
115 }
116
117 pub fn get(&self, node: NodeHandle) -> &Node {
118 &self.nodes[node.0 as usize]
119 }
120
121 pub fn parent(&self, child: NodeHandle) -> Option<NodeHandle> {
122 let mut node = child.0;
123
124 loop {
125 if node == 0 {
126 return None;
127 }
128
129 node -= 1;
130
131 match self.nodes[node as usize] {
132 Node::Token(_) => (),
133 Node::Open { .. } => return Some(NodeHandle(node)),
134 Node::Close { open, .. } => node = open.0,
135 }
136 }
137 }
138
139 pub fn prev_sibling(&self, node: NodeHandle) -> Option<NodeHandle> {
140 let node = node.0;
141
142 if node == 0 {
143 return None;
144 }
145
146 let node = node - 1;
147
148 match &self.nodes[node as usize] {
149 Node::Token(_) => Some(NodeHandle(node)),
150 Node::Open { .. } => None,
151 Node::Close { open, .. } => Some(*open),
152 }
153 }
154
155 pub fn next_sibling(&self, node: NodeHandle) -> Option<NodeHandle> {
156 let node = node.0;
157
158 if node as usize == self.nodes.len() - 1 {
159 return None;
160 }
161
162 let node = node + 1;
163
164 match &self.nodes[node as usize] {
165 Node::Token(_) | Node::Open { .. } => Some(NodeHandle(node)),
166 Node::Close { .. } => None,
167 }
168 }
169
170 pub fn children(&self, parent: NodeHandle) -> Children<'_> {
171 match self.get(parent) {
172 Node::Token(_) | Node::Close { .. } => Children {
173 cst: self,
174 next: NodeHandle(self.nodes.len() as u32 - 1),
175 },
176 Node::Open { .. } => Children {
177 cst: self,
178 next: NodeHandle(parent.0 + 1),
179 },
180 }
181 }
182
183 pub fn is_ancestor(&self, ancestor: NodeHandle, descendant: NodeHandle) -> bool {
184 let Node::Open { close, .. } = self.get(ancestor) else {
185 return false;
186 };
187 let start = ancestor.0;
188 let end = close.0;
189 let descendant = descendant.0;
190
191 start <= descendant && descendant <= end
192 }
193
194 pub fn span(&self, node: NodeHandle, ot: &OriginTable) -> Option<Span> {
195 let range = match self.get(node) {
196 Node::Token(token) => return Some(token.span),
197 Node::Open { close, .. } => node.0 as usize..close.0 as usize,
198 Node::Close { open, .. } => open.0 as usize..node.0 as usize,
199 };
200
201 let start = &self.nodes[range.clone()].iter().find_map(Node::token)?;
202 let end = &self.nodes[range].iter().rev().find_map(Node::token)?;
203
204 let start = start.span.start();
205 let end = end.span.end(ot);
206 let span = Span::cross_origin(start, end, ot);
207
208 Some(span)
209 }
210
211 pub fn iter(&self) -> impl Iterator<Item = NodeRef<'_>> {
212 self.nodes
213 .iter()
214 .enumerate()
215 .map(|(idx, _)| NodeHandle(idx as u32).with_cst(self))
216 }
217
218 pub fn tokens(&self) -> impl Iterator<Item = (NodeHandle, &Token)> {
219 self.iter()
220 .filter_map(|n| Some((n.handle(), self.get(n.handle()).token()?)))
221 }
222}
223
224impl Display for Cst {
225 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
226 let mut depth = 0;
227
228 fn print_indent(f: &mut std::fmt::Formatter<'_>, depth: usize) -> std::fmt::Result {
229 for _ in 0..depth {
230 write!(f, " ")?
231 }
232 Ok(())
233 }
234
235 for node in &self.nodes {
236 match node {
237 Node::Token(token) => {
238 print_indent(f, depth)?;
239 writeln!(f, "{} {:?}", token.kind, token.span)?;
240 }
241 Node::Open { kind, .. } => {
242 print_indent(f, depth)?;
243 writeln!(f, "{kind:?}:")?;
244 depth += 1;
245 }
246 Node::Close { .. } => {
247 depth -= 1;
248 }
249 }
250 }
251
252 Ok(())
253 }
254}
255
256#[derive(Debug, Clone, Copy)]
257pub struct Children<'cst> {
258 cst: &'cst Cst,
259 next: NodeHandle,
260}
261
262impl<'cst> Children<'cst> {
263 pub fn find_kind(mut self, kind: TreeKind) -> Option<NodeRef<'cst>> {
264 self.find(|n| n.is_open(kind))
265 }
266
267 pub fn find_kind_where(mut self, pred: impl Fn(TreeKind) -> bool) -> Option<NodeRef<'cst>> {
268 self.find(|n| n.is_open_where(&pred))
269 }
270
271 pub fn find_tt(mut self, tt: impl Into<Tt>) -> Option<NodeRef<'cst>> {
272 let tt = tt.into();
273 self.find(|n| n.is_tt(tt))
274 }
275}
276
277impl<'cst> Iterator for Children<'cst> {
278 type Item = NodeRef<'cst>;
279
280 fn next(&mut self) -> Option<Self::Item> {
281 let node = self.next;
282 match self.cst.get(node) {
283 Node::Token(_) => {
284 self.next = NodeHandle(node.0 + 1);
285 Some(node.with_cst(self.cst))
286 }
287 Node::Open { close, .. } => {
288 self.next = NodeHandle(close.0 + 1);
289 Some(node.with_cst(self.cst))
290 }
291 Node::Close { .. } => None,
292 }
293 }
294}
295
296#[derive(Clone, Copy)]
297pub struct NodeRef<'cst> {
298 node: NodeHandle,
299 cst: &'cst Cst,
300}
301
302impl Debug for NodeRef<'_> {
303 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
304 f.debug_struct("NodeRef").field("node", &self.node).finish()
305 }
306}
307
308impl PartialEq for NodeRef<'_> {
309 fn eq(&self, other: &Self) -> bool {
310 self.node == other.node
311 }
312}
313
314impl Eq for NodeRef<'_> {}
315
316impl Deref for NodeRef<'_> {
317 type Target = Node;
318
319 fn deref(&self) -> &Self::Target {
320 self.cst.get(self.node)
321 }
322}
323
324impl<'cst> NodeRef<'cst> {
325 pub fn cst(&self) -> &Cst {
326 self.cst
327 }
328
329 pub fn handle(&self) -> NodeHandle {
330 self.node
331 }
332
333 pub fn project(self, f: impl FnOnce(&Cst, NodeHandle) -> NodeHandle) -> Self {
334 NodeRef {
335 node: (f)(self.cst, self.node),
336 cst: self.cst,
337 }
338 }
339
340 pub fn try_project(
341 self,
342 f: impl FnOnce(&Cst, NodeHandle) -> Option<NodeHandle>,
343 ) -> Option<Self> {
344 Some(NodeRef {
345 node: (f)(self.cst, self.node)?,
346 cst: self.cst,
347 })
348 }
349
350 pub fn parent(self) -> Option<Self> {
351 self.try_project(Cst::parent)
352 }
353
354 pub fn next_sibling(self) -> Option<Self> {
355 self.try_project(Cst::next_sibling)
356 }
357
358 pub fn prev_sibling(self) -> Option<Self> {
359 self.try_project(Cst::prev_sibling)
360 }
361
362 pub fn children(self) -> Children<'cst> {
363 self.cst.children(self.node)
364 }
365
366 pub fn span(self, ot: &OriginTable) -> Option<Span> {
367 self.cst.span(self.node, ot)
368 }
369
370 pub fn contains(self, inner: NodeHandle) -> bool {
371 self.cst.is_ancestor(self.handle(), inner)
372 }
373}