Harbor

branch main
showing the latest snapshot on main
parser.odin 46.0 KB · Plain text
gpu/shader/parser.odin 0644 Raw
package shader

import "core:fmt"
import "core:strconv"
import "core:strings"

Parser :: struct {
	tokens:       []Token,
	pos:          int,
	diagnostics:  [dynamic]Diagnostic,
	file:         string,
	panic_mode:   bool,
	error_count:  int,
	block_starts: [dynamic]Block_Start,
}

Block_Start :: struct {
	kind: string, // "function", "if", "for", "while", "struct"
	span: Source_Span,
}

parser_init :: proc(tokens: []Token, file := "", allocator := context.allocator) -> Parser {
	return Parser{
		tokens       = tokens,
		pos          = 0,
		diagnostics  = make([dynamic]Diagnostic, allocator),
		file         = file,
		block_starts = make([dynamic]Block_Start, allocator),
	}
}

parse_module :: proc(p: ^Parser) -> (^Ast_Module, []Diagnostic) {
	mod := new(Ast_Module)
	mod.structs = make([dynamic]^Ast_Struct)
	mod.functions = make([dynamic]^Ast_Function)
	mod.bindings = make([dynamic]^Ast_Binding)
	mod.constants = make([dynamic]^Ast_Const)
	mod.shared_vars = make([dynamic]^Ast_Shared)

	for !is_at_end_p(p) {
		p.panic_mode = false // clear panic mode at top-level
		attrs := parse_attributes(p)

		#partial switch current(p).kind {
		case .KW_Function:
			fn := parse_function(p, attrs)
			if fn != nil do append(&mod.functions, fn)
		case .KW_Struct:
			s := parse_struct(p, attrs)
			if s != nil do append(&mod.structs, s)
		case .KW_Uniform:
			b := parse_binding(p, .Uniform, attrs)
			if b != nil do append(&mod.bindings, b)
		case .KW_Buffer:
			b := parse_binding(p, .Buffer, attrs)
			if b != nil do append(&mod.bindings, b)
		case .KW_Const:
			c := parse_const(p, attrs)
			if c != nil do append(&mod.constants, c)
		case .KW_Shared:
			s := parse_shared(p)
			if s != nil do append(&mod.shared_vars, s)
		case .EOF:
			break
		case:
			if is_stage_starter(p) {
				fn := parse_stage_block(p, attrs)
				if fn != nil do append(&mod.functions, fn)
				continue
			} else if check_identifier_text(p, "group") {
				bindings := parse_group_block(p, attrs)
				for b in bindings {
					if b != nil do append(&mod.bindings, b)
				}
				continue
			}
			error(p, fmt.aprintf("unexpected token %s at top level", token_kind_to_string(current(p).kind)))
			advance_p(p)
		}
	}

	return mod, p.diagnostics[:]
}

// -- Parsing helpers --

@(private = "file")
current :: proc(p: ^Parser) -> Token {
	if p.pos >= len(p.tokens) {
		return Token{kind = .EOF}
	}
	return p.tokens[p.pos]
}

@(private = "file")
peek_p :: proc(p: ^Parser, offset := 1) -> Token {
	idx := p.pos + offset
	if idx >= len(p.tokens) {
		return Token{kind = .EOF}
	}
	return p.tokens[idx]
}

@(private = "file")
advance_p :: proc(p: ^Parser) -> Token {
	tok := current(p)
	if p.pos < len(p.tokens) {
		p.pos += 1
	}
	return tok
}

@(private = "file")
expect :: proc(p: ^Parser, kind: Token_Kind) -> (Token, bool) {
	tok := current(p)
	if tok.kind != kind {
		// Try single-token lookahead recovery: if the expected token is 1 ahead, skip garbage
		if !p.panic_mode && peek_p(p).kind == kind {
			error(p, fmt.aprintf("expected %s, got %s", token_kind_to_string(kind), token_kind_to_string(tok.kind)))
			p.panic_mode = false // recovered from the error we just set
			advance_p(p) // skip garbage
			return advance_p(p), true
		}
		error(p, fmt.aprintf("expected %s, got %s", token_kind_to_string(kind), token_kind_to_string(tok.kind)))
		return tok, false
	}
	return advance_p(p), true
}

@(private = "file")
check :: proc(p: ^Parser, kind: Token_Kind) -> bool {
	return current(p).kind == kind
}

@(private = "file")
match :: proc(p: ^Parser, kinds: ..Token_Kind) -> bool {
	for kind in kinds {
		if current(p).kind == kind {
			advance_p(p)
			return true
		}
	}
	return false
}

@(private = "file")
check_identifier_text :: proc(p: ^Parser, text: string) -> bool {
	tok := current(p)
	return tok.kind == .Identifier && tok.text == text
}

@(private = "file")
is_at_end_p :: proc(p: ^Parser) -> bool {
	return current(p).kind == .EOF
}

@(private = "file")
is_toplevel_keyword :: proc(p: ^Parser) -> bool {
	#partial switch current(p).kind {
	case .KW_Function, .KW_Struct, .KW_Uniform, .KW_Buffer, .KW_Const, .KW_Shared:
		return true
	}
	return false
}

@(private = "file")
error :: proc(p: ^Parser, msg: string) {
	if p.panic_mode do return
	if p.error_count >= MAX_ERRORS {
		if p.error_count == MAX_ERRORS {
			append(&p.diagnostics, Diagnostic{
				level   = .Error,
				message = "too many errors, stopping",
				span    = current(p).span,
			})
			p.error_count += 1
		}
		p.panic_mode = true
		return
	}
	span := current(p).span
	append(&p.diagnostics, Diagnostic{
		level   = .Error,
		message = msg,
		span    = span,
	})
	p.error_count += 1
	p.panic_mode = true
}

@(private = "file")
report_error_at :: proc(p: ^Parser, span: Source_Span, msg: string) {
	if p.error_count >= MAX_ERRORS {
		if p.error_count == MAX_ERRORS {
			append(&p.diagnostics, Diagnostic{
				level   = .Error,
				message = "too many errors, stopping",
				span    = span,
			})
			p.error_count += 1
		}
		return
	}
	append(&p.diagnostics, Diagnostic{
		level   = .Error,
		message = msg,
		span    = span,
	})
	p.error_count += 1
}

@(private = "file")
synchronize :: proc(p: ^Parser) {
	p.panic_mode = false
	for !is_at_end_p(p) {
		#partial switch current(p).kind {
		case .KW_Function, .KW_Struct, .KW_Let, .KW_Var, .KW_Return, .KW_If, .KW_For, .KW_While, .KW_End, .KW_Uniform, .KW_Buffer, .KW_Const, .KW_Shared, .KW_Break, .KW_Continue:
			return
		}
		advance_p(p)
	}
}

@(private = "file")
is_stage_starter :: proc(p: ^Parser) -> bool {
	if current(p).kind != .Identifier do return false
	switch current(p).text {
	case "vertex", "fragment", "compute":
		return true
	}
	return false
}

// -- Attributes --

