Skip to main content

proof_engine/scripting/
compiler.rs

1//! Bytecode compiler — walks the AST and emits `Op` instructions into a `Proto`.
2//!
3//! # Architecture
4//! A single-pass recursive descent over the AST maintains:
5//! - A flat local-variable stack per function (slot numbers are u16).
6//! - A scope depth counter for block exits.
7//! - A list of upvalue descriptors per function (used by `Closure`).
8//! - Per-loop break/continue patch lists.
9//!
10//! Jump offsets are relative (i32): positive = forward, negative = backward.
11
12use std::collections::HashMap;
13use std::sync::Arc;
14use super::ast::*;
15use super::vm::Value as VmValue;
16
17// ── Constant pool ─────────────────────────────────────────────────────────────
18
19/// A compile-time constant value.
20#[derive(Debug, Clone, PartialEq)]
21pub enum Constant {
22    Nil,
23    Bool(bool),
24    Int(i64),
25    Float(f64),
26    Str(String),
27}
28
29// ── Op (bytecode instruction set) ────────────────────────────────────────────
30
31/// VM instruction.  Operands are embedded to allow the interpreter to avoid
32/// secondary table lookups on the hot path.
33#[allow(non_camel_case_types)]
34#[derive(Debug, Clone, PartialEq)]
35pub enum Op {
36    // ── Literals ──────────────────────────────────────────────────────────────
37    /// Push nil.
38    Nil,
39    /// Push true.
40    True,
41    /// Push false.
42    False,
43    /// Push `proto.constants[idx]`.
44    Const(u32),
45
46    // ── Stack ─────────────────────────────────────────────────────────────────
47    Pop,
48    Dup,
49    Swap,
50
51    // ── Locals ────────────────────────────────────────────────────────────────
52    GetLocal(u16),
53    SetLocal(u16),
54
55    // ── Upvalues ──────────────────────────────────────────────────────────────
56    GetUpval(u16),
57    SetUpval(u16),
58
59    // ── Globals (constant-indexed by string name) ─────────────────────────────
60    GetGlobal(u32),
61    SetGlobal(u32),
62
63    // ── Tables ────────────────────────────────────────────────────────────────
64    NewTable,
65    /// `SetField(kidx)`: pop val; table = peek; table\[const_str\] = val.
66    SetField(u32),
67    /// `GetField(kidx)`: pop table; push table\[const_str\].
68    GetField(u32),
69    /// `SetIndex`: pop val, key; table = peek; table\[key\] = val.
70    SetIndex,
71    /// `GetIndex`: pop key, table; push table\[key\].
72    GetIndex,
73    /// `TableAppend`: pop val; table = peek; append val to array part.
74    TableAppend,
75    /// `SetList(n)`: pop n values; table = peek; assign t[1..n].
76    SetList(u16),
77
78    // ── Unary ─────────────────────────────────────────────────────────────────
79    Len,
80    Neg,
81    Not,
82    BitNot,
83
84    // ── Arithmetic ────────────────────────────────────────────────────────────
85    Add, Sub, Mul, Div, IDiv, Mod, Pow,
86    Concat,         // pops 2, pushes concatenated string
87
88    // ── Comparison ────────────────────────────────────────────────────────────
89    Eq, NotEq, Lt, LtEq, Gt, GtEq,
90
91    // ── Bitwise ───────────────────────────────────────────────────────────────
92    BitAnd, BitOr, BitXor, Shl, Shr,
93
94    // ── Control flow ──────────────────────────────────────────────────────────
95    /// Relative unconditional jump.  `ip += offset` (can be negative).
96    Jump(i32),
97    /// Peek top; if truthy jump (no pop).
98    JumpIf(i32),
99    /// Peek top; if falsy jump (no pop).
100    JumpIfNot(i32),
101    /// Pop top; if falsy jump — used for short-circuit `and`.
102    JumpIfNotPop(i32),
103    /// Pop top; if truthy jump — used for short-circuit `or`.
104    JumpIfPop(i32),
105
106    // ── Calls & returns ───────────────────────────────────────────────────────
107    /// `Call(nargs, nret)`: pop nargs + callee; push nret results (0 = all).
108    Call(u8, u8),
109    /// `CallMethod(name_kidx, nargs, nret)`: obj on stack; method = const_str.
110    CallMethod(u32, u8, u8),
111    /// `Return(n)`: pop n values and return (0 = return all). `Return(255)`
112    /// returns everything pushed since the last `MarkReturn`, for
113    /// `return ..., f()` where f can return any number of values.
114    Return(u8),
115    /// Remember the stack height for a following `Return(255)`.
116    MarkReturn,
117    /// Tail-call optimisation.
118    TailCall(u8),
119
120    // ── Closures ──────────────────────────────────────────────────────────────
121    /// Create a closure from `proto.protos[idx]`, capturing upvalues.
122    Closure(u32),
123    /// Close the upvalue at local slot `slot`.
124    Close(u16),
125
126    // ── Iterators ─────────────────────────────────────────────────────────────
127    /// Prepare generic-for: push iterator state.
128    ForPrep(u16),
129    /// Advance generic-for; pop results if exhausted (implied jump offset in
130    /// combination with `ForStepJump`).
131    ForStep,
132    /// Like ForStep but with a jump offset for the exhausted case.
133    ForStepJump(i32),
134    /// Push and validate [start, limit, step] for numeric-for.
135    NumForInit,
136    /// Advance numeric-for; jump by offset if done.
137    NumForStep(i32),
138
139    // ── Varargs ───────────────────────────────────────────────────────────────
140    /// Push `n` vararg values (0 = all).
141    Vararg(u8),
142
143    // ── Debug ─────────────────────────────────────────────────────────────────
144    LineInfo(u32),
145}
146
147// ── Proto (function prototype) ────────────────────────────────────────────────
148
149/// A compiled function — the unit of bytecode.
150#[derive(Debug, Clone)]
151pub struct Proto {
152    pub name:          String,
153    pub code:          Vec<Op>,
154    pub constants:     Vec<Constant>,
155    pub protos:        Vec<Proto>,      // nested closure prototypes
156    pub param_count:   u8,
157    pub is_vararg:     bool,
158    pub upvalue_count: u16,
159    pub max_stack:     u16,
160}
161
162impl Proto {
163    fn new(name: impl Into<String>) -> Self {
164        Proto {
165            name:          name.into(),
166            code:          Vec::new(),
167            constants:     Vec::new(),
168            protos:        Vec::new(),
169            param_count:   0,
170            is_vararg:     false,
171            upvalue_count: 0,
172            max_stack:     0,
173        }
174    }
175
176    /// Add a constant, deduplicating where possible.
177    pub fn add_const(&mut self, c: Constant) -> u32 {
178        for (i, existing) in self.constants.iter().enumerate() {
179            if *existing == c { return i as u32; }
180        }
181        let idx = self.constants.len() as u32;
182        self.constants.push(c);
183        idx
184    }
185
186    fn emit(&mut self, op: Op) -> usize {
187        self.code.push(op);
188        self.code.len() - 1
189    }
190
191    fn patch_jump(&mut self, instr_idx: usize) {
192        let target = self.code.len() as i32;
193        let from   = instr_idx as i32 + 1;
194        let offset = target - from;
195        match &mut self.code[instr_idx] {
196            Op::Jump(o) | Op::JumpIf(o) | Op::JumpIfNot(o)
197            | Op::JumpIfNotPop(o) | Op::JumpIfPop(o)
198            | Op::NumForStep(o) | Op::ForStepJump(o) => *o = offset,
199            _ => {}
200        }
201    }
202}
203
204// ── Instruction (VM-facing bytecode) ─────────────────────────────────────────
205
206/// Runtime instruction set emitted by `Compiler::compile_script`.
207#[derive(Debug, Clone, PartialEq)]
208pub enum Instruction {
209    // Literals
210    LoadNil,
211    LoadBool(bool),
212    LoadInt(i64),
213    LoadFloat(f64),
214    LoadStr(String),
215    LoadConst(usize),
216    // Stack
217    Pop,
218    Dup,
219    Swap,
220    // Locals / upvalues / globals
221    GetLocal(usize),
222    SetLocal(usize),
223    GetUpvalue(usize),
224    SetUpvalue(usize),
225    GetGlobal(String),
226    SetGlobal(String),
227    // Tables
228    NewTable,
229    SetField(String),
230    GetField(String),
231    SetIndex,
232    GetIndex,
233    TableAppend,
234    // Unary
235    Len,
236    Neg,
237    Not,
238    BitNot,
239    // Arithmetic
240    Add, Sub, Mul, Div, IDiv, Mod, Pow,
241    Concat,
242    // Bitwise
243    BitAnd, BitOr, BitXor, Shl, Shr,
244    // Comparison
245    Eq, NotEq, Lt, LtEq, Gt, GtEq,
246    // Control flow
247    Jump(isize),
248    JumpIf(isize),
249    JumpIfNot(isize),
250    /// Peek; if not truthy jump (and leave value); if truthy pop and continue.
251    JumpIfNotPop(isize),
252    /// Peek; if truthy jump (and leave value); if falsy pop and continue.
253    JumpIfPop(isize),
254    JumpAbs(usize),
255    // Calls
256    /// `Call(nargs, nret)`: nret 0 keeps every result, otherwise exactly nret.
257    Call(usize, usize),
258    CallMethod(String, usize, usize),
259    Return(usize),
260    /// Return every value pushed since the last `MarkReturn`.
261    ReturnFromMark,
262    MarkReturn,
263    // Closures
264    MakeFunction(usize),
265    MakeClosure(usize, Vec<(bool, usize)>),
266    CloseUpvalue(usize),
267    // Iterators
268    ForPrep(usize),
269    /// Advance numeric for-loop: `local_idx` = loop var slot, `jump_offset` = exit jump.
270    ForStep(usize, isize),
271    Nop,
272}
273
274// ── Chunk (VM-facing function prototype) ─────────────────────────────────────
275
276/// A compiled function ready for the VM.
277#[derive(Debug, Clone)]
278pub struct Chunk {
279    pub name:         String,
280    pub instructions: Vec<Instruction>,
281    pub constants:    Vec<VmValue>,
282    pub sub_chunks:   Vec<Arc<Chunk>>,
283    pub param_count:  u8,
284    pub is_vararg:    bool,
285}
286
287fn const_to_value(c: &Constant) -> VmValue {
288    match c {
289        Constant::Nil      => VmValue::Nil,
290        Constant::Bool(b)  => VmValue::Bool(*b),
291        Constant::Int(i)   => VmValue::Int(*i),
292        Constant::Float(f) => VmValue::Float(*f),
293        Constant::Str(s)   => VmValue::Str(Arc::new(s.clone())),
294    }
295}
296
297fn proto_to_chunk(proto: &Proto) -> Arc<Chunk> {
298    let instructions = proto.code.iter()
299        .map(|op| op_to_instruction(op, &proto.constants))
300        .collect();
301    let constants = proto.constants.iter().map(const_to_value).collect();
302    let sub_chunks = proto.protos.iter().map(proto_to_chunk).collect();
303    Arc::new(Chunk {
304        name:         proto.name.clone(),
305        instructions,
306        constants,
307        sub_chunks,
308        param_count:  proto.param_count,
309        is_vararg:    proto.is_vararg,
310    })
311}
312
313fn op_to_instruction(op: &Op, constants: &[Constant]) -> Instruction {
314    let get_str = |kidx: u32| -> String {
315        match constants.get(kidx as usize) {
316            Some(Constant::Str(s)) => s.clone(),
317            _ => String::new(),
318        }
319    };
320    match op {
321        Op::Nil               => Instruction::LoadNil,
322        Op::True              => Instruction::LoadBool(true),
323        Op::False             => Instruction::LoadBool(false),
324        Op::Const(idx)        => Instruction::LoadConst(*idx as usize),
325        Op::Pop               => Instruction::Pop,
326        Op::Dup               => Instruction::Dup,
327        Op::Swap              => Instruction::Swap,
328        Op::GetLocal(s)       => Instruction::GetLocal(*s as usize),
329        Op::SetLocal(s)       => Instruction::SetLocal(*s as usize),
330        Op::GetUpval(i)       => Instruction::GetUpvalue(*i as usize),
331        Op::SetUpval(i)       => Instruction::SetUpvalue(*i as usize),
332        Op::GetGlobal(k)      => Instruction::GetGlobal(get_str(*k)),
333        Op::SetGlobal(k)      => Instruction::SetGlobal(get_str(*k)),
334        Op::NewTable          => Instruction::NewTable,
335        Op::SetField(k)       => Instruction::SetField(get_str(*k)),
336        Op::GetField(k)       => Instruction::GetField(get_str(*k)),
337        Op::SetIndex          => Instruction::SetIndex,
338        Op::GetIndex          => Instruction::GetIndex,
339        Op::TableAppend       => Instruction::TableAppend,
340        Op::SetList(_)        => Instruction::Nop,
341        Op::Len               => Instruction::Len,
342        Op::Neg               => Instruction::Neg,
343        Op::Not               => Instruction::Not,
344        Op::BitNot            => Instruction::BitNot,
345        Op::Add               => Instruction::Add,
346        Op::Sub               => Instruction::Sub,
347        Op::Mul               => Instruction::Mul,
348        Op::Div               => Instruction::Div,
349        Op::IDiv              => Instruction::IDiv,
350        Op::Mod               => Instruction::Mod,
351        Op::Pow               => Instruction::Pow,
352        Op::Concat            => Instruction::Concat,
353        Op::Eq                => Instruction::Eq,
354        Op::NotEq             => Instruction::NotEq,
355        Op::Lt                => Instruction::Lt,
356        Op::LtEq              => Instruction::LtEq,
357        Op::Gt                => Instruction::Gt,
358        Op::GtEq              => Instruction::GtEq,
359        Op::BitAnd            => Instruction::BitAnd,
360        Op::BitOr             => Instruction::BitOr,
361        Op::BitXor            => Instruction::BitXor,
362        Op::Shl               => Instruction::Shl,
363        Op::Shr               => Instruction::Shr,
364        Op::Jump(off)         => Instruction::Jump(*off as isize),
365        Op::JumpIf(off)       => Instruction::JumpIf(*off as isize),
366        Op::JumpIfNot(off)    => Instruction::JumpIfNot(*off as isize),
367        Op::JumpIfNotPop(off) => Instruction::JumpIfNotPop(*off as isize),
368        Op::JumpIfPop(off)    => Instruction::JumpIfPop(*off as isize),
369        Op::Call(na, nr)      => Instruction::Call(*na as usize, *nr as usize),
370        Op::CallMethod(k, na, nr) => Instruction::CallMethod(get_str(*k), *na as usize, *nr as usize),
371        Op::Return(255)       => Instruction::ReturnFromMark,
372        Op::Return(n)         => Instruction::Return(*n as usize),
373        Op::MarkReturn        => Instruction::MarkReturn,
374        Op::TailCall(n)       => Instruction::Call(*n as usize, 0),
375        Op::Closure(idx)      => Instruction::MakeFunction(*idx as usize),
376        Op::Close(s)          => Instruction::CloseUpvalue(*s as usize),
377        Op::ForPrep(n)        => Instruction::ForPrep(*n as usize),
378        Op::ForStep           => Instruction::Nop,
379        Op::ForStepJump(off)  => Instruction::ForStep(0, *off as isize),
380        Op::NumForInit        => Instruction::Nop,
381        Op::NumForStep(off)   => Instruction::ForStep(0, *off as isize),
382        Op::Vararg(_)         => Instruction::Nop,
383        Op::LineInfo(_)       => Instruction::Nop,
384    }
385}
386
387// ── Local variable tracking ───────────────────────────────────────────────────
388
389#[derive(Debug, Clone)]
390struct Local {
391    name:  String,
392    slot:  u16,
393    depth: usize,
394}
395
396struct Scope {
397    locals:      Vec<Local>,
398    scope_depth: usize,
399    next_slot:   u16,
400}
401
402impl Scope {
403    fn new() -> Self {
404        Scope { locals: Vec::new(), scope_depth: 0, next_slot: 0 }
405    }
406
407    fn push_scope(&mut self) {
408        self.scope_depth += 1;
409    }
410
411    fn pop_scope(&mut self) -> u16 {
412        let depth = self.scope_depth;
413        let before = self.locals.len();
414        self.locals.retain(|l| l.depth < depth);
415        let removed = (before - self.locals.len()) as u16;
416        self.next_slot -= removed;
417        self.scope_depth -= 1;
418        removed
419    }
420
421    fn add_local(&mut self, name: &str) -> u16 {
422        let slot = self.next_slot;
423        self.locals.push(Local {
424            name: name.to_string(),
425            slot,
426            depth: self.scope_depth,
427        });
428        self.next_slot += 1;
429        slot
430    }
431
432    fn resolve_local(&self, name: &str) -> Option<u16> {
433        self.locals.iter().rev()
434            .find(|l| l.name == name)
435            .map(|l| l.slot)
436    }
437}
438
439// ── Upvalue descriptor ────────────────────────────────────────────────────────
440
441/// How an upvalue is captured by a closure.
442#[derive(Debug, Clone)]
443pub struct UpvalDesc {
444    pub name:     String,
445    /// If true, the upvalue is a local slot in the immediately enclosing scope;
446    /// otherwise it is an upvalue of the enclosing function.
447    pub in_stack: bool,
448    pub index:    u16,
449}
450
451// ── Compiler ─────────────────────────────────────────────────────────────────
452
453/// Single-pass bytecode compiler.
454pub struct Compiler {
455    proto:  Proto,
456    scope:  Scope,
457    breaks: Vec<Vec<usize>>,   // break-patch points indexed by loop nesting
458    /// Set while compiling the last expression of a `return`: a call there
459    /// keeps all of its results instead of exactly one.
460    want_multi: bool,
461    // Note: upvalue handling is simplified — outer-function locals captured
462    // as globals in this basic implementation.
463}
464
465impl Compiler {
466    // ── Public entry points ───────────────────────────────────────────────────
467
468    /// Compile an entire script into a top-level `Arc<Chunk>` for the VM.
469    pub fn compile_script(script: &Script) -> Arc<Chunk> {
470        proto_to_chunk(&Self::compile_to_proto(script))
471    }
472
473    /// Compile to the internal `Proto` representation (used by compiler tests).
474    pub fn compile_to_proto(script: &Script) -> Proto {
475        let mut c = Compiler {
476            proto:  Proto::new(&script.name),
477            scope:  Scope::new(),
478            breaks: Vec::new(),
479            want_multi: false,
480        };
481        c.proto.is_vararg = true;
482        c.compile_block_no_scope(&script.stmts);
483        c.proto.emit(Op::Return(0));
484        c.proto
485    }
486
487    // ── Block compilation ─────────────────────────────────────────────────────
488
489    fn compile_block(&mut self, stmts: &[Stmt]) {
490        self.scope.push_scope();
491        for s in stmts { self.compile_stmt(s); }
492        let popped = self.scope.pop_scope();
493        for _ in 0..popped { self.proto.emit(Op::Pop); }
494    }
495
496    fn compile_block_no_scope(&mut self, stmts: &[Stmt]) {
497        for s in stmts { self.compile_stmt(s); }
498    }
499
500    // ── Statement compilation ─────────────────────────────────────────────────
501
502    fn compile_stmt(&mut self, stmt: &Stmt) {
503        match stmt {
504            Stmt::LocalDecl { name, init } => {
505                if let Some(expr) = init {
506                    self.compile_expr(expr);
507                } else {
508                    self.proto.emit(Op::Nil);
509                }
510                self.scope.add_local(name);
511            }
512
513            Stmt::LocalMulti { names, inits } => {
514                for (i, name) in names.iter().enumerate() {
515                    if i < inits.len() {
516                        self.compile_expr(&inits[i]);
517                    } else {
518                        self.proto.emit(Op::Nil);
519                    }
520                    self.scope.add_local(name);
521                }
522            }
523
524            Stmt::Assign { target, value } => {
525                for (i, t) in target.iter().enumerate() {
526                    // Field and index targets need the table (and key) below
527                    // the value: SetField wants [table, value] and keeps the
528                    // table, SetIndex wants [table, key, value]. The value
529                    // used to be pushed first, so `t.x = v` failed with
530                    // "SetField on non-table" and `t[k] = v` indexed v.
531                    match t {
532                        Expr::Field { table, name } => {
533                            self.compile_expr(table);
534                            self.compile_assign_value(value, i);
535                            let k = self.proto.add_const(Constant::Str(name.clone()));
536                            self.proto.emit(Op::SetField(k));
537                            self.proto.emit(Op::Pop);
538                        }
539                        Expr::Index { table, key } => {
540                            self.compile_expr(table);
541                            self.compile_expr(key);
542                            self.compile_assign_value(value, i);
543                            self.proto.emit(Op::SetIndex);
544                        }
545                        _ => {
546                            self.compile_assign_value(value, i);
547                            self.compile_assign_target(t);
548                        }
549                    }
550                }
551            }
552
553            Stmt::CompoundAssign { target, op, value } => {
554                self.compile_expr(target);
555                self.compile_expr(value);
556                self.compile_binop(*op);
557                self.compile_assign_target(target);
558            }
559
560            Stmt::Call(expr) | Stmt::Expr(expr) => {
561                self.compile_expr(expr);
562                self.proto.emit(Op::Pop);
563            }
564
565            Stmt::Do(body) => {
566                self.compile_block(body);
567            }
568
569            Stmt::While { cond, body } => {
570                let loop_start = self.proto.code.len() as i32;
571                self.compile_expr(cond);
572                let exit = self.proto.emit(Op::JumpIfNot(0));
573                self.breaks.push(Vec::new());
574                self.compile_block(body);
575                let back = loop_start - self.proto.code.len() as i32 - 1;
576                self.proto.emit(Op::Jump(back));
577                self.proto.patch_jump(exit);
578                for b in self.breaks.pop().unwrap_or_default() {
579                    self.proto.patch_jump(b);
580                }
581            }
582
583            Stmt::RepeatUntil { body, cond } => {
584                let loop_start = self.proto.code.len() as i32;
585                self.breaks.push(Vec::new());
586                self.compile_block(body);
587                self.compile_expr(cond);
588                let back = loop_start - self.proto.code.len() as i32 - 1;
589                self.proto.emit(Op::JumpIfNot(back));
590                for b in self.breaks.pop().unwrap_or_default() {
591                    self.proto.patch_jump(b);
592                }
593            }
594
595            Stmt::If { cond, then_body, elseif_branches, else_body } => {
596                self.compile_expr(cond);
597                let skip_then = self.proto.emit(Op::JumpIfNot(0));
598                self.compile_block(then_body);
599
600                let mut end_jumps = Vec::new();
601                if !elseif_branches.is_empty() || else_body.is_some() {
602                    end_jumps.push(self.proto.emit(Op::Jump(0)));
603                }
604                self.proto.patch_jump(skip_then);
605
606                for (ei_cond, ei_body) in elseif_branches {
607                    self.compile_expr(ei_cond);
608                    let skip = self.proto.emit(Op::JumpIfNot(0));
609                    self.compile_block(ei_body);
610                    end_jumps.push(self.proto.emit(Op::Jump(0)));
611                    self.proto.patch_jump(skip);
612                }
613
614                if let Some(eb) = else_body {
615                    self.compile_block(eb);
616                }
617                for j in end_jumps { self.proto.patch_jump(j); }
618            }
619
620            Stmt::NumericFor { var, start, limit, step, body } => {
621                self.compile_expr(start);
622                self.compile_expr(limit);
623                if let Some(s) = step {
624                    self.compile_expr(s);
625                } else {
626                    let k = self.proto.add_const(Constant::Int(1));
627                    self.proto.emit(Op::Const(k));
628                }
629                self.proto.emit(Op::NumForInit);
630                let loop_top = self.proto.code.len();
631                let exit = self.proto.emit(Op::NumForStep(0));
632
633                self.scope.push_scope();
634                let slot = self.scope.add_local(var);
635                self.proto.emit(Op::GetLocal(slot));
636                self.breaks.push(Vec::new());
637                for s in body { self.compile_stmt(s); }
638                self.scope.pop_scope();
639
640                let back = loop_top as i32 - self.proto.code.len() as i32 - 1;
641                self.proto.emit(Op::Jump(back));
642                self.proto.patch_jump(exit);
643                // Pop limit, step, counter
644                for _ in 0..3 { self.proto.emit(Op::Pop); }
645                for b in self.breaks.pop().unwrap_or_default() {
646                    self.proto.patch_jump(b);
647                }
648            }
649
650            Stmt::GenericFor { vars, iter, body } => {
651                for expr in iter { self.compile_expr(expr); }
652                self.proto.emit(Op::ForPrep(vars.len() as u16));
653                let loop_top = self.proto.code.len();
654                let exit = self.proto.emit(Op::ForStepJump(0));
655
656                self.scope.push_scope();
657                for name in vars {
658                    let slot = self.scope.add_local(name);
659                    self.proto.emit(Op::GetLocal(slot));
660                }
661                self.breaks.push(Vec::new());
662                for s in body { self.compile_stmt(s); }
663                self.scope.pop_scope();
664
665                let back = loop_top as i32 - self.proto.code.len() as i32 - 1;
666                self.proto.emit(Op::Jump(back));
667                self.proto.patch_jump(exit);
668                for b in self.breaks.pop().unwrap_or_default() {
669                    self.proto.patch_jump(b);
670                }
671            }
672
673            Stmt::FuncDecl { name, params, vararg, body } => {
674                let fn_proto = self.compile_func(
675                    name.last().map(|s| s.as_str()).unwrap_or("?"),
676                    params, *vararg, body,
677                );
678                let idx = self.proto.protos.len() as u32;
679                self.proto.protos.push(fn_proto);
680                self.proto.emit(Op::Closure(idx));
681
682                if name.len() == 1 {
683                    if let Some(slot) = self.scope.resolve_local(&name[0]) {
684                        self.proto.emit(Op::SetLocal(slot));
685                    } else {
686                        let k = self.proto.add_const(Constant::Str(name[0].clone()));
687                        self.proto.emit(Op::SetGlobal(k));
688                    }
689                } else {
690                    // a.b.fn = closure
691                    self.compile_expr(&Expr::Ident(name[0].clone()));
692                    for part in &name[1..name.len()-1] {
693                        let k = self.proto.add_const(Constant::Str(part.clone()));
694                        self.proto.emit(Op::GetField(k));
695                    }
696                    let last = name.last().unwrap();
697                    let k = self.proto.add_const(Constant::Str(last.clone()));
698                    // Stack is [closure, table]; SetField wants [table, value]
699                    // and leaves the table, so swap first and pop after.
700                    self.proto.emit(Op::Swap);
701                    self.proto.emit(Op::SetField(k));
702                    self.proto.emit(Op::Pop);
703                }
704            }
705
706            Stmt::LocalFunc { name, params, vararg, body } => {
707                let slot = self.scope.add_local(name);
708                self.proto.emit(Op::Nil); // placeholder until closure is made
709                let fn_proto = self.compile_func(name, params, *vararg, body);
710                let idx = self.proto.protos.len() as u32;
711                self.proto.protos.push(fn_proto);
712                self.proto.emit(Op::Closure(idx));
713                self.proto.emit(Op::SetLocal(slot));
714            }
715
716            Stmt::Return(vals) => {
717                // `return f()` / `return a, f()` passes on every value f
718                // returns, as in Lua. Calls used to be fixed at one result,
719                // so `return pcall(g)` lost its leading `true` and
720                // `return table.unpack(t)` gave only the last element.
721                let last_is_call = matches!(
722                    vals.last(),
723                    Some(Expr::Call { .. }) | Some(Expr::MethodCall { .. })
724                );
725                if last_is_call {
726                    self.proto.emit(Op::MarkReturn);
727                    let n = vals.len();
728                    for (i, v) in vals.iter().enumerate() {
729                        self.want_multi = i + 1 == n;
730                        self.compile_expr(v);
731                    }
732                    self.want_multi = false;
733                    self.proto.emit(Op::Return(255));
734                } else {
735                    let n = vals.len() as u8;
736                    for v in vals { self.compile_expr(v); }
737                    self.proto.emit(Op::Return(n));
738                }
739            }
740
741            Stmt::Break => {
742                let j = self.proto.emit(Op::Jump(0));
743                if let Some(list) = self.breaks.last_mut() {
744                    list.push(j);
745                }
746            }
747
748            Stmt::Continue => {
749                // Simplified continue: jump to -1 (loop should handle by re-check)
750                self.proto.emit(Op::Jump(-1));
751            }
752
753            Stmt::Match { expr, arms } => {
754                self.compile_expr(expr);
755                let mut end_jumps = Vec::new();
756
757                for arm in arms {
758                    self.proto.emit(Op::Dup);
759                    match &arm.pattern {
760                        MatchPattern::Wildcard => {
761                            self.proto.emit(Op::Pop);
762                            self.compile_block(&arm.body);
763                            end_jumps.push(self.proto.emit(Op::Jump(0)));
764                            continue;
765                        }
766                        MatchPattern::Ident(bind) => {
767                            let slot = self.scope.add_local(bind);
768                            self.proto.emit(Op::SetLocal(slot));
769                            self.compile_block(&arm.body);
770                            end_jumps.push(self.proto.emit(Op::Jump(0)));
771                            continue;
772                        }
773                        MatchPattern::Literal(lit) => {
774                            self.compile_expr(lit);
775                            self.proto.emit(Op::Eq);
776                        }
777                        MatchPattern::Table(_) => {
778                            self.proto.emit(Op::Pop);
779                            self.proto.emit(Op::True);
780                        }
781                    }
782                    let skip = self.proto.emit(Op::JumpIfNot(0));
783                    self.compile_block(&arm.body);
784                    end_jumps.push(self.proto.emit(Op::Jump(0)));
785                    self.proto.patch_jump(skip);
786                }
787
788                self.proto.emit(Op::Pop);
789                for j in end_jumps { self.proto.patch_jump(j); }
790            }
791
792            Stmt::Import { path, alias } => {
793                let k = self.proto.add_const(Constant::Str(path.clone()));
794                self.proto.emit(Op::Const(k));
795                let rk = self.proto.add_const(Constant::Str("require".to_string()));
796                self.proto.emit(Op::GetGlobal(rk));
797                self.proto.emit(Op::Swap);
798                self.proto.emit(Op::Call(1, 1));
799                let bind = alias.clone().unwrap_or_else(|| {
800                    path.split('/').last().unwrap_or(path).trim_end_matches(".lua").to_string()
801                });
802                let bk = self.proto.add_const(Constant::Str(bind));
803                self.proto.emit(Op::SetGlobal(bk));
804            }
805
806            Stmt::Export(name) => {
807                if let Some(slot) = self.scope.resolve_local(name) {
808                    self.proto.emit(Op::GetLocal(slot));
809                } else {
810                    let k = self.proto.add_const(Constant::Str(name.clone()));
811                    self.proto.emit(Op::GetGlobal(k));
812                }
813                let ek = self.proto.add_const(Constant::Str(name.clone()));
814                let exports_k = self.proto.add_const(Constant::Str("__exports".to_string()));
815                self.proto.emit(Op::GetGlobal(exports_k));
816                self.proto.emit(Op::Swap);
817                self.proto.emit(Op::SetField(ek));
818            }
819        }
820    }
821
822    fn compile_assign_value(&mut self, value: &[Expr], i: usize) {
823        if i < value.len() {
824            self.compile_expr(&value[i]);
825        } else {
826            self.proto.emit(Op::Nil);
827        }
828    }
829
830    fn compile_assign_target(&mut self, target: &Expr) {
831        match target {
832            Expr::Ident(name) => {
833                if let Some(slot) = self.scope.resolve_local(name) {
834                    self.proto.emit(Op::SetLocal(slot));
835                } else {
836                    let k = self.proto.add_const(Constant::Str(name.clone()));
837                    self.proto.emit(Op::SetGlobal(k));
838                }
839            }
840            Expr::Field { table, name } => {
841                self.compile_expr(table);
842                let k = self.proto.add_const(Constant::Str(name.clone()));
843                self.proto.emit(Op::SetField(k));
844            }
845            Expr::Index { table, key } => {
846                self.compile_expr(table);
847                self.compile_expr(key);
848                self.proto.emit(Op::SetIndex);
849            }
850            _ => {}
851        }
852    }
853
854    // ── Expression compilation ────────────────────────────────────────────────
855
856    fn compile_expr(&mut self, expr: &Expr) {
857        match expr {
858            Expr::Nil         => { self.proto.emit(Op::Nil); }
859            Expr::Bool(b)     => { self.proto.emit(if *b { Op::True } else { Op::False }); }
860            Expr::Int(n)      => { let k = self.proto.add_const(Constant::Int(*n)); self.proto.emit(Op::Const(k)); }
861            Expr::Float(f)    => { let k = self.proto.add_const(Constant::Float(*f)); self.proto.emit(Op::Const(k)); }
862            Expr::Str(s)      => { let k = self.proto.add_const(Constant::Str(s.clone())); self.proto.emit(Op::Const(k)); }
863            Expr::Vararg      => { self.proto.emit(Op::Vararg(0)); }
864
865            Expr::Ident(name) => {
866                if let Some(slot) = self.scope.resolve_local(name) {
867                    self.proto.emit(Op::GetLocal(slot));
868                } else {
869                    let k = self.proto.add_const(Constant::Str(name.clone()));
870                    self.proto.emit(Op::GetGlobal(k));
871                }
872            }
873
874            Expr::Field { table, name } => {
875                self.compile_expr(table);
876                let k = self.proto.add_const(Constant::Str(name.clone()));
877                self.proto.emit(Op::GetField(k));
878            }
879
880            Expr::Index { table, key } => {
881                self.compile_expr(table);
882                self.compile_expr(key);
883                self.proto.emit(Op::GetIndex);
884            }
885
886            Expr::Call { callee, args } => {
887                let nret = if std::mem::take(&mut self.want_multi) { 0 } else { 1 };
888                self.compile_expr(callee);
889                let nargs = args.len() as u8;
890                for a in args { self.compile_expr(a); }
891                self.proto.emit(Op::Call(nargs, nret));
892            }
893
894            Expr::MethodCall { obj, method, args } => {
895                let nret = if std::mem::take(&mut self.want_multi) { 0 } else { 1 };
896                self.compile_expr(obj);
897                let k = self.proto.add_const(Constant::Str(method.clone()));
898                let nargs = args.len() as u8;
899                for a in args { self.compile_expr(a); }
900                self.proto.emit(Op::CallMethod(k, nargs, nret));
901            }
902
903            Expr::Unary { op, expr } => {
904                self.compile_expr(expr);
905                match op {
906                    UnOp::Neg    => { self.proto.emit(Op::Neg); }
907                    UnOp::Not    => { self.proto.emit(Op::Not); }
908                    UnOp::Len    => { self.proto.emit(Op::Len); }
909                    UnOp::BitNot => { self.proto.emit(Op::BitNot); }
910                }
911            }
912
913            Expr::Binary { op, lhs, rhs } => {
914                match op {
915                    BinOp::And => {
916                        self.compile_expr(lhs);
917                        let j = self.proto.emit(Op::JumpIfNotPop(0));
918                        self.compile_expr(rhs);
919                        self.proto.patch_jump(j);
920                        return;
921                    }
922                    BinOp::Or => {
923                        self.compile_expr(lhs);
924                        let j = self.proto.emit(Op::JumpIfPop(0));
925                        self.compile_expr(rhs);
926                        self.proto.patch_jump(j);
927                        return;
928                    }
929                    _ => {}
930                }
931                self.compile_expr(lhs);
932                self.compile_expr(rhs);
933                self.compile_binop(*op);
934            }
935
936            Expr::TableCtor(fields) => {
937                self.proto.emit(Op::NewTable);
938                for field in fields {
939                    match field {
940                        TableField::NameKey(name, val) => {
941                            // SetField keeps the table on the stack, so no
942                            // Dup (the Dup left a spare table per field and
943                            // shifted every later local slot).
944                            self.compile_expr(val);
945                            let k = self.proto.add_const(Constant::Str(name.clone()));
946                            self.proto.emit(Op::SetField(k));
947                        }
948                        TableField::ExprKey(key, val) => {
949                            self.proto.emit(Op::Dup);
950                            self.compile_expr(key);
951                            self.compile_expr(val);
952                            self.proto.emit(Op::SetIndex);
953                        }
954                        TableField::Value(val) => {
955                            // Positional values go to the array part. They
956                            // were stored with SetField and an integer
957                            // constant, which SetField reads as a string
958                            // name: every {1, 2, 3} item landed under "".
959                            self.compile_expr(val);
960                            self.proto.emit(Op::TableAppend);
961                        }
962                    }
963                }
964            }
965
966            Expr::FuncExpr { params, vararg, body } => {
967                let fn_proto = self.compile_func("<anon>", params, *vararg, body);
968                let idx = self.proto.protos.len() as u32;
969                self.proto.protos.push(fn_proto);
970                self.proto.emit(Op::Closure(idx));
971            }
972
973            Expr::Ternary { cond, then_val, else_val } => {
974                self.compile_expr(cond);
975                let skip = self.proto.emit(Op::JumpIfNot(0));
976                self.compile_expr(then_val);
977                let end = self.proto.emit(Op::Jump(0));
978                self.proto.patch_jump(skip);
979                self.compile_expr(else_val);
980                self.proto.patch_jump(end);
981            }
982        }
983    }
984
985    fn compile_binop(&mut self, op: BinOp) {
986        let instr = match op {
987            BinOp::Add    => Op::Add,
988            BinOp::Sub    => Op::Sub,
989            BinOp::Mul    => Op::Mul,
990            BinOp::Div    => Op::Div,
991            BinOp::IDiv   => Op::IDiv,
992            BinOp::Mod    => Op::Mod,
993            BinOp::Pow    => Op::Pow,
994            BinOp::Concat => Op::Concat,
995            BinOp::Eq     => Op::Eq,
996            BinOp::NotEq  => Op::NotEq,
997            BinOp::Lt     => Op::Lt,
998            BinOp::LtEq   => Op::LtEq,
999            BinOp::Gt     => Op::Gt,
1000            BinOp::GtEq   => Op::GtEq,
1001            BinOp::And    => Op::BitAnd,
1002            BinOp::Or     => Op::BitOr,
1003            BinOp::BitAnd => Op::BitAnd,
1004            BinOp::BitOr  => Op::BitOr,
1005            BinOp::BitXor => Op::BitXor,
1006            BinOp::Shl    => Op::Shl,
1007            BinOp::Shr    => Op::Shr,
1008        };
1009        self.proto.emit(instr);
1010    }
1011
1012    fn compile_func(&mut self, name: &str, params: &[String], vararg: bool, body: &[Stmt]) -> Proto {
1013        let mut child = Compiler {
1014            proto:  Proto::new(name),
1015            scope:  Scope::new(),
1016            breaks: Vec::new(),
1017            want_multi: false,
1018        };
1019        child.proto.param_count = params.len() as u8;
1020        child.proto.is_vararg   = vararg;
1021        child.scope.push_scope();
1022        for p in params { child.scope.add_local(p); }
1023        for s in body   { child.compile_stmt(s); }
1024        child.scope.pop_scope();
1025        child.proto.emit(Op::Return(0));
1026        child.proto
1027    }
1028}
1029
1030// ── Tests ─────────────────────────────────────────────────────────────────────
1031
1032#[cfg(test)]
1033mod tests {
1034    use super::*;
1035    use crate::scripting::parser;
1036
1037    fn compile_src(src: &str) -> Proto {
1038        let script = parser::parse(src, "test").expect("parse failed");
1039        Compiler::compile_to_proto(&script)
1040    }
1041
1042    #[test]
1043    fn test_compile_nil() {
1044        let p = compile_src("local x");
1045        assert!(p.code.iter().any(|op| *op == Op::Nil));
1046    }
1047
1048    #[test]
1049    fn test_compile_int_const() {
1050        let p = compile_src("local x = 42");
1051        assert!(p.constants.iter().any(|c| *c == Constant::Int(42)));
1052    }
1053
1054    #[test]
1055    fn test_compile_float_const() {
1056        let p = compile_src("local pi = 3.14");
1057        assert!(p.constants.iter().any(|c| matches!(c, Constant::Float(f) if (*f - 3.14).abs() < 1e-6)));
1058    }
1059
1060    #[test]
1061    fn test_compile_string_const() {
1062        let p = compile_src(r#"local s = "hello""#);
1063        assert!(p.constants.iter().any(|c| *c == Constant::Str("hello".to_string())));
1064    }
1065
1066    #[test]
1067    fn test_compile_add() {
1068        let p = compile_src("local z = 1 + 2");
1069        assert!(p.code.iter().any(|op| *op == Op::Add));
1070    }
1071
1072    #[test]
1073    fn test_compile_bool_true() {
1074        let p = compile_src("local b = true");
1075        assert!(p.code.iter().any(|op| *op == Op::True));
1076    }
1077
1078    #[test]
1079    fn test_compile_while_has_back_jump() {
1080        let p = compile_src("local i = 0 while i < 10 do i = i + 1 end");
1081        let has_exit = p.code.iter().any(|op| matches!(op, Op::JumpIfNot(_)));
1082        let has_back = p.code.iter().any(|op| matches!(op, Op::Jump(n) if *n < 0));
1083        assert!(has_exit, "expected JumpIfNot");
1084        assert!(has_back, "expected backward Jump");
1085    }
1086
1087    #[test]
1088    fn test_compile_if_else() {
1089        let p = compile_src("if x then return 1 else return 2 end");
1090        assert!(p.code.iter().any(|op| matches!(op, Op::JumpIfNot(_))));
1091        assert!(p.code.iter().any(|op| matches!(op, Op::Jump(_))));
1092    }
1093
1094    #[test]
1095    fn test_compile_function_creates_proto() {
1096        let p = compile_src("function add(a, b) return a + b end");
1097        assert!(!p.protos.is_empty());
1098        assert_eq!(p.protos[0].param_count, 2);
1099    }
1100
1101    #[test]
1102    fn test_compile_local_function() {
1103        let p = compile_src("local function square(x) return x * x end");
1104        assert!(p.code.iter().any(|op| matches!(op, Op::Closure(_))));
1105        assert!(!p.protos.is_empty());
1106    }
1107
1108    #[test]
1109    fn test_compile_table_ctor() {
1110        let p = compile_src("local t = {x = 1, y = 2}");
1111        assert!(p.code.iter().any(|op| *op == Op::NewTable));
1112        assert!(p.code.iter().any(|op| matches!(op, Op::SetField(_))));
1113    }
1114
1115    #[test]
1116    fn test_compile_method_call() {
1117        let p = compile_src("obj:doThing(1, 2)");
1118        assert!(p.code.iter().any(|op| matches!(op, Op::CallMethod(..))));
1119    }
1120
1121    #[test]
1122    fn test_compile_for_numeric() {
1123        let p = compile_src("for i = 1, 10, 2 do end");
1124        assert!(p.code.iter().any(|op| *op == Op::NumForInit));
1125        assert!(p.code.iter().any(|op| matches!(op, Op::NumForStep(_))));
1126    }
1127
1128    #[test]
1129    fn test_compile_for_generic() {
1130        let p = compile_src("for k, v in pairs(t) do end");
1131        assert!(p.code.iter().any(|op| matches!(op, Op::ForPrep(_))));
1132    }
1133
1134    #[test]
1135    fn test_compile_and_short_circuit() {
1136        let p = compile_src("local r = a and b");
1137        assert!(p.code.iter().any(|op| matches!(op, Op::JumpIfNotPop(_))));
1138    }
1139
1140    #[test]
1141    fn test_compile_or_short_circuit() {
1142        let p = compile_src("local r = a or b");
1143        assert!(p.code.iter().any(|op| matches!(op, Op::JumpIfPop(_))));
1144    }
1145
1146    #[test]
1147    fn test_compile_ternary() {
1148        let p = compile_src("local x = cond ? 1 : 2");
1149        assert!(p.code.iter().any(|op| matches!(op, Op::JumpIfNot(_))));
1150    }
1151
1152    #[test]
1153    fn test_compile_nested_function() {
1154        let p = compile_src("
1155            function outer(x)
1156                local function inner(y) return x + y end
1157                return inner(10)
1158            end
1159        ");
1160        assert!(!p.protos.is_empty());
1161        let outer = &p.protos[0];
1162        assert!(!outer.protos.is_empty(), "expected inner proto");
1163    }
1164
1165    #[test]
1166    fn test_compile_concat() {
1167        let p = compile_src(r#"local s = "hello" .. " " .. "world""#);
1168        assert!(p.code.iter().filter(|op| **op == Op::Concat).count() >= 1);
1169    }
1170
1171    #[test]
1172    fn test_compile_repeat_until() {
1173        let p = compile_src("local i = 0 repeat i = i + 1 until i >= 10");
1174        assert!(p.code.iter().any(|op| matches!(op, Op::JumpIfNot(n) if *n < 0)));
1175    }
1176
1177    #[test]
1178    fn test_compile_match() {
1179        let p = compile_src("match x { case 1 => return 1, case 2 => return 2 }");
1180        assert!(p.code.iter().any(|op| *op == Op::Dup));
1181        assert!(p.code.iter().any(|op| *op == Op::Eq));
1182    }
1183
1184    #[test]
1185    fn test_compile_import() {
1186        let p = compile_src(r#"import "math" as m"#);
1187        assert!(p.constants.iter().any(|c| *c == Constant::Str("math".to_string())));
1188        assert!(p.constants.iter().any(|c| *c == Constant::Str("require".to_string())));
1189    }
1190
1191    #[test]
1192    fn test_add_const_deduplication() {
1193        let mut p = Proto::new("test");
1194        let i1 = p.add_const(Constant::Int(42));
1195        let i2 = p.add_const(Constant::Int(42));
1196        assert_eq!(i1, i2, "deduplication failed");
1197        assert_eq!(p.constants.len(), 1);
1198    }
1199}