Skip to main content

wazabin_qcode_parser/
parser.rs

1use crate::ast::{
2    Atom, BlockParamDecl, Callee, CastOp, ExprNode, ExtractField, FnDecl, FnKind, GepField, Label,
3    Program, ProgramKind, SourcePosition, SourceSpan, Statement, StructDecl, StructFieldDecl,
4    StructFieldType, TupleField, TypedAtom,
5};
6use pest::Parser;
7use pest::iterators::Pair;
8use pest_derive::Parser;
9use std::fmt;
10
11#[derive(Debug, Clone)]
12pub struct ParseError {
13    message: String,
14}
15
16impl ParseError {
17    fn new(message: impl Into<String>) -> Self {
18        Self {
19            message: message.into(),
20        }
21    }
22}
23
24impl fmt::Display for ParseError {
25    fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
26        f.write_str(&self.message)
27    }
28}
29
30impl std::error::Error for ParseError {}
31
32#[derive(Parser)]
33#[grammar = "qcode.pest"]
34struct QCodeParser;
35
36pub fn parse_program(program: &str) -> Result<Program, ParseError> {
37    let mut parsed = QCodeParser::parse(Rule::program, program).map_err(to_parse_error)?;
38    let root = parsed
39        .next()
40        .ok_or_else(|| ParseError::new("missing program"))?;
41
42    let mut fn_decls: Vec<FnDecl> = Vec::new();
43    let mut statements: Vec<Statement> = Vec::new();
44    let mut top_varnodes: Vec<Statement> = Vec::new();
45    let mut structs: Vec<StructDecl> = Vec::new();
46    let mut is_fn_program = false;
47
48    // A comment may appear at the program level (before statement_list or fn_decl)
49    // or between elements inside those rules.
50    let mut pending_comment: Option<String> = None;
51
52    for pair in root.into_inner() {
53        match pair.as_rule() {
54            Rule::COMMENT => {
55                pending_comment = Some(comment_text(pair.as_str()));
56            }
57            Rule::struct_decl => {
58                pending_comment = None;
59                structs.push(parse_struct_decl(pair)?);
60            }
61            Rule::top_varnode_list => {
62                for part in pair.into_inner() {
63                    if part.as_rule() == Rule::local_decl {
64                        top_varnodes.push(parse_local_decl(part)?);
65                    }
66                }
67            }
68            Rule::fn_decl => {
69                is_fn_program = true;
70                pending_comment = None; // comments before fn_decl are not yet attached
71                fn_decls.push(parse_fn_decl(pair)?);
72            }
73            Rule::statement_list => {
74                // The pending_comment (if any) was outside statement_list; treat it as
75                // preceding the first compound_stmt.
76                let mut inner_comment = pending_comment.take();
77                for item in pair.into_inner() {
78                    match item.as_rule() {
79                        Rule::COMMENT => {
80                            inner_comment = Some(comment_text(item.as_str()));
81                        }
82                        Rule::compound_stmt => {
83                            let mut compound_stmts = Vec::new();
84                            parse_compound(item, &mut compound_stmts)?;
85                            attach_comment(&mut inner_comment, &mut compound_stmts);
86                            statements.extend(compound_stmts);
87                        }
88                        _ => {}
89                    }
90                }
91            }
92            _ => {}
93        }
94    }
95
96    let kind = if is_fn_program {
97        ProgramKind::Functions {
98            varnodes: top_varnodes,
99            fns: fn_decls,
100        }
101    } else {
102        ProgramKind::Statements(statements)
103    };
104    Ok(Program { structs, kind })
105}
106
107fn parse_struct_decl(pair: Pair<'_, Rule>) -> Result<StructDecl, ParseError> {
108    let span = source_span(pair.as_span());
109    let mut inner = pair.into_inner();
110    let name = inner
111        .next()
112        .ok_or_else(|| ParseError::new("missing struct name"))?
113        .as_str()
114        .to_owned();
115    let mut fields = Vec::new();
116    for field in inner {
117        if field.as_rule() != Rule::struct_field {
118            continue;
119        }
120        let mut parts = field.into_inner();
121        let field_name = parts
122            .next()
123            .ok_or_else(|| ParseError::new("missing struct field name"))?
124            .as_str()
125            .to_owned();
126        let ty_pair = parts
127            .next()
128            .ok_or_else(|| ParseError::new("missing struct field type"))?;
129        let ty = match ty_pair.as_rule() {
130            Rule::struct_ptr_ty => {
131                StructFieldType::StructPtr(ty_pair.as_str().trim_end_matches('*').to_owned())
132            }
133            Rule::integer => StructFieldType::Int(parse_integer(ty_pair.as_str())? as usize),
134            _ => return Err(ParseError::new("invalid struct field type")),
135        };
136        fields.push(StructFieldDecl {
137            name: field_name,
138            ty,
139        });
140    }
141    Ok(StructDecl { name, fields, span })
142}
143
144fn parse_fn_decl(pair: Pair<'_, Rule>) -> Result<FnDecl, ParseError> {
145    let span = source_span(pair.as_span());
146    let mut inner = pair.into_inner();
147    let kind_pair = inner
148        .next()
149        .ok_or_else(|| ParseError::new("missing function kind"))?;
150    let kind = match kind_pair.as_str() {
151        "fn" => FnKind::Machine,
152        "lambda" => FnKind::Lambda,
153        _ => return Err(ParseError::new("invalid function kind")),
154    };
155    let name_pair = inner
156        .next()
157        .ok_or_else(|| ParseError::new("missing function name"))?;
158    let name_span = source_span(name_pair.as_span());
159    let name = name_pair.as_str().to_owned();
160
161    let mut statements = Vec::new();
162    for part in inner {
163        if part.as_rule() == Rule::fn_body {
164            let mut pending_comment: Option<String> = None;
165            for fn_stmt in part.into_inner() {
166                match fn_stmt.as_rule() {
167                    Rule::COMMENT => {
168                        pending_comment = Some(comment_text(fn_stmt.as_str()));
169                    }
170                    Rule::fn_stmt => {
171                        for compound in fn_stmt.into_inner() {
172                            if compound.as_rule() != Rule::compound_stmt {
173                                continue;
174                            }
175                            let mut compound_stmts = Vec::new();
176                            parse_compound(compound, &mut compound_stmts)?;
177                            attach_comment(&mut pending_comment, &mut compound_stmts);
178                            statements.extend(compound_stmts);
179                        }
180                    }
181                    _ => {}
182                }
183            }
184        }
185    }
186
187    Ok(FnDecl {
188        kind,
189        name,
190        name_span,
191        span,
192        statements,
193    })
194}
195
196fn parse_compound(pair: Pair<'_, Rule>, out: &mut Vec<Statement>) -> Result<(), ParseError> {
197    let mut pending_comment: Option<String> = None;
198    for part in pair.into_inner() {
199        match part.as_rule() {
200            Rule::label_decl => out.push(parse_label_decl(part)?),
201            Rule::inner_stmt => {
202                let mut stmts = Vec::new();
203                parse_inner_stmt(part, &mut stmts)?;
204                attach_comment(&mut pending_comment, &mut stmts);
205                out.extend(stmts);
206            }
207            Rule::COMMENT => {
208                pending_comment = Some(comment_text(part.as_str()));
209            }
210            _ => return Err(ParseError::new("unexpected compound statement")),
211        }
212    }
213    Ok(())
214}
215
216fn parse_label_decl(pair: Pair<'_, Rule>) -> Result<Statement, ParseError> {
217    let span = source_span(pair.as_span());
218    let mut inner = pair.into_inner();
219    let name_pair = inner
220        .next()
221        .ok_or_else(|| ParseError::new("missing label name"))?;
222    let label_span = source_span(name_pair.as_span());
223    let label = match name_pair.as_rule() {
224        Rule::ident => {
225            let name = name_pair.as_str().to_owned();
226            let params: Vec<BlockParamDecl> = inner
227                .filter(|p| p.as_rule() == Rule::block_param_decl)
228                .map(parse_block_param_decl)
229                .collect::<Result<_, _>>()?;
230            Label::Named {
231                name,
232                params,
233                span: label_span,
234            }
235        }
236        Rule::integer => {
237            let addr = parse_integer(name_pair.as_str())?;
238            Label::Address {
239                value: addr,
240                span: label_span,
241            }
242        }
243        _ => return Err(ParseError::new("invalid label declaration")),
244    };
245    Ok(Statement::LabelDecl { label, span })
246}
247
248fn parse_block_param_decl(pair: Pair<'_, Rule>) -> Result<BlockParamDecl, ParseError> {
249    let mut name = None;
250    let mut size_bytes = None;
251
252    for part in pair.into_inner() {
253        match part.as_rule() {
254            Rule::block_param_name => {
255                name = Some(
256                    part.as_str()
257                        .strip_prefix('@')
258                        .unwrap_or(part.as_str())
259                        .to_owned(),
260                );
261            }
262            Rule::ty => {
263                size_bytes = Some(parse_size_bytes(part.as_str(), "block parameter")?);
264            }
265            _ => {}
266        }
267    }
268
269    Ok(BlockParamDecl {
270        name: name.ok_or_else(|| ParseError::new("missing block parameter name"))?,
271        size_bytes,
272    })
273}
274
275fn parse_inner_stmt(pair: Pair<'_, Rule>, out: &mut Vec<Statement>) -> Result<(), ParseError> {
276    let inner = pair
277        .into_inner()
278        .next()
279        .ok_or_else(|| ParseError::new("empty inner statement"))?;
280
281    let stmt = match inner.as_rule() {
282        Rule::local_decl => parse_local_decl(inner)?,
283        Rule::assignment_ssa => parse_assignment_ssa(inner)?,
284        Rule::terminator => parse_terminator(inner)?,
285        Rule::assert_stmt => parse_assert_stmt(inner)?,
286        Rule::expr => Statement::Expr(parse_expr(inner)?),
287        _ => return Err(ParseError::new("unexpected inner statement")),
288    };
289
290    out.push(stmt);
291    Ok(())
292}
293
294fn parse_local_decl(pair: Pair<'_, Rule>) -> Result<Statement, ParseError> {
295    let span = source_span(pair.as_span());
296    let mut inner = pair.into_inner();
297    let size_bytes = parse_size_bytes(
298        inner
299            .next()
300            .ok_or_else(|| ParseError::new("missing local type"))?
301            .as_str(),
302        "local declaration",
303    )?;
304    let name_pair = inner
305        .next()
306        .ok_or_else(|| ParseError::new("missing local name"))?;
307    let name_span = source_span(name_pair.as_span());
308    let name = name_pair.as_str().to_owned();
309    let display_name = inner
310        .next()
311        .map(|pair| pair.as_str().to_owned())
312        .unwrap_or_else(|| name.clone());
313
314    Ok(Statement::LocalDecl {
315        name,
316        name_span,
317        display_name,
318        size_bytes,
319        span,
320    })
321}
322
323/// Parses the content of a `<...>` label into a `Label`.
324fn label_value(pair: Pair<'_, Rule>) -> Result<Label, ParseError> {
325    let span = source_span(pair.as_span());
326    match pair.as_rule() {
327        Rule::ident => Ok(Label::Named {
328            name: pair.as_str().to_owned(),
329            params: vec![],
330            span,
331        }),
332        Rule::integer => {
333            let addr = parse_integer(pair.as_str())?;
334            Ok(Label::Address { value: addr, span })
335        }
336        _ => Err(ParseError::new("invalid label content")),
337    }
338}
339
340/// Parses an `edge_hint` rule (`// -> <a>, <b>`) into its list of target labels.
341fn parse_edge_hint(pair: Pair<'_, Rule>) -> Result<Vec<Label>, ParseError> {
342    let mut targets = Vec::new();
343    for label_pair in pair.into_inner() {
344        if label_pair.as_rule() == Rule::label {
345            let content = label_pair
346                .into_inner()
347                .next()
348                .ok_or_else(|| ParseError::new("missing edge-hint label content"))?;
349            targets.push(label_value(content)?);
350        }
351    }
352    Ok(targets)
353}
354
355/// Parses a `branch_label` rule (`<(ident|int) branch_arg*>`) into a `Label` and args.
356fn parse_branch_label(
357    pair: Pair<'_, Rule>,
358) -> Result<(Label, Vec<(String, TypedAtom)>), ParseError> {
359    let mut inner = pair.into_inner();
360    let target_pair = inner
361        .next()
362        .ok_or_else(|| ParseError::new("missing branch label target"))?;
363    let label = label_value(target_pair)?;
364    let mut args = Vec::new();
365    for part in inner {
366        if part.as_rule() == Rule::branch_arg {
367            let mut arg_inner = part.into_inner();
368            let name_pair = arg_inner
369                .next()
370                .ok_or_else(|| ParseError::new("missing branch arg name"))?;
371            let name = name_pair
372                .as_str()
373                .strip_prefix('@')
374                .unwrap_or(name_pair.as_str())
375                .to_owned();
376            let value_pair = arg_inner
377                .next()
378                .ok_or_else(|| ParseError::new("missing branch arg value"))?;
379            args.push((name, parse_typed_atom(value_pair)?));
380        }
381    }
382    Ok((label, args))
383}
384
385fn parse_terminator(pair: Pair<'_, Rule>) -> Result<Statement, ParseError> {
386    let span = source_span(pair.as_span());
387    let specific = pair
388        .into_inner()
389        .next()
390        .ok_or_else(|| ParseError::new("missing terminator kind"))?;
391
392    match specific.as_rule() {
393        Rule::branch_stmt => {
394            let branch_label_pair = specific
395                .into_inner()
396                .next()
397                .ok_or_else(|| ParseError::new("missing branch label"))?;
398            let (target, args) = parse_branch_label(branch_label_pair)?;
399            Ok(Statement::Branch { target, args, span })
400        }
401
402        Rule::branchind_stmt => {
403            let mut inner = specific.into_inner();
404            let ptr = parse_typed_atom(
405                inner
406                    .next()
407                    .ok_or_else(|| ParseError::new("missing branchind pointer"))?,
408            )?;
409            let targets = match inner.next() {
410                Some(hint) if hint.as_rule() == Rule::edge_hint => parse_edge_hint(hint)?,
411                _ => Vec::new(),
412            };
413            Ok(Statement::BranchInd { ptr, targets, span })
414        }
415
416        Rule::switch_stmt => {
417            let mut inner = specific.into_inner();
418            let scrutinee = parse_typed_atom(
419                inner
420                    .next()
421                    .ok_or_else(|| ParseError::new("missing switch scrutinee"))?,
422            )?;
423            let mut cases = Vec::new();
424            let mut default = None;
425            for arm in inner {
426                let arm = arm
427                    .into_inner()
428                    .next()
429                    .ok_or_else(|| ParseError::new("empty switch arm"))?;
430                match arm.as_rule() {
431                    Rule::switch_case => {
432                        let mut parts = arm.into_inner();
433                        let value = parse_integer(
434                            parts
435                                .next()
436                                .ok_or_else(|| ParseError::new("missing switch case value"))?
437                                .as_str(),
438                        )?;
439                        let (target, args) = parse_branch_label(
440                            parts
441                                .next()
442                                .ok_or_else(|| ParseError::new("missing switch case target"))?,
443                        )?;
444                        cases.push((value, target, args));
445                    }
446                    Rule::switch_default => {
447                        let target = arm
448                            .into_inner()
449                            .next()
450                            .ok_or_else(|| ParseError::new("missing switch default target"))?;
451                        default = Some(parse_branch_label(target)?);
452                    }
453                    other => {
454                        return Err(ParseError::new(format!("unexpected switch arm {other:?}")));
455                    }
456                }
457            }
458            Ok(Statement::Switch {
459                scrutinee,
460                cases,
461                default,
462                span,
463            })
464        }
465
466        Rule::cbranch_stmt => {
467            let mut inner = specific.into_inner();
468            let condition = parse_typed_atom(
469                inner
470                    .next()
471                    .ok_or_else(|| ParseError::new("missing cbranch condition"))?,
472            )?;
473            let target_pair = inner
474                .next()
475                .ok_or_else(|| ParseError::new("missing cbranch target"))?;
476            let (target, target_args) = parse_branch_label(target_pair)?;
477            let fallthrough_pair = inner
478                .next()
479                .ok_or_else(|| ParseError::new("missing cbranch fallthrough"))?;
480            let (fallthrough, fallthrough_args) = parse_branch_label(fallthrough_pair)?;
481            Ok(Statement::CBranch {
482                condition,
483                target,
484                target_args,
485                fallthrough,
486                fallthrough_args,
487                span,
488            })
489        }
490
491        Rule::call_stmt => {
492            let mut specific_inner = specific.into_inner();
493            let form = specific_inner
494                .next()
495                .ok_or_else(|| ParseError::new("missing call form"))?;
496            let mut target = None;
497            let mut args = Vec::new();
498            match form.as_rule() {
499                Rule::call_direct => {
500                    for part in form.into_inner() {
501                        match part.as_rule() {
502                            Rule::ident | Rule::minted_callee => target = Some(parse_callee(part)?),
503                            Rule::call_arg => {
504                                let mut inner = part.into_inner();
505                                let name = inner
506                                    .next()
507                                    .ok_or_else(|| ParseError::new("missing call arg name"))?
508                                    .as_str()
509                                    .to_owned();
510                                let atom =
511                                    parse_typed_atom(inner.next().ok_or_else(|| {
512                                        ParseError::new("missing call arg value")
513                                    })?)?;
514                                args.push((name, atom));
515                            }
516                            _ => {}
517                        }
518                    }
519                }
520                Rule::call_legacy => {
521                    let label = label_value(
522                        form.into_inner()
523                            .next()
524                            .ok_or_else(|| ParseError::new("missing call target"))?
525                            .into_inner()
526                            .next()
527                            .ok_or_else(|| ParseError::new("missing call target content"))?,
528                    )?;
529                    match label {
530                        Label::Named { name, .. } => target = Some(Callee::Named(name)),
531                        Label::Address { .. } => {
532                            return Err(ParseError::new(
533                                "call with address target is not supported",
534                            ));
535                        }
536                    }
537                }
538                _ => {}
539            }
540            let targets = match specific_inner.next() {
541                Some(hint) if hint.as_rule() == Rule::edge_hint => parse_edge_hint(hint)?,
542                _ => Vec::new(),
543            };
544            Ok(Statement::Call {
545                target: target.ok_or_else(|| ParseError::new("missing call target"))?,
546                tail: false,
547                args,
548                targets,
549                span,
550            })
551        }
552
553        Rule::tailcall_stmt => {
554            let mut inner = specific.into_inner();
555            let target = parse_callee(
556                inner
557                    .next()
558                    .ok_or_else(|| ParseError::new("missing tailcall target"))?,
559            )?;
560            let args = inner
561                .filter(|part| part.as_rule() == Rule::typed_atom)
562                .map(parse_typed_atom)
563                .collect::<Result<Vec<_>, _>>()?
564                .into_iter()
565                .enumerate()
566                .map(|(i, atom)| (format!("arg{i}"), atom))
567                .collect();
568            Ok(Statement::Call {
569                target,
570                tail: true,
571                args,
572                targets: Vec::new(),
573                span,
574            })
575        }
576
577        Rule::callind_stmt => {
578            let mut inner = specific.into_inner();
579            let ptr = parse_typed_atom(
580                inner
581                    .next()
582                    .ok_or_else(|| ParseError::new("missing callind pointer"))?,
583            )?;
584            let mut args = Vec::new();
585            let mut targets = Vec::new();
586            for part in inner {
587                match part.as_rule() {
588                    Rule::callind_args => {
589                        for atom in part.into_inner() {
590                            args.push(parse_typed_atom(atom)?);
591                        }
592                    }
593                    Rule::edge_hint => targets = parse_edge_hint(part)?,
594                    _ => {}
595                }
596            }
597            Ok(Statement::CallInd {
598                ptr,
599                args,
600                targets,
601                span,
602            })
603        }
604
605        Rule::badinsn_stmt => Ok(Statement::BadInsn { span }),
606
607        Rule::return_stmt => {
608            let mut inner = specific.into_inner();
609            let ret = inner
610                .next()
611                .ok_or_else(|| ParseError::new("missing return body"))?;
612            match ret.as_rule() {
613                Rule::return_at_stmt => {
614                    let ptr = parse_typed_atom(
615                        ret.into_inner()
616                            .find(|p| p.as_rule() == Rule::typed_atom)
617                            .ok_or_else(|| ParseError::new("missing return pointer"))?,
618                    )?;
619                    Ok(Statement::Return {
620                        ptr,
621                        value: None,
622                        span,
623                    })
624                }
625                Rule::return_value_at_stmt => {
626                    let mut atoms = ret.into_inner().filter(|p| p.as_rule() == Rule::typed_atom);
627                    let value = parse_typed_atom(
628                        atoms
629                            .next()
630                            .ok_or_else(|| ParseError::new("missing return value"))?,
631                    )?;
632                    let ptr = parse_typed_atom(
633                        atoms
634                            .next()
635                            .ok_or_else(|| ParseError::new("missing return pointer"))?,
636                    )?;
637                    Ok(Statement::Return {
638                        ptr,
639                        value: Some(value),
640                        span,
641                    })
642                }
643                Rule::return_value_stmt => {
644                    let value = parse_typed_atom(
645                        ret.into_inner()
646                            .find(|p| p.as_rule() == Rule::typed_atom)
647                            .ok_or_else(|| ParseError::new("missing return value"))?,
648                    )?;
649                    Ok(Statement::ReturnValue { value, span })
650                }
651                _ => Err(ParseError::new("invalid return")),
652            }
653        }
654
655        _ => Err(ParseError::new("invalid terminator")),
656    }
657}
658
659fn parse_assert_stmt(pair: Pair<'_, Rule>) -> Result<Statement, ParseError> {
660    let span = source_span(pair.as_span());
661    let mut inner = pair.into_inner();
662    let condition = parse_typed_atom(
663        inner
664            .next()
665            .ok_or_else(|| ParseError::new("missing assert condition"))?,
666    )?;
667    Ok(Statement::Assert { condition, span })
668}
669
670fn parse_assignment_ssa(pair: Pair<'_, Rule>) -> Result<Statement, ParseError> {
671    let span = source_span(pair.as_span());
672    let mut inner = pair.into_inner();
673    // Optional declared type (`decl_ty`): an `iN`/`fN` size, or a `Foo*` struct
674    // pointer. Only the struct-pointer form is recorded (it retypes the result).
675    let name_or_ty = inner
676        .next()
677        .ok_or_else(|| ParseError::new("missing ssa assignment name"))?;
678    let (decl_struct_ptr, ssa_pair) = if name_or_ty.as_rule() == Rule::decl_ty {
679        let decl = name_or_ty
680            .into_inner()
681            .next()
682            .ok_or_else(|| ParseError::new("empty declared type"))?;
683        let decl_struct_ptr = match decl.as_rule() {
684            Rule::struct_ptr_ty => Some(decl.as_str().trim_end_matches('*').to_owned()),
685            _ => None,
686        };
687        (
688            decl_struct_ptr,
689            inner
690                .next()
691                .ok_or_else(|| ParseError::new("missing ssa name after type"))?,
692        )
693    } else {
694        (None, name_or_ty)
695    };
696    let name_span = source_span(ssa_pair.as_span());
697    let name = ssa_pair
698        .as_str()
699        .strip_prefix('%')
700        .ok_or_else(|| ParseError::new("ssa name missing % prefix"))?
701        .to_owned();
702    let expr_pair = inner
703        .find(|p| p.as_rule() == Rule::expr)
704        .ok_or_else(|| ParseError::new("missing ssa assignment expression"))?;
705    Ok(Statement::Assign {
706        name,
707        name_span,
708        expr: parse_expr(expr_pair)?,
709        decl_struct_ptr,
710        span,
711    })
712}
713
714fn parse_expr(pair: Pair<'_, Rule>) -> Result<ExprNode, ParseError> {
715    let inner = pair
716        .into_inner()
717        .next()
718        .ok_or_else(|| ParseError::new("missing expression"))?;
719
720    match inner.as_rule() {
721        Rule::atom_expr => {
722            let typed_atom = inner
723                .into_inner()
724                .next()
725                .ok_or_else(|| ParseError::new("missing atom"))?;
726            Ok(ExprNode::Atom(parse_typed_atom(typed_atom)?))
727        }
728        Rule::unop => parse_unop(inner),
729        Rule::func_unop => parse_func_unop(inner),
730        Rule::func_call => parse_func_call(inner),
731        Rule::intrinsic_call => parse_intrinsic_call(inner),
732        Rule::apply => parse_apply(inner),
733        Rule::scan => parse_scan(inner),
734        Rule::map => parse_map(inner),
735        Rule::binary => parse_binary(inner),
736        Rule::memory => parse_memory(inner),
737        Rule::cast => parse_cast(inner),
738        Rule::tuple => parse_tuple(inner),
739        Rule::extract => parse_extract(inner),
740        Rule::gep => parse_gep(inner),
741        Rule::range => parse_range(inner),
742        _ => Err(ParseError::new("invalid expression")),
743    }
744}
745
746fn parse_apply(pair: Pair<'_, Rule>) -> Result<ExprNode, ParseError> {
747    let mut target = None;
748    let mut args = Vec::new();
749
750    for part in pair.into_inner() {
751        match part.as_rule() {
752            Rule::ident | Rule::minted_callee => target = Some(parse_callee(part)?),
753            Rule::typed_atom => args.push(parse_typed_atom(part)?),
754            _ => {}
755        }
756    }
757
758    Ok(ExprNode::Apply {
759        target: target.ok_or_else(|| ParseError::new("missing apply target"))?,
760        args,
761    })
762}
763
764fn parse_callee(pair: Pair<'_, Rule>) -> Result<Callee, ParseError> {
765    match pair.as_rule() {
766        Rule::ident => Ok(Callee::Named(pair.as_str().to_owned())),
767        Rule::minted_callee => {
768            let slot = pair
769                .as_str()
770                .strip_prefix("<minted:")
771                .and_then(|text| text.strip_suffix('>'))
772                .ok_or_else(|| ParseError::new("invalid minted callee"))?
773                .parse::<u32>()
774                .map_err(|_| ParseError::new("minted callee slot exceeds u32"))?;
775            Ok(Callee::Minted(slot))
776        }
777        _ => Err(ParseError::new("invalid callee")),
778    }
779}
780
781fn parse_range(pair: Pair<'_, Rule>) -> Result<ExprNode, ParseError> {
782    let mut src = None;
783    let mut start = None;
784    let mut end = None;
785    for part in pair.into_inner() {
786        match part.as_rule() {
787            Rule::typed_atom => src = Some(parse_typed_atom(part)?),
788            Rule::range_start => start = Some(parse_integer(part.as_str().trim())?),
789            Rule::range_end => end = Some(parse_integer(part.as_str().trim())?),
790            _ => {}
791        }
792    }
793    let src = src.ok_or_else(|| ParseError::new("missing range source"))?;
794    Ok(ExprNode::Range { src, start, end })
795}
796
797fn parse_gep(pair: Pair<'_, Rule>) -> Result<ExprNode, ParseError> {
798    let inner = pair
799        .into_inner()
800        .next()
801        .ok_or_else(|| ParseError::new("missing gep body"))?;
802    let rule = inner.as_rule();
803    let mut parts = inner.into_inner();
804    let base = parse_typed_atom(
805        parts
806            .find(|p| p.as_rule() == Rule::typed_atom)
807            .ok_or_else(|| ParseError::new("missing gep base"))?,
808    )?;
809    let field = match rule {
810        Rule::named_gep => GepField::Name(
811            parts
812                .find(|p| p.as_rule() == Rule::ident)
813                .ok_or_else(|| ParseError::new("missing gep field name"))?
814                .as_str()
815                .to_owned(),
816        ),
817        Rule::offset_gep => GepField::Offset(parse_integer(
818            parts
819                .find(|p| p.as_rule() == Rule::integer)
820                .ok_or_else(|| ParseError::new("missing gep offset"))?
821                .as_str(),
822        )?),
823        _ => return Err(ParseError::new("invalid gep")),
824    };
825    Ok(ExprNode::Gep { base, field })
826}
827
828fn parse_tuple(pair: Pair<'_, Rule>) -> Result<ExprNode, ParseError> {
829    let inner = pair
830        .into_inner()
831        .next()
832        .ok_or_else(|| ParseError::new("missing tuple body"))?;
833    let fields = match inner.as_rule() {
834        Rule::pack_tuple => inner
835            .into_inner()
836            .filter(|p| p.as_rule() == Rule::tuple_field)
837            .map(|field| {
838                let mut parts = field.into_inner();
839                let name = parts
840                    .find(|p| p.as_rule() == Rule::ident)
841                    .ok_or_else(|| ParseError::new("missing tuple field name"))?
842                    .as_str()
843                    .to_owned();
844                let value = parse_typed_atom(
845                    parts
846                        .find(|p| p.as_rule() == Rule::typed_atom)
847                        .ok_or_else(|| ParseError::new("missing tuple field value"))?,
848                )?;
849                Ok(TupleField {
850                    name: Some(name),
851                    value,
852                })
853            })
854            .collect::<Result<Vec<_>, _>>()?,
855        Rule::positional_tuple => inner
856            .into_inner()
857            .filter(|p| p.as_rule() == Rule::typed_atom)
858            .enumerate()
859            .map(|(i, value)| {
860                Ok(TupleField {
861                    name: Some(format!("field{}", i + 1)),
862                    value: parse_typed_atom(value)?,
863                })
864            })
865            .collect::<Result<Vec<_>, _>>()?,
866        _ => return Err(ParseError::new("invalid tuple")),
867    };
868    if fields.is_empty() {
869        return Err(ParseError::new("tuple must have at least one field"));
870    }
871    Ok(ExprNode::Tuple { fields })
872}
873
874fn parse_extract(pair: Pair<'_, Rule>) -> Result<ExprNode, ParseError> {
875    let inner = pair
876        .into_inner()
877        .next()
878        .ok_or_else(|| ParseError::new("missing extract body"))?;
879    let rule = inner.as_rule();
880    let mut parts = inner.into_inner();
881    let agg = parse_typed_atom(
882        parts
883            .find(|p| p.as_rule() == Rule::typed_atom)
884            .ok_or_else(|| ParseError::new("missing extract aggregate"))?,
885    )?;
886    let field = match rule {
887        Rule::named_extract => ExtractField::Name(
888            parts
889                .find(|p| p.as_rule() == Rule::ident)
890                .ok_or_else(|| ParseError::new("missing extract field name"))?
891                .as_str()
892                .to_owned(),
893        ),
894        Rule::indexed_extract => ExtractField::Index(parse_integer(
895            parts
896                .find(|p| p.as_rule() == Rule::integer)
897                .ok_or_else(|| ParseError::new("missing extract index"))?
898                .as_str(),
899        )?),
900        _ => return Err(ParseError::new("invalid extract")),
901    };
902    Ok(ExprNode::Extract { agg, field })
903}
904
905fn parse_unop(pair: Pair<'_, Rule>) -> Result<ExprNode, ParseError> {
906    let mut inner = pair.into_inner();
907    let op = inner
908        .next()
909        .ok_or_else(|| ParseError::new("missing unary operator"))?
910        .as_str()
911        .to_owned();
912    let src = parse_typed_atom(
913        inner
914            .next()
915            .ok_or_else(|| ParseError::new("missing unary source"))?,
916    )?;
917
918    Ok(ExprNode::Unop { op, src })
919}
920
921fn parse_func_unop(pair: Pair<'_, Rule>) -> Result<ExprNode, ParseError> {
922    let mut inner = pair.into_inner();
923    let op = inner
924        .next()
925        .ok_or_else(|| ParseError::new("missing unary function"))?
926        .as_str()
927        .to_owned();
928    let src = parse_typed_atom(
929        inner
930            .next()
931            .ok_or_else(|| ParseError::new("missing unary function source"))?,
932    )?;
933
934    Ok(ExprNode::Unop { op, src })
935}
936
937fn parse_func_call(pair: Pair<'_, Rule>) -> Result<ExprNode, ParseError> {
938    let mut op = None;
939    let mut args = Vec::new();
940
941    for part in pair.into_inner() {
942        match part.as_rule() {
943            Rule::func_ident => op = Some(part.as_str().to_owned()),
944            Rule::typed_atom => args.push(parse_typed_atom(part)?),
945            _ => {}
946        }
947    }
948
949    Ok(ExprNode::FuncCall {
950        op: op.ok_or_else(|| ParseError::new("missing function name"))?,
951        args,
952    })
953}
954
955fn parse_intrinsic_call(pair: Pair<'_, Rule>) -> Result<ExprNode, ParseError> {
956    let mut name = None;
957    let mut args = Vec::new();
958
959    for part in pair.into_inner() {
960        match part.as_rule() {
961            // Strip the leading `$` sigil.
962            Rule::intrinsic_name => name = Some(part.as_str()[1..].to_owned()),
963            Rule::typed_atom => args.push(parse_typed_atom(part)?),
964            _ => {}
965        }
966    }
967
968    Ok(ExprNode::Intrinsic {
969        name: name.ok_or_else(|| ParseError::new("missing intrinsic name"))?,
970        args,
971    })
972}
973
974fn parse_map(pair: Pair<'_, Rule>) -> Result<ExprNode, ParseError> {
975    let mut inner = pair.into_inner();
976    // `map_app` holds the body callee and any parenthesized captures.
977    let app = inner
978        .next()
979        .filter(|p| p.as_rule() == Rule::map_app)
980        .ok_or_else(|| ParseError::new("missing map body"))?;
981    let mut app_parts = app.into_inner();
982    let body = app_parts
983        .next()
984        .ok_or_else(|| ParseError::new("missing map body function"))?;
985    let body = parse_callee(body)?;
986    let captures = app_parts
987        .filter(|p| p.as_rule() == Rule::typed_atom)
988        .map(parse_typed_atom)
989        .collect::<Result<Vec<_>, _>>()?;
990    let src = parse_typed_atom(
991        inner
992            .find(|p| p.as_rule() == Rule::typed_atom)
993            .ok_or_else(|| ParseError::new("missing map source"))?,
994    )?;
995    Ok(ExprNode::Map {
996        body,
997        src,
998        captures,
999    })
1000}
1001
1002fn parse_scan(pair: Pair<'_, Rule>) -> Result<ExprNode, ParseError> {
1003    let mut inner = pair.into_inner();
1004    // `scan_app` holds the `@body` symbol and any parenthesized captures.
1005    let app = inner
1006        .next()
1007        .filter(|p| p.as_rule() == Rule::scan_app)
1008        .ok_or_else(|| ParseError::new("missing scan body"))?;
1009    let mut app_parts = app.into_inner();
1010    let body = app_parts
1011        .next()
1012        .ok_or_else(|| ParseError::new("missing scan body function"))?;
1013    let body = match body.as_rule() {
1014        Rule::scan_body => Callee::Named(body.as_str().trim_start_matches('@').to_owned()),
1015        Rule::minted_callee => parse_callee(body)?,
1016        _ => return Err(ParseError::new("invalid scan body function")),
1017    };
1018    let captures = app_parts
1019        .filter(|p| p.as_rule() == Rule::typed_atom)
1020        .map(parse_typed_atom)
1021        .collect::<Result<Vec<_>, _>>()?;
1022    // After `scan_app` come two `typed_atom`s: the initial accumulator and the
1023    // scanned source array, in that order.
1024    let mut atoms = inner.filter(|p| p.as_rule() == Rule::typed_atom);
1025    let init = parse_typed_atom(
1026        atoms
1027            .next()
1028            .ok_or_else(|| ParseError::new("missing scan init"))?,
1029    )?;
1030    let src = parse_typed_atom(
1031        atoms
1032            .next()
1033            .ok_or_else(|| ParseError::new("missing scan source"))?,
1034    )?;
1035    Ok(ExprNode::Scan {
1036        body,
1037        init,
1038        src,
1039        captures,
1040    })
1041}
1042
1043fn parse_binary(pair: Pair<'_, Rule>) -> Result<ExprNode, ParseError> {
1044    let mut inner = pair.into_inner();
1045    let lhs = parse_typed_atom(inner.next().ok_or_else(|| ParseError::new("missing lhs"))?)?;
1046    let op = inner
1047        .next()
1048        .ok_or_else(|| ParseError::new("missing operator"))?
1049        .as_str()
1050        .to_owned();
1051    let rhs = parse_typed_atom(inner.next().ok_or_else(|| ParseError::new("missing rhs"))?)?;
1052
1053    Ok(ExprNode::Binary { lhs, op, rhs })
1054}
1055
1056fn parse_memory(pair: Pair<'_, Rule>) -> Result<ExprNode, ParseError> {
1057    let inner = pair
1058        .into_inner()
1059        .next()
1060        .ok_or_else(|| ParseError::new("missing memory expression"))?;
1061
1062    match inner.as_rule() {
1063        Rule::load => {
1064            let mut space = None;
1065            let mut size_bytes = None;
1066            let mut ptr = None;
1067            for part in inner.into_inner() {
1068                match part.as_rule() {
1069                    Rule::mem_loc => {
1070                        let (name, bytes) = parse_mem_loc(part)?;
1071                        space = Some(name);
1072                        size_bytes = Some(bytes);
1073                    }
1074                    Rule::typed_atom => ptr = Some(parse_typed_atom(part)?),
1075                    _ => {}
1076                }
1077            }
1078            Ok(ExprNode::Load {
1079                space: space.ok_or_else(|| ParseError::new("missing load space"))?,
1080                size_bytes: size_bytes.ok_or_else(|| ParseError::new("missing load size"))?,
1081                ptr: ptr.ok_or_else(|| ParseError::new("missing load pointer"))?,
1082            })
1083        }
1084        Rule::store => {
1085            let mut space = None;
1086            let mut size_bytes = None;
1087            let mut atoms = Vec::new();
1088            for part in inner.into_inner() {
1089                match part.as_rule() {
1090                    Rule::mem_loc => {
1091                        let (name, bytes) = parse_mem_loc(part)?;
1092                        space = Some(name);
1093                        size_bytes = Some(bytes);
1094                    }
1095                    Rule::typed_atom => atoms.push(parse_typed_atom(part)?),
1096                    _ => {}
1097                }
1098            }
1099            let mut atoms = atoms.into_iter();
1100            let ptr = atoms
1101                .next()
1102                .ok_or_else(|| ParseError::new("missing store pointer"))?;
1103            let src = atoms
1104                .next()
1105                .ok_or_else(|| ParseError::new("missing store source"))?;
1106            Ok(ExprNode::Store {
1107                space: space.ok_or_else(|| ParseError::new("missing store space"))?,
1108                size_bytes: size_bytes.ok_or_else(|| ParseError::new("missing store size"))?,
1109                ptr,
1110                src,
1111            })
1112        }
1113        _ => Err(ParseError::new("invalid memory expression")),
1114    }
1115}
1116
1117/// Parse a `space:bytes` memory location into `(space_name, byte_size)`.
1118fn parse_mem_loc(pair: Pair<'_, Rule>) -> Result<(String, usize), ParseError> {
1119    let mut name = None;
1120    let mut bytes = None;
1121    for part in pair.into_inner() {
1122        match part.as_rule() {
1123            Rule::mem_space => name = Some(part.as_str().to_owned()),
1124            Rule::integer => bytes = Some(parse_integer(part.as_str())? as usize),
1125            _ => {}
1126        }
1127    }
1128    Ok((
1129        name.ok_or_else(|| ParseError::new("missing space name"))?,
1130        bytes.ok_or_else(|| ParseError::new("missing space size"))?,
1131    ))
1132}
1133
1134fn parse_cast(pair: Pair<'_, Rule>) -> Result<ExprNode, ParseError> {
1135    let mut op = None;
1136    let mut ty = None;
1137    let mut src = None;
1138
1139    for part in pair.into_inner() {
1140        match part.as_rule() {
1141            Rule::cast_op => {
1142                op = Some(match part.as_str() {
1143                    "zext" => CastOp::Zext,
1144                    "sext" => CastOp::Sext,
1145                    "int2float" => CastOp::IntToFloat,
1146                    "float2float" => CastOp::FloatToFloat,
1147                    "trunc" => CastOp::Trunc,
1148                    _ => return Err(ParseError::new("invalid cast operation")),
1149                })
1150            }
1151            Rule::ty => ty = Some(parse_size_bytes(part.as_str(), "cast")?),
1152            Rule::typed_atom => src = Some(parse_typed_atom(part)?),
1153            _ => {}
1154        }
1155    }
1156
1157    Ok(ExprNode::Cast {
1158        op: op.ok_or_else(|| ParseError::new("missing cast op"))?,
1159        size_bytes: ty.ok_or_else(|| ParseError::new("missing cast type"))?,
1160        src: src.ok_or_else(|| ParseError::new("missing cast source"))?,
1161    })
1162}
1163
1164fn parse_typed_atom(pair: Pair<'_, Rule>) -> Result<TypedAtom, ParseError> {
1165    let mut size_bytes = None;
1166    let mut atom = None;
1167    let mut span = None;
1168
1169    for part in pair.into_inner() {
1170        match part.as_rule() {
1171            Rule::ty => size_bytes = Some(parse_size_bytes(part.as_str(), "typed atom")?),
1172            Rule::atom => {
1173                span = Some(source_span(part.as_span()));
1174                atom = Some(parse_atom(part)?);
1175            }
1176            _ => {}
1177        }
1178    }
1179
1180    Ok(TypedAtom {
1181        size_bytes,
1182        atom: atom.ok_or_else(|| ParseError::new("missing atom"))?,
1183        span: span.ok_or_else(|| ParseError::new("missing atom span"))?,
1184    })
1185}
1186
1187fn parse_atom(pair: Pair<'_, Rule>) -> Result<Atom, ParseError> {
1188    let inner = pair
1189        .into_inner()
1190        .next()
1191        .ok_or_else(|| ParseError::new("invalid atom"))?;
1192
1193    match inner.as_rule() {
1194        Rule::capture => {
1195            let ident = inner
1196                .into_inner()
1197                .next()
1198                .ok_or_else(|| ParseError::new("invalid capture identifier"))?
1199                .as_str()
1200                .to_owned();
1201            Ok(Atom::External(ident))
1202        }
1203        Rule::block_param_name => Ok(Atom::BlockParam(
1204            inner
1205                .as_str()
1206                .strip_prefix('@')
1207                .unwrap_or(inner.as_str())
1208                .to_owned(),
1209        )),
1210        Rule::ssa_name => Ok(Atom::Ssa(
1211            inner
1212                .as_str()
1213                .strip_prefix('%')
1214                .unwrap_or(inner.as_str())
1215                .to_owned(),
1216        )),
1217        Rule::ident => Ok(Atom::Varnode(inner.as_str().to_owned())),
1218        Rule::addressof => {
1219            let name = inner
1220                .into_inner()
1221                .next()
1222                .ok_or_else(|| ParseError::new("invalid addressof: missing identifier"))?
1223                .as_str()
1224                .to_owned();
1225            Ok(Atom::AddressOf(name))
1226        }
1227        Rule::integer => parse_integer(inner.as_str()).map(Atom::Int),
1228        Rule::bool_lit => Ok(Atom::Bool(inner.as_str() == "true")),
1229        _ => Err(ParseError::new("invalid atom")),
1230    }
1231}
1232
1233fn parse_integer(text: &str) -> Result<u64, ParseError> {
1234    if let Some(hex) = text.strip_prefix("0x").or_else(|| text.strip_prefix("0X")) {
1235        return u64::from_str_radix(hex, 16)
1236            .map_err(|_| ParseError::new("failed to parse integer literal"));
1237    }
1238
1239    text.parse::<u64>()
1240        .map_err(|_| ParseError::new("failed to parse integer literal"))
1241}
1242
1243fn parse_size_bytes(text: &str, context: &str) -> Result<usize, ParseError> {
1244    if text == "bool" {
1245        return Ok(1);
1246    }
1247    let bits = parse_size_bits(text).ok_or_else(|| {
1248        ParseError::new(format!(
1249            "invalid {context} size `{text}`; expected iNN or fNN"
1250        ))
1251    })?;
1252
1253    if bits % 8 != 0 {
1254        return Err(ParseError::new(format!(
1255            "{context} size {text} is not byte-aligned"
1256        )));
1257    }
1258
1259    Ok(bits / 8)
1260}
1261
1262fn parse_size_bits(ident: &str) -> Option<usize> {
1263    if ident.len() < 2 {
1264        return None;
1265    }
1266
1267    let mut chars = ident.chars();
1268    let prefix = chars.next()?;
1269    if prefix != 'i' && prefix != 'f' {
1270        return None;
1271    }
1272
1273    let bits = chars.as_str();
1274    if bits.chars().all(|ch| ch.is_ascii_digit()) {
1275        bits.parse::<usize>().ok()
1276    } else {
1277        None
1278    }
1279}
1280
1281fn comment_text(raw: &str) -> String {
1282    raw.strip_prefix('#').unwrap_or("").trim().to_owned()
1283}
1284
1285fn attach_comment(pending: &mut Option<String>, stmts: &mut Vec<Statement>) {
1286    if let Some(comment) = pending.take()
1287        && !stmts.is_empty()
1288    {
1289        let first = stmts.remove(0);
1290        stmts.insert(
1291            0,
1292            Statement::Commented {
1293                comment,
1294                inner: Box::new(first),
1295            },
1296        );
1297    }
1298}
1299
1300fn to_parse_error(error: pest::error::Error<Rule>) -> ParseError {
1301    ParseError::new(format!("qcode parse error: {error}"))
1302}
1303
1304fn source_span(span: pest::Span<'_>) -> SourceSpan {
1305    let start_pos = span.start_pos();
1306    let end_pos = span.end_pos();
1307    let (start_line, start_column) = start_pos.line_col();
1308    let (end_line, end_column) = end_pos.line_col();
1309
1310    SourceSpan {
1311        start: SourcePosition {
1312            offset: span.start(),
1313            line: start_line,
1314            column: start_column,
1315        },
1316        end: SourcePosition {
1317            offset: span.end(),
1318            line: end_line,
1319            column: end_column,
1320        },
1321    }
1322}
1323
1324#[cfg(test)]
1325mod tests {
1326    use super::parse_program;
1327    use crate::ast::{Atom, Callee, CastOp, ExprNode, Label, ProgramKind, Statement};
1328
1329    fn stmts(program: &str) -> Vec<Statement> {
1330        match parse_program(program).expect("parse should succeed").kind {
1331            ProgramKind::Statements(s) => s,
1332            ProgramKind::Functions { .. } => panic!("expected statements, got functions"),
1333        }
1334    }
1335
1336    #[test]
1337    fn parses_ssa_assignment_and_use() {
1338        let statements = stmts("%tmp = {v1} + 3; %tmp + 2");
1339        assert_eq!(statements.len(), 2);
1340
1341        match &statements[0] {
1342            Statement::Assign { name, expr, .. } => {
1343                assert_eq!(name, "tmp");
1344                match expr {
1345                    ExprNode::Binary { lhs, op, rhs } => {
1346                        assert_eq!(op, "+");
1347                        match &lhs.atom {
1348                            Atom::External(name) => assert_eq!(name, "v1"),
1349                            _ => panic!("expected external lhs"),
1350                        }
1351                        match &rhs.atom {
1352                            Atom::Int(value) => assert_eq!(*value, 3),
1353                            _ => panic!("expected integer rhs"),
1354                        }
1355                    }
1356                    _ => panic!("expected binary expression"),
1357                }
1358            }
1359            _ => panic!("expected assignment statement"),
1360        }
1361
1362        match &statements[1] {
1363            Statement::Expr(ExprNode::Binary { lhs, op, rhs }) => {
1364                assert_eq!(op, "+");
1365                match &lhs.atom {
1366                    Atom::Ssa(name) => assert_eq!(name, "tmp"),
1367                    _ => panic!("expected ssa lhs"),
1368                }
1369                match &rhs.atom {
1370                    Atom::Int(value) => assert_eq!(*value, 2),
1371                    _ => panic!("expected integer rhs"),
1372                }
1373            }
1374            _ => panic!("expected expression statement"),
1375        }
1376    }
1377
1378    #[test]
1379    fn parses_range_with_and_without_default_bounds() {
1380        // Explicit `[start:end]`, defaulted start `[:end]`, defaulted end
1381        // `[start:]`, and fully defaulted `[:]`.
1382        let cases = [
1383            ("%r = %a[1:4]", Some(1), Some(4)),
1384            ("%r = %a[:4]", None, Some(4)),
1385            ("%r = %a[1:]", Some(1), None),
1386            ("%r = %a[:]", None, None),
1387        ];
1388        for (src, want_start, want_end) in cases {
1389            match &stmts(src)[0] {
1390                Statement::Assign {
1391                    expr:
1392                        ExprNode::Range {
1393                            start,
1394                            end,
1395                            src: atom,
1396                        },
1397                    ..
1398                } => {
1399                    assert_eq!(*start, want_start, "start for `{src}`");
1400                    assert_eq!(*end, want_end, "end for `{src}`");
1401                    match &atom.atom {
1402                        Atom::Ssa(name) => assert_eq!(name, "a"),
1403                        _ => panic!("expected ssa range source"),
1404                    }
1405                }
1406                other => panic!("expected range expression for `{src}`, got {other:?}"),
1407            }
1408        }
1409    }
1410
1411    #[test]
1412    fn parses_map_without_captures() {
1413        let statements = stmts("%m = inc <$> %src");
1414        match &statements[0] {
1415            Statement::Assign {
1416                expr:
1417                    ExprNode::Map {
1418                        body,
1419                        src,
1420                        captures,
1421                    },
1422                ..
1423            } => {
1424                assert_eq!(body, &Callee::Named("inc".into()));
1425                assert!(captures.is_empty());
1426                match &src.atom {
1427                    Atom::Ssa(name) => assert_eq!(name, "src"),
1428                    _ => panic!("expected ssa src"),
1429                }
1430            }
1431            _ => panic!("expected map expression"),
1432        }
1433    }
1434
1435    #[test]
1436    fn parses_map_with_captures() {
1437        let statements = stmts("%m = (addk %k0 %k1) <$> %src");
1438        match &statements[0] {
1439            Statement::Assign {
1440                expr:
1441                    ExprNode::Map {
1442                        body,
1443                        src,
1444                        captures,
1445                    },
1446                ..
1447            } => {
1448                assert_eq!(body, &Callee::Named("addk".into()));
1449                assert_eq!(captures.len(), 2);
1450                match &src.atom {
1451                    Atom::Ssa(name) => assert_eq!(name, "src"),
1452                    _ => panic!("expected ssa src"),
1453                }
1454            }
1455            _ => panic!("expected map expression"),
1456        }
1457    }
1458
1459    #[test]
1460    fn parses_local_declaration() {
1461        let statements = stmts("varnode i64 ptr; ptr");
1462        assert_eq!(statements.len(), 2);
1463
1464        match &statements[0] {
1465            Statement::LocalDecl {
1466                name,
1467                display_name,
1468                size_bytes,
1469                name_span,
1470                span,
1471            } => {
1472                assert_eq!(name, "ptr");
1473                assert_eq!(display_name, "ptr");
1474                assert_eq!(*size_bytes, 8);
1475                assert_eq!(name_span.start.column, 13);
1476                assert_eq!(name_span.end.column, 16);
1477                assert_eq!(span.start.column, 1);
1478            }
1479            _ => panic!("expected local declaration"),
1480        }
1481
1482        match &statements[1] {
1483            Statement::Expr(ExprNode::Atom(atom)) => match &atom.atom {
1484                Atom::Varnode(name) => assert_eq!(name, "ptr"),
1485                _ => panic!("expected varnode atom"),
1486            },
1487            _ => panic!("expected local expression"),
1488        }
1489    }
1490
1491    #[test]
1492    fn parses_local_declaration_with_display_name() {
1493        let statements = stmts("varnode i64 ptr as PTR; ptr");
1494        assert_eq!(statements.len(), 2);
1495
1496        match &statements[0] {
1497            Statement::LocalDecl {
1498                name,
1499                display_name,
1500                size_bytes,
1501                ..
1502            } => {
1503                assert_eq!(name, "ptr");
1504                assert_eq!(display_name, "PTR");
1505                assert_eq!(*size_bytes, 8);
1506            }
1507            _ => panic!("expected local declaration"),
1508        }
1509    }
1510
1511    #[test]
1512    fn parses_cast_expression() {
1513        let statements = stmts("zext(i32, {v1})");
1514        assert_eq!(statements.len(), 1);
1515
1516        match &statements[0] {
1517            Statement::Expr(ExprNode::Cast {
1518                op,
1519                size_bytes,
1520                src,
1521            }) => {
1522                assert!(matches!(op, CastOp::Zext));
1523                assert_eq!(*size_bytes, 4);
1524                match &src.atom {
1525                    Atom::External(name) => assert_eq!(name, "v1"),
1526                    _ => panic!("expected external cast source"),
1527                }
1528            }
1529            _ => panic!("expected cast expression"),
1530        }
1531    }
1532
1533    #[test]
1534    fn parses_load_and_store_statements() {
1535        let statements = stmts("load(ram:4, {ptr}); store(ram:4, {ptr} <- {src})");
1536        assert_eq!(statements.len(), 2);
1537
1538        match &statements[0] {
1539            Statement::Expr(ExprNode::Load {
1540                space,
1541                size_bytes,
1542                ptr,
1543            }) => {
1544                assert_eq!(space, "ram");
1545                assert_eq!(*size_bytes, 4);
1546                match &ptr.atom {
1547                    Atom::External(name) => assert_eq!(name, "ptr"),
1548                    _ => panic!("expected external load pointer"),
1549                }
1550            }
1551            _ => panic!("expected load statement"),
1552        }
1553
1554        match &statements[1] {
1555            Statement::Expr(ExprNode::Store {
1556                space,
1557                size_bytes,
1558                ptr,
1559                src,
1560            }) => {
1561                assert_eq!(space, "ram");
1562                assert_eq!(*size_bytes, 4);
1563                match &ptr.atom {
1564                    Atom::External(name) => assert_eq!(name, "ptr"),
1565                    _ => panic!("expected external store pointer"),
1566                }
1567                match &src.atom {
1568                    Atom::External(name) => assert_eq!(name, "src"),
1569                    _ => panic!("expected external store source"),
1570                }
1571            }
1572            _ => panic!("expected store statement"),
1573        }
1574    }
1575
1576    #[test]
1577    fn parses_typed_binary_expression_statement() {
1578        let statements = stmts("i32 {v0} + i32 0x2");
1579        assert_eq!(statements.len(), 1);
1580
1581        match &statements[0] {
1582            Statement::Expr(ExprNode::Binary { lhs, op, rhs }) => {
1583                assert_eq!(op, "+");
1584                assert_eq!(lhs.size_bytes, Some(4));
1585                assert_eq!(rhs.size_bytes, Some(4));
1586                match &lhs.atom {
1587                    Atom::External(name) => assert_eq!(name, "v0"),
1588                    _ => panic!("expected external lhs"),
1589                }
1590                match &rhs.atom {
1591                    Atom::Int(value) => assert_eq!(*value, 2),
1592                    _ => panic!("expected integer rhs"),
1593                }
1594            }
1595            _ => panic!("expected binary expression statement"),
1596        }
1597    }
1598
1599    #[test]
1600    fn parses_float_negate_expression_statement() {
1601        let statements = stmts("f-{v0}");
1602        assert_eq!(statements.len(), 1);
1603
1604        match &statements[0] {
1605            Statement::Expr(ExprNode::Unop { op, src }) => {
1606                assert_eq!(op, "f-");
1607                match &src.atom {
1608                    Atom::External(name) => assert_eq!(name, "v0"),
1609                    _ => panic!("expected external unary source"),
1610                }
1611            }
1612            _ => panic!("expected unary expression statement"),
1613        }
1614    }
1615
1616    #[test]
1617    fn parses_float_abs_expression_statement() {
1618        let statements = stmts("abs({v0})");
1619        assert_eq!(statements.len(), 1);
1620
1621        match &statements[0] {
1622            Statement::Expr(ExprNode::Unop { op, src }) => {
1623                assert_eq!(op, "abs");
1624                match &src.atom {
1625                    Atom::External(name) => assert_eq!(name, "v0"),
1626                    _ => panic!("expected external unary source"),
1627                }
1628            }
1629            _ => panic!("expected unary expression statement"),
1630        }
1631    }
1632
1633    #[test]
1634    fn parses_float_binary_expression_statement() {
1635        let statements = stmts("{v0} f+ {v1}");
1636        assert_eq!(statements.len(), 1);
1637
1638        match &statements[0] {
1639            Statement::Expr(ExprNode::Binary { op, .. }) => assert_eq!(op, "f+"),
1640            _ => panic!("expected binary expression statement"),
1641        }
1642    }
1643
1644    #[test]
1645    fn parses_int_unary_negate_expression_statement() {
1646        let statements = stmts("-{v0}");
1647        assert_eq!(statements.len(), 1);
1648
1649        match &statements[0] {
1650            Statement::Expr(ExprNode::Unop { op, src }) => {
1651                assert_eq!(op, "-");
1652                match &src.atom {
1653                    Atom::External(name) => assert_eq!(name, "v0"),
1654                    _ => panic!("expected external unary source"),
1655                }
1656            }
1657            _ => panic!("expected unary expression statement"),
1658        }
1659    }
1660
1661    #[test]
1662    fn parses_int_signed_binary_expression_statement() {
1663        let statements = stmts("{v0} s< {v1}");
1664        assert_eq!(statements.len(), 1);
1665
1666        match &statements[0] {
1667            Statement::Expr(ExprNode::Binary { op, .. }) => assert_eq!(op, "s<"),
1668            _ => panic!("expected binary expression statement"),
1669        }
1670    }
1671
1672    #[test]
1673    fn parses_misc_single_arg_function_call() {
1674        let statements = stmts("popcount({v0})");
1675        assert_eq!(statements.len(), 1);
1676
1677        match &statements[0] {
1678            Statement::Expr(ExprNode::FuncCall { op, args }) => {
1679                assert_eq!(op, "popcount");
1680                assert_eq!(args.len(), 1);
1681            }
1682            _ => panic!("expected function call expression statement"),
1683        }
1684    }
1685
1686    #[test]
1687    fn parses_misc_two_arg_function_call() {
1688        let statements = stmts("carry({v0}, {v1})");
1689        assert_eq!(statements.len(), 1);
1690
1691        match &statements[0] {
1692            Statement::Expr(ExprNode::FuncCall { op, args }) => {
1693                assert_eq!(op, "carry");
1694                assert_eq!(args.len(), 2);
1695            }
1696            _ => panic!("expected function call expression statement"),
1697        }
1698    }
1699
1700    #[test]
1701    fn parses_ssa_assignment_chain() {
1702        let statements = stmts("i64 %a = i64 1 + i64 2; i64 %b = i64 %a + i64 5");
1703        assert_eq!(statements.len(), 2);
1704
1705        match &statements[0] {
1706            Statement::Assign {
1707                name,
1708                expr,
1709                name_span,
1710                span,
1711                ..
1712            } => {
1713                assert_eq!(name, "a");
1714                assert!(matches!(expr, ExprNode::Binary { .. }));
1715                assert_eq!(name_span.start.column, 5);
1716                assert_eq!(name_span.end.column, 7);
1717                assert_eq!(span.start.column, 1);
1718            }
1719            _ => panic!("expected ssa assignment"),
1720        }
1721
1722        match &statements[1] {
1723            Statement::Assign { name, expr, .. } => {
1724                assert_eq!(name, "b");
1725                match expr {
1726                    ExprNode::Binary { lhs, .. } => match &lhs.atom {
1727                        Atom::Ssa(name) => assert_eq!(name, "a"),
1728                        _ => panic!("expected ssa lhs referencing %a"),
1729                    },
1730                    _ => panic!("expected binary expression"),
1731                }
1732            }
1733            _ => panic!("expected ssa assignment"),
1734        }
1735    }
1736
1737    #[test]
1738    fn parses_label_decl_named() {
1739        let statements = stmts("<entry>");
1740        assert_eq!(statements.len(), 1);
1741
1742        match &statements[0] {
1743            Statement::LabelDecl {
1744                label:
1745                    Label::Named {
1746                        name,
1747                        span: label_span,
1748                        ..
1749                    },
1750                span,
1751            } => {
1752                assert_eq!(name, "entry");
1753                assert_eq!(label_span.start.column, 2);
1754                assert_eq!(label_span.end.column, 7);
1755                assert_eq!(span.start.column, 1);
1756            }
1757            _ => panic!("expected named label declaration"),
1758        }
1759    }
1760
1761    #[test]
1762    fn parses_label_decl_address() {
1763        let statements = stmts("<0x1000>");
1764        assert_eq!(statements.len(), 1);
1765
1766        match &statements[0] {
1767            Statement::LabelDecl {
1768                label: Label::Address { value: addr, .. },
1769                ..
1770            } => assert_eq!(*addr, 0x1000),
1771            _ => panic!("expected address label declaration"),
1772        }
1773    }
1774
1775    #[test]
1776    fn parses_branch_named() {
1777        let statements = stmts("goto <done>");
1778        assert_eq!(statements.len(), 1);
1779
1780        match &statements[0] {
1781            Statement::Branch {
1782                target: Label::Named { name, span, .. },
1783                ..
1784            } => {
1785                assert_eq!(name, "done");
1786                assert_eq!(span.start.column, 7);
1787            }
1788            _ => panic!("expected named branch statement"),
1789        }
1790    }
1791
1792    #[test]
1793    fn parses_branch_address() {
1794        let statements = stmts("goto <0x1001>");
1795        assert_eq!(statements.len(), 1);
1796
1797        match &statements[0] {
1798            Statement::Branch {
1799                target: Label::Address { value: addr, .. },
1800                ..
1801            } => assert_eq!(*addr, 0x1001),
1802            _ => panic!("expected address branch statement"),
1803        }
1804    }
1805
1806    #[test]
1807    fn parses_branchind() {
1808        let statements = stmts("goto [{ptr}]");
1809        assert_eq!(statements.len(), 1);
1810
1811        match &statements[0] {
1812            Statement::BranchInd { ptr, .. } => match &ptr.atom {
1813                Atom::External(name) => assert_eq!(name, "ptr"),
1814                _ => panic!("expected external pointer"),
1815            },
1816            _ => panic!("expected branchind statement"),
1817        }
1818    }
1819
1820    #[test]
1821    fn parses_cbranch() {
1822        let statements = stmts("if {cond} goto <then_lbl> else goto <else_lbl>");
1823        assert_eq!(statements.len(), 1);
1824
1825        match &statements[0] {
1826            Statement::CBranch {
1827                condition,
1828                target,
1829                fallthrough,
1830                ..
1831            } => {
1832                match &condition.atom {
1833                    Atom::External(name) => assert_eq!(name, "cond"),
1834                    _ => panic!("expected external condition"),
1835                }
1836                assert!(matches!(target, Label::Named { name, .. } if name == "then_lbl"));
1837                assert!(matches!(fallthrough, Label::Named { name, .. } if name == "else_lbl"));
1838            }
1839            _ => panic!("expected cbranch statement"),
1840        }
1841    }
1842
1843    #[test]
1844    fn parses_call() {
1845        let statements = stmts("call <target>");
1846        assert_eq!(statements.len(), 1);
1847
1848        match &statements[0] {
1849            Statement::Call { target, args, .. } => {
1850                assert_eq!(target, &Callee::Named("target".into()));
1851                assert!(args.is_empty());
1852            }
1853            _ => panic!("expected call statement"),
1854        }
1855    }
1856
1857    #[test]
1858    fn parses_call_with_args() {
1859        let statements = stmts("call fn callee(@arg0={x}, @arg1={y})");
1860        assert_eq!(statements.len(), 1);
1861
1862        match &statements[0] {
1863            Statement::Call { target, args, .. } => {
1864                assert_eq!(target, &Callee::Named("callee".into()));
1865                assert_eq!(args.len(), 2);
1866                assert_eq!(args[0].0, "@arg0");
1867                assert_eq!(args[1].0, "@arg1");
1868            }
1869            _ => panic!("expected call statement"),
1870        }
1871    }
1872
1873    #[test]
1874    fn parses_call_with_edge_hint() {
1875        let statements = stmts("call fn callee(@p={x}) // -> <resume>");
1876        assert_eq!(statements.len(), 1);
1877
1878        match &statements[0] {
1879            Statement::Call {
1880                target, targets, ..
1881            } => {
1882                assert_eq!(target, &Callee::Named("callee".into()));
1883                assert_eq!(targets.len(), 1);
1884                assert!(matches!(&targets[0], Label::Named { name, .. } if name == "resume"));
1885            }
1886            _ => panic!("expected call statement"),
1887        }
1888    }
1889
1890    #[test]
1891    fn parses_branchind_with_edge_hint() {
1892        let statements = stmts("goto [{p}] // -> <a>, <b>");
1893        assert_eq!(statements.len(), 1);
1894
1895        match &statements[0] {
1896            Statement::BranchInd { targets, .. } => {
1897                let names: Vec<&str> = targets
1898                    .iter()
1899                    .map(|t| match t {
1900                        Label::Named { name, .. } => name.as_str(),
1901                        _ => panic!("expected named target"),
1902                    })
1903                    .collect();
1904                assert_eq!(names, ["a", "b"]);
1905            }
1906            _ => panic!("expected branchind statement"),
1907        }
1908    }
1909
1910    #[test]
1911    fn parses_callind_with_args() {
1912        let statements = stmts("call [{ptr}]({a}, {b})");
1913        assert_eq!(statements.len(), 1);
1914
1915        match &statements[0] {
1916            Statement::CallInd { args, .. } => assert_eq!(args.len(), 2),
1917            _ => panic!("expected callind statement"),
1918        }
1919    }
1920
1921    #[test]
1922    fn parses_callind() {
1923        let statements = stmts("call [{ptr}]");
1924        assert_eq!(statements.len(), 1);
1925
1926        match &statements[0] {
1927            Statement::CallInd { ptr, .. } => match &ptr.atom {
1928                Atom::External(name) => assert_eq!(name, "ptr"),
1929                _ => panic!("expected external pointer"),
1930            },
1931            _ => panic!("expected callind statement"),
1932        }
1933    }
1934
1935    #[test]
1936    fn parses_return() {
1937        let statements = stmts("return at {ptr}");
1938        assert_eq!(statements.len(), 1);
1939
1940        match &statements[0] {
1941            Statement::Return { ptr, .. } => match &ptr.atom {
1942                Atom::External(name) => assert_eq!(name, "ptr"),
1943                _ => panic!("expected external pointer"),
1944            },
1945            _ => panic!("expected return statement"),
1946        }
1947    }
1948
1949    #[test]
1950    fn parses_multi_block_program() {
1951        // Label prefixes the next statement without a `;` between them.
1952        let program = "%tmp = {v1} + 3; goto <done>; <done> %tmp + 2";
1953        let statements = stmts(program);
1954        assert_eq!(statements.len(), 4);
1955
1956        match &statements[0] {
1957            Statement::Assign { name, .. } => assert_eq!(name, "tmp"),
1958            _ => panic!("expected assignment"),
1959        }
1960        match &statements[1] {
1961            Statement::Branch {
1962                target: Label::Named { name, .. },
1963                ..
1964            } => assert_eq!(name, "done"),
1965            _ => panic!("expected branch"),
1966        }
1967        match &statements[2] {
1968            Statement::LabelDecl {
1969                label: Label::Named { name, .. },
1970                ..
1971            } => assert_eq!(name, "done"),
1972            _ => panic!("expected label declaration"),
1973        }
1974        match &statements[3] {
1975            Statement::Expr(ExprNode::Binary { op, .. }) => assert_eq!(op, "+"),
1976            _ => panic!("expected expression"),
1977        }
1978    }
1979
1980    #[test]
1981    fn parses_standalone_label() {
1982        // A label with no following statement is valid (e.g. as a terminal label).
1983        let statements = stmts("goto <end>; <end>");
1984        assert_eq!(statements.len(), 2);
1985
1986        match &statements[1] {
1987            Statement::LabelDecl {
1988                label: Label::Named { name, .. },
1989                ..
1990            } => assert_eq!(name, "end"),
1991            _ => panic!("expected label declaration"),
1992        }
1993    }
1994
1995    #[test]
1996    fn parses_fn_decl() {
1997        let program = parse_program("fn f: <entry> varnode i64 a; goto <done>; <done> a + 1")
1998            .expect("parse should succeed");
1999        match program.kind {
2000            ProgramKind::Functions { fns, .. } => {
2001                assert_eq!(fns.len(), 1);
2002                let f = &fns[0];
2003                assert_eq!(f.name, "f");
2004                assert_eq!(f.name_span.start.column, 4);
2005                assert_eq!(f.statements.len(), 5);
2006                assert!(
2007                    matches!(&f.statements[0], Statement::LabelDecl { label: Label::Named { name, .. }, .. } if name == "entry")
2008                );
2009                assert!(
2010                    matches!(&f.statements[1], Statement::LocalDecl { name, .. } if name == "a")
2011                );
2012                assert!(
2013                    matches!(&f.statements[2], Statement::Branch { target: Label::Named { name, .. }, .. } if name == "done")
2014                );
2015                assert!(
2016                    matches!(&f.statements[3], Statement::LabelDecl { label: Label::Named { name, .. }, .. } if name == "done")
2017                );
2018                assert!(
2019                    matches!(&f.statements[4], Statement::Expr(ExprNode::Binary { op, .. }) if op == "+")
2020                );
2021            }
2022            _ => panic!("expected function program"),
2023        }
2024    }
2025
2026    #[test]
2027    fn parses_fn_decl_with_address_labels() {
2028        let program = parse_program("fn f: <entry> varnode i64 a; goto <0x1001>")
2029            .expect("parse should succeed");
2030        match program.kind {
2031            ProgramKind::Functions { fns, .. } => {
2032                assert_eq!(fns.len(), 1);
2033                let f = &fns[0];
2034                assert!(matches!(
2035                    &f.statements[2],
2036                    Statement::Branch {
2037                        target: Label::Address { value: 0x1001, .. },
2038                        ..
2039                    }
2040                ));
2041            }
2042            _ => panic!("expected function program"),
2043        }
2044    }
2045
2046    #[test]
2047    fn parses_label_decl_with_params() {
2048        let statements = stmts("<entry @v1 @v2>");
2049        assert_eq!(statements.len(), 1);
2050        match &statements[0] {
2051            Statement::LabelDecl {
2052                label: Label::Named { name, params, .. },
2053                ..
2054            } => {
2055                assert_eq!(name, "entry");
2056                assert_eq!(params[0].name, "v1");
2057                assert_eq!(params[0].size_bytes, None);
2058                assert_eq!(params[1].name, "v2");
2059                assert_eq!(params[1].size_bytes, None);
2060            }
2061            _ => panic!("expected named label declaration with params"),
2062        }
2063    }
2064
2065    #[test]
2066    fn parses_label_decl_with_typed_params() {
2067        let statements = stmts("<entry @v1:i64 @v2:i32>");
2068        assert_eq!(statements.len(), 1);
2069        match &statements[0] {
2070            Statement::LabelDecl {
2071                label: Label::Named { name, params, .. },
2072                ..
2073            } => {
2074                assert_eq!(name, "entry");
2075                assert_eq!(params[0].name, "v1");
2076                assert_eq!(params[0].size_bytes, Some(8));
2077                assert_eq!(params[1].name, "v2");
2078                assert_eq!(params[1].size_bytes, Some(4));
2079            }
2080            _ => panic!("expected named label declaration with typed params"),
2081        }
2082    }
2083
2084    #[test]
2085    fn parses_branch_with_args() {
2086        let statements = stmts("goto <done @v1=1 @v2=%x>");
2087        assert_eq!(statements.len(), 1);
2088        match &statements[0] {
2089            Statement::Branch {
2090                target: Label::Named { name, .. },
2091                args,
2092                ..
2093            } => {
2094                assert_eq!(name, "done");
2095                assert_eq!(args.len(), 2);
2096                assert_eq!(args[0].0, "v1");
2097                assert!(matches!(args[0].1.atom, Atom::Int(1)));
2098                assert_eq!(args[1].0, "v2");
2099                assert!(matches!(&args[1].1.atom, Atom::Ssa(n) if n == "x"));
2100            }
2101            _ => panic!("expected branch with args"),
2102        }
2103    }
2104
2105    #[test]
2106    fn parses_cbranch_with_args() {
2107        let statements = stmts("if %c goto <then_lbl @x=1> else goto <else_lbl @y=%v>");
2108        assert_eq!(statements.len(), 1);
2109        match &statements[0] {
2110            Statement::CBranch {
2111                condition,
2112                target,
2113                target_args,
2114                fallthrough,
2115                fallthrough_args,
2116                ..
2117            } => {
2118                assert!(matches!(&condition.atom, Atom::Ssa(n) if n == "c"));
2119                assert!(matches!(target, Label::Named { name, .. } if name == "then_lbl"));
2120                assert_eq!(target_args.len(), 1);
2121                assert_eq!(target_args[0].0, "x");
2122                assert!(matches!(target_args[0].1.atom, Atom::Int(1)));
2123                assert!(matches!(fallthrough, Label::Named { name, .. } if name == "else_lbl"));
2124                assert_eq!(fallthrough_args.len(), 1);
2125                assert_eq!(fallthrough_args[0].0, "y");
2126                assert!(matches!(&fallthrough_args[0].1.atom, Atom::Ssa(n) if n == "v"));
2127            }
2128            _ => panic!("expected cbranch with args"),
2129        }
2130    }
2131
2132    #[test]
2133    fn parses_switch_with_cases_and_default() {
2134        let statements =
2135            stmts("switch %idx { 0x0 => <a_lbl>, 0x3 => <b_lbl @p=%v>, default => <d_lbl> }");
2136        assert_eq!(statements.len(), 1);
2137        match &statements[0] {
2138            Statement::Switch {
2139                scrutinee,
2140                cases,
2141                default,
2142                ..
2143            } => {
2144                assert!(matches!(&scrutinee.atom, Atom::Ssa(n) if n == "idx"));
2145                assert_eq!(cases.len(), 2);
2146
2147                assert_eq!(cases[0].0, 0);
2148                assert!(matches!(&cases[0].1, Label::Named { name, .. } if name == "a_lbl"));
2149                assert!(cases[0].2.is_empty());
2150
2151                // A case arm carries block arguments just as a `goto` does.
2152                assert_eq!(cases[1].0, 3);
2153                assert!(matches!(&cases[1].1, Label::Named { name, .. } if name == "b_lbl"));
2154                assert_eq!(cases[1].2.len(), 1);
2155                assert_eq!(cases[1].2[0].0, "p");
2156                assert!(matches!(&cases[1].2[0].1.atom, Atom::Ssa(n) if n == "v"));
2157
2158                let (default_label, default_args) = default.as_ref().expect("default arm");
2159                assert!(matches!(default_label, Label::Named { name, .. } if name == "d_lbl"));
2160                assert!(default_args.is_empty());
2161            }
2162            other => panic!("expected switch, got {other:?}"),
2163        }
2164    }
2165
2166    /// A jump table guarded by a bounds check is total over the values it lists,
2167    /// so the default arm is optional.
2168    #[test]
2169    fn parses_switch_without_default() {
2170        let statements = stmts("switch %idx { 0x0 => <a_lbl>, 0x1 => <b_lbl> }");
2171        match &statements[0] {
2172            Statement::Switch { cases, default, .. } => {
2173                assert_eq!(cases.len(), 2);
2174                assert!(default.is_none());
2175            }
2176            other => panic!("expected switch, got {other:?}"),
2177        }
2178    }
2179
2180    #[test]
2181    fn parses_comment_attached_to_next_statement() {
2182        let statements = stmts("# add one\n%x = {v} + 1");
2183        assert_eq!(statements.len(), 1);
2184        match &statements[0] {
2185            Statement::Commented { comment, inner } => {
2186                assert_eq!(comment, "add one");
2187                assert!(matches!(inner.as_ref(), Statement::Assign { name, .. } if name == "x"));
2188            }
2189            _ => panic!("expected commented statement"),
2190        }
2191    }
2192
2193    #[test]
2194    fn parses_comment_stripped_for_standalone_expression() {
2195        let statements = stmts("# first\n%a = 1; # second\n%b = 2");
2196        assert_eq!(statements.len(), 2);
2197        match &statements[0] {
2198            Statement::Commented { comment, .. } => assert_eq!(comment, "first"),
2199            _ => panic!("expected first statement to be commented"),
2200        }
2201        match &statements[1] {
2202            Statement::Commented { comment, .. } => assert_eq!(comment, "second"),
2203            _ => panic!("expected second statement to be commented"),
2204        }
2205    }
2206
2207    #[test]
2208    fn parses_uncommented_statements_unaffected() {
2209        let statements = stmts("%x = 1 + 2");
2210        assert_eq!(statements.len(), 1);
2211        assert!(matches!(&statements[0], Statement::Assign { name, .. } if name == "x"));
2212    }
2213
2214    #[test]
2215    fn parses_comment_in_fn_body() {
2216        let program = parse_program("fn f: <entry> # load value\n%x = {v} + 1; return at %x")
2217            .expect("parse should succeed");
2218        match program.kind {
2219            ProgramKind::Functions { fns, .. } => {
2220                let stmts = &fns[0].statements;
2221                // stmts[0] = LabelDecl <entry>, stmts[1] = Commented(%x = ...), stmts[2] = return
2222                assert!(
2223                    matches!(&stmts[1], Statement::Commented { comment, .. } if comment == "load value")
2224                );
2225            }
2226            _ => panic!("expected function program"),
2227        }
2228    }
2229
2230    #[test]
2231    fn parses_top_level_varnode_before_fn() {
2232        let program = parse_program("varnode i64 ptr; fn f: <entry> return at ptr")
2233            .expect("parse should succeed");
2234        match program.kind {
2235            ProgramKind::Functions { varnodes, fns } => {
2236                assert_eq!(varnodes.len(), 1);
2237                assert!(
2238                    matches!(&varnodes[0], Statement::LocalDecl { name, size_bytes, .. } if name == "ptr" && *size_bytes == 8)
2239                );
2240                assert_eq!(fns.len(), 1);
2241            }
2242            _ => panic!("expected function program"),
2243        }
2244    }
2245
2246    #[test]
2247    fn parses_lambda_apply_and_value_returns() {
2248        let program = parse_program(
2249            "lambda rec: <entry @s:i64> %next = @s + 1; %out = apply rec(%next); return %out",
2250        )
2251        .expect("parse should succeed");
2252        let ProgramKind::Functions { fns, .. } = program.kind else {
2253            panic!("expected function program")
2254        };
2255        assert_eq!(fns.len(), 1);
2256        assert_eq!(fns[0].kind, crate::ast::FnKind::Lambda);
2257        assert!(matches!(
2258            &fns[0].statements[2],
2259            Statement::Assign { expr: ExprNode::Apply { target, args }, .. }
2260                if target == &Callee::Named("rec".into()) && args.len() == 1
2261        ));
2262        assert!(matches!(
2263            &fns[0].statements[3],
2264            Statement::ReturnValue { value, .. } if matches!(&value.atom, Atom::Ssa(name) if name == "out")
2265        ));
2266    }
2267
2268    #[test]
2269    fn parses_machine_return_at_forms() {
2270        let statements = stmts("return at %ptr; return %v at %ptr");
2271        assert_eq!(statements.len(), 2);
2272        assert!(matches!(
2273            &statements[0],
2274            Statement::Return { value: None, ptr, .. }
2275                if matches!(&ptr.atom, Atom::Ssa(name) if name == "ptr")
2276        ));
2277        assert!(matches!(
2278            &statements[1],
2279            Statement::Return { value: Some(value), ptr, .. }
2280                if matches!(&value.atom, Atom::Ssa(name) if name == "v")
2281                    && matches!(&ptr.atom, Atom::Ssa(name) if name == "ptr")
2282        ));
2283    }
2284
2285    #[test]
2286    fn parses_struct_decl_offsets_and_padding() {
2287        use crate::ast::{ExprNode, GepField, StructFieldType};
2288        let program =
2289            parse_program("type Foo { a: 4, _: 5, b: 2, p: Bar* }; %x + 1").expect("parse");
2290        assert_eq!(program.structs.len(), 1);
2291        let foo = &program.structs[0];
2292        assert_eq!(foo.name, "Foo");
2293        // 4 named fields minus padding = a, b, p.
2294        let named: Vec<&str> = foo
2295            .fields
2296            .iter()
2297            .filter(|f| !f.is_padding())
2298            .map(|f| f.name.as_str())
2299            .collect();
2300        assert_eq!(named, ["a", "b", "p"]);
2301        // Padding `_` advances the running offset (a@0, then 4 + 5 = 9 → b@9).
2302        assert!(matches!(foo.fields[1].ty, StructFieldType::Int(5)));
2303        assert!(foo.fields[1].is_padding());
2304        // Pointer-to-struct field type is captured.
2305        assert!(matches!(&foo.fields[3].ty, StructFieldType::StructPtr(name) if name == "Bar"));
2306
2307        // gep parses as its own expression node.
2308        let geps = parse_program("%y = gep(%p.field)").expect("parse").kind;
2309        let ProgramKind::Statements(stmts) = geps else {
2310            panic!("expected statements")
2311        };
2312        assert!(matches!(
2313            &stmts[0],
2314            Statement::Assign { expr: ExprNode::Gep { field: GepField::Name(n), .. }, .. } if n == "field"
2315        ));
2316    }
2317
2318    #[test]
2319    fn parses_canonical_minted_callees_in_all_direct_forms() {
2320        let cases = [
2321            ("%x = apply <minted:1>()", 1),
2322            ("%x = <minted:2> <$> i64 0", 2),
2323            ("%x = scanl <minted:3> i64 0 i64 1", 3),
2324            ("call fn <minted:4>()", 4),
2325            ("tailcall fn <minted:5>()", 5),
2326        ];
2327
2328        for (source, expected) in cases {
2329            let statement = stmts(source).pop().expect("one statement");
2330            let callee = match statement {
2331                Statement::Assign {
2332                    expr: ExprNode::Apply { target, .. },
2333                    ..
2334                } => target,
2335                Statement::Assign {
2336                    expr: ExprNode::Map { body, .. },
2337                    ..
2338                }
2339                | Statement::Assign {
2340                    expr: ExprNode::Scan { body, .. },
2341                    ..
2342                } => body,
2343                Statement::Call { target, .. } => target,
2344                other => panic!("unexpected parse for {source}: {other:?}"),
2345            };
2346            assert_eq!(callee, Callee::Minted(expected));
2347        }
2348
2349        assert!(matches!(
2350            stmts("tailcall fn <minted:9>()").pop(),
2351            Some(Statement::Call { tail: true, .. })
2352        ));
2353    }
2354}