@(private = "file")
parse_attributes :: proc(p: ^Parser) -> []Ast_Attribute {
	attrs := make([dynamic]Ast_Attribute)
	for check(p, .At) {
		advance_p(p) // @
		name_tok, ok := expect(p, .Identifier)
		if !ok {
			synchronize(p)
			continue
		}

		args := make([dynamic]string)
		if check(p, .Lparen) {
			advance_p(p)
			for !check(p, .Rparen) && !is_at_end_p(p) {
				arg_tok := advance_p(p)
				append(&args, arg_tok.text)
				if !check(p, .Rparen) {
					if !match(p, .Comma) {
						// Allow args without comma separation for simple cases
					}
				}
			}
			expect(p, .Rparen)
		}

		append(&attrs, Ast_Attribute{
			name = name_tok.text,
			args = args[:],
			span = name_tok.span,
		})
	}
	return attrs[:]
}

// -- Top-level declarations --

@(private = "file")
parse_function :: proc(p: ^Parser, attrs: []Ast_Attribute) -> ^Ast_Function {
	span_start := current(p).span
	expect(p, .KW_Function)

	name_tok, ok := expect(p, .Identifier)
	if !ok {
		synchronize(p)
		return nil
	}

	// Parameters
	expect(p, .Lparen)
	params := parse_params(p)
	expect(p, .Rparen)

	// Return type (optional) — can be a type name, array, or tuple (name: Type, ...)
	ret_type: ^Type_Expr
	if match(p, .Arrow) {
		if check(p, .Lparen) {
			// Tuple return: -> (name: Type, ...)
			te := parse_tuple_type(p)
			ret_type = new(Type_Expr)
			ret_type^ = te
		} else {
			te := parse_type_expr(p)
			ret_type = new(Type_Expr)
			ret_type^ = te
		}
	}

	// Body
	append(&p.block_starts, Block_Start{kind = "function", span = span_start})
	body := parse_body(p)

	// End
	if !match(p, .KW_End) {
		bs := pop(&p.block_starts) if len(p.block_starts) > 0 else Block_Start{}
		error(p, fmt.aprintf("missing 'end' to close 'function' started at line %d", bs.span.line_start))
	} else {
		if len(p.block_starts) > 0 do pop(&p.block_starts)
	}

	fn := new(Ast_Function)
	fn^ = Ast_Function{
		name        = name_tok.text,
		params      = params,
		return_type = ret_type,
		body        = body,
		attributes  = attrs,
		span        = span_start,
	}
	return fn
}

@(private = "file")
parse_stage_block :: proc(p: ^Parser, attrs: []Ast_Attribute) -> ^Ast_Function {
	span_start := current(p).span
	stage_tok := advance_p(p) // contextual vertex, fragment, compute

	name_tok, ok := expect(p, .Identifier)
	if !ok {
		synchronize(p)
		return nil
	}

	header_attrs := parse_attributes(p)
	inputs := make([dynamic]Ast_Param)
	outputs := make([dynamic]Ast_Struct_Field)

	for !check(p, .KW_Do) && !check(p, .KW_End) && !is_at_end_p(p) {
		start_pos := p.pos
		if check(p, .KW_In) {
			input := parse_stage_input(p)
			append(&inputs, input)
		} else if check_identifier_text(p, "out") {
			output := parse_stage_output(p)
			append(&outputs, output)
		} else {
			error(p, "expected 'in', 'out', or 'do' in stage block")
			synchronize(p)
		}
		if p.pos == start_pos do advance_p(p)
	}

	if stage_tok.text == "compute" && len(outputs) > 0 {
		report_error_at(p, outputs[0].span, "compute stages cannot declare out slots")
	}

	expect(p, .KW_Do)
	append(&p.block_starts, Block_Start{kind = stage_tok.text, span = span_start})
	body := parse_body(p)
	if !match(p, .KW_End) {
		bs := pop(&p.block_starts) if len(p.block_starts) > 0 else Block_Start{}
		error(p, fmt.aprintf("missing 'end' to close '%s' started at line %d", stage_tok.text, bs.span.line_start))
	} else {
		if len(p.block_starts) > 0 do pop(&p.block_starts)
	}

	rewrite_workflow_builtin_namespaces(p, stage_tok.text, &inputs, body[:])
	body = lower_workflow_outputs(p, stage_tok.text, name_tok.text, outputs[:], body)

	return_type: ^Type_Expr
	if len(outputs) > 0 {
		te := Type_Tuple{
			fields = outputs[:],
			span   = span_start,
		}
		return_type = new(Type_Expr)
		return_type^ = te
	}

	fn_attrs := make([dynamic]Ast_Attribute)
	for attr in attrs do append(&fn_attrs, attr)
	for attr in header_attrs do append(&fn_attrs, attr)
	append(&fn_attrs, make_string_attribute("entry", stage_tok.text, stage_tok.span))
	append(&fn_attrs, Ast_Attribute{name = "luma_workflow", span = stage_tok.span})

	fn := new(Ast_Function)
	fn^ = Ast_Function{
		name        = name_tok.text,
		params      = inputs[:],
		return_type = return_type,
		body        = body,
		attributes  = fn_attrs[:],
		span        = span_start,
	}
	return fn
}

@(private = "file")
parse_stage_input :: proc(p: ^Parser) -> Ast_Param {
	span_start := current(p).span
	advance_p(p) // in

	name_tok, ok := expect(p, .Identifier)
	if !ok {
		return Ast_Param{span = span_start}
	}
	expect(p, .Colon)
	type_expr := parse_type_expr(p)
	attrs := parse_attributes(p)

	return Ast_Param{
		name       = name_tok.text,
		type       = type_expr,
		attributes = attrs,
		span       = name_tok.span,
	}
}

@(private = "file")
parse_stage_output :: proc(p: ^Parser) -> Ast_Struct_Field {
	span_start := current(p).span
	advance_p(p) // contextual out

	name_tok, ok := expect(p, .Identifier)
	if !ok {
		return Ast_Struct_Field{span = span_start}
	}
	expect(p, .Colon)
	type_expr := parse_type_expr(p)
	attrs := parse_attributes(p)

	return Ast_Struct_Field{
		name       = name_tok.text,
		type       = type_expr,
		attributes = attrs,
		span       = name_tok.span,
	}
}

@(private = "file")
make_string_attribute :: proc(name, arg: string, span: Source_Span) -> Ast_Attribute {
	args := make([]string, 1)
	args[0] = arg
	return Ast_Attribute{
		name = name,
		args = args,
		span = span,
	}
}

@(private = "file")
rewrite_workflow_builtin_namespaces :: proc(
	p: ^Parser,
	stage: string,
	inputs: ^[dynamic]Ast_Param,
	body: []^Ast_Node,
) {
	for stmt in body {
		rewrite_workflow_builtins_in_node(p, stage, inputs, stmt)
	}
}

