Skip to main content

valua_codegen/
lib.rs

1//! Lua 5.1 code emitter — walks a transformed AST and produces formatted source.
2// emit_statement and emit_expr_inner are large dispatch functions — splitting them would
3// scatter the grammar and hurt readability.
4#![allow(clippy::too_many_lines)]
5
6use valua_ast::{BinaryOp, Block, Call, Expression, FunctionBody, Statement, TableField, UnaryOp};
7
8pub use error::CodeGenError;
9
10mod error;
11
12// ── Configuration ─────────────────────────────────────────────────────────────
13
14// `LuaTarget` is defined in `valua-ast` (the common dep layer) and re-exported
15// here so that existing `valua_codegen::LuaTarget` paths continue to compile.
16pub use valua_ast::LuaTarget;
17
18/// Options that control how the emitted Lua source is formatted.
19#[derive(Debug, Clone)]
20pub struct EmitOptions {
21    /// String used for each indentation level (default: four spaces).
22    pub indent: String,
23    /// Target runtime.
24    pub target: LuaTarget,
25    /// Emit a `-- Generated by valua` header comment.
26    pub emit_header_comment: bool,
27}
28
29impl Default for EmitOptions {
30    fn default() -> Self {
31        Self {
32            indent: "    ".to_string(),
33            target: LuaTarget::default(),
34            emit_header_comment: false,
35        }
36    }
37}
38
39// ── Trait ─────────────────────────────────────────────────────────────────────
40
41/// A code-generation backend that can emit a `Block` to a `String`.
42pub trait CodeGen {
43    /// Emit the entire `block` and return the formatted Lua source.
44    ///
45    /// # Errors
46    /// Returns `CodeGenError` if the AST contains a node that cannot be
47    /// represented in the target Lua version.
48    fn emit(&self, block: &Block) -> Result<String, CodeGenError>;
49}
50
51// ── Emitter ───────────────────────────────────────────────────────────────────
52
53/// Concrete Lua 5.1 / `LuaJIT` emitter.
54pub struct LuaEmitter {
55    pub options: EmitOptions,
56}
57
58impl LuaEmitter {
59    #[must_use]
60    pub fn new(options: EmitOptions) -> Self {
61        Self { options }
62    }
63
64    #[must_use]
65    pub fn lua51() -> Self {
66        Self::new(EmitOptions {
67            target: LuaTarget::Lua51,
68            ..EmitOptions::default()
69        })
70    }
71
72    #[must_use]
73    pub fn luajit() -> Self {
74        Self::new(EmitOptions {
75            target: LuaTarget::LuaJIT,
76            ..EmitOptions::default()
77        })
78    }
79}
80
81impl CodeGen for LuaEmitter {
82    fn emit(&self, block: &Block) -> Result<String, CodeGenError> {
83        let mut ctx = EmitContext::new(&self.options);
84        if self.options.emit_header_comment {
85            ctx.push("-- Generated by valua v");
86            ctx.push(env!("CARGO_PKG_VERSION"));
87            ctx.push("\n");
88        }
89        ctx.emit_block(block)?;
90        Ok(ctx.finish())
91    }
92}
93
94// ── Operator helpers ──────────────────────────────────────────────────────────
95
96fn binary_op_str(op: BinaryOp) -> &'static str {
97    match op {
98        BinaryOp::Add => "+",
99        BinaryOp::Sub => "-",
100        BinaryOp::Mul => "*",
101        BinaryOp::Div => "/",
102        BinaryOp::Mod => "%",
103        BinaryOp::Pow => "^",
104        BinaryOp::IDiv => "//",
105        BinaryOp::Concat => "..",
106        BinaryOp::Lt => "<",
107        BinaryOp::Le => "<=",
108        BinaryOp::Gt => ">",
109        BinaryOp::Ge => ">=",
110        BinaryOp::Eq => "==",
111        BinaryOp::Ne => "~=",
112        BinaryOp::And => "and",
113        BinaryOp::Or => "or",
114        BinaryOp::BitwiseAnd => "&",
115        BinaryOp::BitwiseOr => "|",
116        BinaryOp::BitwiseXor => "~",
117        BinaryOp::Shl => "<<",
118        BinaryOp::Shr => ">>",
119    }
120}
121
122fn binary_op_prec(op: BinaryOp) -> u8 {
123    match op {
124        BinaryOp::Or => 1,
125        BinaryOp::And => 3,
126        BinaryOp::Lt | BinaryOp::Le | BinaryOp::Gt | BinaryOp::Ge | BinaryOp::Eq | BinaryOp::Ne => {
127            5
128        }
129        BinaryOp::BitwiseOr => 7,
130        BinaryOp::BitwiseXor => 9,
131        BinaryOp::BitwiseAnd => 11,
132        BinaryOp::Shl | BinaryOp::Shr => 13,
133        BinaryOp::Concat => 16,
134        BinaryOp::Add | BinaryOp::Sub => 17,
135        BinaryOp::Mul | BinaryOp::Div | BinaryOp::IDiv | BinaryOp::Mod => 19,
136        BinaryOp::Pow => 24,
137    }
138}
139
140fn is_right_assoc(op: BinaryOp) -> bool {
141    matches!(op, BinaryOp::Pow | BinaryOp::Concat)
142}
143
144fn expr_outer_prec(expr: &Expression) -> u8 {
145    match expr {
146        Expression::BinOp(_, op, _, _) => binary_op_prec(*op),
147        Expression::UnOp(_, _, _) => 21,
148        _ => u8::MAX,
149    }
150}
151
152// ── Emit context ──────────────────────────────────────────────────────────────
153
154pub(crate) struct EmitContext<'opts> {
155    options: &'opts EmitOptions,
156    buf: String,
157    depth: usize,
158}
159
160impl<'opts> EmitContext<'opts> {
161    pub(crate) fn new(options: &'opts EmitOptions) -> Self {
162        Self {
163            options,
164            buf: String::new(),
165            depth: 0,
166        }
167    }
168
169    pub(crate) fn finish(self) -> String {
170        self.buf
171    }
172
173    fn push(&mut self, s: &str) {
174        self.buf.push_str(s);
175    }
176
177    fn push_indent(&mut self) {
178        self.buf.push_str(&self.options.indent.repeat(self.depth));
179    }
180
181    pub(crate) fn emit_block(&mut self, block: &Block) -> Result<(), CodeGenError> {
182        for stmt in &block.stmts {
183            self.emit_statement(stmt)?;
184        }
185        Ok(())
186    }
187
188    pub(crate) fn emit_statement(&mut self, stmt: &Statement) -> Result<(), CodeGenError> {
189        match stmt {
190            Statement::LocalDecl(d) => {
191                self.push_indent();
192                self.push("local ");
193                for (i, name) in d.names.iter().enumerate() {
194                    if i > 0 {
195                        self.push(", ");
196                    }
197                    self.push(&name.name);
198                }
199                if !d.values.is_empty() {
200                    self.push(" = ");
201                    for (i, val) in d.values.iter().enumerate() {
202                        if i > 0 {
203                            self.push(", ");
204                        }
205                        self.emit_expression(val)?;
206                    }
207                }
208                self.push("\n");
209            }
210            Statement::Assign(a) => {
211                self.push_indent();
212                for (i, target) in a.targets.iter().enumerate() {
213                    if i > 0 {
214                        self.push(", ");
215                    }
216                    self.emit_expression(target)?;
217                }
218                self.push(" = ");
219                for (i, val) in a.values.iter().enumerate() {
220                    if i > 0 {
221                        self.push(", ");
222                    }
223                    self.emit_expression(val)?;
224                }
225                self.push("\n");
226            }
227            Statement::ExprStmt(e) => {
228                self.push_indent();
229                self.emit_expression(e)?;
230                self.push("\n");
231            }
232            Statement::Return(r) => {
233                self.push_indent();
234                self.push("return");
235                if !r.values.is_empty() {
236                    self.push(" ");
237                    for (i, val) in r.values.iter().enumerate() {
238                        if i > 0 {
239                            self.push(", ");
240                        }
241                        self.emit_expression(val)?;
242                    }
243                }
244                self.push("\n");
245            }
246            Statement::Break(_) => {
247                self.push_indent();
248                self.push("break\n");
249            }
250            Statement::Goto(g) => {
251                self.push_indent();
252                self.push("goto ");
253                self.push(&g.label);
254                self.push("\n");
255            }
256            Statement::Label(l) => {
257                self.push_indent();
258                self.push("::");
259                self.push(&l.name);
260                self.push("::\n");
261            }
262            Statement::Do(d) => {
263                self.push_indent();
264                self.push("do\n");
265                self.depth += 1;
266                self.emit_block(&d.body)?;
267                self.depth -= 1;
268                self.push_indent();
269                self.push("end\n");
270            }
271            Statement::While(w) => {
272                self.push_indent();
273                self.push("while ");
274                self.emit_expression(&w.condition)?;
275                self.push(" do\n");
276                self.depth += 1;
277                self.emit_block(&w.body)?;
278                self.depth -= 1;
279                self.push_indent();
280                self.push("end\n");
281            }
282            Statement::Repeat(r) => {
283                self.push_indent();
284                self.push("repeat\n");
285                self.depth += 1;
286                self.emit_block(&r.body)?;
287                self.depth -= 1;
288                self.push_indent();
289                self.push("until ");
290                self.emit_expression(&r.condition)?;
291                self.push("\n");
292            }
293            Statement::If(i) => {
294                self.push_indent();
295                self.push("if ");
296                self.emit_expression(&i.condition)?;
297                self.push(" then\n");
298                self.depth += 1;
299                self.emit_block(&i.then_block)?;
300                self.depth -= 1;
301                for elseif in &i.elseif_clauses {
302                    self.push_indent();
303                    self.push("elseif ");
304                    self.emit_expression(&elseif.condition)?;
305                    self.push(" then\n");
306                    self.depth += 1;
307                    self.emit_block(&elseif.body)?;
308                    self.depth -= 1;
309                }
310                if let Some(ref else_block) = i.else_block {
311                    self.push_indent();
312                    self.push("else\n");
313                    self.depth += 1;
314                    self.emit_block(else_block)?;
315                    self.depth -= 1;
316                }
317                self.push_indent();
318                self.push("end\n");
319            }
320            Statement::NumericFor(f) => {
321                self.push_indent();
322                self.push("for ");
323                self.push(&f.var);
324                self.push(" = ");
325                self.emit_expression(&f.start)?;
326                self.push(", ");
327                self.emit_expression(&f.limit)?;
328                if let Some(ref step) = f.step {
329                    self.push(", ");
330                    self.emit_expression(step)?;
331                }
332                self.push(" do\n");
333                self.depth += 1;
334                self.emit_block(&f.body)?;
335                self.depth -= 1;
336                self.push_indent();
337                self.push("end\n");
338            }
339            Statement::GenericFor(f) => {
340                self.push_indent();
341                self.push("for ");
342                for (i, var) in f.vars.iter().enumerate() {
343                    if i > 0 {
344                        self.push(", ");
345                    }
346                    self.push(var);
347                }
348                self.push(" in ");
349                for (i, iter) in f.iterators.iter().enumerate() {
350                    if i > 0 {
351                        self.push(", ");
352                    }
353                    self.emit_expression(iter)?;
354                }
355                self.push(" do\n");
356                self.depth += 1;
357                self.emit_block(&f.body)?;
358                self.depth -= 1;
359                self.push_indent();
360                self.push("end\n");
361            }
362            Statement::FunctionDecl(f) => {
363                self.push_indent();
364                self.push("function ");
365                let parts = f.name.parts.join(".");
366                self.push(&parts);
367                if let Some(ref method) = f.name.method {
368                    self.push(":");
369                    self.push(method);
370                }
371                self.emit_function_body(&f.func)?;
372                self.push("\n");
373            }
374            Statement::LocalFunctionDecl(f) => {
375                self.push_indent();
376                self.push("local function ");
377                self.push(&f.name);
378                self.emit_function_body(&f.func)?;
379                self.push("\n");
380            }
381        }
382        Ok(())
383    }
384
385    fn emit_function_body(&mut self, func: &FunctionBody) -> Result<(), CodeGenError> {
386        self.push("(");
387        for (i, param) in func.params.iter().enumerate() {
388            if i > 0 {
389                self.push(", ");
390            }
391            self.push(&param.name);
392        }
393        if func.is_vararg {
394            if !func.params.is_empty() {
395                self.push(", ");
396            }
397            self.push("...");
398        }
399        self.push(")\n");
400        self.depth += 1;
401        self.emit_block(&func.body)?;
402        self.depth -= 1;
403        self.push_indent();
404        self.push("end");
405        Ok(())
406    }
407
408    pub(crate) fn emit_expression(&mut self, expr: &Expression) -> Result<(), CodeGenError> {
409        self.emit_expr_prec(expr, 0)
410    }
411
412    fn emit_expr_prec(&mut self, expr: &Expression, min_prec: u8) -> Result<(), CodeGenError> {
413        let ep = expr_outer_prec(expr);
414        let needs_parens = ep < min_prec;
415        if needs_parens {
416            self.push("(");
417        }
418        self.emit_expr_inner(expr)?;
419        if needs_parens {
420            self.push(")");
421        }
422        Ok(())
423    }
424
425    fn emit_expr_inner(&mut self, expr: &Expression) -> Result<(), CodeGenError> {
426        match expr {
427            Expression::Nil(_) => self.push("nil"),
428            Expression::True(_) => self.push("true"),
429            Expression::False(_) => self.push("false"),
430            Expression::Vararg(_) => self.push("..."),
431            Expression::Integer(v, _) => {
432                let s = v.to_string();
433                self.push(&s);
434            }
435            Expression::Float(f, span) => {
436                // NaN and Inf have no Lua literal syntax. Emitting "NaN" or "inf"
437                // produces a bare identifier that Lua treats as a variable name, not
438                // a numeric value — silent correctness failure. Reject early.
439                if f.is_nan() {
440                    return Err(CodeGenError::UnsupportedNode {
441                        target: "Lua 5.1",
442                        detail: "NaN has no literal representation; use `0/0` or avoid NaN"
443                            .to_string(),
444                        span: *span,
445                    });
446                }
447                if f.is_infinite() {
448                    return Err(CodeGenError::UnsupportedNode {
449                        target: "Lua 5.1",
450                        detail: "Infinity has no literal representation; use `math.huge`"
451                            .to_string(),
452                        span: *span,
453                    });
454                }
455                let s = f.to_string();
456                self.push(&s);
457                if !s.contains('.') && !s.contains('e') && !s.contains('E') {
458                    self.push(".0");
459                }
460            }
461            Expression::String(s, _) => {
462                self.push("\"");
463                let escaped = emit_string_content(s);
464                self.push(&escaped);
465                self.push("\"");
466            }
467            Expression::Name(n, _) => self.push(n),
468            Expression::Index(base, field, _) => {
469                // Base may need parens if it's a binary/unary op (unlikely but correct).
470                self.emit_expr_prec(base, u8::MAX)?;
471                self.push(".");
472                self.push(field);
473            }
474            Expression::IndexExpr(base, key, _) => {
475                self.emit_expr_prec(base, u8::MAX)?;
476                self.push("[");
477                self.emit_expression(key)?;
478                self.push("]");
479            }
480            Expression::BinOp(lhs, op, rhs, _) => {
481                let prec = binary_op_prec(*op);
482                let ra = is_right_assoc(*op);
483                let lhs_min = if ra { prec + 1 } else { prec };
484                let rhs_min = if ra { prec } else { prec + 1 };
485                self.emit_expr_prec(lhs, lhs_min)?;
486                self.push(" ");
487                self.push(binary_op_str(*op));
488                self.push(" ");
489                self.emit_expr_prec(rhs, rhs_min)?;
490            }
491            Expression::UnOp(op, operand, _) => {
492                match op {
493                    UnaryOp::Neg => self.push("- "),
494                    UnaryOp::Not => self.push("not "),
495                    UnaryOp::Len => self.push("#"),
496                    UnaryOp::BitwiseNot => self.push("~ "),
497                }
498                // Unary prec is 21; operand needs >= 21 to omit parens.
499                self.emit_expr_prec(operand, 21)?;
500            }
501            Expression::Call(call) => match call {
502                Call::Call { func, args, .. } => {
503                    self.emit_expr_prec(func, u8::MAX)?;
504                    self.push("(");
505                    for (i, arg) in args.iter().enumerate() {
506                        if i > 0 {
507                            self.push(", ");
508                        }
509                        self.emit_expression(arg)?;
510                    }
511                    self.push(")");
512                }
513                Call::MethodCall {
514                    obj, method, args, ..
515                } => {
516                    self.emit_expr_prec(obj, u8::MAX)?;
517                    self.push(":");
518                    self.push(method);
519                    self.push("(");
520                    for (i, arg) in args.iter().enumerate() {
521                        if i > 0 {
522                            self.push(", ");
523                        }
524                        self.emit_expression(arg)?;
525                    }
526                    self.push(")");
527                }
528            },
529            Expression::Function(func) => {
530                self.push("function");
531                self.emit_function_body(func)?;
532            }
533            Expression::Table(t) => {
534                self.push("{");
535                for (i, field) in t.fields.iter().enumerate() {
536                    if i > 0 {
537                        self.push(", ");
538                    }
539                    match field {
540                        TableField::ExprKey { key, value, .. } => {
541                            self.push("[");
542                            self.emit_expression(key)?;
543                            self.push("] = ");
544                            self.emit_expression(value)?;
545                        }
546                        TableField::NameKey { key, value, .. } => {
547                            self.push(key);
548                            self.push(" = ");
549                            self.emit_expression(value)?;
550                        }
551                        TableField::Positional(val) => {
552                            self.emit_expression(val)?;
553                        }
554                    }
555                }
556                self.push("}");
557            }
558        }
559        Ok(())
560    }
561}
562
563fn emit_string_content(s: &str) -> String {
564    let mut out = String::with_capacity(s.len());
565    for ch in s.chars() {
566        match ch {
567            '"' => out.push_str("\\\""),
568            '\\' => out.push_str("\\\\"),
569            '\n' => out.push_str("\\n"),
570            '\r' => out.push_str("\\r"),
571            '\t' => out.push_str("\\t"),
572            '\x07' => out.push_str("\\a"),
573            '\x08' => out.push_str("\\b"),
574            '\x0C' => out.push_str("\\f"),
575            '\x0B' => out.push_str("\\v"),
576            c if (c as u32) < 32 => {
577                out.push('\\');
578                out.push_str(&(c as u32).to_string());
579            }
580            c => out.push(c),
581        }
582    }
583    out
584}
585
586#[cfg(test)]
587mod tests {
588    use super::*;
589
590    #[test]
591    fn test_emit_options_default() {
592        let opts = EmitOptions::default();
593        assert_eq!(opts.indent, "    ");
594        assert_eq!(opts.target, LuaTarget::Lua51);
595        assert!(!opts.emit_header_comment);
596    }
597
598    #[test]
599    fn test_emit_empty_block() {
600        let block = valua_ast::Block {
601            stmts: vec![],
602            span: valua_diagnostics::Span::dummy(),
603        };
604        let out = LuaEmitter::lua51().emit(&block).unwrap();
605        assert_eq!(out, "");
606    }
607
608    #[test]
609    fn test_luajit_target_assumed() {
610        let emitter = LuaEmitter::luajit();
611        assert_eq!(emitter.options.target, LuaTarget::LuaJIT);
612    }
613
614    #[test]
615    fn test_emit_string_content_escapes() {
616        assert_eq!(emit_string_content("hello"), "hello");
617        assert_eq!(emit_string_content("a\"b"), "a\\\"b");
618        assert_eq!(emit_string_content("a\\b"), "a\\\\b");
619        assert_eq!(emit_string_content("a\nb"), "a\\nb");
620    }
621
622    // ── Numeric emission invariants ───────────────────────────────────────────
623
624    fn emit_expr(expr: Expression) -> Result<String, CodeGenError> {
625        let opts = EmitOptions::default();
626        let mut ctx = EmitContext::new(&opts);
627        ctx.emit_expression(&expr)?;
628        Ok(ctx.finish())
629    }
630
631    fn int_expr(v: i64) -> Expression {
632        Expression::Integer(v, valua_diagnostics::Span::dummy())
633    }
634
635    fn float_expr(v: f64) -> Expression {
636        Expression::Float(v, valua_diagnostics::Span::dummy())
637    }
638
639    #[test]
640    fn integer_emitted_as_plain_decimal() {
641        assert_eq!(emit_expr(int_expr(0)).unwrap(), "0");
642        assert_eq!(emit_expr(int_expr(42)).unwrap(), "42");
643        assert_eq!(emit_expr(int_expr(-1)).unwrap(), "-1");
644        assert_eq!(emit_expr(int_expr(255)).unwrap(), "255");
645        // Source hex literal 0xFF → stored as i64(255) → emitted as "255", never "0xff".
646        // This is the decimal-only invariant (PRD §TD1).
647    }
648
649    #[test]
650    fn integer_max_emitted_as_decimal_not_hex() {
651        let s = emit_expr(int_expr(i64::MAX)).unwrap();
652        assert!(
653            !s.contains("0x") && !s.contains("0X"),
654            "must not emit hex: {s}"
655        );
656        assert_eq!(s, i64::MAX.to_string());
657    }
658
659    #[test]
660    fn float_with_fractional_part_emitted_verbatim() {
661        assert_eq!(emit_expr(float_expr(1.5)).unwrap(), "1.5");
662        assert_eq!(emit_expr(float_expr(3.14)).unwrap(), "3.14");
663    }
664
665    #[test]
666    fn float_without_fractional_part_gets_dot_zero_suffix() {
667        // Ensures Lua treats the value as float, not integer.
668        let s = emit_expr(float_expr(1.0)).unwrap();
669        assert!(
670            s.contains('.') || s.contains('e') || s.contains('E'),
671            "float 1.0 must have fractional marker: {s}"
672        );
673        assert_eq!(s, "1.0");
674    }
675
676    #[test]
677    fn float_every_emitted_value_has_decimal_marker() {
678        // Rust Display for f64 uses full decimal notation (not 'e'). For any
679        // finite float that lacks a '.' in its Display string, the emitter
680        // appends ".0". This test verifies no finite float escapes without a
681        // decimal marker — which would cause Lua to treat it as an identifier.
682        let cases = [0.0_f64, 1.0, -1.0, 42.0, 1e10, 1e15, 1e-10, 1e-300];
683        for v in cases {
684            let s = emit_expr(float_expr(v)).unwrap();
685            assert!(
686                s.contains('.') || s.contains('e') || s.contains('E'),
687                "float {v} emitted without decimal marker: {s}"
688            );
689        }
690    }
691
692    #[test]
693    fn float_nan_returns_unsupported_node_error() {
694        let err = emit_expr(float_expr(f64::NAN)).unwrap_err();
695        match err {
696            CodeGenError::UnsupportedNode { ref detail, .. } => {
697                assert!(detail.contains("NaN"), "error must mention NaN: {detail}");
698            }
699            other => panic!("expected UnsupportedNode, got: {other}"),
700        }
701    }
702
703    #[test]
704    fn float_positive_infinity_returns_unsupported_node_error() {
705        let err = emit_expr(float_expr(f64::INFINITY)).unwrap_err();
706        assert!(matches!(err, CodeGenError::UnsupportedNode { .. }));
707    }
708
709    #[test]
710    fn float_negative_infinity_returns_unsupported_node_error() {
711        let err = emit_expr(float_expr(f64::NEG_INFINITY)).unwrap_err();
712        assert!(matches!(err, CodeGenError::UnsupportedNode { .. }));
713    }
714
715    #[test]
716    fn float_no_runtime_type_wrapper_in_output() {
717        // Emitting a float must produce a plain literal — no math.type, no
718        // runtime dispatch, no wrapping call. This is the Case C compliance check.
719        let s = emit_expr(float_expr(3.14)).unwrap();
720        assert!(
721            !s.contains("math"),
722            "emitted float must not reference math.*: {s}"
723        );
724        assert!(
725            !s.contains("type"),
726            "emitted float must not call type(): {s}"
727        );
728        assert!(
729            !s.contains("("),
730            "emitted float literal must not contain a call: {s}"
731        );
732    }
733
734    #[test]
735    fn integer_no_runtime_type_wrapper_in_output() {
736        let s = emit_expr(int_expr(42)).unwrap();
737        assert!(
738            !s.contains("math"),
739            "emitted integer must not reference math.*: {s}"
740        );
741        assert!(
742            !s.contains("("),
743            "emitted integer literal must not contain a call: {s}"
744        );
745    }
746}