Skip to main content

blues_lsp/syntax/cst/
tree.rs

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    // TODO: make this not pub?
61    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        // TODO: naive linear impl
75        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        // TODO: naive linear impl
97        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}