@(private = "file")
rewrite_workflow_builtins_in_node :: proc(
	p: ^Parser,
	stage: string,
	inputs: ^[dynamic]Ast_Param,
	node: ^Ast_Node,
) {
	if node == nil do return

	#partial switch d in node.derived {
	case ^Ast_Let:
		rewrite_workflow_builtins_in_expr(p, stage, inputs, &d.value)
	case ^Ast_Assign:
		rewrite_workflow_builtins_in_expr(p, stage, inputs, &d.target)
		rewrite_workflow_builtins_in_expr(p, stage, inputs, &d.value)
	case ^Ast_Return:
		rewrite_workflow_builtins_in_expr(p, stage, inputs, &d.value)
	case ^Ast_If:
		rewrite_workflow_builtins_in_expr(p, stage, inputs, &d.condition)
		for stmt in d.then_body do rewrite_workflow_builtins_in_node(p, stage, inputs, stmt)
		for &clause in d.elseif_clauses {
			rewrite_workflow_builtins_in_expr(p, stage, inputs, &clause.condition)
			for stmt in clause.body do rewrite_workflow_builtins_in_node(p, stage, inputs, stmt)
		}
		for stmt in d.else_body do rewrite_workflow_builtins_in_node(p, stage, inputs, stmt)
	case ^Ast_For:
		rewrite_workflow_builtins_in_expr(p, stage, inputs, &d.start)
		rewrite_workflow_builtins_in_expr(p, stage, inputs, &d.stop)
		if d.step != nil do rewrite_workflow_builtins_in_expr(p, stage, inputs, &d.step)
		for stmt in d.body do rewrite_workflow_builtins_in_node(p, stage, inputs, stmt)
	case ^Ast_While:
		rewrite_workflow_builtins_in_expr(p, stage, inputs, &d.condition)
		for stmt in d.body do rewrite_workflow_builtins_in_node(p, stage, inputs, stmt)
	case ^Ast_Call, ^Ast_Field_Access, ^Ast_Index, ^Ast_Swizzle, ^Ast_Unary, ^Ast_Binary, ^Ast_Struct_Literal:
		expr := node
		rewrite_workflow_builtins_in_expr(p, stage, inputs, &expr)
	}
}

@(private = "file")
rewrite_workflow_builtins_in_expr :: proc(
	p: ^Parser,
	stage: string,
	inputs: ^[dynamic]Ast_Param,
	node_ptr: ^^Ast_Node,
) {
	node := node_ptr^
	if node == nil do return

	if field, ok := node.derived.(^Ast_Field_Access); ok {
		if obj, obj_ok := field.object.derived.(^Ast_Ident); obj_ok {
			if is_workflow_builtin_namespace(obj.name) {
				field_name, ok := workflow_builtin_input_name(p, stage, inputs, obj.name, field.field, node.span)
				if ok {
					ident := new(Ast_Ident)
					ident^ = Ast_Ident{name = field_name, span = node.span}
					node_ptr^ = make_node(.Ident, ident, node.span)
				}
				return
			}
		}
		rewrite_workflow_builtins_in_expr(p, stage, inputs, &field.object)
		return
	}

	#partial switch d in node.derived {
	case ^Ast_Call:
		rewrite_workflow_builtins_in_expr(p, stage, inputs, &d.callee)
		for &arg in d.args do rewrite_workflow_builtins_in_expr(p, stage, inputs, &arg)
	case ^Ast_Index:
		rewrite_workflow_builtins_in_expr(p, stage, inputs, &d.object)
		rewrite_workflow_builtins_in_expr(p, stage, inputs, &d.index)
	case ^Ast_Swizzle:
		rewrite_workflow_builtins_in_expr(p, stage, inputs, &d.object)
	case ^Ast_Unary:
		rewrite_workflow_builtins_in_expr(p, stage, inputs, &d.operand)
	case ^Ast_Binary:
		rewrite_workflow_builtins_in_expr(p, stage, inputs, &d.left)
		rewrite_workflow_builtins_in_expr(p, stage, inputs, &d.right)
	case ^Ast_Struct_Literal:
		for &f in d.fields do rewrite_workflow_builtins_in_expr(p, stage, inputs, &f.value)
	}
}

@(private = "file")
is_workflow_builtin_namespace :: proc(name: string) -> bool {
	switch name {
	case "vertex", "frag", "work":
		return true
	}
	return false
}

@(private = "file")
workflow_builtin_input_name :: proc(
	p: ^Parser,
	stage: string,
	inputs: ^[dynamic]Ast_Param,
	namespace: string,
	member: string,
	span: Source_Span,
) -> (string, bool) {
	required_stage, builtin, type_name, ok := workflow_builtin_info(namespace, member)
	if !ok {
		report_error_at(p, span, fmt.aprintf("unknown builtin '%s.%s'", namespace, member))
		return "", false
	}
	if stage != required_stage {
		report_error_at(p, span, fmt.aprintf("builtin '%s.%s' is only valid in %s stages", namespace, member, required_stage))
		return "", false
	}

	if existing, found := workflow_input_name_for_builtin(inputs^[:], builtin); found {
		return existing, true
	}

	name := fmt.aprintf("__luma_%s_%s", namespace, member)
	append(inputs, Ast_Param{
		name       = name,
		type       = Type_Named{name = type_name, span = span},
		attributes = make_single_attribute_slice(make_string_attribute("builtin", builtin, span)),
		span       = span,
	})
	return name, true
}

@(private = "file")
make_single_attribute_slice :: proc(attr: Ast_Attribute) -> []Ast_Attribute {
	attrs := make([]Ast_Attribute, 1)
	attrs[0] = attr
	return attrs
}

@(private = "file")
workflow_builtin_info :: proc(namespace, member: string) -> (stage: string, builtin: string, type_name: string, ok: bool) {
	if namespace == "vertex" {
		stage = "vertex"
		switch member {
		case "index":
			return stage, "vertex_id", "uint", true
		case "instance":
			return stage, "instance_id", "uint", true
		}
	}
	if namespace == "frag" {
		stage = "fragment"
		switch member {
		case "coord":
			return stage, "frag_coord", "vec4", true
		case "front_facing":
			return stage, "front_facing", "bool", true
		}
	}
	if namespace == "work" {
		stage = "compute"
		switch member {
		case "global_id":
			return stage, "global_invocation_id", "uvec3", true
		case "local_id":
			return stage, "local_invocation_id", "uvec3", true
		case "group_id":
			return stage, "workgroup_id", "uvec3", true
		case "local_index":
			return stage, "local_invocation_index", "uint", true
		}
	}
	return "", "", "", false
}

@(private = "file")
workflow_input_name_for_builtin :: proc(inputs: []Ast_Param, builtin: string) -> (string, bool) {
	for input in inputs {
		if get_builtin_name(input.attributes) == builtin {
			return input.name, true
		}
	}
	return "", false
}

@(private = "file")
lower_workflow_outputs :: proc(
	p: ^Parser,
	stage: string,
	fn_name: string,
	outputs: []Ast_Struct_Field,
	body: [dynamic]^Ast_Node,
) -> [dynamic]^Ast_Node {
	if stage == "compute" {
		validate_no_workflow_returns(p, stage, body[:])
		return body
	}

	state := workflow_assign_state_make(len(outputs))
	flags := Workflow_Output_Flags{}
	new_body := workflow_transform_body(p, stage, outputs, body[:], state, &flags, true, false).body
	return new_body
}

Workflow_Assign_State :: struct {
	assigned:   []bool,
	terminated: bool,
}

