diff --git a/src/astgen.cpp b/src/astgen.cpp index d851dc4..6bfd9c7 100644 --- a/src/astgen.cpp +++ b/src/astgen.cpp @@ -3036,10 +3036,34 @@ static bool collectSwitchCases(const NodeStore &ns, JamCodegenContext &ctx, return true; } case AstTag::PatEnumVariant: { + bool hasBindings = (p.flags & 1) != 0; + bool inferReceiver = (p.flags & 4) != 0; + // Constant pattern: a bare identifier (infer-receiver, no bindings) + // naming a module-level const. comptime-evaluate it to an integer and + // emit a switch case — the scalar-const -> SwitchInt path rustc takes + // via Const::try_eval_bits. Unfoldable / non-int consts return false + // and fall back to the per-arm compare chain (astgenPatternCompare). + if (!scrutIsEnum && inferReceiver && !hasBindings) { + const std::string &cname = + ctx.getStringPool().get(static_cast(p.rhs)); + if (const auto *mc = ctx.getModuleConst(cname)) { + jam::ComptimeEvaluator ev(ns, ctx.getStringPool(), + ctx.getTypePool()); + jam::ComptimeValue v = + ev.eval(mc->initExpr, jam::ComptimeScope{}); + if (v.isInt()) { + const TypeKey &sk = ctx.getTypePool().get(scrutTy); + bool signedCmp = sk.kind == TypeKind::Int && sk.b != 0; + out.push_back( + SwitchCase{v.intVal.bits, signedCmp, armBlock}); + return true; + } + } + return false; + } if (!scrutIsEnum || einfo == nullptr) return false; // Reject pattern shapes that need a binding block before the // arm body — Switch can only jump straight to `armBlock`. - bool hasBindings = (p.flags & 1) != 0; if (hasBindings) return false; StringIdx variantNameId = static_cast(p.rhs); const std::string &variantName = ctx.getStringPool().get(variantNameId); @@ -3190,6 +3214,21 @@ static void astgenPatternCompare(AstGenCtx &gctx, NodeIdx patIdx, JirRef scrut, einfo = gctx.ctx.getEnum(recvName); } if (einfo == nullptr) { + // Not an enum. A bare-identifier pattern (infer-receiver, no + // bindings) that names a module-level `const` is a CONSTANT + // pattern: compare the scrutinee against the const's value, + // exactly like a literal arm. Mirrors Rust, where an identifier + // pattern that resolves to a const is the const, not a binding. + if (inferReceiver && !hasBindings) { + const std::string &constName = + gctx.ctx.getStringPool().get(variantNameId); + if (const auto *mc = gctx.ctx.getModuleConst(constName)) { + JirRef k = astgenExpr(gctx, mc->initExpr, scrutTy); + JirRef cmp = emitCmp(JirTag::ICmpEq, scrut, k); + emitCondBr(gctx, cmp, armBlock, nextBlock); + return; + } + } failHere(gctx, "astgen: pattern receiver doesn't resolve to an enum"); } diff --git a/tests/unit/test_match_const.jam b/tests/unit/test_match_const.jam new file mode 100644 index 0000000..b94c7f0 --- /dev/null +++ b/tests/unit/test_match_const.jam @@ -0,0 +1,68 @@ +// Constant patterns in match: a bare identifier that names a module-level +// `const` matches the scrutinee against the const's value (like a literal +// arm), not a new binding. Mirrors Rust's const-pattern resolution. + +const { assert } = import("test"); + +const A: u32 = 10; +const B: u32 = 20; +const C: u32 = 30; +const D: u32 = 40; +const E: u32 = 50; +const F: u32 = 60; +const G: u32 = 70; + +fn classify(x: u32) u32 { + match (x) { + A { return 1; } + B { return 2; } + C { return 3; } + _ { return 0; } + } +} + +// const and literal arms in the same match. +fn mixed(x: u32) u32 { + match (x) { + A { return 1; } + 5 { return 5; } + _ { return 0; } + } +} + +tfn matchConstPatterns() { + assert(classify(10), 1); + assert(classify(20), 2); + assert(classify(30), 3); + assert(classify(99), 0); +} + +tfn matchMixedLiteralAndConst() { + assert(mixed(10), 1); + assert(mixed(5), 5); + assert(mixed(0), 0); +} + +// All-const arms: the match lowers to a jump-table switch (collectSwitchCases +// comptime-evaluates each const), and must give the same result as the +// compare-chain it would otherwise use. +fn grade(x: u32) u32 { + match (x) { + A { return 1; } + B { return 2; } + C { return 3; } + D { return 4; } + E { return 5; } + F { return 6; } + G { return 7; } + _ { return 0; } + } +} + +tfn matchManyConstArmsSwitch() { + assert(grade(10), 1); + assert(grade(40), 4); + assert(grade(70), 7); + assert(grade(0), 0); + assert(grade(55), 0); +}