From 981ccb047a8d75a534a549edbd60ea16469ad0cf Mon Sep 17 00:00:00 2001 From: hl-valdemar Date: Sun, 28 Jun 2026 22:54:20 +0200 Subject: [PATCH] tagged unions --- TODO.md | 120 +++++++++++++++++++++++- compiler/checker/checker.odin | 19 ++++ compiler/llvm/llvm.odin | 40 +++++++- compiler/parser/parser.odin | 46 ++++++++- compiler/types/types.odin | 64 ++++++++++++- compiler_tests.odin | 99 +++++++++++++++++++ examples/programs/tagged_union/main.bro | 22 +++++ 7 files changed, 396 insertions(+), 14 deletions(-) create mode 100644 examples/programs/tagged_union/main.bro diff --git a/TODO.md b/TODO.md index 55fae48..10364a0 100644 --- a/TODO.md +++ b/TODO.md @@ -412,7 +412,7 @@ a Animal } - # tagged unions (constrained to the backing enum) — DEFERRED to 21.5 + # tagged unions (constrained to the backing enum) — see 21.5 AnimalNameOrHeight :: union(Animal) { # use just `enum` instead of `Animal` for unconstrained tagged union dog []u8 cat []u8 @@ -440,10 +440,33 @@ runtime `{tag, payload}` layout, tag init at construction, and tag extraction — the real codegen work and the foundation for milestone 22 (`match` with payload unwrapping) -21.5. tagged unions with a runtime tag (DEFERRED) - - `union(Enum)` (variant names constrained to the backing enum's members) and `union(enum)` - (auto-generated tag, one member per variant); both store an enum tag beside the payload - (`{tag, payload}`), switchable in milestone 22 +21.5. tagged unions with a runtime tag (implemented; see below) + - `union(Enum)` (variant names constrained to an existing enum's members) and `union(enum)` + (compiler-synthesized anonymous tag enum, one dense 0-based member per variant); both store + the discriminant beside the payload as `{tag, payload}` + - a tagged union is a first-class `.Union` type node whose `child` holds the tag enum (untagged + unions and structs leave `child` INVALID). its fields are the variants (payload types), still + looked up by name. `union(enum)` synthesizes its tag enum up front (`types.enum_anonymous`, + variant names as members), so both forms converge on one representation + - **key reuse:** the per-variant tag value is *derivable* from `(union type, active variant + index)` — the variant's field name matches a member of the tag enum, whose value is the tag. + so codegen computes the tag itself and the existing `.Struct`/`.Field` HIR (carrying `type` + + active/field `integer`) is sufficient. **no HIR/IR/lowering change** (like 18/19/20); the delta + is parser + types layout + one checker validation + three LLVM emit sites + - runtime layout `{tag, payload-carrier}`: tag at offset 0, payload carrier at + `payload_offset = round_up(sizeof(tag), payload_align)` (`types.union_payload_offset`, shared by + `size` and the emitter). field access uses byte-offset GEPs, so offsets stay self-consistent + - construction `T{ variant = value }` reuses the union-literal path and additionally stores the + derived tag; payload read `x.variant` reuses field access, reading at the payload offset + (unchecked reinterpret, Zig-style). safe tag dispatch + payload capture is milestone 22 (`match`) + - changes: `parse_struct` parses `union(...)` (`enum` → synthesize; else an existing enum); + `types.define_record` takes a `tag` param stored in `child`, plus `is_tagged_union` / + `union_tag_enum` / `union_payload_offset` helpers and tagged `size`/`alignment`; the checker + record-decl pass requires the tag to be an enum and each variant to name a member of it; the + LLVM emitter lays out `{tag, [pad], carrier, [pad]}`, stores the tag in construction, and offsets + payload field access + - deferred to milestone 22 (`match`): safe tag dispatch + payload capture (`match x { .bird |v| … }`), + exhaustiveness checking, and any first-class tag-read accessor 22. match statements with tagged unions payload unwrapping @@ -741,6 +764,93 @@ for 0..10 |i| blk: { # bad, no name binds returned value } ``` +## A word on match statements + +``` +# matching on enums +match status { + .ok: print("success") + .error: print("failure") + .pending: { + log("still waiting") + retry() + } +} + +# matching on integers and other values +match code { + 0: print("zero") + 1: print("one") + 2: print("two") + else: print("other") # needed - missing variants +} + +Status :: enum { ok, error, pending } + +status :: get_status() # returns a `Status` +match status { + .ok: handle_ok() + .error: handle_error() + .pending: handle_pending() + # no else needed - all variants covered +} + +# matching on tagged unions +Result :: union(enum) { + success Data + failure struct { + msg []u8 + code i32 + } + pending void +} + +match result { + .success |data|: { # use `|name|` to capture the variant's payload + process(data) + } + .failure |info|: { + print("error {d}: {s}", { info.code, info.msg }) + } + .pending: { + # void payload - no capture needed + wait() + } +} + +# when a union variant has a `void` payload, omit the capture +Event :: union(enum) { + click struct { x i32, y i32 } + keypress KeyCode + quit void +} + +event :: get_event() # returns an `Event` +match event { + .click |pos|: handle_click(pos.x, pos.y) + .keypress |key|: handle_key(key) + .quit: should_exit = true +} + +# single-expression arms yield value implicitly +label :: match priority { + .critical: "CRIT" + .high: "HIGH" + .normal: "NORM" + .low: " LOW" +} + +# multi-statement arms use `yield` +message :: match code { + 0: "success" + 1: { + log("warning encountered") + yield "warning" + } + else: "unknown" +} +``` + ## A word on memory allocation (NOTE THAT SYNTAX MAY NOT MATCH BROLANG EXACTLY AND SHOULD BE TAKEN WITH A GRAIN OF SALT - INSPIRATION ONLY) diff --git a/compiler/checker/checker.odin b/compiler/checker/checker.odin index 2c8679f..d82398f 100644 --- a/compiler/checker/checker.odin +++ b/compiler/checker/checker.odin @@ -1048,6 +1048,25 @@ validate_type_nodes :: proc(checker: ^Checker) { ) } } + // A tagged union (`item.child` set) must tag with an enum, and each variant + // must name a member of it. The synthesized `union(enum)` tag satisfies this + // by construction; the check guards the explicit `union(Enum)` form. + if item.kind == .Union && types.is_valid(item.child) { + if !types.is_enum(item.child, &checker.module.types) { + source.add(checker.diagnostics, source.Span{}, "a tagged union's tag must be an enum") + } else { + for field in types.fields_for(&checker.module.types, id) { + if _, ok := find_enum_member(checker, item.child, symbol.Id(field.name)); !ok { + source.addf( + checker.diagnostics, + source.Span{}, + "union variant '%s' is not a member of the tag enum", + symbol_text(checker, symbol.Id(field.name)), + ) + } + } + } + } } if item.kind == .Function { if !item.c_abi { diff --git a/compiler/llvm/llvm.odin b/compiler/llvm/llvm.odin index cd3acf2..3fcf178 100644 --- a/compiler/llvm/llvm.odin +++ b/compiler/llvm/llvm.odin @@ -711,9 +711,27 @@ emit_instruction_stream :: proc( } fmt.sbprintf(&emitter.builder, " %%union_slot%d = alloca %s, align %d\n", instruction_index, type_name, types.alignment_of(instruction.type, &emitter.module.types, emitter.module.target)) fmt.sbprintf(&emitter.builder, " store %s zeroinitializer, ptr %%union_slot%d\n", type_name, instruction_index) + // A tagged union writes the discriminant (the tag enum member matching the + // active variant's name) at offset 0, then the payload after it; untagged + // unions write the payload at offset 0. + payload_ptr := fmt.tprintf("%%union_slot%d", instruction_index) + tag_enum := types.union_tag_enum(instruction.type, &emitter.module.types) + if types.is_valid(tag_enum) { + tag_value: i128 = 0 + for member in types.enum_members_for(&emitter.module.types, tag_enum) { + if member.name == fields[field_index].name { + tag_value = member.value + break + } + } + fmt.sbprintf(&emitter.builder, " store %s %d, ptr %%union_slot%d\n", llvm_type(tag_enum, &emitter.module.types), tag_value, instruction_index) + offset := types.union_payload_offset(instruction.type, &emitter.module.types, emitter.module.target) + fmt.sbprintf(&emitter.builder, " %%union_payload%d = getelementptr i8, ptr %%union_slot%d, i64 %d\n", instruction_index, instruction_index, offset) + payload_ptr = fmt.tprintf("%%union_payload%d", instruction_index) + } fmt.sbprintf(&emitter.builder, " store %s ", llvm_type(fields[field_index].type, &emitter.module.types)) write_operand(&emitter.builder, instructions, instruction.args[0], fields[field_index].type, &emitter.module.types) - fmt.sbprintf(&emitter.builder, ", ptr %%union_slot%d\n", instruction_index) + fmt.sbprintf(&emitter.builder, ", ptr %s\n", payload_ptr) fmt.sbprintf(&emitter.builder, " %%v%d = load %s, ptr %%union_slot%d\n", instruction_index, type_name, instruction_index) continue } @@ -941,7 +959,7 @@ emit_instruction_stream :: proc( continue } if types.is_union(base_type, &emitter.module.types) { - fmt.sbprintf(&emitter.builder, " %%v%d = getelementptr i8, ptr %%v%d, i64 0\n", instruction_index, instruction.a) + fmt.sbprintf(&emitter.builder, " %%v%d = getelementptr i8, ptr %%v%d, i64 %d\n", instruction_index, instruction.a, types.union_payload_offset(base_type, &emitter.module.types, emitter.module.target)) } else { fmt.sbprintf( &emitter.builder, @@ -1804,10 +1822,24 @@ emit_types :: proc(emitter: ^Emitter) { fmt.sbprintf(&emitter.builder, "[%d x i8]\n", total_size) continue } + // A tagged union lays out `{ tag, [pad], carrier, [pad] }`; the explicit i8 + // padding makes the LLVM type's size and field offsets match the byte offsets + // used by construction and field access. Untagged unions have offset 0 and no tag. + payload_offset := types.union_payload_offset(id, &emitter.module.types, emitter.module.target) strings.write_string(&emitter.builder, "{ ") + if types.is_tagged_union(id, &emitter.module.types) { + tag_type := types.union_tag_enum(id, &emitter.module.types) + tag_size := types.size(tag_type, &emitter.module.types, emitter.module.target) + strings.write_string(&emitter.builder, llvm_type(tag_type, &emitter.module.types)) + if payload_offset > tag_size { + fmt.sbprintf(&emitter.builder, ", [%d x i8]", payload_offset-tag_size) + } + strings.write_string(&emitter.builder, ", ") + } strings.write_string(&emitter.builder, llvm_type(carrier, &emitter.module.types)) - if carrier_size < total_size { - fmt.sbprintf(&emitter.builder, ", [%d x i8]", total_size-carrier_size) + used := payload_offset + carrier_size + if used < total_size { + fmt.sbprintf(&emitter.builder, ", [%d x i8]", total_size-used) } strings.write_string(&emitter.builder, " }\n") continue diff --git a/compiler/parser/parser.odin b/compiler/parser/parser.odin index 5c9dcf2..bc41f99 100644 --- a/compiler/parser/parser.odin +++ b/compiler/parser/parser.odin @@ -1730,6 +1730,22 @@ parse_function :: proc(parser: ^Parser, name: token.Token, c_abi: bool) { parse_struct :: proc(parser: ^Parser, name: token.Token, c_layout: bool, is_union := false) { start := advance(parser) id := types.named(&parser.module.type_store, u32(parser.pkg), u32(name.symbol)) + // A tagged union spells its discriminant in parens: `union(Enum)` reuses an existing + // enum; `union(enum)` synthesizes one from the variant names after the body is parsed. + tag := types.INVALID + inferred_tag := false + if is_union { + if _, ok := allow(parser, .Left_Paren); ok { + if _, enum_ok := allow(parser, .Keyword_Enum); enum_ok { + inferred_tag = true + } else { + tag = parse_type(parser) + } + if _, close_ok := allow(parser, .Right_Paren); !close_ok { + source.add(parser.diagnostics, current(parser).span, "expected ')' after union tag") + } + } + } ended_by_newline := current(parser).kind == .Newline skip_newlines(parser) if current(parser).kind != .Left_Brace { @@ -1739,7 +1755,7 @@ parse_struct :: proc(parser: ^Parser, name: token.Token, c_layout: bool, is_unio "native union declarations require a body" if is_union else "native struct declarations require a body", ) } - if !types.define_record(&parser.module.type_store, id, nil, c_layout, true, is_union) { + if !types.define_record(&parser.module.type_store, id, nil, c_layout, true, is_union, tag=tag) { source.addf(parser.diagnostics, name.span, "duplicate type declaration '%s'", token_text(parser, name)) } if !ended_by_newline { @@ -1775,12 +1791,38 @@ parse_struct :: proc(parser: ^Parser, name: token.Token, c_layout: bool, is_unio if _, ok := allow(parser, .Right_Brace); !ok { source.add(parser.diagnostics, current(parser).span, "expected '}' after struct fields") } - if !types.define_record(&parser.module.type_store, id, fields[:], c_layout, false, is_union) { + if inferred_tag { + tag = synthesize_union_tag(parser, fields[:]) + } + if !types.define_record(&parser.module.type_store, id, fields[:], c_layout, false, is_union, tag=tag) { source.addf(parser.diagnostics, name.span, "duplicate type declaration '%s'", token_text(parser, name)) } _ = finish_statement(parser) } +// synthesize_union_tag builds the anonymous discriminant enum for a `union(enum)`: +// one member per variant, dense 0-based, in the smallest fitting unsigned backing +// (mirroring `parse_enum`'s unbacked pick). +synthesize_union_tag :: proc(parser: ^Parser, fields: []types.Field) -> types.Type { + members := make([]types.Enum_Member, len(fields), parser.module.allocator) + defer delete(members, parser.module.allocator) + for field, index in fields { + members[index] = types.Enum_Member{name=field.name, value=i128(index)} + } + max_value := u64(max(len(fields)-1, 0)) + backing := types.U8 + if max_value > 0xff { + backing = types.U16 + } + if max_value > 0xffff { + backing = types.U32 + } + if max_value > 0xffff_ffff { + backing = types.U64 + } + return types.enum_anonymous(&parser.module.type_store, members, backing) +} + parse_distinct :: proc(parser: ^Parser, name: token.Token) { start := advance(parser) child := parse_type(parser) diff --git a/compiler/types/types.odin b/compiler/types/types.odin index f98a813..a9cc752 100644 --- a/compiler/types/types.odin +++ b/compiler/types/types.odin @@ -229,6 +229,7 @@ define_record :: proc( is_union := false, explicit_size: u64 = 0, explicit_alignment: u32 = 0, + tag: Type = INVALID, ) -> bool { existing, ok := node(store, id) if !ok || (existing.kind != .Named && existing.kind != .Struct && existing.kind != .Union) || @@ -242,12 +243,31 @@ define_record :: proc( store.nodes[index].declared = true store.nodes[index].explicit_size = explicit_size store.nodes[index].explicit_alignment = explicit_alignment + // A tagged union stashes its discriminant enum in `child` (untagged unions and + // structs leave it INVALID); the per-variant tag value is derived from the enum + // member whose name matches the variant, so no extra storage is needed. + store.nodes[index].child = tag store.nodes[index].field_start = u32(len(store.fields)) store.nodes[index].field_count = u32(len(fields)) append(&store.fields, ..fields) return true } +// enum_anonymous interns an unnamed enum (used as the synthesized discriminant of a +// `union(enum)` tagged union). Members carry the variant names with dense 0-based +// values, so the same name→member→value lookup used for `union(Enum)` resolves tags. +enum_anonymous :: proc(store: ^Store, members: []Enum_Member, backing: Type) -> Type { + start := u32(len(store.enum_members)) + append(&store.enum_members, ..members) + return intern(store, Node{ + kind=.Enum, + child=backing, + field_start=start, + field_count=u32(len(members)), + declared=true, + }) +} + define_struct :: proc(store: ^Store, id: Type, fields: []Field, c_layout, opaque: bool) -> bool { return define_record(store, id, fields, c_layout, opaque) } @@ -541,6 +561,36 @@ is_union :: proc(value: Type, store: ^Store) -> bool { return kind(value, store) == .Union } +// A tagged union is a `.Union` whose `child` is a valid enum (the discriminant). +is_tagged_union :: proc(value: Type, store: ^Store) -> bool { + item, ok := node(store, value) + return ok && item.kind == .Union && is_enum(item.child, store) +} + +union_tag_enum :: proc(value: Type, store: ^Store) -> Type { + if !is_tagged_union(value, store) { + return INVALID + } + item, _ := node(store, value) + return item.child +} + +// union_payload_offset is the byte offset of a tagged union's payload carrier (after +// the discriminant), shared by `size` and the LLVM emitter so construction, field +// access, and layout agree. Zero for untagged unions. +union_payload_offset :: proc(value: Type, store: ^Store, selected := target.DEFAULT) -> u64 { + if !is_tagged_union(value, store) { + return 0 + } + item, _ := node(store, value) + tag_size := size(item.child, store, selected) + payload_align: u64 = 1 + for field in fields_for(store, value) { + payload_align = max(payload_align, u64(alignment_of(field.type, store, selected))) + } + return (tag_size+payload_align-1)/payload_align*payload_align +} + is_distinct :: proc(value: Type, store: ^Store) -> bool { return kind(value, store) == .Distinct } @@ -1048,13 +1098,18 @@ size :: proc(value: Type, store: ^Store, selected := target.DEFAULT) -> u64 { if item.explicit_size > 0 { return item.explicit_size } - result: u64 + carrier_size: u64 max_align: u64 = 1 for field in fields_for(store, value) { - result = max(result, size(field.type, store, selected)) + carrier_size = max(carrier_size, size(field.type, store, selected)) max_align = max(max_align, u64(alignment_of(field.type, store, selected))) } - return (result+max_align-1)/max_align*max_align + if is_enum(item.child, store) { + payload_offset := union_payload_offset(value, store, selected) + total_align := max(max_align, u64(alignment_of(item.child, store, selected))) + return (payload_offset+carrier_size+total_align-1)/total_align*total_align + } + return (carrier_size+max_align-1)/max_align*max_align case: return 0 } @@ -1091,6 +1146,9 @@ alignment_of :: proc(value: Type, store: ^Store, selected := target.DEFAULT) -> for field in fields_for(store, value) { result = max(result, alignment_of(field.type, store, selected)) } + if is_enum(item.child, store) { + result = max(result, alignment_of(item.child, store, selected)) + } return result case: return 1 diff --git a/compiler_tests.odin b/compiler_tests.odin index bd81cfd..4831c92 100644 --- a/compiler_tests.odin +++ b/compiler_tests.odin @@ -2152,6 +2152,105 @@ native_union_compiles_and_runs :: proc(t: ^testing.T) { testing.expect_value(t, state.exit_code, 42) } +@(test) +tagged_union_compiles_and_runs :: proc(t: ^testing.T) { + output := "/tmp/brolang-test-tagged-union" + defer _ = os.remove(output) + status := compiler_core.compile_package("examples/programs/tagged_union", output) + testing.expect_value(t, status, 0) + state := run_executable(output) + // Both tagged forms — `union(Animal)` (existing enum tag) and `union(enum)` (synthesized + // tag) — constructed via keyed literals, stored as `{tag, payload}`, with the active + // payload read back at its post-tag offset: 37 + 5 = 42. Non-zero payloads make a wrong + // payload offset (e.g. overlapping the tag) fail the exit code. + testing.expect_value(t, state.exit_code, 42) +} + +@(test) +tagged_union_stores_the_discriminant :: proc(t: ^testing.T) { + // The runtime test observes only the payload; this one checks the *tag* is written. + // `Animal` is unbacked/dense (dog=0, cat=1, bird=2) in a u8 backing, so `Data{ bird = 99 }` + // lays out as `{ i8, [3 x i8], i32 }` (tag at 0, i32 payload at offset 4) and writes the + // discriminant `store i8 2` beside the payload `store i32 99`. + text := `Animal :: enum { + dog + cat + bird +} +Data :: union(Animal) { + dog i32 + bird i32 +} +main :: func() i32 { + x Data = Data{ bird = 99 } + return x.bird +} +` + source_file := source.Source{path="test.bro", text=text} + diagnostics := source.init_diagnostics(&source_file) + defer source.destroy_diagnostics(&diagnostics) + symbols := symbol.init_table() + defer symbol.destroy_table(&symbols) + stream := lexer.lex(&source_file, &diagnostics, &symbols) + defer delete(stream.items) + ast_module := parser.parse(&stream, &source_file, &diagnostics) + defer ast.destroy_module(&ast_module) + hir_module := checker.check(&ast_module, &diagnostics, &symbols) + defer hir.destroy_module(&hir_module) + ir_module := lower.lower(&hir_module) + defer ir.destroy_module(&ir_module) + llvm_text := llvm.emit(&ir_module, &diagnostics, &symbols) + defer delete(llvm_text) + + testing.expect_value(t, len(diagnostics.items), 0) + testing.expect(t, strings.contains(llvm_text, "{ i8, [3 x i8], i32 }")) + testing.expect(t, strings.contains(llvm_text, "store i8 2,")) + testing.expect(t, strings.contains(llvm_text, "store i32 99,")) +} + +@(test) +tagged_union_validation_is_diagnosed :: proc(t: ^testing.T) { + // A `union(T)` tag must be an enum, and every variant of a `union(Enum)` must name a + // member of that enum. + text := `Color :: struct { + r u8 +} +Animal :: enum { + dog + cat +} +BadTag :: union(Color) { + dog i32 +} +BadVariant :: union(Animal) { + snake i32 +} +main :: func() i32 { + return 0 +} +` + source_file := source.Source{path="test.bro", text=text} + diagnostics := source.init_diagnostics(&source_file) + defer source.destroy_diagnostics(&diagnostics) + symbols := symbol.init_table() + defer symbol.destroy_table(&symbols) + stream := lexer.lex(&source_file, &diagnostics, &symbols) + defer delete(stream.items) + ast_module := parser.parse(&stream, &source_file, &diagnostics) + defer ast.destroy_module(&ast_module) + hir_module := checker.check(&ast_module, &diagnostics, &symbols) + defer hir.destroy_module(&hir_module) + + found_tag := false + found_variant := false + for diagnostic in diagnostics.items { + found_tag = found_tag || strings.contains(diagnostic.message, "tagged union's tag must be an enum") + found_variant = found_variant || strings.contains(diagnostic.message, "'snake' is not a member of the tag enum") + } + testing.expect(t, found_tag) + testing.expect(t, found_variant) +} + @(test) yield_misuse_is_diagnosed :: proc(t: ^testing.T) { // A value block that does not end in `yield`, and a `yield` nested inside an diff --git a/examples/programs/tagged_union/main.bro b/examples/programs/tagged_union/main.bro new file mode 100644 index 0000000..1d862e5 --- /dev/null +++ b/examples/programs/tagged_union/main.bro @@ -0,0 +1,22 @@ +Animal :: enum { + dog + cat + bird + lizard +} + +Data :: union(Animal) { + dog i32 + bird i32 +} + +Thing :: union(enum) { + a i32 + b f64 +} + +main :: func() i32 { + x Data = Data{ bird = 37 } + y Thing = Thing{ a = 5 } + return x.bird + y.a +}