Workflow_Output_Flags :: struct {
	explicit_write_seen: bool,
	return_sugar_seen:   bool,
}

Workflow_Body_Result :: struct {
	body:  [dynamic]^Ast_Node,
	state: Workflow_Assign_State,
}

@(private = "file")
workflow_assign_state_make :: proc(count: int) -> Workflow_Assign_State {
	return {assigned = make([]bool, count)}
}

@(private = "file")
workflow_assign_state_clone :: proc(state: Workflow_Assign_State) -> Workflow_Assign_State {
	assigned := make([]bool, len(state.assigned))
	for value, i in state.assigned do assigned[i] = value
	return {assigned = assigned, terminated = state.terminated}
}

@(private = "file")
workflow_transform_body :: proc(
	p: ^Parser,
	stage: string,
	outputs: []Ast_Struct_Field,
	stmts: []^Ast_Node,
	state: Workflow_Assign_State,
	flags: ^Workflow_Output_Flags,
	top_level: bool,
	in_loop: bool,
) -> Workflow_Body_Result {
	result := make([dynamic]^Ast_Node)
	current_state := workflow_assign_state_clone(state)

	for stmt in stmts {
		if stmt == nil do continue
		transformed := workflow_transform_stmt(p, stage, outputs, stmt, current_state, flags, top_level, in_loop)
		current_state = transformed.state
		append(&result, ..transformed.body[:])
	}

	if top_level && !current_state.terminated {
		workflow_report_missing_outputs(p, outputs, current_state)
	}

	return {body = result, state = current_state}
}

@(private = "file")
workflow_transform_stmt :: proc(
	p: ^Parser,
	stage: string,
	outputs: []Ast_Struct_Field,
	stmt: ^Ast_Node,
	state: Workflow_Assign_State,
	flags: ^Workflow_Output_Flags,
	top_level: bool,
	in_loop: bool,
) -> Workflow_Body_Result {
	body := make([dynamic]^Ast_Node)
	current_state := workflow_assign_state_clone(state)

	if assign, ok := stmt.derived.(^Ast_Assign); ok {
		if slot, slot_ok := workflow_output_assignment_slot(p, assign, outputs); slot_ok {
			if flags.return_sugar_seen {
				report_error_at(p, assign.span, "cannot mix fragment return sugar with explicit out writes")
			}
			flags.explicit_write_seen = true
			if !in_loop && !current_state.terminated {
				current_state.assigned[slot] = true
			}
			append(&body, make_workflow_output_assign(outputs[slot], slot, assign.value, assign.span))
			return {body = body, state = current_state}
		}
	}

	if ret, ok := stmt.derived.(^Ast_Return); ok {
		if stage == "fragment" && len(outputs) == 1 && ret.value != nil && top_level && !in_loop {
			if flags.explicit_write_seen {
				report_error_at(p, ret.span, "cannot mix explicit out writes with fragment return sugar")
			}
			flags.return_sugar_seen = true
			if !current_state.terminated {
				current_state.assigned[0] = true
			}
			append(&body, make_workflow_output_assign(outputs[0], 0, ret.value, ret.span))
			return {body = body, state = current_state}
		}
		report_error_at(p, ret.span, fmt.aprintf("'%s' workflow stage return sugar is only supported as a top-level single-output fragment statement", stage))
		return {body = body, state = current_state}
	}

	#partial switch d in stmt.derived {
	case ^Ast_If:
		if_result := workflow_transform_if(p, stage, outputs, d, current_state, flags, in_loop)
		append(&body, ..if_result.body[:])
		current_state = if_result.state
	case ^Ast_For:
		body_result := workflow_transform_body(p, stage, outputs, d.body[:], current_state, flags, false, true)
		d.body = body_result.body
		append(&body, stmt)
	case ^Ast_While:
		body_result := workflow_transform_body(p, stage, outputs, d.body[:], current_state, flags, false, true)
		d.body = body_result.body
		append(&body, stmt)
	case ^Ast_Discard:
		current_state.terminated = true
		append(&body, stmt)
	case:
		append(&body, stmt)
	}

	return {body = body, state = current_state}
}

@(private = "file")
workflow_transform_if :: proc(
	p: ^Parser,
	stage: string,
	outputs: []Ast_Struct_Field,
	if_node: ^Ast_If,
	in_state: Workflow_Assign_State,
	flags: ^Workflow_Output_Flags,
	in_loop: bool,
) -> Workflow_Body_Result {
	body := make([dynamic]^Ast_Node)

	then_result := workflow_transform_body(p, stage, outputs, if_node.then_body[:], in_state, flags, false, in_loop)
	if_node.then_body = then_result.body

	branch_states := make([dynamic]Workflow_Assign_State)
	defer delete(branch_states)
	append(&branch_states, then_result.state)

	for &clause in if_node.elseif_clauses {
		clause_result := workflow_transform_body(p, stage, outputs, clause.body[:], in_state, flags, false, in_loop)
		clause.body = clause_result.body
		append(&branch_states, clause_result.state)
	}

	if len(if_node.else_body) > 0 {
		else_result := workflow_transform_body(p, stage, outputs, if_node.else_body[:], in_state, flags, false, in_loop)
		if_node.else_body = else_result.body
		append(&branch_states, else_result.state)
	} else {
		append(&branch_states, workflow_assign_state_clone(in_state))
	}

	merged := workflow_assign_state_clone(in_state)
	if in_loop {
		append(&body, make_node(.If, if_node, if_node.span))
		return {body = body, state = merged}
	}

	merged = workflow_assign_state_make(len(outputs))
	all_terminated := true
	for branch in branch_states {
		if !branch.terminated {
			all_terminated = false
			break
		}
	}
	merged.terminated = all_terminated
	if all_terminated {
		append(&body, make_node(.If, if_node, if_node.span))
		return {body = body, state = merged}
	}

	for _, i in outputs {
		definitely_assigned := true
		for branch in branch_states {
			if branch.terminated do continue
			if !branch.assigned[i] {
				definitely_assigned = false
				break
			}
		}
		merged.assigned[i] = definitely_assigned
	}
	append(&body, make_node(.If, if_node, if_node.span))
	return {body = body, state = merged}
}

@(private = "file")
workflow_report_missing_outputs :: proc(p: ^Parser, outputs: []Ast_Struct_Field, state: Workflow_Assign_State) {
	for out, i in outputs {
		if !state.assigned[i] {
			report_error_at(p, out.span, fmt.aprintf("out slot '%s' is not definitely assigned on every path", out.name))
		}
	}
}

@(private = "file")
make_workflow_output_assign :: proc(out: Ast_Struct_Field, io_index: int, value: ^Ast_Node, span: Source_Span) -> ^Ast_Node {
	assign := new(Ast_Output_Assign)
	assign^ = Ast_Output_Assign{
		name     = out.name,
		io_index = io_index,
		type     = out.type,
		value    = value,
		span     = span,
	}
	return make_node(.Output_Assign, assign, span)
}

