Skip to main content

asm_rs/
preprocessor.rs

1//! Preprocessor for assembly source text.
2//!
3//! Handles macro definitions (`.macro`/`.endm`), repeat loops (`.rept`, `.irp`,
4//! `.irpc`), and conditional assembly (`.if`/`.ifdef`/`.ifndef`/`.else`/`.endif`)
5//! before the source reaches the parser.
6//!
7//! The preprocessor operates on raw text, expanding directives in-place so that
8//! the downstream lexer and parser see only ordinary assembly statements.
9
10use alloc::borrow::Cow;
11use alloc::collections::BTreeMap;
12use alloc::format;
13use alloc::rc::Rc;
14use alloc::string::String;
15use alloc::vec::Vec;
16
17use crate::error::{AsmError, Span};
18
19/// Default maximum macro expansion recursion depth.
20///
21/// Every level of `.rept`/macro nesting costs a native stack frame, so this is
22/// deliberately shallow: a limit deep enough to overflow the stack would abort
23/// the process instead of returning [`AsmError::ResourceLimitExceeded`], and an
24/// abort cannot be caught by a host embedding the assembler.
25const DEFAULT_MAX_RECURSION_DEPTH: usize = 32;
26
27/// Maximum total iterations across all `.rept`/`.irp`/`.irpc` blocks.
28const DEFAULT_MAX_ITERATION_COUNT: usize = 100_000;
29
30/// Default ceiling on total bytes of expanded text.
31const DEFAULT_MAX_EXPANDED_BYTES: usize = 64 * 1024 * 1024;
32
33/// Append `body` to `out`, replacing every occurrence of `placeholder` with `value`.
34/// Avoids allocating an intermediate `String` — writes directly into `out`.
35fn replace_single_param(out: &mut String, body: &str, placeholder: &str, value: &str) {
36    let ph_bytes = placeholder.as_bytes();
37    let body_bytes = body.as_bytes();
38    let ph_len = ph_bytes.len();
39    let mut start = 0;
40    while start < body_bytes.len() {
41        if let Some(pos) = body[start..].find(placeholder) {
42            out.push_str(&body[start..start + pos]);
43            out.push_str(value);
44            start += pos + ph_len;
45        } else {
46            out.push_str(&body[start..]);
47            break;
48        }
49    }
50}
51
52/// A macro definition.
53///
54/// Stored behind an [`Rc`](alloc::rc::Rc) in [`Preprocessor::macros`] so that
55/// invoking a macro clones a refcount rather than its whole body text.
56#[derive(Debug, Clone)]
57struct MacroDef {
58    /// Parameter names (without leading `\`).
59    params: Vec<MacroParam>,
60    /// The raw body text (lines between `.macro` and `.endm`).
61    body: String,
62}
63
64/// A macro parameter with optional default value.
65#[derive(Debug, Clone)]
66struct MacroParam {
67    name: String,
68    default: Option<String>,
69    is_vararg: bool,
70}
71
72/// Preprocessor state.
73#[derive(Debug)]
74pub struct Preprocessor {
75    /// Defined macros: name → definition.
76    macros: BTreeMap<String, Rc<MacroDef>>,
77    /// Defined symbols for `.ifdef`/`.ifndef` (name → value).
78    symbols: BTreeMap<String, i128>,
79    /// Counter for `\@` unique label generation.
80    expansion_counter: usize,
81    /// Current recursion depth for macro expansion.
82    recursion_depth: usize,
83    /// Maximum recursion depth (configurable).
84    max_recursion_depth: usize,
85    /// Maximum total iteration count (configurable).
86    max_iteration_count: usize,
87    /// Total iteration count across all loops (bounds check).
88    iteration_count: usize,
89    /// Maximum total bytes of expanded text (configurable).
90    max_expanded_bytes: usize,
91    /// Bytes of expanded text produced so far in this `process()` call.
92    expanded_bytes: usize,
93}
94
95impl Preprocessor {
96    /// Create a new preprocessor.
97    pub fn new() -> Self {
98        Self {
99            macros: BTreeMap::new(),
100            symbols: BTreeMap::new(),
101            expansion_counter: 0,
102            recursion_depth: 0,
103            max_recursion_depth: DEFAULT_MAX_RECURSION_DEPTH,
104            max_iteration_count: DEFAULT_MAX_ITERATION_COUNT,
105            iteration_count: 0,
106            max_expanded_bytes: DEFAULT_MAX_EXPANDED_BYTES,
107            expanded_bytes: 0,
108        }
109    }
110
111    /// Set the maximum total size, in bytes, of expanded preprocessor output.
112    pub fn set_max_expanded_bytes(&mut self, bytes: usize) {
113        self.max_expanded_bytes = bytes;
114    }
115
116    /// Charge `n` bytes against the expansion budget.
117    ///
118    /// Iteration and recursion counters bound how *often* the preprocessor
119    /// loops, not how much text each pass emits, so expansion size is metered
120    /// separately.
121    fn charge_expansion(&mut self, n: usize) -> Result<(), AsmError> {
122        self.expanded_bytes = self.expanded_bytes.saturating_add(n);
123        if self.expanded_bytes > self.max_expanded_bytes {
124            return Err(AsmError::ResourceLimitExceeded {
125                resource: String::from("preprocessor expanded bytes"),
126                limit: self.max_expanded_bytes,
127            });
128        }
129        Ok(())
130    }
131
132    /// Set the maximum recursion depth for macro expansion.
133    pub fn set_max_recursion_depth(&mut self, depth: usize) {
134        self.max_recursion_depth = depth;
135    }
136
137    /// Set the maximum total iteration count for `.rept`/`.irp`/`.irpc`.
138    pub fn set_max_iterations(&mut self, count: usize) {
139        self.max_iteration_count = count;
140    }
141
142    /// Define a symbol for conditional assembly.
143    pub fn define_symbol(&mut self, name: &str, value: i128) {
144        self.symbols.insert(String::from(name), value);
145    }
146
147    /// Process source text, expanding all preprocessor directives.
148    ///
149    /// Returns the expanded source text ready for lexing/parsing.
150    /// When no preprocessor directives are present and no macros are defined,
151    /// returns a borrowed reference to the original source (zero allocation).
152    ///
153    /// # Errors
154    ///
155    /// Returns `AsmError` on malformed directives, recursion limit, or
156    /// iteration limit exceeded.
157    pub fn process<'a>(&mut self, source: &'a str) -> Result<Cow<'a, str>, AsmError> {
158        // Reset iteration count per process() call so long-lived assemblers
159        // don't accumulate towards the limit across multiple emit() calls.
160        self.iteration_count = 0;
161        self.expanded_bytes = 0;
162        if !self.needs_expansion(source) {
163            return Ok(Cow::Borrowed(source));
164        }
165        self.expand_text(source).map(Cow::Owned)
166    }
167
168    /// Check whether the source text requires preprocessing.
169    ///
170    /// Returns `false` when no macros are defined and the source contains
171    /// no preprocessor directives, meaning the text can be passed straight
172    /// to the lexer without any transformation.
173    fn needs_expansion(&self, source: &str) -> bool {
174        // If macros are defined, any line could be an invocation.
175        if !self.macros.is_empty() {
176            return true;
177        }
178        // If symbols are defined, we still don't need expansion unless
179        // the source actually references them via .ifdef/.ifndef.
180        // Scan for directive prefixes.  We look for lines whose first
181        // non-whitespace content matches a preprocessor directive.
182        // This is cheaper than the full expansion: no allocation, no
183        // string building — just a linear scan of the source bytes.
184        //
185        // The scan runs over every line of every `emit()`, so it starts with
186        // the cheapest possible discriminator: every directive it looks for
187        // begins with '.', and almost no line in real assembly does. That one
188        // byte comparison skips the ~20 string comparisons below for the
189        // overwhelming majority of lines.
190        if !source.as_bytes().contains(&b'.') {
191            return false;
192        }
193        for line in source.lines() {
194            let trimmed = line.trim_start();
195            if !trimmed.starts_with('.') {
196                continue;
197            }
198            let trimmed = trimmed.trim_end();
199            if trimmed.starts_with(".macro ")
200                || trimmed.starts_with(".macro\t")
201                || trimmed.starts_with(".rept ")
202                || trimmed.starts_with(".rept\t")
203                || trimmed.starts_with(".irp ")
204                || trimmed.starts_with(".irp\t")
205                || trimmed.starts_with(".irpc ")
206                || trimmed.starts_with(".irpc\t")
207                || trimmed.starts_with(".if ")
208                || trimmed.starts_with(".if\t")
209                || trimmed == ".if"
210                || trimmed.starts_with(".ifdef ")
211                || trimmed.starts_with(".ifdef\t")
212                || trimmed.starts_with(".ifndef ")
213                || trimmed.starts_with(".ifndef\t")
214                || trimmed == ".exitm"
215            {
216                return true;
217            }
218        }
219        false
220    }
221
222    /// Core expansion loop — processes one level of text.
223    fn expand_text(&mut self, source: &str) -> Result<String, AsmError> {
224        self.recursion_depth += 1;
225        if self.recursion_depth > self.max_recursion_depth {
226            self.recursion_depth -= 1;
227            return Err(AsmError::ResourceLimitExceeded {
228                resource: String::from("macro recursion depth"),
229                limit: self.max_recursion_depth,
230            });
231        }
232
233        let lines: Vec<&str> = source.lines().collect();
234        let mut output = String::new();
235        let mut i = 0;
236
237        let result = self.expand_text_inner(&lines, &mut output, &mut i);
238        self.recursion_depth -= 1;
239        result?;
240        Ok(output)
241    }
242
243    /// Inner expansion logic (separated to ensure recursion_depth cleanup).
244    fn expand_text_inner(
245        &mut self,
246        lines: &[&str],
247        output: &mut String,
248        i: &mut usize,
249    ) -> Result<(), AsmError> {
250        while *i < lines.len() {
251            let line = lines[*i];
252            let trimmed = line.trim();
253
254            // --- Macro definition ---
255            if trimmed.starts_with(".macro ") || trimmed.starts_with(".macro\t") {
256                let (macro_def, end_idx) = self.parse_macro_def(lines, *i)?;
257                let name = parse_macro_name(trimmed, *i)?;
258                self.macros.insert(name, Rc::new(macro_def));
259                *i = end_idx + 1;
260                continue;
261            }
262
263            // --- .rept ---
264            if trimmed.starts_with(".rept ") || trimmed.starts_with(".rept\t") {
265                let (body, end_idx) = collect_block(lines, *i, ".rept", ".endr")?;
266                let count = parse_rept_count(trimmed, *i)?;
267                let expanded = self.expand_rept(count, &body)?;
268                output.push_str(&expanded);
269                *i = end_idx + 1;
270                continue;
271            }
272
273            // --- .irp ---
274            if trimmed.starts_with(".irp ") || trimmed.starts_with(".irp\t") {
275                let (body, end_idx) = collect_block(lines, *i, ".irp", ".endr")?;
276                let (sym, values) = parse_irp_args(trimmed, *i)?;
277                let expanded = self.expand_irp(&sym, &values, &body)?;
278                output.push_str(&expanded);
279                *i = end_idx + 1;
280                continue;
281            }
282
283            // --- .irpc ---
284            if trimmed.starts_with(".irpc ") || trimmed.starts_with(".irpc\t") {
285                let (body, end_idx) = collect_block(lines, *i, ".irpc", ".endr")?;
286                let (sym, chars) = parse_irpc_args(trimmed, *i)?;
287                let expanded = self.expand_irpc(&sym, &chars, &body)?;
288                output.push_str(&expanded);
289                *i = end_idx + 1;
290                continue;
291            }
292
293            // --- Conditional assembly ---
294            if trimmed.starts_with(".if ")
295                || trimmed.starts_with(".if\t")
296                || trimmed == ".if"
297                || trimmed.starts_with(".ifdef ")
298                || trimmed.starts_with(".ifdef\t")
299                || trimmed.starts_with(".ifndef ")
300                || trimmed.starts_with(".ifndef\t")
301            {
302                let (selected_body, end_idx) = self.process_conditional(lines, *i)?;
303                if !selected_body.is_empty() {
304                    let expanded = self.expand_text(&selected_body)?;
305                    output.push_str(&expanded);
306                }
307                *i = end_idx + 1;
308                continue;
309            }
310
311            // --- .exitm (only meaningful inside macro expansion) ---
312            if trimmed == ".exitm" {
313                if self.recursion_depth <= 1 {
314                    // At the top level, .exitm is meaningless — warn the user.
315                    return Err(AsmError::Syntax {
316                        msg: String::from(".exitm outside of macro expansion"),
317                        span: crate::error::Span::new((*i + 1) as u32, 1, 0, trimmed.len()),
318                    });
319                }
320                // Inside macro expansion, stop expanding this level
321                break;
322            }
323
324            // --- Macro invocation ---
325            if let Some(expanded) = self.try_expand_macro(trimmed)? {
326                // Mutually-invoking macros grow the text exponentially with
327                // depth, so meter each expansion against the byte budget
328                // before recursing into it.
329                self.charge_expansion(expanded.len())?;
330                // Recursively expand the result
331                let re_expanded = self.expand_text(&expanded)?;
332                output.push_str(&re_expanded);
333                *i += 1;
334                continue;
335            }
336
337            // --- .equ / .set / NAME = expr: track symbols for .ifdef ---
338            if let Some((name, val)) = try_parse_symbol_def(trimmed) {
339                self.symbols.insert(name, val);
340            }
341
342            // Ordinary line — pass through
343            output.push_str(line);
344            output.push('\n');
345            *i += 1;
346        }
347
348        Ok(())
349    }
350
351    /// Parse a `.macro name [params...]` ... `.endm` definition.
352    fn parse_macro_def(&self, lines: &[&str], start: usize) -> Result<(MacroDef, usize), AsmError> {
353        let header = lines[start].trim();
354        let params = parse_macro_params(header)?;
355
356        let mut body_lines = Vec::new();
357        let mut depth = 1usize;
358        let mut i = start + 1;
359
360        while i < lines.len() {
361            let trimmed = lines[i].trim();
362            if trimmed.starts_with(".macro ") || trimmed.starts_with(".macro\t") {
363                depth += 1;
364            } else if trimmed == ".endm" {
365                depth -= 1;
366                if depth == 0 {
367                    let body = body_lines.join("\n");
368                    return Ok((MacroDef { params, body }, i));
369                }
370            }
371            body_lines.push(lines[i]);
372            i += 1;
373        }
374
375        Err(AsmError::Syntax {
376            msg: String::from("unterminated .macro (missing .endm)"),
377            span: line_span(start),
378        })
379    }
380
381    /// Try to expand a line as a macro invocation. Returns `None` if no macro matches.
382    fn try_expand_macro(&mut self, line: &str) -> Result<Option<String>, AsmError> {
383        let trimmed = line.trim();
384        if trimmed.is_empty() || trimmed.starts_with('#') || trimmed.starts_with('.') {
385            return Ok(None);
386        }
387
388        // Extract first word as potential macro name
389        let first_word = trimmed.split_whitespace().next().unwrap_or("");
390
391        // Also check if it ends with ':' (label definition) — skip
392        if first_word.ends_with(':') {
393            // Could be `label: macro_name args` — check remainder
394            let rest = trimmed[first_word.len()..].trim();
395            if rest.is_empty() {
396                return Ok(None);
397            }
398            let macro_name = rest.split_whitespace().next().unwrap_or("");
399            if let Some(def) = self.macros.get(macro_name).cloned() {
400                let args_str = rest[macro_name.len()..].trim();
401                let args = parse_macro_args(args_str);
402                let expanded = self.substitute_macro(&def, &args);
403                // Preserve the label
404                return Ok(Some(format!("{}\n{}", first_word, expanded)));
405            }
406            return Ok(None);
407        }
408
409        if let Some(def) = self.macros.get(first_word).cloned() {
410            let args_str = trimmed[first_word.len()..].trim();
411            let args = parse_macro_args(args_str);
412            let expanded = self.substitute_macro(&def, &args);
413            return Ok(Some(expanded));
414        }
415
416        Ok(None)
417    }
418
419    /// Substitute macro parameters and `\@` counter into body text.
420    ///
421    /// Uses a single-pass scan: walks the body once, and at each `\` checks
422    /// for parameter names or `@`.  This is O(M × log N) where M = body length
423    /// and N = parameter count, versus the prior O(N × M) multi-pass approach.
424    fn substitute_macro(&mut self, def: &MacroDef, args: &[String]) -> String {
425        let counter = self.expansion_counter;
426        self.expansion_counter += 1;
427
428        // Pre-compute replacement strings for each parameter
429        let replacements: Vec<(&str, String)> = def
430            .params
431            .iter()
432            .enumerate()
433            .map(|(idx, param)| {
434                let value = if param.is_vararg {
435                    if idx < args.len() {
436                        args[idx..].join(", ")
437                    } else {
438                        param.default.clone().unwrap_or_default()
439                    }
440                } else if idx < args.len() {
441                    args[idx].clone()
442                } else {
443                    param.default.clone().unwrap_or_default()
444                };
445                (param.name.as_str(), value)
446            })
447            .collect();
448
449        let body = &def.body;
450        let mut result = String::with_capacity(body.len());
451        let bytes = body.as_bytes();
452        let len = bytes.len();
453        let mut i = 0;
454
455        while i < len {
456            if bytes[i] == b'\\' && i + 1 < len {
457                // Check for \@ (unique counter)
458                if bytes[i + 1] == b'@' {
459                    use core::fmt::Write;
460                    let _ = write!(result, "{}", counter);
461                    i += 2;
462                    continue;
463                }
464                // Check for \param_name
465                let rest = &body[i + 1..];
466                let mut matched = false;
467                for &(name, ref value) in &replacements {
468                    if rest.starts_with(name) {
469                        // Ensure we match the full token — the char after the
470                        // name must NOT be alphanumeric or '_' (otherwise
471                        // \foo would partially match \foobar).
472                        let end = name.len();
473                        let boundary = end >= rest.len()
474                            || !rest.as_bytes()[end].is_ascii_alphanumeric()
475                                && rest.as_bytes()[end] != b'_';
476                        if boundary {
477                            result.push_str(value);
478                            i += 1 + name.len();
479                            matched = true;
480                            break;
481                        }
482                    }
483                }
484                if !matched {
485                    result.push('\\');
486                    i += 1;
487                }
488            } else {
489                // Copy one character (handles multi-byte UTF-8)
490                let ch = body[i..].chars().next().unwrap_or('\0');
491                result.push(ch);
492                i += ch.len_utf8();
493            }
494        }
495
496        result
497    }
498
499    /// Expand `.rept count` block.
500    fn expand_rept(&mut self, count: usize, body: &str) -> Result<String, AsmError> {
501        let mut raw = String::new();
502        for _ in 0..count {
503            self.iteration_count += 1;
504            if self.iteration_count > self.max_iteration_count {
505                return Err(AsmError::ResourceLimitExceeded {
506                    resource: String::from("preprocessor iterations"),
507                    limit: self.max_iteration_count,
508                });
509            }
510            self.charge_expansion(body.len() + 1)?;
511            raw.push_str(body);
512            raw.push('\n');
513        }
514        // Re-expand to handle nested .rept/.irp/.irpc/macros
515        self.expand_text(&raw)
516    }
517
518    /// Expand `.irp sym, val1, val2, ...` block.
519    fn expand_irp(&mut self, sym: &str, values: &[String], body: &str) -> Result<String, AsmError> {
520        let placeholder = format!("\\{}", sym);
521        let mut raw = String::new();
522        for val in values {
523            self.iteration_count += 1;
524            if self.iteration_count > self.max_iteration_count {
525                return Err(AsmError::ResourceLimitExceeded {
526                    resource: String::from("preprocessor iterations"),
527                    limit: self.max_iteration_count,
528                });
529            }
530            self.charge_expansion(body.len() + val.len() + 1)?;
531            replace_single_param(&mut raw, body, &placeholder, val);
532            raw.push('\n');
533        }
534        self.expand_text(&raw)
535    }
536
537    /// Expand `.irpc sym, string` block.
538    fn expand_irpc(&mut self, sym: &str, chars: &str, body: &str) -> Result<String, AsmError> {
539        let placeholder = format!("\\{}", sym);
540        let mut raw = String::new();
541        let mut ch_buf = [0u8; 4];
542        for ch in chars.chars() {
543            self.iteration_count += 1;
544            if self.iteration_count > self.max_iteration_count {
545                return Err(AsmError::ResourceLimitExceeded {
546                    resource: String::from("preprocessor iterations"),
547                    limit: self.max_iteration_count,
548                });
549            }
550            self.charge_expansion(body.len() + ch.len_utf8() + 1)?;
551            let ch_str = ch.encode_utf8(&mut ch_buf);
552            replace_single_param(&mut raw, body, &placeholder, ch_str);
553            raw.push('\n');
554        }
555        self.expand_text(&raw)
556    }
557
558    /// Process a conditional block (`.if`/`.ifdef`/`.ifndef`).
559    /// Returns the selected body text and the line index of `.endif`.
560    fn process_conditional(
561        &self,
562        lines: &[&str],
563        start: usize,
564    ) -> Result<(String, usize), AsmError> {
565        let header = lines[start].trim();
566
567        // Determine the initial condition result
568        let condition = evaluate_condition(header, &self.symbols, start)?;
569
570        let mut branches: Vec<(bool, Vec<&str>)> = Vec::new();
571        let mut current_cond = condition;
572        let mut current_body: Vec<&str> = Vec::new();
573        let mut depth = 1usize;
574        let mut i = start + 1;
575
576        while i < lines.len() {
577            let trimmed = lines[i].trim();
578
579            // Nested conditional
580            if trimmed.starts_with(".if ")
581                || trimmed.starts_with(".if\t")
582                || trimmed == ".if"
583                || trimmed.starts_with(".ifdef ")
584                || trimmed.starts_with(".ifdef\t")
585                || trimmed.starts_with(".ifndef ")
586                || trimmed.starts_with(".ifndef\t")
587            {
588                depth += 1;
589                current_body.push(lines[i]);
590                i += 1;
591                continue;
592            }
593
594            if trimmed == ".endif" {
595                depth -= 1;
596                if depth == 0 {
597                    branches.push((current_cond, current_body));
598                    // Select first true branch
599                    for (cond, body) in &branches {
600                        if *cond {
601                            return Ok((body.join("\n"), i));
602                        }
603                    }
604                    return Ok((String::new(), i));
605                }
606                current_body.push(lines[i]);
607                i += 1;
608                continue;
609            }
610
611            if depth == 1
612                && (trimmed == ".else"
613                    || trimmed.starts_with(".elseif ")
614                    || trimmed.starts_with(".elseif\t"))
615            {
616                branches.push((current_cond, core::mem::take(&mut current_body)));
617                if trimmed == ".else" {
618                    // .else is true if no prior branch was taken
619                    current_cond = !branches.iter().any(|(c, _)| *c);
620                } else {
621                    // .elseif expr
622                    let expr_str = trimmed.strip_prefix(".elseif").unwrap().trim();
623                    current_cond = if branches.iter().any(|(c, _)| *c) {
624                        false // A prior branch was already taken
625                    } else {
626                        eval_simple_expr(expr_str, &self.symbols) != 0
627                    };
628                }
629                i += 1;
630                continue;
631            }
632
633            current_body.push(lines[i]);
634            i += 1;
635        }
636
637        Err(AsmError::Syntax {
638            msg: String::from("unterminated conditional (missing .endif)"),
639            span: line_span(start),
640        })
641    }
642}
643
644impl Default for Preprocessor {
645    fn default() -> Self {
646        Self::new()
647    }
648}
649
650// --- Helper functions ---
651
652/// Parse the macro name from `.macro name ...`.
653fn parse_macro_name(header: &str, line: usize) -> Result<String, AsmError> {
654    let rest = header.strip_prefix(".macro").unwrap_or(header).trim_start();
655    let name = rest
656        .split(|c: char| c.is_whitespace() || c == ',')
657        .next()
658        .unwrap_or("");
659    if name.is_empty() {
660        return Err(AsmError::Syntax {
661            msg: String::from(".macro directive requires a name"),
662            span: Span::new((line + 1) as u32, 1, 0, header.len()),
663        });
664    }
665    Ok(String::from(name))
666}
667
668/// Parse macro parameters from `.macro name param1, param2=default, rest:vararg`.
669fn parse_macro_params(header: &str) -> Result<Vec<MacroParam>, AsmError> {
670    let rest = header.strip_prefix(".macro").unwrap_or(header).trim_start();
671
672    // Skip the macro name
673    let after_name = rest
674        .split_once(|c: char| c.is_whitespace() || c == ',')
675        .map(|(_, p)| p.trim_start_matches(',').trim())
676        .unwrap_or("");
677
678    if after_name.is_empty() {
679        return Ok(Vec::new());
680    }
681
682    let mut params = Vec::new();
683    for part in after_name.split(',') {
684        let part = part.trim();
685        if part.is_empty() {
686            continue;
687        }
688        if let Some((name, rest)) = part.split_once(':') {
689            let name = name.trim();
690            let rest = rest.trim();
691            if rest == "vararg" {
692                params.push(MacroParam {
693                    name: String::from(name),
694                    default: None,
695                    is_vararg: true,
696                });
697            } else {
698                params.push(MacroParam {
699                    name: String::from(part),
700                    default: None,
701                    is_vararg: false,
702                });
703            }
704        } else if let Some((name, default)) = part.split_once('=') {
705            params.push(MacroParam {
706                name: String::from(name.trim()),
707                default: Some(String::from(default.trim())),
708                is_vararg: false,
709            });
710        } else {
711            params.push(MacroParam {
712                name: String::from(part),
713                default: None,
714                is_vararg: false,
715            });
716        }
717    }
718    Ok(params)
719}
720
721/// Parse arguments passed to a macro invocation.
722fn parse_macro_args(args_str: &str) -> Vec<String> {
723    if args_str.is_empty() {
724        return Vec::new();
725    }
726    args_str
727        .split(',')
728        .map(|s| String::from(s.trim()))
729        .collect()
730}
731
732/// Parse `.rept count` header.
733fn parse_rept_count(header: &str, line: usize) -> Result<usize, AsmError> {
734    let rest = header.strip_prefix(".rept").unwrap_or(header).trim();
735    rest.parse::<usize>().map_err(|_| AsmError::Syntax {
736        msg: format!("invalid .rept count: '{}'", rest),
737        span: Span::new((line + 1) as u32, 1, 0, header.len()),
738    })
739}
740
741/// Parse `.irp sym, val1, val2, ...` header.
742fn parse_irp_args(header: &str, line: usize) -> Result<(String, Vec<String>), AsmError> {
743    let rest = header.strip_prefix(".irp").unwrap_or(header).trim();
744    let (sym, vals_str) = rest.split_once(',').ok_or_else(|| AsmError::Syntax {
745        msg: String::from(".irp requires a symbol and a comma-separated value list"),
746        span: Span::new((line + 1) as u32, 1, 0, header.len()),
747    })?;
748    let sym = sym.trim();
749    let values: Vec<String> = vals_str
750        .split(',')
751        .map(|s| String::from(s.trim()))
752        .filter(|s| !s.is_empty())
753        .collect();
754    Ok((String::from(sym), values))
755}
756
757/// Parse `.irpc sym, chars` header.
758fn parse_irpc_args(header: &str, line: usize) -> Result<(String, String), AsmError> {
759    let rest = header.strip_prefix(".irpc").unwrap_or(header).trim();
760    let (sym, chars) = rest.split_once(',').ok_or_else(|| AsmError::Syntax {
761        msg: String::from(".irpc requires a symbol and a string"),
762        span: Span::new((line + 1) as u32, 1, 0, header.len()),
763    })?;
764    Ok((String::from(sym.trim()), String::from(chars.trim())))
765}
766
767/// Collect lines of a block between `open_directive` and `close_directive`,
768/// handling nesting.
769fn collect_block(
770    lines: &[&str],
771    start: usize,
772    open_kw: &str,
773    close_kw: &str,
774) -> Result<(String, usize), AsmError> {
775    let mut depth = 1usize;
776    let mut body_lines = Vec::new();
777    let mut i = start + 1;
778
779    // All directives that share `.endr` as their terminator.
780    let endr_openers: &[&str] = &[".rept", ".irp", ".irpc"];
781
782    while i < lines.len() {
783        let trimmed = lines[i].trim();
784
785        // Check for nested open — if `.endr` is the terminator we must
786        // count *any* `.rept`/`.irp`/`.irpc` as nesting, not just the
787        // exact `open_kw`.
788        if close_kw == ".endr" {
789            for &opener in endr_openers {
790                if trimmed.starts_with(opener)
791                    && (trimmed.len() == opener.len()
792                        || trimmed.as_bytes().get(opener.len()) == Some(&b' ')
793                        || trimmed.as_bytes().get(opener.len()) == Some(&b'\t'))
794                {
795                    depth += 1;
796                    break;
797                }
798            }
799        } else if trimmed.starts_with(open_kw)
800            && (trimmed.len() == open_kw.len()
801                || trimmed.as_bytes().get(open_kw.len()) == Some(&b' ')
802                || trimmed.as_bytes().get(open_kw.len()) == Some(&b'\t'))
803        {
804            depth += 1;
805        }
806
807        if trimmed == close_kw {
808            depth -= 1;
809            if depth == 0 {
810                return Ok((body_lines.join("\n"), i));
811            }
812        }
813
814        body_lines.push(lines[i]);
815        i += 1;
816    }
817
818    Err(AsmError::Syntax {
819        msg: format!("unterminated {} (missing {})", open_kw, close_kw),
820        span: line_span(start),
821    })
822}
823
824/// Evaluate a conditional directive header.
825fn evaluate_condition(
826    header: &str,
827    symbols: &BTreeMap<String, i128>,
828    line: usize,
829) -> Result<bool, AsmError> {
830    let trimmed = header.trim();
831
832    if let Some(rest) = trimmed.strip_prefix(".ifdef") {
833        let name = rest.trim();
834        return Ok(symbols.contains_key(name));
835    }
836
837    if let Some(rest) = trimmed.strip_prefix(".ifndef") {
838        let name = rest.trim();
839        return Ok(!symbols.contains_key(name));
840    }
841
842    if let Some(rest) = trimmed.strip_prefix(".if") {
843        let expr = rest.trim();
844        return Ok(eval_simple_expr(expr, symbols) != 0);
845    }
846
847    Err(AsmError::Syntax {
848        msg: format!("unrecognized conditional directive: {}", trimmed),
849        span: Span::new((line + 1) as u32, 1, 0, header.len()),
850    })
851}
852
853/// Recursive-descent expression evaluator with proper C-like operator precedence.
854///
855/// Precedence (lowest → highest):
856///  1. `||`  logical OR
857///  2. `&&`  logical AND
858///  3. `|`   bitwise OR
859///  4. `^`   bitwise XOR
860///  5. `&`   bitwise AND
861///  6. `==` `!=`  equality
862///  7. `<` `>` `<=` `>=`  relational
863///  8. `<<` `>>`  shift
864///  9. `+` `-`  additive
865/// 10. `*` `/` `%`  multiplicative
866/// 11. `!` `-` `~`  unary prefix
867/// 12. literals, symbols, `defined()`, `(expr)`
868struct ExprEval<'a> {
869    src: &'a [u8],
870    pos: usize,
871    symbols: &'a BTreeMap<String, i128>,
872}
873
874impl<'a> ExprEval<'a> {
875    fn new(expr: &'a str, symbols: &'a BTreeMap<String, i128>) -> Self {
876        Self {
877            src: expr.as_bytes(),
878            pos: 0,
879            symbols,
880        }
881    }
882
883    fn eval(mut self) -> i128 {
884        self.skip_ws();
885        if self.pos >= self.src.len() {
886            return 0;
887        }
888        self.parse_logical_or()
889    }
890
891    fn skip_ws(&mut self) {
892        while self.pos < self.src.len() && self.src[self.pos].is_ascii_whitespace() {
893            self.pos += 1;
894        }
895    }
896
897    /// Try to consume a two-byte operator token. Returns `true` on match.
898    fn eat2(&mut self, c1: u8, c2: u8) -> bool {
899        self.skip_ws();
900        if self.pos + 1 < self.src.len() && self.src[self.pos] == c1 && self.src[self.pos + 1] == c2
901        {
902            self.pos += 2;
903            true
904        } else {
905            false
906        }
907    }
908
909    // ── precedence 1: || ──────────────────────────────────────────────
910    fn parse_logical_or(&mut self) -> i128 {
911        let mut v = self.parse_logical_and();
912        while self.eat2(b'|', b'|') {
913            let r = self.parse_logical_and();
914            v = if v != 0 || r != 0 { 1 } else { 0 };
915        }
916        v
917    }
918
919    // ── precedence 2: && ──────────────────────────────────────────────
920    fn parse_logical_and(&mut self) -> i128 {
921        let mut v = self.parse_bitwise_or();
922        while self.eat2(b'&', b'&') {
923            let r = self.parse_bitwise_or();
924            v = if v != 0 && r != 0 { 1 } else { 0 };
925        }
926        v
927    }
928
929    // ── precedence 3: | (but not ||) ─────────────────────────────────
930    fn parse_bitwise_or(&mut self) -> i128 {
931        let mut v = self.parse_bitwise_xor();
932        loop {
933            self.skip_ws();
934            if self.pos < self.src.len() && self.src[self.pos] == b'|' {
935                // Distinguish | from ||
936                if self.pos + 1 < self.src.len() && self.src[self.pos + 1] == b'|' {
937                    break;
938                }
939                self.pos += 1;
940                v |= self.parse_bitwise_xor();
941            } else {
942                break;
943            }
944        }
945        v
946    }
947
948    // ── precedence 4: ^ ──────────────────────────────────────────────
949    fn parse_bitwise_xor(&mut self) -> i128 {
950        let mut v = self.parse_bitwise_and();
951        loop {
952            self.skip_ws();
953            if self.pos < self.src.len() && self.src[self.pos] == b'^' {
954                self.pos += 1;
955                v ^= self.parse_bitwise_and();
956            } else {
957                break;
958            }
959        }
960        v
961    }
962
963    // ── precedence 5: & (but not &&) ─────────────────────────────────
964    fn parse_bitwise_and(&mut self) -> i128 {
965        let mut v = self.parse_equality();
966        loop {
967            self.skip_ws();
968            if self.pos < self.src.len() && self.src[self.pos] == b'&' {
969                if self.pos + 1 < self.src.len() && self.src[self.pos + 1] == b'&' {
970                    break;
971                }
972                self.pos += 1;
973                v &= self.parse_equality();
974            } else {
975                break;
976            }
977        }
978        v
979    }
980
981    // ── precedence 6: == != ──────────────────────────────────────────
982    fn parse_equality(&mut self) -> i128 {
983        let mut v = self.parse_relational();
984        loop {
985            if self.eat2(b'=', b'=') {
986                let r = self.parse_relational();
987                v = if v == r { 1 } else { 0 };
988            } else if self.eat2(b'!', b'=') {
989                let r = self.parse_relational();
990                v = if v == r { 0 } else { 1 };
991            } else {
992                break;
993            }
994        }
995        v
996    }
997
998    // ── precedence 7: < > <= >= ──────────────────────────────────────
999    fn parse_relational(&mut self) -> i128 {
1000        let mut v = self.parse_shift();
1001        loop {
1002            if self.eat2(b'<', b'=') {
1003                v = if v <= self.parse_shift() { 1 } else { 0 };
1004            } else if self.eat2(b'>', b'=') {
1005                v = if v >= self.parse_shift() { 1 } else { 0 };
1006            } else {
1007                self.skip_ws();
1008                if self.pos < self.src.len() && self.src[self.pos] == b'<' {
1009                    // Not << or <=
1010                    if self.pos + 1 < self.src.len()
1011                        && (self.src[self.pos + 1] == b'<' || self.src[self.pos + 1] == b'=')
1012                    {
1013                        break;
1014                    }
1015                    self.pos += 1;
1016                    v = if v < self.parse_shift() { 1 } else { 0 };
1017                } else if self.pos < self.src.len() && self.src[self.pos] == b'>' {
1018                    if self.pos + 1 < self.src.len()
1019                        && (self.src[self.pos + 1] == b'>' || self.src[self.pos + 1] == b'=')
1020                    {
1021                        break;
1022                    }
1023                    self.pos += 1;
1024                    v = if v > self.parse_shift() { 1 } else { 0 };
1025                } else {
1026                    break;
1027                }
1028            }
1029        }
1030        v
1031    }
1032
1033    // ── precedence 8: << >> ──────────────────────────────────────────
1034    fn parse_shift(&mut self) -> i128 {
1035        let mut v = self.parse_additive();
1036        loop {
1037            if self.eat2(b'<', b'<') {
1038                let r = self.parse_additive();
1039                v = if (0..128).contains(&r) {
1040                    v.wrapping_shl(r as u32)
1041                } else {
1042                    0
1043                };
1044            } else if self.eat2(b'>', b'>') {
1045                let r = self.parse_additive();
1046                v = if (0..128).contains(&r) {
1047                    v.wrapping_shr(r as u32)
1048                } else {
1049                    0
1050                };
1051            } else {
1052                break;
1053            }
1054        }
1055        v
1056    }
1057
1058    // ── precedence 9: + - ────────────────────────────────────────────
1059    fn parse_additive(&mut self) -> i128 {
1060        let mut v = self.parse_multiplicative();
1061        loop {
1062            self.skip_ws();
1063            if self.pos < self.src.len() && self.src[self.pos] == b'+' {
1064                self.pos += 1;
1065                v = v.wrapping_add(self.parse_multiplicative());
1066            } else if self.pos < self.src.len() && self.src[self.pos] == b'-' {
1067                self.pos += 1;
1068                v = v.wrapping_sub(self.parse_multiplicative());
1069            } else {
1070                break;
1071            }
1072        }
1073        v
1074    }
1075
1076    // ── precedence 10: * / % ─────────────────────────────────────────
1077    fn parse_multiplicative(&mut self) -> i128 {
1078        let mut v = self.parse_unary();
1079        loop {
1080            self.skip_ws();
1081            if self.pos < self.src.len() && self.src[self.pos] == b'*' {
1082                self.pos += 1;
1083                v = v.wrapping_mul(self.parse_unary());
1084            } else if self.pos < self.src.len() && self.src[self.pos] == b'/' {
1085                self.pos += 1;
1086                let r = self.parse_unary();
1087                v = if r != 0 { v / r } else { 0 };
1088            } else if self.pos < self.src.len() && self.src[self.pos] == b'%' {
1089                self.pos += 1;
1090                let r = self.parse_unary();
1091                v = if r != 0 { v % r } else { 0 };
1092            } else {
1093                break;
1094            }
1095        }
1096        v
1097    }
1098
1099    // ── precedence 11: unary ! - ~ ───────────────────────────────────
1100    fn parse_unary(&mut self) -> i128 {
1101        self.skip_ws();
1102        if self.pos < self.src.len() {
1103            match self.src[self.pos] {
1104                // Logical NOT (but not !=)
1105                b'!' if self.pos + 1 >= self.src.len() || self.src[self.pos + 1] != b'=' => {
1106                    self.pos += 1;
1107                    let v = self.parse_unary();
1108                    return if v == 0 { 1 } else { 0 };
1109                }
1110                b'-' => {
1111                    self.pos += 1;
1112                    return self.parse_unary().wrapping_neg();
1113                }
1114                b'~' => {
1115                    self.pos += 1;
1116                    return !self.parse_unary();
1117                }
1118                _ => {}
1119            }
1120        }
1121        self.parse_primary()
1122    }
1123
1124    // ── precedence 12: atoms ─────────────────────────────────────────
1125    fn parse_primary(&mut self) -> i128 {
1126        self.skip_ws();
1127        if self.pos >= self.src.len() {
1128            return 0;
1129        }
1130        let ch = self.src[self.pos];
1131
1132        // Parenthesised sub-expression
1133        if ch == b'(' {
1134            self.pos += 1;
1135            let v = self.parse_logical_or();
1136            self.skip_ws();
1137            if self.pos < self.src.len() && self.src[self.pos] == b')' {
1138                self.pos += 1;
1139            }
1140            return v;
1141        }
1142
1143        // Numeric literal (decimal, 0x, 0b, 0o)
1144        if ch.is_ascii_digit() {
1145            return self.parse_number();
1146        }
1147
1148        // Character literal 'c'
1149        if ch == b'\'' && self.pos + 2 < self.src.len() && self.src[self.pos + 2] == b'\'' {
1150            let c = self.src[self.pos + 1];
1151            self.pos += 3;
1152            return c as i128;
1153        }
1154
1155        // Identifier: symbol name or `defined()`
1156        if ch.is_ascii_alphabetic() || ch == b'_' || ch == b'.' {
1157            let start = self.pos;
1158            while self.pos < self.src.len() {
1159                let c = self.src[self.pos];
1160                if c.is_ascii_alphanumeric() || c == b'_' || c == b'.' {
1161                    self.pos += 1;
1162                } else {
1163                    break;
1164                }
1165            }
1166            let name = core::str::from_utf8(&self.src[start..self.pos]).unwrap_or("");
1167
1168            // `defined(sym)` pseudo-function
1169            if name == "defined" {
1170                self.skip_ws();
1171                if self.pos < self.src.len() && self.src[self.pos] == b'(' {
1172                    self.pos += 1;
1173                    self.skip_ws();
1174                    let s = self.pos;
1175                    while self.pos < self.src.len() {
1176                        let c = self.src[self.pos];
1177                        if c.is_ascii_alphanumeric() || c == b'_' || c == b'.' {
1178                            self.pos += 1;
1179                        } else {
1180                            break;
1181                        }
1182                    }
1183                    let sym = core::str::from_utf8(&self.src[s..self.pos]).unwrap_or("");
1184                    self.skip_ws();
1185                    if self.pos < self.src.len() && self.src[self.pos] == b')' {
1186                        self.pos += 1;
1187                    }
1188                    return if self.symbols.contains_key(sym) { 1 } else { 0 };
1189                }
1190            }
1191
1192            if let Some(&val) = self.symbols.get(name) {
1193                return val;
1194            }
1195            return 0; // unknown symbol → 0
1196        }
1197
1198        0
1199    }
1200
1201    /// Parse a numeric literal at the current position.
1202    fn parse_number(&mut self) -> i128 {
1203        if self.src[self.pos] == b'0' && self.pos + 1 < self.src.len() {
1204            match self.src[self.pos + 1] {
1205                b'x' | b'X' => {
1206                    self.pos += 2;
1207                    let start = self.pos;
1208                    while self.pos < self.src.len() && self.src[self.pos].is_ascii_hexdigit() {
1209                        self.pos += 1;
1210                    }
1211                    let s = core::str::from_utf8(&self.src[start..self.pos]).unwrap_or("0");
1212                    return i128::from_str_radix(s, 16).unwrap_or(0);
1213                }
1214                b'b' | b'B' => {
1215                    self.pos += 2;
1216                    let start = self.pos;
1217                    while self.pos < self.src.len() && matches!(self.src[self.pos], b'0' | b'1') {
1218                        self.pos += 1;
1219                    }
1220                    let s = core::str::from_utf8(&self.src[start..self.pos]).unwrap_or("0");
1221                    return i128::from_str_radix(s, 2).unwrap_or(0);
1222                }
1223                b'o' | b'O' => {
1224                    self.pos += 2;
1225                    let start = self.pos;
1226                    while self.pos < self.src.len() && matches!(self.src[self.pos], b'0'..=b'7') {
1227                        self.pos += 1;
1228                    }
1229                    let s = core::str::from_utf8(&self.src[start..self.pos]).unwrap_or("0");
1230                    return i128::from_str_radix(s, 8).unwrap_or(0);
1231                }
1232                _ => {}
1233            }
1234        }
1235        // Plain decimal
1236        let start = self.pos;
1237        while self.pos < self.src.len() && self.src[self.pos].is_ascii_digit() {
1238            self.pos += 1;
1239        }
1240        let s = core::str::from_utf8(&self.src[start..self.pos]).unwrap_or("0");
1241        s.parse::<i128>().unwrap_or(0)
1242    }
1243}
1244
1245/// Evaluate a simple integer expression (with symbol lookup).
1246///
1247/// Uses a recursive-descent parser with full C-like operator precedence.
1248/// Supports all arithmetic, bitwise, shift, logical, and comparison operators,
1249/// parenthesised sub-expressions, `defined()`, and numeric/symbol atoms.
1250fn eval_simple_expr(expr: &str, symbols: &BTreeMap<String, i128>) -> i128 {
1251    ExprEval::new(expr.trim(), symbols).eval()
1252}
1253
1254/// Try to parse a symbol definition from `.equ name, value` or `name = value`.
1255fn try_parse_symbol_def(line: &str) -> Option<(String, i128)> {
1256    let trimmed = line.trim();
1257
1258    // `.equ name, value` or `.set name, value`
1259    for prefix in &[".equ ", ".set "] {
1260        if let Some(rest) = trimmed.strip_prefix(prefix) {
1261            let rest = rest.trim();
1262            if let Some((name, val_str)) = rest.split_once(',') {
1263                if let Ok(val) = parse_int_literal(val_str.trim()) {
1264                    return Some((String::from(name.trim()), val));
1265                }
1266            }
1267        }
1268    }
1269
1270    // `name = value`
1271    if let Some((name, val_str)) = trimmed.split_once('=') {
1272        let name = name.trim();
1273        let val_str = val_str.trim();
1274        // Must not start with '=' (that would be '==')
1275        if !val_str.is_empty()
1276            && !val_str.starts_with('=')
1277            && name.chars().all(|c| c.is_alphanumeric() || c == '_')
1278        {
1279            if let Ok(val) = parse_int_literal(val_str) {
1280                return Some((String::from(name), val));
1281            }
1282        }
1283    }
1284
1285    None
1286}
1287
1288/// Parse an integer literal (decimal, hex, octal, binary).
1289fn parse_int_literal(s: &str) -> Result<i128, ()> {
1290    let s = s.trim();
1291    if let Some(hex) = s.strip_prefix("0x").or_else(|| s.strip_prefix("0X")) {
1292        i128::from_str_radix(hex, 16).map_err(|_| ())
1293    } else if let Some(bin) = s.strip_prefix("0b").or_else(|| s.strip_prefix("0B")) {
1294        i128::from_str_radix(bin, 2).map_err(|_| ())
1295    } else if let Some(oct) = s.strip_prefix("0o").or_else(|| s.strip_prefix("0O")) {
1296        i128::from_str_radix(oct, 8).map_err(|_| ())
1297    } else {
1298        s.parse::<i128>().map_err(|_| ())
1299    }
1300}
1301
1302/// Create a dummy span for a given line index.
1303fn line_span(line: usize) -> Span {
1304    Span::new((line + 1) as u32, 1, 0, 0)
1305}
1306
1307#[cfg(test)]
1308mod tests {
1309    use super::*;
1310
1311    // === Macro definition and expansion ===
1312
1313    #[test]
1314    fn macro_simple_expansion() {
1315        let mut pp = Preprocessor::new();
1316        let source = "\
1317.macro push_pair r1, r2
1318    push \\r1
1319    push \\r2
1320.endm
1321push_pair rax, rbx
1322";
1323        let result = pp.process(source).unwrap();
1324        assert!(result.contains("push rax"));
1325        assert!(result.contains("push rbx"));
1326    }
1327
1328    #[test]
1329    fn macro_with_defaults() {
1330        let mut pp = Preprocessor::new();
1331        let source = "\
1332.macro load_imm reg=rax, val=0
1333    mov \\reg, \\val
1334.endm
1335load_imm
1336load_imm rcx, 42
1337";
1338        let result = pp.process(source).unwrap();
1339        assert!(result.contains("mov rax, 0"));
1340        assert!(result.contains("mov rcx, 42"));
1341    }
1342
1343    #[test]
1344    fn macro_unique_labels() {
1345        let mut pp = Preprocessor::new();
1346        let source = "\
1347.macro my_loop
1348    jmp label_\\@
1349label_\\@:
1350.endm
1351my_loop
1352my_loop
1353";
1354        let result = pp.process(source).unwrap();
1355        assert!(result.contains("label_0"));
1356        assert!(result.contains("label_1"));
1357    }
1358
1359    #[test]
1360    fn macro_recursion_limit() {
1361        let mut pp = Preprocessor::new();
1362        let source = "\
1363.macro recurse
1364    nop
1365    recurse
1366.endm
1367recurse
1368";
1369        let err = pp.process(source).unwrap_err();
1370        match err {
1371            AsmError::ResourceLimitExceeded { resource, .. } => {
1372                assert!(resource.contains("recursion"));
1373            }
1374            _ => panic!("expected ResourceLimitExceeded, got {:?}", err),
1375        }
1376    }
1377
1378    #[test]
1379    fn macro_vararg() {
1380        let mut pp = Preprocessor::new();
1381        let source = "\
1382.macro pushall regs:vararg
1383    # push \\regs
1384.endm
1385pushall rax, rbx, rcx
1386";
1387        let result = pp.process(source).unwrap();
1388        assert!(result.contains("rax, rbx, rcx"));
1389    }
1390
1391    #[test]
1392    fn macro_nested_endm() {
1393        let mut pp = Preprocessor::new();
1394        // Macro containing a nested macro definition
1395        let source = "\
1396.macro outer
1397    nop
1398.endm
1399outer
1400";
1401        let result = pp.process(source).unwrap();
1402        assert!(result.contains("nop"));
1403    }
1404
1405    // === .rept ===
1406
1407    #[test]
1408    fn rept_basic() {
1409        let mut pp = Preprocessor::new();
1410        let source = "\
1411.rept 3
1412    nop
1413.endr
1414";
1415        let result = pp.process(source).unwrap();
1416        let nop_count = result.matches("nop").count();
1417        assert_eq!(nop_count, 3);
1418    }
1419
1420    #[test]
1421    fn rept_zero() {
1422        let mut pp = Preprocessor::new();
1423        let source = "\
1424.rept 0
1425    nop
1426.endr
1427";
1428        let result = pp.process(source).unwrap();
1429        assert!(!result.contains("nop"));
1430    }
1431
1432    #[test]
1433    fn rept_nested() {
1434        let mut pp = Preprocessor::new();
1435        let source = "\
1436.rept 2
1437.rept 3
1438    nop
1439.endr
1440.endr
1441";
1442        let result = pp.process(source).unwrap();
1443        let nop_count = result.matches("nop").count();
1444        assert_eq!(nop_count, 6);
1445    }
1446
1447    // === .irp ===
1448
1449    #[test]
1450    fn irp_basic() {
1451        let mut pp = Preprocessor::new();
1452        let source = "\
1453.irp reg, rax, rbx, rcx
1454    push \\reg
1455.endr
1456";
1457        let result = pp.process(source).unwrap();
1458        assert!(result.contains("push rax"));
1459        assert!(result.contains("push rbx"));
1460        assert!(result.contains("push rcx"));
1461    }
1462
1463    // === .irpc ===
1464
1465    #[test]
1466    fn irpc_basic() {
1467        let mut pp = Preprocessor::new();
1468        let source = "\
1469.irpc c, abc
1470    .byte '\\c'
1471.endr
1472";
1473        let result = pp.process(source).unwrap();
1474        assert!(result.contains("'a'"));
1475        assert!(result.contains("'b'"));
1476        assert!(result.contains("'c'"));
1477    }
1478
1479    // === Conditional assembly ===
1480
1481    #[test]
1482    fn if_true() {
1483        let mut pp = Preprocessor::new();
1484        let source = "\
1485.if 1
1486    nop
1487.endif
1488";
1489        let result = pp.process(source).unwrap();
1490        assert!(result.contains("nop"));
1491    }
1492
1493    #[test]
1494    fn if_false() {
1495        let mut pp = Preprocessor::new();
1496        let source = "\
1497.if 0
1498    nop
1499.endif
1500";
1501        let result = pp.process(source).unwrap();
1502        assert!(!result.contains("nop"));
1503    }
1504
1505    #[test]
1506    fn if_else() {
1507        let mut pp = Preprocessor::new();
1508        let source = "\
1509.if 0
1510    mov rax, 1
1511.else
1512    mov rax, 2
1513.endif
1514";
1515        let result = pp.process(source).unwrap();
1516        assert!(!result.contains("mov rax, 1"));
1517        assert!(result.contains("mov rax, 2"));
1518    }
1519
1520    #[test]
1521    fn if_elseif() {
1522        let mut pp = Preprocessor::new();
1523        let source = "\
1524.if 0
1525    mov rax, 1
1526.elseif 1
1527    mov rax, 2
1528.else
1529    mov rax, 3
1530.endif
1531";
1532        let result = pp.process(source).unwrap();
1533        assert!(!result.contains("mov rax, 1"));
1534        assert!(result.contains("mov rax, 2"));
1535        assert!(!result.contains("mov rax, 3"));
1536    }
1537
1538    #[test]
1539    fn ifdef_defined() {
1540        let mut pp = Preprocessor::new();
1541        pp.define_symbol("MY_FLAG", 1);
1542        let source = "\
1543.ifdef MY_FLAG
1544    nop
1545.endif
1546";
1547        let result = pp.process(source).unwrap();
1548        assert!(result.contains("nop"));
1549    }
1550
1551    #[test]
1552    fn ifdef_undefined() {
1553        let mut pp = Preprocessor::new();
1554        let source = "\
1555.ifdef UNDEFINED_FLAG
1556    nop
1557.endif
1558";
1559        let result = pp.process(source).unwrap();
1560        assert!(!result.contains("nop"));
1561    }
1562
1563    #[test]
1564    fn ifndef_undefined() {
1565        let mut pp = Preprocessor::new();
1566        let source = "\
1567.ifndef MY_FLAG
1568    nop
1569.endif
1570";
1571        let result = pp.process(source).unwrap();
1572        assert!(result.contains("nop"));
1573    }
1574
1575    #[test]
1576    fn nested_conditionals() {
1577        let mut pp = Preprocessor::new();
1578        pp.define_symbol("OUTER", 1);
1579        pp.define_symbol("INNER", 1);
1580        let source = "\
1581.ifdef OUTER
1582    .ifdef INNER
1583        nop
1584    .endif
1585.endif
1586";
1587        let result = pp.process(source).unwrap();
1588        assert!(result.contains("nop"));
1589    }
1590
1591    #[test]
1592    fn if_expression_with_symbols() {
1593        let mut pp = Preprocessor::new();
1594        pp.define_symbol("X", 5);
1595        let source = "\
1596.if X > 3
1597    nop
1598.endif
1599";
1600        let result = pp.process(source).unwrap();
1601        assert!(result.contains("nop"));
1602    }
1603
1604    #[test]
1605    fn equ_tracks_symbols() {
1606        let mut pp = Preprocessor::new();
1607        let source = "\
1608.equ MY_CONST, 42
1609.ifdef MY_CONST
1610    nop
1611.endif
1612";
1613        let result = pp.process(source).unwrap();
1614        assert!(result.contains("nop"));
1615        // The .equ line also passes through for the parser
1616        assert!(result.contains(".equ MY_CONST, 42"));
1617    }
1618
1619    #[test]
1620    fn if_defined_function() {
1621        let mut pp = Preprocessor::new();
1622        pp.define_symbol("X", 1);
1623        let source = "\
1624.if defined(X)
1625    nop
1626.endif
1627";
1628        let result = pp.process(source).unwrap();
1629        assert!(result.contains("nop"));
1630    }
1631
1632    // === Error cases ===
1633
1634    #[test]
1635    fn unterminated_macro() {
1636        let mut pp = Preprocessor::new();
1637        let source = ".macro foo\n    nop\n";
1638        let err = pp.process(source).unwrap_err();
1639        match err {
1640            AsmError::Syntax { msg, .. } => {
1641                assert!(msg.contains("unterminated .macro"));
1642            }
1643            _ => panic!("expected Syntax error"),
1644        }
1645    }
1646
1647    #[test]
1648    fn unterminated_rept() {
1649        let mut pp = Preprocessor::new();
1650        let source = ".rept 3\n    nop\n";
1651        let err = pp.process(source).unwrap_err();
1652        match err {
1653            AsmError::Syntax { msg, .. } => {
1654                assert!(msg.contains("unterminated"));
1655            }
1656            _ => panic!("expected Syntax error"),
1657        }
1658    }
1659
1660    #[test]
1661    fn unterminated_conditional() {
1662        let mut pp = Preprocessor::new();
1663        let source = ".if 1\n    nop\n";
1664        let err = pp.process(source).unwrap_err();
1665        match err {
1666            AsmError::Syntax { msg, .. } => {
1667                assert!(msg.contains("unterminated"));
1668            }
1669            _ => panic!("expected Syntax error"),
1670        }
1671    }
1672
1673    #[test]
1674    fn iteration_limit() {
1675        let mut pp = Preprocessor::new();
1676        let source = ".rept 200000\n    nop\n.endr\n";
1677        let err = pp.process(source).unwrap_err();
1678        match err {
1679            AsmError::ResourceLimitExceeded { resource, .. } => {
1680                assert!(resource.contains("iteration"));
1681            }
1682            _ => panic!("expected ResourceLimitExceeded"),
1683        }
1684    }
1685
1686    // === Expression evaluator tests ===
1687
1688    /// Helper: evaluate expression with given symbols.
1689    fn eval(expr: &str) -> i128 {
1690        let syms = BTreeMap::new();
1691        super::eval_simple_expr(expr, &syms)
1692    }
1693
1694    fn eval_with(expr: &str, syms: &BTreeMap<String, i128>) -> i128 {
1695        super::eval_simple_expr(expr, syms)
1696    }
1697
1698    #[test]
1699    fn expr_decimal_literals() {
1700        assert_eq!(eval("0"), 0);
1701        assert_eq!(eval("42"), 42);
1702        assert_eq!(eval("123456789"), 123_456_789);
1703    }
1704
1705    #[test]
1706    fn expr_hex_literals() {
1707        assert_eq!(eval("0xFF"), 255);
1708        assert_eq!(eval("0x10"), 16);
1709        assert_eq!(eval("0XAB"), 0xAB);
1710    }
1711
1712    #[test]
1713    fn expr_binary_literals() {
1714        assert_eq!(eval("0b1010"), 10);
1715        assert_eq!(eval("0B11111111"), 255);
1716    }
1717
1718    #[test]
1719    fn expr_octal_literals() {
1720        assert_eq!(eval("0o77"), 63);
1721        assert_eq!(eval("0O10"), 8);
1722    }
1723
1724    #[test]
1725    fn expr_char_literal() {
1726        assert_eq!(eval("'A'"), 65);
1727        assert_eq!(eval("'0'"), 48);
1728    }
1729
1730    #[test]
1731    fn expr_addition() {
1732        assert_eq!(eval("1 + 2"), 3);
1733        assert_eq!(eval("10+20+30"), 60);
1734    }
1735
1736    #[test]
1737    fn expr_subtraction() {
1738        assert_eq!(eval("10 - 3"), 7);
1739        assert_eq!(eval("100 - 50 - 25"), 25);
1740    }
1741
1742    #[test]
1743    fn expr_multiplication() {
1744        assert_eq!(eval("3 * 4"), 12);
1745        assert_eq!(eval("2 * 3 * 5"), 30);
1746    }
1747
1748    #[test]
1749    fn expr_division() {
1750        assert_eq!(eval("12 / 4"), 3);
1751        assert_eq!(eval("100 / 10 / 2"), 5);
1752        // Division by zero → 0
1753        assert_eq!(eval("42 / 0"), 0);
1754    }
1755
1756    #[test]
1757    fn expr_modulo() {
1758        assert_eq!(eval("10 % 3"), 1);
1759        assert_eq!(eval("17 % 5"), 2);
1760        assert_eq!(eval("42 % 0"), 0);
1761    }
1762
1763    #[test]
1764    fn expr_precedence_mul_over_add() {
1765        assert_eq!(eval("2 + 3 * 4"), 14);
1766        assert_eq!(eval("3 * 4 + 2"), 14);
1767        assert_eq!(eval("10 - 2 * 3"), 4);
1768    }
1769
1770    #[test]
1771    fn expr_parentheses() {
1772        assert_eq!(eval("(2 + 3) * 4"), 20);
1773        assert_eq!(eval("((1 + 2) * (3 + 4))"), 21);
1774        assert_eq!(eval("(10)"), 10);
1775    }
1776
1777    #[test]
1778    fn expr_nested_parentheses() {
1779        assert_eq!(eval("((2 + 3) * (4 - 1))"), 15);
1780        assert_eq!(eval("(((5)))"), 5);
1781    }
1782
1783    #[test]
1784    fn expr_bitwise_and() {
1785        assert_eq!(eval("0xFF & 0x0F"), 0x0F);
1786        assert_eq!(eval("0b1010 & 0b1100"), 0b1000);
1787    }
1788
1789    #[test]
1790    fn expr_bitwise_or() {
1791        assert_eq!(eval("0x0F | 0xF0"), 0xFF);
1792        assert_eq!(eval("0b1010 | 0b0101"), 0b1111);
1793    }
1794
1795    #[test]
1796    fn expr_bitwise_xor() {
1797        assert_eq!(eval("0xFF ^ 0x0F"), 0xF0);
1798        assert_eq!(eval("0b1010 ^ 0b1100"), 0b0110);
1799    }
1800
1801    #[test]
1802    fn expr_bitwise_not() {
1803        // ~0 in i128 is all ones = -1
1804        assert_eq!(eval("~0"), -1);
1805        assert_eq!(eval("~0xFF & 0xFF"), 0);
1806    }
1807
1808    #[test]
1809    fn expr_shift_left() {
1810        assert_eq!(eval("1 << 8"), 256);
1811        assert_eq!(eval("0xFF << 4"), 0xFF0);
1812    }
1813
1814    #[test]
1815    fn expr_shift_right() {
1816        assert_eq!(eval("256 >> 8"), 1);
1817        assert_eq!(eval("0xFF0 >> 4"), 0xFF);
1818    }
1819
1820    #[test]
1821    fn expr_logical_and() {
1822        assert_eq!(eval("1 && 1"), 1);
1823        assert_eq!(eval("1 && 0"), 0);
1824        assert_eq!(eval("0 && 1"), 0);
1825        assert_eq!(eval("0 && 0"), 0);
1826    }
1827
1828    #[test]
1829    fn expr_logical_or() {
1830        assert_eq!(eval("1 || 1"), 1);
1831        assert_eq!(eval("1 || 0"), 1);
1832        assert_eq!(eval("0 || 1"), 1);
1833        assert_eq!(eval("0 || 0"), 0);
1834    }
1835
1836    #[test]
1837    fn expr_logical_not() {
1838        assert_eq!(eval("!0"), 1);
1839        assert_eq!(eval("!1"), 0);
1840        assert_eq!(eval("!42"), 0);
1841    }
1842
1843    #[test]
1844    fn expr_equality() {
1845        assert_eq!(eval("5 == 5"), 1);
1846        assert_eq!(eval("5 == 6"), 0);
1847        assert_eq!(eval("5 != 6"), 1);
1848        assert_eq!(eval("5 != 5"), 0);
1849    }
1850
1851    #[test]
1852    fn expr_relational() {
1853        assert_eq!(eval("3 < 5"), 1);
1854        assert_eq!(eval("5 < 3"), 0);
1855        assert_eq!(eval("5 > 3"), 1);
1856        assert_eq!(eval("3 > 5"), 0);
1857        assert_eq!(eval("5 <= 5"), 1);
1858        assert_eq!(eval("5 <= 6"), 1);
1859        assert_eq!(eval("6 <= 5"), 0);
1860        assert_eq!(eval("5 >= 5"), 1);
1861        assert_eq!(eval("6 >= 5"), 1);
1862        assert_eq!(eval("5 >= 6"), 0);
1863    }
1864
1865    #[test]
1866    fn expr_unary_minus() {
1867        assert_eq!(eval("-1"), -1);
1868        assert_eq!(eval("-(-5)"), 5);
1869        assert_eq!(eval("3 + -2"), 1);
1870        assert_eq!(eval("3 - -2"), 5);
1871    }
1872
1873    #[test]
1874    fn expr_mixed_precedence() {
1875        // Shift lower than add: 1 + 2 << 3 == (1+2) << 3 ... NO
1876        // Actually: << is higher than +: 1 + (2 << 3) = 1 + 16 = 17
1877        assert_eq!(eval("1 + 2 << 3"), 24); // (1+2)<<3, since shift is HIGHER than add... wait
1878                                            // Let me think about this. In C, << is higher precedence (binds tighter) than +.
1879                                            // But in our parser, additive is level 9 and shift is level 8 (lower number = lower precedence).
1880                                            // additive calls parse_multiplicative, shift calls parse_additive.
1881                                            // Wait - that's wrong. Let me re-check.
1882                                            // Actually: parse_additive calls parse_multiplicative, and parse_shift calls parse_additive.
1883                                            // So shift calls additive which calls multiplicative. This means additive binds tighter than shift.
1884                                            // That matches C precedence where + binds tighter than <<.
1885                                            // So 1 + 2 << 3 = (1+2) << 3 = 3 << 3 = 24. Correct for C.
1886        assert_eq!(eval("1 + 2 << 3"), 24);
1887
1888        // Comparison: == is lower than +
1889        assert_eq!(eval("2 + 3 == 5"), 1);
1890        assert_eq!(eval("2 + 3 == 6"), 0);
1891
1892        // Logical: && is lower than ==
1893        assert_eq!(eval("1 == 1 && 2 == 2"), 1);
1894        assert_eq!(eval("1 == 1 && 2 == 3"), 0);
1895
1896        // || is the lowest
1897        assert_eq!(eval("0 && 1 || 1"), 1);
1898        assert_eq!(eval("1 || 0 && 0"), 1);
1899    }
1900
1901    #[test]
1902    fn expr_complex_bitwise() {
1903        // Page-align: addr & ~0xFFF
1904        // Can't test with large addresses easily, but logic works
1905        assert_eq!(eval("0x1234 & ~0xFFF & 0xFFFF"), 0x1000);
1906        // Flag test
1907        assert_eq!(eval("(0x03 & 0x01) != 0"), 1);
1908        assert_eq!(eval("(0x02 & 0x01) != 0"), 0);
1909    }
1910
1911    #[test]
1912    fn expr_symbols() {
1913        let mut syms = BTreeMap::new();
1914        syms.insert(String::from("X"), 10);
1915        syms.insert(String::from("Y"), 20);
1916        assert_eq!(eval_with("X + Y", &syms), 30);
1917        assert_eq!(eval_with("X * Y", &syms), 200);
1918        assert_eq!(eval_with("(X + Y) * 2", &syms), 60);
1919    }
1920
1921    #[test]
1922    fn expr_defined_function() {
1923        let mut syms = BTreeMap::new();
1924        syms.insert(String::from("FOO"), 1);
1925        assert_eq!(eval_with("defined(FOO)", &syms), 1);
1926        assert_eq!(eval_with("defined(BAR)", &syms), 0);
1927        assert_eq!(eval_with("defined(FOO) && defined(BAR)", &syms), 0);
1928        assert_eq!(eval_with("defined(FOO) || defined(BAR)", &syms), 1);
1929    }
1930
1931    #[test]
1932    fn expr_regression_0x_minus() {
1933        // Old parser choked on "0x10 - 1" because rsplit_once('-') split at "0x10"
1934        assert_eq!(eval("0x10 - 1"), 15);
1935        assert_eq!(eval("0xFF - 0xF0"), 15);
1936    }
1937
1938    #[test]
1939    fn expr_whitespace_tolerance() {
1940        assert_eq!(eval("  42  "), 42);
1941        assert_eq!(eval("  1  +  2  "), 3);
1942        assert_eq!(eval(" ( 1 + 2 ) * 3 "), 9);
1943    }
1944
1945    #[test]
1946    fn expr_empty() {
1947        assert_eq!(eval(""), 0);
1948        assert_eq!(eval("   "), 0);
1949    }
1950
1951    #[test]
1952    fn expr_if_mul_integrated() {
1953        // This was the data-corruption bug: `.if 2 * 3 == 6` silently failed
1954        let mut pp = Preprocessor::new();
1955        let source = "\
1956.if 2 * 3 == 6
1957    nop
1958.endif
1959";
1960        let result = pp.process(source).unwrap();
1961        assert!(result.contains("nop"), "2*3==6 should be true");
1962    }
1963
1964    #[test]
1965    fn expr_if_parenthesised_integrated() {
1966        let mut pp = Preprocessor::new();
1967        let source = "\
1968.if (1 + 2) * 4 == 12
1969    mov eax, 1
1970.endif
1971";
1972        let result = pp.process(source).unwrap();
1973        assert!(result.contains("mov eax, 1"));
1974    }
1975
1976    #[test]
1977    fn expr_if_shift_integrated() {
1978        let mut pp = Preprocessor::new();
1979        let source = "\
1980.if 1 << 4 == 16
1981    nop
1982.endif
1983";
1984        let result = pp.process(source).unwrap();
1985        assert!(result.contains("nop"), "1<<4 should equal 16");
1986    }
1987
1988    #[test]
1989    fn expr_if_bitwise_and_integrated() {
1990        let mut pp = Preprocessor::new();
1991        let source = "\
1992.equ FLAGS, 0x07
1993.if FLAGS & 0x02
1994    nop
1995.endif
1996";
1997        let result = pp.process(source).unwrap();
1998        assert!(result.contains("nop"), "0x07 & 0x02 should be non-zero");
1999    }
2000
2001    #[test]
2002    fn expr_if_logical_and_integrated() {
2003        let mut pp = Preprocessor::new();
2004        pp.define_symbol("A", 1);
2005        pp.define_symbol("B", 1);
2006        let source = "\
2007.if defined(A) && defined(B)
2008    nop
2009.endif
2010";
2011        let result = pp.process(source).unwrap();
2012        assert!(result.contains("nop"), "both A and B defined");
2013    }
2014
2015    #[test]
2016    fn expr_if_logical_or_integrated() {
2017        let mut pp = Preprocessor::new();
2018        pp.define_symbol("A", 1);
2019        let source = "\
2020.if defined(A) || defined(B)
2021    nop
2022.endif
2023";
2024        let result = pp.process(source).unwrap();
2025        assert!(result.contains("nop"), "A is defined so OR should be true");
2026    }
2027
2028    #[test]
2029    fn expr_elseif_with_operators() {
2030        let mut pp = Preprocessor::new();
2031        pp.define_symbol("MODE", 2);
2032        let source = "\
2033.if MODE * 2 == 2
2034    wrong
2035.elseif MODE * 2 == 4
2036    correct
2037.else
2038    also_wrong
2039.endif
2040";
2041        let result = pp.process(source).unwrap();
2042        assert!(!result.contains("wrong"));
2043        assert!(result.contains("correct"));
2044    }
2045}