1use std::collections::HashMap;
13use std::sync::Arc;
14use super::ast::*;
15use super::vm::Value as VmValue;
16
17#[derive(Debug, Clone, PartialEq)]
21pub enum Constant {
22 Nil,
23 Bool(bool),
24 Int(i64),
25 Float(f64),
26 Str(String),
27}
28
29#[allow(non_camel_case_types)]
34#[derive(Debug, Clone, PartialEq)]
35pub enum Op {
36 Nil,
39 True,
41 False,
43 Const(u32),
45
46 Pop,
48 Dup,
49 Swap,
50
51 GetLocal(u16),
53 SetLocal(u16),
54
55 GetUpval(u16),
57 SetUpval(u16),
58
59 GetGlobal(u32),
61 SetGlobal(u32),
62
63 NewTable,
65 SetField(u32),
67 GetField(u32),
69 SetIndex,
71 GetIndex,
73 TableAppend,
75 SetList(u16),
77
78 Len,
80 Neg,
81 Not,
82 BitNot,
83
84 Add, Sub, Mul, Div, IDiv, Mod, Pow,
86 Concat, Eq, NotEq, Lt, LtEq, Gt, GtEq,
90
91 BitAnd, BitOr, BitXor, Shl, Shr,
93
94 Jump(i32),
97 JumpIf(i32),
99 JumpIfNot(i32),
101 JumpIfNotPop(i32),
103 JumpIfPop(i32),
105
106 Call(u8, u8),
109 CallMethod(u32, u8, u8),
111 Return(u8),
115 MarkReturn,
117 TailCall(u8),
119
120 Closure(u32),
123 Close(u16),
125
126 ForPrep(u16),
129 ForStep,
132 ForStepJump(i32),
134 NumForInit,
136 NumForStep(i32),
138
139 Vararg(u8),
142
143 LineInfo(u32),
145}
146
147#[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>, 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 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#[derive(Debug, Clone, PartialEq)]
208pub enum Instruction {
209 LoadNil,
211 LoadBool(bool),
212 LoadInt(i64),
213 LoadFloat(f64),
214 LoadStr(String),
215 LoadConst(usize),
216 Pop,
218 Dup,
219 Swap,
220 GetLocal(usize),
222 SetLocal(usize),
223 GetUpvalue(usize),
224 SetUpvalue(usize),
225 GetGlobal(String),
226 SetGlobal(String),
227 NewTable,
229 SetField(String),
230 GetField(String),
231 SetIndex,
232 GetIndex,
233 TableAppend,
234 Len,
236 Neg,
237 Not,
238 BitNot,
239 Add, Sub, Mul, Div, IDiv, Mod, Pow,
241 Concat,
242 BitAnd, BitOr, BitXor, Shl, Shr,
244 Eq, NotEq, Lt, LtEq, Gt, GtEq,
246 Jump(isize),
248 JumpIf(isize),
249 JumpIfNot(isize),
250 JumpIfNotPop(isize),
252 JumpIfPop(isize),
254 JumpAbs(usize),
255 Call(usize, usize),
258 CallMethod(String, usize, usize),
259 Return(usize),
260 ReturnFromMark,
262 MarkReturn,
263 MakeFunction(usize),
265 MakeClosure(usize, Vec<(bool, usize)>),
266 CloseUpvalue(usize),
267 ForPrep(usize),
269 ForStep(usize, isize),
271 Nop,
272}
273
274#[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#[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#[derive(Debug, Clone)]
443pub struct UpvalDesc {
444 pub name: String,
445 pub in_stack: bool,
448 pub index: u16,
449}
450
451pub struct Compiler {
455 proto: Proto,
456 scope: Scope,
457 breaks: Vec<Vec<usize>>, want_multi: bool,
461 }
464
465impl Compiler {
466 pub fn compile_script(script: &Script) -> Arc<Chunk> {
470 proto_to_chunk(&Self::compile_to_proto(script))
471 }
472
473 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 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 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 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 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 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 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); 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 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 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 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 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 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#[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}