@(private = "file")
workflow_output_assignment_slot :: proc(p: ^Parser, assign: ^Ast_Assign, outputs: []Ast_Struct_Field) -> (int, bool) {
	if ident, ok := assign.target.derived.(^Ast_Ident); ok {
		return workflow_output_slot_index(outputs, ident.name)
	}
	if field, ok := assign.target.derived.(^Ast_Field_Access); ok {
		if obj, obj_ok := field.object.derived.(^Ast_Ident); obj_ok && obj.name == "out" {
			idx, found := workflow_output_slot_index(outputs, field.field)
			if !found {
				report_error_at(p, assign.span, fmt.aprintf("unknown out slot '%s'", field.field))
			}
			return idx, found
		}
	}
	return -1, false
}

@(private = "file")
workflow_output_slot_index :: proc(outputs: []Ast_Struct_Field, name: string) -> (int, bool) {
	for out, i in outputs {
		if out.name == name do return i, true
	}
	return -1, false
}

@(private = "file")
validate_no_workflow_returns :: proc(p: ^Parser, stage: string, stmts: []^Ast_Node) {
	for stmt in stmts {
		if stmt == nil do continue
		if ret, ok := stmt.derived.(^Ast_Return); ok {
			report_error_at(p, ret.span, fmt.aprintf("'%s' workflow stages cannot use return", stage))
		}
		#partial switch d in stmt.derived {
		case ^Ast_If:
			validate_no_workflow_returns(p, stage, d.then_body[:])
			for clause in d.elseif_clauses do validate_no_workflow_returns(p, stage, clause.body[:])
			validate_no_workflow_returns(p, stage, d.else_body[:])
		case ^Ast_For:
			validate_no_workflow_returns(p, stage, d.body[:])
		case ^Ast_While:
			validate_no_workflow_returns(p, stage, d.body[:])
		}
	}
}

@(private = "file")
validate_no_nested_workflow_output_effects :: proc(p: ^Parser, stage: string, stmt: ^Ast_Node, outputs: []Ast_Struct_Field) {
	#partial switch d in stmt.derived {
	case ^Ast_If:
		validate_no_nested_workflow_output_effects_in_body(p, stage, d.then_body[:], outputs)
		for clause in d.elseif_clauses do validate_no_nested_workflow_output_effects_in_body(p, stage, clause.body[:], outputs)
		validate_no_nested_workflow_output_effects_in_body(p, stage, d.else_body[:], outputs)
	case ^Ast_For:
		validate_no_nested_workflow_output_effects_in_body(p, stage, d.body[:], outputs)
	case ^Ast_While:
		validate_no_nested_workflow_output_effects_in_body(p, stage, d.body[:], outputs)
	}
}

@(private = "file")
validate_no_nested_workflow_output_effects_in_body :: proc(p: ^Parser, stage: string, stmts: []^Ast_Node, outputs: []Ast_Struct_Field) {
	for stmt in stmts {
		if stmt == nil do continue
		if assign, ok := stmt.derived.(^Ast_Assign); ok {
			if _, slot_ok := workflow_output_assignment_slot(p, assign, outputs); slot_ok {
				report_error_at(p, assign.span, "out slot assignments inside control flow are not supported yet")
			}
		}
		if ret, ok := stmt.derived.(^Ast_Return); ok {
			report_error_at(p, ret.span, fmt.aprintf("'%s' workflow stage return sugar is only supported at top level", stage))
		}
		validate_no_nested_workflow_output_effects(p, stage, stmt, outputs)
	}
}

@(private = "file")
make_workflow_output_return :: proc(
	fn_name: string,
	outputs: []Ast_Struct_Field,
	values: []^Ast_Node,
	span: Source_Span,
) -> ^Ast_Node {
	lit_fields := make([]Ast_Struct_Literal_Field, len(outputs))
	for out, i in outputs {
		lit_fields[i] = Ast_Struct_Literal_Field{
			name  = out.name,
			value = values[i],
			span  = out.span,
		}
	}

	sl := new(Ast_Struct_Literal)
	sl^ = Ast_Struct_Literal{
		type_name = fmt.aprintf("_%s_Output", fn_name),
		fields    = lit_fields,
		span      = span,
	}
	ret := new(Ast_Return)
	ret^ = Ast_Return{
		value = make_node(.Struct_Literal, sl, span),
		span  = span,
	}
	return make_node(.Return, ret, span)
}

@(private = "file")
parse_params :: proc(p: ^Parser) -> []Ast_Param {
	params := make([dynamic]Ast_Param)
	if check(p, .Rparen) do return params[:]

	for {
		param_attrs := parse_attributes(p)
		name_tok, ok := expect(p, .Identifier)
		if !ok do break
		expect(p, .Colon)
		type_expr := parse_type_expr(p)
		append(&params, Ast_Param{
			name       = name_tok.text,
			type       = type_expr,
			attributes = param_attrs,
			span       = name_tok.span,
		})
		if !match(p, .Comma) do break
	}
	return params[:]
}

@(private = "file")
parse_struct :: proc(p: ^Parser, attrs: []Ast_Attribute = nil) -> ^Ast_Struct {
	span_start := current(p).span
	expect(p, .KW_Struct)

	name_tok, ok := expect(p, .Identifier)
	if !ok {
		synchronize(p)
		return nil
	}

	append(&p.block_starts, Block_Start{kind = "struct", span = span_start})
	fields := make([dynamic]Ast_Struct_Field)
	for !check(p, .KW_End) && !is_at_end_p(p) && !is_toplevel_keyword(p) {
		start_pos := p.pos

		// If we see @, peek ahead: if the token after the attribute is a top-level keyword,
		// this attribute belongs to the next declaration, not a struct field
		if check(p, .At) {
			saved_pos := p.pos
			parse_attributes(p) // consume speculatively
			if is_toplevel_keyword(p) {
				p.pos = saved_pos // rewind — let parse_module handle this
				break
			}
			p.pos = saved_pos // rewind and parse normally
		}

		field_attrs := parse_attributes(p)

		field_name, field_ok := expect(p, .Identifier)
		if !field_ok {
			synchronize(p)
			if p.pos == start_pos do advance_p(p) // prevent infinite loop on stuck keywords
			continue
		}
		expect(p, .Colon)
		field_type := parse_type_expr(p)
		append(&fields, Ast_Struct_Field{
			name       = field_name.text,
			type       = field_type,
			attributes = field_attrs,
			span       = field_name.span,
		})
	}

	if !match(p, .KW_End) {
		bs := pop(&p.block_starts) if len(p.block_starts) > 0 else Block_Start{}
		error(p, fmt.aprintf("missing 'end' to close 'struct' started at line %d", bs.span.line_start))
	} else {
		if len(p.block_starts) > 0 do pop(&p.block_starts)
	}

	s := new(Ast_Struct)
	s^ = Ast_Struct{
		name       = name_tok.text,
		fields     = fields[:],
		attributes = attrs,
		span       = span_start,
	}
	return s
}

