diff --git a/cmd/wcc/check.c b/cmd/wcc/check.c index 579eee83..86926ddc 100644 --- a/cmd/wcc/check.c +++ b/cmd/wcc/check.c @@ -909,6 +909,28 @@ cexpr(Checker *c, Node *n) cs->type = vt; for (Node *alt = cs->list; alt; alt = alt->next) alt->type = resolve_type(c, alt); + /* Validity: every `case T =>` pattern must + * refer to a variant of the scrutinee's + * tagged union. Mirrors the existing is/as + * check; `match (u) { case f64 => ... }` + * where f64 isn't a variant of u is dead code + * the dispatch never reaches, so refuse it. */ + if (vt && vt != ty_err && + !variant_present(u->params, vt)) + err(c, cs->pos, + "case: %s is not a variant of %s", + type_name(c->a, vt), + type_name(c->a, st)); + for (Node *alt = cs->list; alt; alt = alt->next) { + if (alt->type == NULL || + alt->type == ty_err) continue; + if (!variant_present(u->params, + alt->type)) + err(c, cs->pos, + "case: %s is not a variant of %s", + type_name(c->a, alt->type), + type_name(c->a, st)); + } if (cs->str && cs->str[0]) scope_define(c->cur, cs->str, SK_VAR, vt, cs); } diff --git a/selfhost/cmd/w6c/main.combined.ww b/selfhost/cmd/w6c/main.combined.ww index 329515b9..f1b1877c 100644 --- a/selfhost/cmd/w6c/main.combined.ww +++ b/selfhost/cmd/w6c/main.combined.ww @@ -4079,6 +4079,31 @@ fn err_match_variant(c: *checker, n: *node, vname: *node) void = { c.errs += 1; }; +// case_variant_in — true iff `pat` (a `case T` pattern, including +// each alt of a multi-pattern) names a variant of the tagged +// union `tagged`. +fn case_variant_in(tagged: *node, pat: *node) bool = { + let v: *node = tagged.list; + for (v != nil) { + if (type_eq_ast(v, pat)) { return true; }; + v = v.next; + }; + return false; +}; + +fn err_bad_case_variant(c: *checker, pat: *node) void = { + os.write(2, "case: not a variant of scrutinee".ptr, 32u64); + if (pat != nil) { + if (pat.kind == N_TNAME) { + os.write(2, " (".ptr, 2u64); + os.write(2, pat.str.ptr, pat.str.len: u64); + os.write(2, ")".ptr, 1u64); + }; + }; + os.write(2, "\n".ptr, 1u64); + c.errs += 1; +}; + fn check_match_exhaustive(c: *checker, n: *node) void = { if (n == nil) { return; }; if (n.lhs == nil) { return; }; @@ -4086,7 +4111,26 @@ fn check_match_exhaustive(c: *checker, n: *node) void = { let u: *node = resolvealias(c, unwrapbang(st)); if (u == nil) { return; }; if (u.kind != N_TTAGGED) { return; }; - // Default arm absorbs anything; skip. + // Validity: every `case T` pattern (and multi-pattern alts) + // must name a variant of u. Catches typos and dead arms that + // the dispatch would never reach. + let cs0: *node = n.list; + for (cs0 != nil) { + if (cs0.lhs != nil) { + if (!case_variant_in(u, cs0.lhs)) { + err_bad_case_variant(c, cs0.lhs); + }; + let alt: *node = cs0.list; + for (alt != nil) { + if (!case_variant_in(u, alt)) { + err_bad_case_variant(c, alt); + }; + alt = alt.next; + }; + }; + cs0 = cs0.next; + }; + // Default arm absorbs anything; skip exhaustiveness. let cs: *node = n.list; for (cs != nil) { if (cs.lhs == nil) { return; }; // default diff --git a/selfhost/cmd/wcc/check.ww b/selfhost/cmd/wcc/check.ww index 92bd67ed..feff0a0e 100644 --- a/selfhost/cmd/wcc/check.ww +++ b/selfhost/cmd/wcc/check.ww @@ -362,6 +362,31 @@ fn err_match_variant(c: *checker, n: *node, vname: *node) void = { c.errs += 1; }; +// case_variant_in — true iff `pat` (a `case T` pattern, including +// each alt of a multi-pattern) names a variant of the tagged +// union `tagged`. +fn case_variant_in(tagged: *node, pat: *node) bool = { + let v: *node = tagged.list; + for (v != nil) { + if (type_eq_ast(v, pat)) { return true; }; + v = v.next; + }; + return false; +}; + +fn err_bad_case_variant(c: *checker, pat: *node) void = { + os.write(2, "case: not a variant of scrutinee".ptr, 32u64); + if (pat != nil) { + if (pat.kind == N_TNAME) { + os.write(2, " (".ptr, 2u64); + os.write(2, pat.str.ptr, pat.str.len: u64); + os.write(2, ")".ptr, 1u64); + }; + }; + os.write(2, "\n".ptr, 1u64); + c.errs += 1; +}; + fn check_match_exhaustive(c: *checker, n: *node) void = { if (n == nil) { return; }; if (n.lhs == nil) { return; }; @@ -369,7 +394,26 @@ fn check_match_exhaustive(c: *checker, n: *node) void = { let u: *node = resolvealias(c, unwrapbang(st)); if (u == nil) { return; }; if (u.kind != N_TTAGGED) { return; }; - // Default arm absorbs anything; skip. + // Validity: every `case T` pattern (and multi-pattern alts) + // must name a variant of u. Catches typos and dead arms that + // the dispatch would never reach. + let cs0: *node = n.list; + for (cs0 != nil) { + if (cs0.lhs != nil) { + if (!case_variant_in(u, cs0.lhs)) { + err_bad_case_variant(c, cs0.lhs); + }; + let alt: *node = cs0.list; + for (alt != nil) { + if (!case_variant_in(u, alt)) { + err_bad_case_variant(c, alt); + }; + alt = alt.next; + }; + }; + cs0 = cs0.next; + }; + // Default arm absorbs anything; skip exhaustiveness. let cs: *node = n.list; for (cs != nil) { if (cs.lhs == nil) { return; }; // default diff --git a/selfhost/cmd/wwdump/main.combined.ww b/selfhost/cmd/wwdump/main.combined.ww index 68feb3ab..6de333ca 100644 --- a/selfhost/cmd/wwdump/main.combined.ww +++ b/selfhost/cmd/wwdump/main.combined.ww @@ -4079,6 +4079,31 @@ fn err_match_variant(c: *checker, n: *node, vname: *node) void = { c.errs += 1; }; +// case_variant_in — true iff `pat` (a `case T` pattern, including +// each alt of a multi-pattern) names a variant of the tagged +// union `tagged`. +fn case_variant_in(tagged: *node, pat: *node) bool = { + let v: *node = tagged.list; + for (v != nil) { + if (type_eq_ast(v, pat)) { return true; }; + v = v.next; + }; + return false; +}; + +fn err_bad_case_variant(c: *checker, pat: *node) void = { + os.write(2, "case: not a variant of scrutinee".ptr, 32u64); + if (pat != nil) { + if (pat.kind == N_TNAME) { + os.write(2, " (".ptr, 2u64); + os.write(2, pat.str.ptr, pat.str.len: u64); + os.write(2, ")".ptr, 1u64); + }; + }; + os.write(2, "\n".ptr, 1u64); + c.errs += 1; +}; + fn check_match_exhaustive(c: *checker, n: *node) void = { if (n == nil) { return; }; if (n.lhs == nil) { return; }; @@ -4086,7 +4111,26 @@ fn check_match_exhaustive(c: *checker, n: *node) void = { let u: *node = resolvealias(c, unwrapbang(st)); if (u == nil) { return; }; if (u.kind != N_TTAGGED) { return; }; - // Default arm absorbs anything; skip. + // Validity: every `case T` pattern (and multi-pattern alts) + // must name a variant of u. Catches typos and dead arms that + // the dispatch would never reach. + let cs0: *node = n.list; + for (cs0 != nil) { + if (cs0.lhs != nil) { + if (!case_variant_in(u, cs0.lhs)) { + err_bad_case_variant(c, cs0.lhs); + }; + let alt: *node = cs0.list; + for (alt != nil) { + if (!case_variant_in(u, alt)) { + err_bad_case_variant(c, alt); + }; + alt = alt.next; + }; + }; + cs0 = cs0.next; + }; + // Default arm absorbs anything; skip exhaustiveness. let cs: *node = n.list; for (cs != nil) { if (cs.lhs == nil) { return; }; // default diff --git a/test/wcc/300_check.c b/test/wcc/300_check.c index 745dde29..56a3c0a8 100644 --- a/test/wcc/300_check.c +++ b/test/wcc/300_check.c @@ -149,6 +149,19 @@ static const struct row rows[] = { /* nullable pointer folding accepts `(*T | void)` */ { "fn lookup(p: *i32) (*i32 | void) = { return p; };", "ok" }, { "fn lookup() (*i32 | void) = { return; };", "ok" }, /* bare return → null */ + + /* `case T` must name a variant of the scrutinee */ + { "fn pick() (i32 | str) = { return 1; }; " + "fn caller() void = { let v: (i32 | str) = pick(); " + "match (v) { case let n: i32 => { }; case let s: str => { }; " + "case let f: f64 => { }; }; };", + "case: f64 is not a variant" }, + { "fn pick() (i32 | str | bool) = { return 1; }; " + "fn caller() i32 = { let v: (i32 | str | bool) = pick(); " + "match (v) { case let n: i32 => return n; " + "case str | f64 => return 9; case let b: bool => return 1; }; " + "return 0; };", + "case: f64 is not a variant" }, }; int diff --git a/test/wcc/950_selfcheck.c b/test/wcc/950_selfcheck.c index fd4efa2f..eba30de4 100644 --- a/test/wcc/950_selfcheck.c +++ b/test/wcc/950_selfcheck.c @@ -92,6 +92,17 @@ static const struct row rows[] = { " return v + 1;\n" "};\n", "?: error variant not in enclosing return" }, + /* case T => names a non-variant */ + { "fn pick() (i32 | str) = { return 1; };\n" + "fn caller() void = {\n" + " let v: (i32 | str) = pick();\n" + " match (v) {\n" + " case let n: i32 => { };\n" + " case let s: str => { };\n" + " case let f: f64 => { };\n" + " };\n" + "};\n", + "case: not a variant" }, }; int @@ -123,6 +134,8 @@ main(void) "?: enclosing fn return is not tagged") != NULL); err_present = err_present || (err && strstr(err, "?: enclosing fn has no tagged-union return") != NULL); + err_present = err_present || (err && strstr(err, + "case: not a variant") != NULL); int ok; if (expected_no_err) ok = !err_present; else ok = got_match;