conditional optional unwrapping

This commit is contained in:
2026-06-21 20:38:12 +02:00
parent c90ada608e
commit f4194492cc
10 changed files with 391 additions and 12 deletions
+69 -8
View File
@@ -58,7 +58,6 @@ Build_Ctx :: struct {
global_reads: ^[dynamic]hir.Global_Id,
calls: ^[dynamic]hir.Function_Id,
problematic: ^bool,
has_return: ^bool,
}
Constant_Kind :: enum {
@@ -1412,9 +1411,19 @@ infer_statements :: proc(
}
}
case .If:
_ = infer_expr(checker, statement.expr, locals^[:], pkg, file, demanded)
infer_statements(checker, statement.body, locals, pkg, file, demanded, result)
infer_statements(checker, statement.else_body, locals, pkg, file, demanded, result)
value_type := infer_expr(checker, statement.expr, locals^[:], pkg, file, demanded)
if statement.name != symbol.INVALID {
// Conditional unwrap: `v` is in scope (with the unwrapped type) only inside the then-block.
child := types.child_type(value_type, &checker.module.types) if types.is_optional(value_type, &checker.module.types) else types.INVALID
binding_start := len(locals^)
append(locals, Infer_Local{name = statement.name, type = child})
infer_statements(checker, statement.body, locals, pkg, file, demanded, result)
resize(locals, binding_start)
infer_statements(checker, statement.else_body, locals, pkg, file, demanded, result)
} else {
infer_statements(checker, statement.body, locals, pkg, file, demanded, result)
infer_statements(checker, statement.else_body, locals, pkg, file, demanded, result)
}
}
}
resize(locals, scope_start)
@@ -2949,7 +2958,6 @@ build_block :: proc(ctx: ^Build_Ctx, statements: []ast.Stmt_Id) -> []hir.Stmt_Id
})
ctx.problematic^ = ctx.problematic^ || checker.module.exprs[value].kind == .Invalid
case .Return:
ctx.has_return^ = true
if statement.expr == ast.INVALID_EXPR {
if !types.is_void(ctx.result) {
id := source.add(checker.diagnostics, statement.span, "'return _' is only valid in a void function")
@@ -3007,6 +3015,39 @@ build_block :: proc(ctx: ^Build_Ctx, statements: []ast.Stmt_Id) -> []hir.Stmt_Id
})
}
case .If:
if statement.name != symbol.INVALID {
// Conditional unwrap `if opt |v| { ... }`: `expr` is the optional, `v`
// binds the unwrapped value (immutable) for the duration of the then-block.
value := build_expr(checker, statement.expr, ctx.locals^[:], ctx.global_reads, ctx.calls, types.INVALID, ctx.pkg, ctx.file)
value_type := checker.module.exprs[value].type
child := types.INVALID
if checker.module.exprs[value].kind != .Invalid && !types.is_optional(value_type, &checker.module.types) {
id := source.add(checker.diagnostics, statement.span, "'if' unwrap requires an optional value")
value = invalid_hir_expr(checker, statement.span, id)
ctx.problematic^ = true
} else if checker.module.exprs[value].kind != .Invalid {
child = types.child_type(value_type, &checker.module.types)
}
binding := hir.local_id(len(ctx.hir_locals^))
append(ctx.hir_locals, hir.Local{name = statement.name, type = child, mutable = false})
locals_before := len(ctx.locals^)
append(ctx.locals, Build_Local{name = statement.name, type = child, mutable = false, id = binding})
then_body := build_block(ctx, statement.body)
resize(ctx.locals, locals_before)
else_body: []hir.Stmt_Id = nil
if statement.else_body != nil {
else_body = build_block(ctx, statement.else_body)
}
append(&body, hir.stmt_id(len(checker.module.statements)))
append(&checker.module.statements, hir.Stmt{
kind = .If, span = statement.span, expr = value,
then_body = then_body, else_body = else_body,
local = binding, target = hir.INVALID_EXPR,
diagnostic = source.INVALID_DIAGNOSTIC,
})
ctx.problematic^ = ctx.problematic^ || checker.module.exprs[value].kind == .Invalid
continue
}
condition := build_expr(checker, statement.expr, ctx.locals^[:], ctx.global_reads, ctx.calls, types.BOOL, ctx.pkg, ctx.file)
if checker.module.exprs[condition].kind != .Invalid && !types.is_bool(checker.module.exprs[condition].type) {
id := source.add(checker.diagnostics, statement.span, "'if' condition must be a bool")
@@ -3039,6 +3080,27 @@ build_block :: proc(ctx: ^Build_Ctx, statements: []ast.Stmt_Id) -> []hir.Stmt_Id
return body[:]
}
// Reports whether every control-flow path through `stmts` terminates (returns or traps),
// so the end of the block is unreachable. A `.Return` or `.Trap` terminates outright; an
// `.If` terminates only when it has an `else` and both arms terminate. Recursion into the
// `then_body`/`else_body` slices handles nested ifs and `else if` chains.
all_paths_return :: proc(module: ^hir.Module, stmts: []hir.Stmt_Id) -> bool {
for id in stmts {
statement := module.statements[id]
#partial switch statement.kind {
case .Return, .Trap:
return true
case .If:
if statement.else_body != nil &&
all_paths_return(module, statement.then_body) &&
all_paths_return(module, statement.else_body) {
return true
}
}
}
return false
}
build_function :: proc(checker: ^Checker, id: Spec_Id) {
spec := checker.specs[id]
function := checker.ast_module.functions[spec.template]
@@ -3118,7 +3180,6 @@ build_function :: proc(checker: ^Checker, id: Spec_Id) {
return
}
has_return := false
if signature_diagnostic != source.INVALID_DIAGNOSTIC {
append(&body, hir.stmt_id(len(checker.module.statements)))
append(
@@ -3142,15 +3203,15 @@ build_function :: proc(checker: ^Checker, id: Spec_Id) {
global_reads = &global_reads,
calls = &calls,
problematic = &problematic,
has_return = &has_return,
}
block := build_block(&ctx, function.body)
returns := all_paths_return(&checker.module, block)
for block_stmt in block {
append(&body, block_stmt)
}
delete(block, checker.allocator)
if !types.is_void(spec.result) && !has_return {
if !types.is_void(spec.result) && !returns {
id := source.addf(
checker.diagnostics,
function.span,
+2
View File
@@ -85,6 +85,8 @@ Opcode :: enum u8 {
Length,
Slice_Ptr,
Unwrap,
Optional_Is_Some,
Optional_Value,
Orelse_Begin,
Orelse,
Widen,
+3
View File
@@ -204,6 +204,9 @@ lex :: proc(
case ',':
append_token(&stream, source_file, .Comma, cursor, cursor+1)
cursor += 1
case '|':
append_token(&stream, source_file, .Pipe, cursor, cursor+1)
cursor += 1
case '"':
start := cursor
cursor += 1
+26 -1
View File
@@ -232,7 +232,8 @@ valid_value :: proc(
}
switch instructions[value_id].op {
case .Param, .Const, .String, .Aggregate, .None, .Optional_Some,
.Load_Global, .Function_Address, .Address_Of, .Load, .Slice, .Length, .Slice_Ptr, .Unwrap, .Orelse,
.Load_Global, .Function_Address, .Address_Of, .Load, .Slice, .Length, .Slice_Ptr, .Unwrap,
.Optional_Is_Some, .Optional_Value, .Orelse,
.Widen, .C_Vararg_Promote, .Weaken_Pointer, .Weaken_Slice, .Decay_Array_Pointer,
.Neg_Checked, .Add_Checked, .Pointer_Add, .Not, .Compare, .Call:
return true
@@ -928,6 +929,30 @@ emit_instruction_stream :: proc(
} else {
fmt.sbprintf(&emitter.builder, " %%v%d = extractvalue %s %%v%d, 1\n", instruction_index, llvm_type(optional_type, &emitter.module.types), instruction.a)
}
case .Optional_Is_Some:
optional_type := instructions[instruction.a].type if valid_instruction(instructions, instruction.a) else types.INVALID
item, ok := types.node(&emitter.module.types, optional_type)
if !ok || item.kind != .Optional {
emit_recovery_value(emitter, instruction_index, instruction, "invalid optional presence test")
continue
}
if types.is_pointer(item.child, &emitter.module.types) {
fmt.sbprintf(&emitter.builder, " %%v%d = icmp ne ptr %%v%d, null\n", instruction_index, instruction.a)
} else {
fmt.sbprintf(&emitter.builder, " %%v%d = extractvalue %s %%v%d, 0\n", instruction_index, llvm_type(optional_type, &emitter.module.types), instruction.a)
}
case .Optional_Value:
optional_type := instructions[instruction.a].type if valid_instruction(instructions, instruction.a) else types.INVALID
item, ok := types.node(&emitter.module.types, optional_type)
if !ok || item.kind != .Optional || !types.equal(item.child, instruction.type) {
emit_recovery_value(emitter, instruction_index, instruction, "invalid optional value")
continue
}
if types.is_pointer(item.child, &emitter.module.types) {
fmt.sbprintf(&emitter.builder, " %%v%d = select i1 true, ptr %%v%d, ptr null\n", instruction_index, instruction.a)
} else {
fmt.sbprintf(&emitter.builder, " %%v%d = extractvalue %s %%v%d, 1\n", instruction_index, llvm_type(optional_type, &emitter.module.types), instruction.a)
}
case .Orelse_Begin:
optional_type := instructions[instruction.a].type if valid_instruction(instructions, instruction.a) else types.INVALID
item, ok := types.node(&emitter.module.types, optional_type)
+36 -1
View File
@@ -633,7 +633,22 @@ lower_statements :: proc(state: ^State, statements: []hir.Stmt_Id) {
diagnostic=statement.diagnostic,
})
case .If:
cond := lower_expr(state, statement.expr)
// A conditional unwrap (`if opt |v| { ... }`) carries the binding local id in
// `statement.local`; `expr` is then the optional, not a bool condition. Test it
// for presence, and inside the then-block bind the unwrapped value to the local.
is_unwrap := statement.local != hir.INVALID_LOCAL
opt := ir.INVALID_INSTRUCTION
cond: ir.Instruction_Id
if is_unwrap {
opt = lower_expr(state, statement.expr)
cond = append_instruction(state, ir.Instruction{
op=.Optional_Is_Some, span=statement.span, type=types.BOOL,
target=ir.INVALID_REF, a=opt, b=ir.INVALID_INSTRUCTION,
diagnostic=source.INVALID_DIAGNOSTIC,
})
} else {
cond = lower_expr(state, statement.expr)
}
has_else := statement.else_body != nil
then_lbl := fresh_label(state)
else_lbl := fresh_label(state) if has_else else then_lbl
@@ -649,6 +664,26 @@ lower_statements :: proc(state: ^State, statements: []hir.Stmt_Id) {
target=ir.INVALID_REF, a=ir.INVALID_INSTRUCTION, b=ir.INVALID_INSTRUCTION,
diagnostic=source.INVALID_DIAGNOSTIC,
})
if is_unwrap && int(statement.local) < len(state.func_locals) {
local := state.func_locals[statement.local]
slot := append_instruction(state, ir.Instruction{
op=.Alloca, span=statement.span, type=local.type,
target=ir.local_ref(ir.Local_Id(statement.local)),
a=ir.INVALID_INSTRUCTION, b=ir.INVALID_INSTRUCTION,
diagnostic=source.INVALID_DIAGNOSTIC,
})
state.local_slots[statement.local] = slot
inner := append_instruction(state, ir.Instruction{
op=.Optional_Value, span=statement.span, type=local.type,
target=ir.INVALID_REF, a=opt, b=ir.INVALID_INSTRUCTION,
diagnostic=source.INVALID_DIAGNOSTIC,
})
append_instruction(state, ir.Instruction{
op=.Store, span=statement.span, type=local.type,
target=ir.INVALID_REF, a=slot, b=inner,
diagnostic=source.INVALID_DIAGNOSTIC,
})
}
lower_statements(state, statement.then_body)
append_instruction(state, ir.Instruction{
op=.Br, span=statement.span, type=types.VOID, integer=merge_lbl,
+13
View File
@@ -1079,6 +1079,18 @@ parse_if :: proc(parser: ^Parser) -> ast.Stmt_Id {
parser.no_struct_literal = true
condition := parse_expression(parser)
parser.no_struct_literal = saved
binding := symbol.INVALID
if _, ok := allow(parser, .Pipe); ok {
name_tok, name_ok := allow(parser, .Identifier)
if name_ok {
binding = name_tok.symbol
} else {
source.add(parser.diagnostics, current(parser).span, "expected a binding name after '|'")
}
if _, close_ok := allow(parser, .Pipe); !close_ok {
source.add(parser.diagnostics, current(parser).span, "expected '|' to close the unwrap binding")
}
}
skip_newlines(parser)
then_body := parse_block(parser)
else_body: []ast.Stmt_Id = nil
@@ -1102,6 +1114,7 @@ parse_if :: proc(parser: ^Parser) -> ast.Stmt_Id {
append(&parser.module.statements, ast.Stmt{
kind=.If,
span=span_from(start.span, previous(parser).span),
name=binding,
expr=condition,
body=then_body,
else_body=else_body,
+1
View File
@@ -40,6 +40,7 @@ Kind :: enum u8 {
Left_Brace,
Right_Brace,
Comma,
Pipe,
Keyword_Func,
Keyword_C_Func,
Keyword_Struct,