Skip to main content

dry4rust/
node.rs

1// Copyright (c) 2026 Matjaz Domen Pecan
2// Copyright 2026 Umberto Gotti <umberto.gotti@umbertogotti.dev>
3// Licensed under the MIT License
4// SPDX-License-Identifier: MIT
5
6use std::collections::HashMap;
7
8use crate::placeholder_order_collector::PlaceholderOrderCollector;
9
10/// Kinds of literals — preserves type but erases value.
11#[derive(Debug, Clone, PartialEq, Eq, Hash)]
12pub enum LiteralKind {
13    Int,
14    Float,
15    Str,
16    ByteStr,
17    CStr,
18    Byte,
19    Char,
20    Bool,
21    Null,
22}
23
24/// Kinds of placeholders — what the original identifier referred to.
25#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)]
26pub enum PlaceholderKind {
27    Variable,
28    Function,
29    Type,
30    Lifetime,
31    Label,
32}
33
34/// Binary operators.
35#[derive(Debug, Clone, PartialEq, Eq, Hash)]
36pub enum BinOpKind {
37    Add,
38    Sub,
39    Mul,
40    Div,
41    Rem,
42    And,
43    Or,
44    BitXor,
45    BitAnd,
46    BitOr,
47    Shl,
48    Shr,
49    Eq,
50    Lt,
51    Le,
52    Ne,
53    Ge,
54    Gt,
55    AddAssign,
56    SubAssign,
57    MulAssign,
58    DivAssign,
59    RemAssign,
60    BitXorAssign,
61    BitAndAssign,
62    BitOrAssign,
63    ShlAssign,
64    ShrAssign,
65    FloorDiv,
66    Pow,
67    In,
68    NotIn,
69    Is,
70    IsNot,
71    FloorDivAssign,
72    PowAssign,
73    Other,
74}
75
76/// Unary operators.
77#[derive(Debug, Clone, PartialEq, Eq, Hash)]
78pub enum UnOpKind {
79    Deref,
80    Not,
81    Neg,
82    Other,
83}
84
85/// The kind of a normalized AST node. Carries only non-child data
86/// (operator kinds, literal kinds, placeholder indices, mutability flags, macro names).
87#[derive(Debug, Clone, PartialEq, Eq, Hash)]
88pub enum NodeKind {
89    // Blocks and statements
90    Block,
91    LetBinding,
92    Semi,
93    Paren,
94
95    // Literals and identifiers
96    Literal(LiteralKind),
97    Placeholder(PlaceholderKind, usize),
98
99    // Operations
100    BinaryOp(BinOpKind),
101    UnaryOp(UnOpKind),
102    Range,
103
104    // Calls and access
105    Call,
106    MethodCall,
107    FieldAccess,
108    Index,
109    Path,
110
111    // Closures and functions
112    Closure,
113    FnSignature,
114
115    // Control flow
116    Return,
117    Break,
118    Continue,
119    Assign,
120
121    // References and pointers
122    Reference {
123        mutable: bool,
124    },
125
126    // Compound types
127    Tuple,
128    Array,
129    Set,
130    Repeat,
131
132    // Type operations
133    Cast,
134    StructInit,
135
136    // Async/error
137    Await,
138    Yield,
139    Try,
140
141    // Control flow structures
142    If,
143    Match,
144    MatchArm,
145    Loop,
146    While,
147    ForLoop,
148    LetExpr,
149
150    // Patterns
151    PatWild,
152    PatPlaceholder(PlaceholderKind, usize),
153    PatTuple,
154    PatStruct,
155    PatOr,
156    PatLiteral,
157    PatReference {
158        mutable: bool,
159    },
160    PatSlice,
161    PatRest,
162    PatRange,
163
164    // Types
165    TypePlaceholder(PlaceholderKind, usize),
166    TypeReference {
167        mutable: bool,
168    },
169    TypeTuple,
170    TypeSlice,
171    TypeArray,
172    TypePath,
173    TypeImplTrait,
174    TypeInfer,
175    TypeUnit,
176    TypeNever,
177
178    // Field initializer (name = value)
179    FieldValue,
180
181    // Macro invocations
182    MacroCall {
183        name: String,
184    },
185
186    // Opaque — unsupported constructs
187    Opaque,
188
189    /// Sentinel for absent optional children, ensuring fixed child positions
190    /// for correct zip alignment in similarity comparison.
191    None,
192}
193
194/// A normalized AST node. Uses a data-driven `{ kind, children }` representation
195/// instead of a large enum with differently-shaped variants. This allows generic
196/// traversal algorithms (count_nodes, reindex, count_matching, extract) to work
197/// without exhaustive matching on every variant.
198///
199/// ## Child ordering conventions
200///
201/// - **Fixed with None sentinels** (always same child count):
202///   - `If` -> [condition, then_branch, else_or_None]
203///   - `LetBinding` -> [pattern, type_or_None, init_or_None, diverge_or_None]
204///   - `Range` / `PatRange` -> [from_or_None, to_or_None]
205///   - `MatchArm` -> [pattern, guard_or_None, body]
206/// - **Fixed children first, variable after** (for zip alignment):
207///   - `Call` -> [func, arg0, arg1, ...]
208///   - `MethodCall` -> [receiver, method, arg0, ...]
209///   - `Closure` -> [body, param0, ...]
210///   - `FnSignature` -> [return_type_or_None, param0, ...]
211///   - `Match` -> [expr, arm0, arm1, ...]
212///   - `StructInit` -> [rest_or_None, field0, field1, ...]
213///   - `MacroCall` -> [arg0, arg1, ...]
214/// - **Variable-length (0 or 1)**: `Return`, `Break` -> `[]` or `[value]`
215/// - **Homogeneous**: `Block`, `Tuple`, `Array`, `Path`, `PatTuple`, etc. -> [elem0, ...]
216/// - **All other fixed**: e.g. `BinaryOp` -> [left, right], `ForLoop` -> [pat, iter, body]
217#[derive(Debug, Clone, PartialEq, Eq, Hash)]
218pub struct NormalizedNode {
219    pub kind: NodeKind,
220    pub children: Vec<Self>,
221}
222
223impl NormalizedNode {
224    /// Create a leaf node (no children).
225    #[must_use]
226    pub const fn leaf(kind: NodeKind) -> Self {
227        Self {
228            kind,
229            children: vec![],
230        }
231    }
232
233    /// Create a node with children.
234    #[must_use]
235    pub const fn with_children(kind: NodeKind, children: Vec<Self>) -> Self {
236        Self { kind, children }
237    }
238
239    /// Create a None sentinel node.
240    #[must_use]
241    pub const fn none() -> Self {
242        Self::leaf(NodeKind::None)
243    }
244
245    /// Convert an `Option<NormalizedNode>` to a node, using the [`Self::none`] sentinel for
246    /// absent values.
247    pub fn opt(node: Option<Self>) -> Self {
248        node.unwrap_or_else(Self::none)
249    }
250
251    /// Check if this is a None sentinel node.
252    #[must_use]
253    pub const fn is_none(&self) -> bool {
254        matches!(self.kind, NodeKind::None)
255    }
256}
257
258// -- Placeholder re-indexing --------------------------------------------------
259
260/// Applies the reindex mapping to a node, returning a new node with remapped indices.
261fn apply_reindex(
262    node: &NormalizedNode,
263    mapping: &HashMap<(PlaceholderKind, usize), usize>,
264) -> NormalizedNode {
265    let kind = match &node.kind {
266        NodeKind::Placeholder(kind, idx) => {
267            let new_idx = mapping.get(&(*kind, *idx)).copied().unwrap_or(*idx);
268            NodeKind::Placeholder(*kind, new_idx)
269        }
270        NodeKind::PatPlaceholder(kind, idx) => {
271            let new_idx = mapping.get(&(*kind, *idx)).copied().unwrap_or(*idx);
272            NodeKind::PatPlaceholder(*kind, new_idx)
273        }
274        NodeKind::TypePlaceholder(kind, idx) => {
275            let new_idx = mapping.get(&(*kind, *idx)).copied().unwrap_or(*idx);
276            NodeKind::TypePlaceholder(*kind, new_idx)
277        }
278        other => other.clone(),
279    };
280    let children = node
281        .children
282        .iter()
283        .map(|c| apply_reindex(c, mapping))
284        .collect();
285    NormalizedNode { kind, children }
286}
287
288/// Re-index all placeholders in a sub-tree so that indices start from 0
289/// per kind, assigned by first-occurrence depth-first order.
290/// This allows comparing sub-trees extracted from different function contexts.
291#[must_use]
292pub fn reindex_placeholders(node: &NormalizedNode) -> NormalizedNode {
293    let mut collector = PlaceholderOrderCollector::new();
294    collector.collect(node);
295    let order = collector.into_order();
296
297    // Build mapping: (kind, old_index) -> new sequential index per kind
298    let mut counters: HashMap<PlaceholderKind, usize> = HashMap::new();
299    let mut mapping: HashMap<(PlaceholderKind, usize), usize> = HashMap::new();
300    for (kind, old_idx) in order {
301        let counter = counters.entry(kind).or_insert(0);
302        mapping.insert((kind, old_idx), *counter);
303        *counter += 1;
304    }
305
306    apply_reindex(node, &mapping)
307}
308
309/// Count the number of nodes in a normalized tree.
310/// None sentinel nodes are not counted.
311pub fn count_nodes(node: &NormalizedNode) -> usize {
312    if node.is_none() {
313        return 0;
314    }
315    1 + node.children.iter().map(count_nodes).sum::<usize>()
316}