@(private = "file")
parse_binding :: proc(p: ^Parser, kind: Binding_Kind, attrs: []Ast_Attribute) -> ^Ast_Binding {
	span_start := current(p).span
	advance_p(p) // uniform or buffer

	// Check for optional qualifier (e.g. "buffer storage")
	qualifier := ""
	if kind == .Buffer && check(p, .KW_Storage) {
		qualifier = "storage"
		advance_p(p)
	}

	name_tok, ok := expect(p, .Identifier)
	if !ok {
		synchronize(p)
		return nil
	}

	expect(p, .Colon)
	type_expr := parse_type_expr(p)

	b := new(Ast_Binding)
	b^ = Ast_Binding{
		kind       = kind,
		name       = name_tok.text,
		type_expr  = type_expr,
		attributes = attrs,
		qualifier  = qualifier,
		span       = span_start,
	}
	return b
}

@(private = "file")
parse_group_block :: proc(p: ^Parser, attrs: []Ast_Attribute) -> []^Ast_Binding {
	span_start := current(p).span
	advance_p(p) // contextual "group"

	if len(attrs) > 0 {
		error(p, "attributes are not supported on group blocks")
		p.panic_mode = false
	}

	group_name, ok := expect(p, .Identifier)
	if !ok {
		synchronize(p)
		return nil
	}

	expect(p, .Eq)
	group_tok, group_ok := expect(p, .Integer)
	group_num := 0
	if group_ok {
		if parsed, parse_ok := strconv.parse_i64_of_base(group_tok.text, 10); parse_ok {
			group_num = int(parsed)
		} else {
			error(p, fmt.aprintf("invalid group index '%s'", group_tok.text))
			p.panic_mode = false
		}
	}

	append(&p.block_starts, Block_Start{kind = "group", span = span_start})
	bindings := make([dynamic]^Ast_Binding)
	for !check(p, .KW_End) && !is_at_end_p(p) {
		start_pos := p.pos
		binding_attrs := parse_attributes(p)
		binding_attrs = append_group_attribute(p, binding_attrs, group_num, group_name.span)

		if check(p, .KW_Uniform) {
			b := parse_binding(p, .Uniform, binding_attrs)
			if b != nil do append(&bindings, b)
		} else if check(p, .KW_Buffer) {
			b := parse_binding(p, .Buffer, binding_attrs)
			if b != nil do append(&bindings, b)
		} else if check_identifier_text(p, "texture") {
			b := parse_binding(p, .Uniform, binding_attrs)
			if b != nil do append(&bindings, b)
		} else {
			error(p, fmt.aprintf("expected binding declaration in group '%s'", group_name.text))
			synchronize(p)
		}

		if p.pos == start_pos do advance_p(p)
	}

	if !match(p, .KW_End) {
		bs := pop(&p.block_starts) if len(p.block_starts) > 0 else Block_Start{}
		error(p, fmt.aprintf("missing 'end' to close 'group' started at line %d", bs.span.line_start))
	} else {
		if len(p.block_starts) > 0 do pop(&p.block_starts)
	}

	return bindings[:]
}

@(private = "file")
append_group_attribute :: proc(p: ^Parser, attrs: []Ast_Attribute, group_num: int, span: Source_Span) -> []Ast_Attribute {
	if has_attribute(attrs, "group") {
		error(p, "binding inside a group block cannot also declare @group")
		p.panic_mode = false
	}

	result := make([dynamic]Ast_Attribute)
	for attr in attrs {
		append(&result, attr)
	}

	args := make([]string, 1)
	args[0] = fmt.aprintf("%d", group_num)
	append(&result, Ast_Attribute{
		name = "group",
		args = args,
		span = span,
	})
	return result[:]
}

@(private = "file")
parse_const :: proc(p: ^Parser, attrs: []Ast_Attribute = {}) -> ^Ast_Const {
	span_start := current(p).span
	advance_p(p) // const

	name_tok, ok := expect(p, .Identifier)
	if !ok {
		synchronize(p)
		return nil
	}

	type_ptr: ^Type_Expr
	if match(p, .Colon) {
		te := parse_type_expr(p)
		type_ptr = new(Type_Expr)
		type_ptr^ = te
	}

	expect(p, .Eq)
	value := parse_expr(p)

	c := new(Ast_Const)
	c^ = Ast_Const{
		name       = name_tok.text,
		type       = type_ptr,
		value      = value,
		attributes = attrs,
		span       = span_start,
	}
	return c
}

@(private = "file")
parse_shared :: proc(p: ^Parser) -> ^Ast_Shared {
	span_start := current(p).span
	advance_p(p) // shared

	name_tok, ok := expect(p, .Identifier)
	if !ok {
		synchronize(p)
		return nil
	}

	expect(p, .Colon)
	type_expr := parse_type_expr(p)

	s := new(Ast_Shared)
	s^ = Ast_Shared{
		name      = name_tok.text,
		type_expr = type_expr,
		span      = span_start,
	}
	return s
}

// -- Type expressions --

@(private = "file")
parse_type_expr :: proc(p: ^Parser) -> Type_Expr {
	// Check for array type: []T or [N]T
	if check(p, .Lbracket) {
		return parse_array_type(p)
	}

	name_tok, ok := expect(p, .Identifier)
	if !ok {
		return Type_Named{name = "error", span = current(p).span}
	}
	return Type_Named{name = name_tok.text, span = name_tok.span}
}

@(private = "file")
parse_tuple_type :: proc(p: ^Parser) -> Type_Expr {
	span_start := current(p).span
	expect(p, .Lparen)

	fields := make([dynamic]Ast_Struct_Field)
	for !check(p, .Rparen) && !is_at_end_p(p) {
		field_attrs := parse_attributes(p)
		field_name, ok := expect(p, .Identifier)
		if !ok do break
		expect(p, .Colon)
		field_type := parse_type_expr(p)
		append(&fields, Ast_Struct_Field{
			name       = field_name.text,
			type       = field_type,
			attributes = field_attrs,
			span       = field_name.span,
		})
		if !match(p, .Comma) do break
	}
	expect(p, .Rparen)

	if len(fields) == 0 {
		error(p, "empty tuple return type")
	}

	return Type_Tuple{
		fields = fields[:],
		span   = span_start,
	}
}

@(private = "file")
parse_array_type :: proc(p: ^Parser) -> Type_Expr {
	span_start := current(p).span
	expect(p, .Lbracket)

	size: ^Ast_Node
	if !check(p, .Rbracket) {
		size = parse_expr(p)
	}

	expect(p, .Rbracket)

	elem_te := parse_type_expr(p)
	elem_ptr := new(Type_Expr)
	elem_ptr^ = elem_te

	return Type_Array{
		elem = elem_ptr,
		size = size,
		span = span_start,
	}
}

// -- Statements --

@(private = "file")
parse_body :: proc(p: ^Parser) -> [dynamic]^Ast_Node {
	stmts := make([dynamic]^Ast_Node)
	for !check(p, .KW_End) && !check(p, .KW_Else) && !check(p, .KW_Elseif) && !is_at_end_p(p) {
		start_pos := p.pos
		stmt := parse_statement(p)
		if stmt != nil {
			append(&stmts, stmt)
		}
		if p.pos == start_pos do advance_p(p) // prevent infinite loop
	}
	return stmts
}

@(private = "file")
parse_statement :: proc(p: ^Parser) -> ^Ast_Node {
	#partial switch current(p).kind {
	case .KW_Let, .KW_Var:
		return parse_let(p)
	case .KW_Return:
		return parse_return(p)
	case .KW_If:
		return parse_if(p)
	case .KW_For:
		return parse_for(p)
	case .KW_While:
		return parse_while(p)
	case .KW_Discard:
		span := current(p).span
		advance_p(p)
		d := new(Ast_Discard)
		d.span = span
		return make_node(.Discard, d, span)
	case .KW_Break:
		span := current(p).span
		advance_p(p)
		b := new(Ast_Break)
		b.span = span
		return make_node(.Break, b, span)
	case .KW_Continue:
		span := current(p).span
		advance_p(p)
		c := new(Ast_Continue)
		c.span = span
		return make_node(.Continue, c, span)
	case:
		return parse_assign_or_expr(p)
	}
}

@(private = "file")
parse_let :: proc(p: ^Parser) -> ^Ast_Node {
	span_start := current(p).span
	is_var := current(p).kind == .KW_Var
	advance_p(p) // let or var

	name_tok, ok := expect(p, .Identifier)
	if !ok {
		synchronize(p)
		return nil
	}

	type_ptr: ^Type_Expr
	if match(p, .Colon) {
		te := parse_type_expr(p)
		type_ptr = new(Type_Expr)
		type_ptr^ = te
	}

	expect(p, .Eq)
	value := parse_expr(p)

	let := new(Ast_Let)
	let^ = Ast_Let{
		name      = name_tok.text,
		type_expr = type_ptr,
		value     = value,
		span      = span_start,
		mutable   = is_var,
	}
	return make_node(.Let, let, span_start)
}

@(private = "file")
parse_return :: proc(p: ^Parser) -> ^Ast_Node {
	span_start := current(p).span
	advance_p(p) // return

	value: ^Ast_Node
	// Check if there's an expression following (not a block-ender)
	if !check(p, .KW_End) && !check(p, .KW_Else) && !check(p, .KW_Elseif) && !is_at_end_p(p) {
		value = parse_expr(p)
	}

	ret := new(Ast_Return)
	ret^ = Ast_Return{
		value = value,
		span  = span_start,
	}
	return make_node(.Return, ret, span_start)
}

@(private = "file")
parse_if :: proc(p: ^Parser) -> ^Ast_Node {
	span_start := current(p).span
	advance_p(p) // if

	append(&p.block_starts, Block_Start{kind = "if", span = span_start})
	condition := parse_expr(p)
	expect(p, .KW_Then)
	then_body := parse_body(p)

	elseifs := make([dynamic]Ast_Elseif)
	for check(p, .KW_Elseif) {
		advance_p(p) // elseif
		ei_cond := parse_expr(p)
		expect(p, .KW_Then)
		ei_body := parse_body(p)
		append(&elseifs, Ast_Elseif{
			condition = ei_cond,
			body      = ei_body,
		})
	}

	else_body: [dynamic]^Ast_Node
	if match(p, .KW_Else) {
		else_body = parse_body(p)
	}

	if !match(p, .KW_End) {
		bs := pop(&p.block_starts) if len(p.block_starts) > 0 else Block_Start{}
		error(p, fmt.aprintf("missing 'end' to close 'if' started at line %d", bs.span.line_start))
	} else {
		if len(p.block_starts) > 0 do pop(&p.block_starts)
	}

	if_node := new(Ast_If)
	if_node^ = Ast_If{
		condition      = condition,
		then_body      = then_body,
		elseif_clauses = elseifs[:],
		else_body      = else_body,
		span           = span_start,
	}
	return make_node(.If, if_node, span_start)
}

@(private = "file")
parse_for :: proc(p: ^Parser) -> ^Ast_Node {
	span_start := current(p).span
	advance_p(p) // for

	var_tok, ok := expect(p, .Identifier)
	if !ok {
		synchronize(p)
		return nil
	}

	append(&p.block_starts, Block_Start{kind = "for", span = span_start})
	expect(p, .Eq)
	start_expr := parse_expr(p)
	expect(p, .Comma)
	stop_expr := parse_expr(p)

	step: ^Ast_Node
	if match(p, .Comma) {
		step = parse_expr(p)
	}

	expect(p, .KW_Do)
	body := parse_body(p)
	if !match(p, .KW_End) {
		bs := pop(&p.block_starts) if len(p.block_starts) > 0 else Block_Start{}
		error(p, fmt.aprintf("missing 'end' to close 'for' started at line %d", bs.span.line_start))
	} else {
		if len(p.block_starts) > 0 do pop(&p.block_starts)
	}

	for_node := new(Ast_For)
	for_node^ = Ast_For{
		var_name = var_tok.text,
		start    = start_expr,
		stop     = stop_expr,
		step     = step,
		body     = body,
		span     = span_start,
	}
	return make_node(.For, for_node, span_start)
}

@(private = "file")
parse_while :: proc(p: ^Parser) -> ^Ast_Node {
	span_start := current(p).span
	advance_p(p) // while

	append(&p.block_starts, Block_Start{kind = "while", span = span_start})
	condition := parse_expr(p)
	expect(p, .KW_Do)
	body := parse_body(p)
	if !match(p, .KW_End) {
		bs := pop(&p.block_starts) if len(p.block_starts) > 0 else Block_Start{}
		error(p, fmt.aprintf("missing 'end' to close 'while' started at line %d", bs.span.line_start))
	} else {
		if len(p.block_starts) > 0 do pop(&p.block_starts)
	}

	while_node := new(Ast_While)
	while_node^ = Ast_While{
		condition = condition,
		body      = body,
		span      = span_start,
	}
	return make_node(.While, while_node, span_start)
}

@(private = "file")
parse_assign_or_expr :: proc(p: ^Parser) -> ^Ast_Node {
	expr := parse_expr(p)
	if expr == nil do return nil

	// Check if this is an assignment
	if check(p, .Eq) {
		advance_p(p) // =
		value := parse_expr(p)
		assign := new(Ast_Assign)
		assign^ = Ast_Assign{
			target = expr,
			value  = value,
			span   = expr.span,
		}
		return make_node(.Assign, assign, expr.span)
	}

	return expr
}

// -- Expressions (precedence climbing via recursive descent) --

@(private = "file")
parse_expr :: proc(p: ^Parser) -> ^Ast_Node {
	return parse_or(p)
}

@(private = "file")
parse_or :: proc(p: ^Parser) -> ^Ast_Node {
	left := parse_and(p)
	for check(p, .KW_Or) {
		advance_p(p)
		right := parse_and(p)
		bin := new(Ast_Binary)
		bin^ = Ast_Binary{op = .Or, left = left, right = right, span = left.span}
		left = make_node(.Binary, bin, left.span)
	}
	return left
}

@(private = "file")
parse_and :: proc(p: ^Parser) -> ^Ast_Node {
	left := parse_comparison(p)
	for check(p, .KW_And) {
		advance_p(p)
		right := parse_comparison(p)
		bin := new(Ast_Binary)
		bin^ = Ast_Binary{op = .And, left = left, right = right, span = left.span}
		left = make_node(.Binary, bin, left.span)
	}
	return left
}

@(private = "file")
parse_comparison :: proc(p: ^Parser) -> ^Ast_Node {
	left := parse_addition(p)

	for {
		op: Binary_Op
		#partial switch current(p).kind {
		case .Eq_Eq:  op = .Eq
		case .Not_Eq: op = .Neq
		case .Lt:     op = .Lt
		case .Gt:     op = .Gt
		case .Lt_Eq:  op = .Lte
		case .Gt_Eq:  op = .Gte
		case:         return left
		}
		advance_p(p)
		right := parse_addition(p)
		bin := new(Ast_Binary)
		bin^ = Ast_Binary{op = op, left = left, right = right, span = left.span}
		left = make_node(.Binary, bin, left.span)
	}
}

@(private = "file")
parse_addition :: proc(p: ^Parser) -> ^Ast_Node {
	left := parse_multiplication(p)

	for {
		op: Binary_Op
		#partial switch current(p).kind {
		case .Plus:  op = .Add
		case .Minus: op = .Sub
		case:        return left
		}
		advance_p(p)
		right := parse_multiplication(p)
		bin := new(Ast_Binary)
		bin^ = Ast_Binary{op = op, left = left, right = right, span = left.span}
		left = make_node(.Binary, bin, left.span)
	}
}

@(private = "file")
parse_multiplication :: proc(p: ^Parser) -> ^Ast_Node {
	left := parse_unary(p)

	for {
		op: Binary_Op
		#partial switch current(p).kind {
		case .Star:    op = .Mul
		case .Slash:   op = .Div
		case .Percent: op = .Mod
		case:          return left
		}
		advance_p(p)
		right := parse_unary(p)
		bin := new(Ast_Binary)
		bin^ = Ast_Binary{op = op, left = left, right = right, span = left.span}
		left = make_node(.Binary, bin, left.span)
	}
}

@(private = "file")
parse_unary :: proc(p: ^Parser) -> ^Ast_Node {
	if check(p, .Minus) {
		span := current(p).span
		advance_p(p)
		operand := parse_unary(p)
		un := new(Ast_Unary)
		un^ = Ast_Unary{op = .Neg, operand = operand, span = span}
		return make_node(.Unary, un, span)
	}
	if check(p, .KW_Not) {
		span := current(p).span
		advance_p(p)
		operand := parse_unary(p)
		un := new(Ast_Unary)
		un^ = Ast_Unary{op = .Not, operand = operand, span = span}
		return make_node(.Unary, un, span)
	}
	return parse_postfix(p)
}

@(private = "file")
parse_postfix :: proc(p: ^Parser) -> ^Ast_Node {
	node := parse_primary(p)

	for {
		#partial switch current(p).kind {
		case .Dot:
			advance_p(p)
			field_tok, ok := expect(p, .Identifier)
			if !ok do return node
			fa := new(Ast_Field_Access)
			fa^ = Ast_Field_Access{
				object = node,
				field  = field_tok.text,
				span   = node.span,
			}
			node = make_node(.Field_Access, fa, node.span)

		case .Lbracket:
			advance_p(p)
			index_expr := parse_expr(p)
			expect(p, .Rbracket)
			idx := new(Ast_Index)
			idx^ = Ast_Index{
				object = node,
				index  = index_expr,
				span   = node.span,
			}
			node = make_node(.Index, idx, node.span)

		case .Lparen:
			advance_p(p)
			args := make([dynamic]^Ast_Node)
			if !check(p, .Rparen) {
				for {
					arg := parse_expr(p)
					append(&args, arg)
					if !match(p, .Comma) do break
				}
			}
			expect(p, .Rparen)
			call := new(Ast_Call)
			call^ = Ast_Call{
				callee = node,
				args   = args[:],
				span   = node.span,
			}
			node = make_node(.Call, call, node.span)

		case .Lbrace:
			// Struct literal: only if previous node is an identifier
			if node.kind == .Ident {
				ident := node.derived.(^Ast_Ident)
				sl := parse_struct_literal_body(p, ident.name, node.span)
				node = sl
			} else {
				return node
			}

		case:
			return node
		}
	}
}

@(private = "file")
parse_primary :: proc(p: ^Parser) -> ^Ast_Node {
	tok := current(p)

	#partial switch tok.kind {
	case .Integer:
		advance_p(p)
		val, ok := strconv.parse_i64_of_base(tok.text, 10)
		lit := new(Ast_Literal)
		lit^ = Ast_Literal{value = val, span = tok.span}
		return make_node(.Literal, lit, tok.span)

	case .Float:
		advance_p(p)
		val, ok := strconv.parse_f64(tok.text)
		lit := new(Ast_Literal)
		lit^ = Ast_Literal{value = val, span = tok.span}
		return make_node(.Literal, lit, tok.span)

	case .KW_True:
		advance_p(p)
		lit := new(Ast_Literal)
		lit^ = Ast_Literal{value = true, span = tok.span}
		return make_node(.Literal, lit, tok.span)

	case .KW_False:
		advance_p(p)
		lit := new(Ast_Literal)
		lit^ = Ast_Literal{value = false, span = tok.span}
		return make_node(.Literal, lit, tok.span)

	case .String:
		advance_p(p)
		// Strip quotes
		text := tok.text[1:len(tok.text)-1]
		lit := new(Ast_Literal)
		lit^ = Ast_Literal{value = text, span = tok.span}
		return make_node(.Literal, lit, tok.span)

	case .Identifier:
		advance_p(p)
		ident := new(Ast_Ident)
		ident^ = Ast_Ident{name = tok.text, span = tok.span}
		return make_node(.Ident, ident, tok.span)

	case .Lparen:
		advance_p(p)
		expr := parse_expr(p)
		expect(p, .Rparen)
		return expr

	case:
		error(p, fmt.aprintf("unexpected token %s in expression", token_kind_to_string(tok.kind)))
		advance_p(p)
		// Return a dummy node
		lit := new(Ast_Literal)
		lit^ = Ast_Literal{value = i64(0), span = tok.span}
		return make_node(.Literal, lit, tok.span)
	}
}

@(private = "file")
parse_struct_literal_body :: proc(p: ^Parser, type_name: string, span: Source_Span) -> ^Ast_Node {
	expect(p, .Lbrace)
	fields := make([dynamic]Ast_Struct_Literal_Field)

	for !check(p, .Rbrace) && !is_at_end_p(p) {
		start_pos := p.pos
		field_name_tok, ok := expect(p, .Identifier)
		if !ok {
			if p.pos == start_pos do advance_p(p) // prevent infinite loop
			break
		}
		expect(p, .Eq)
		value := parse_expr(p)
		append(&fields, Ast_Struct_Literal_Field{
			name  = field_name_tok.text,
			value = value,
			span  = field_name_tok.span,
		})
		match(p, .Comma) // optional trailing comma
	}

	expect(p, .Rbrace)

	sl := new(Ast_Struct_Literal)
	sl^ = Ast_Struct_Literal{
		type_name = type_name,
		fields    = fields[:],
		span      = span,
	}
	return make_node(.Struct_Literal, sl, span)
}