Harbor

branch main
showing the latest snapshot on main
emit_dxbc.odin 91.2 KB · Plain text
shader/emit_dxbc.odin 0644 Raw
package shader

import "core:fmt"
import "core:math"
import "core:strings"

// DXBC SM5.0 binary backend — emits from IR_Module

DXBC_Reg :: struct {
	kind:       u32,  // DXBC_OPERAND_* type
	index:      u32,  // register number
	index2:     u32,  // second index (for cb[N][M])
	components: int,  // 1-4 components this value occupies
	mask:       u32,  // write mask (bitmask: 0x1=x, 0x2=y, etc.)
	swizzle:    u32,  // explicit source swizzle (0 = derive from mask)
	has_swizzle: bool, // if true, use swizzle field directly for source encoding
	negate:     bool, // source modifier: negate
	abs:        bool, // source modifier: absolute value
}

DXBC_Binding_Info :: struct {
	kind:    IR_Binding_Kind,
	reg_num: u32,
	space:   u32,
	size:    u32, // cbuffer size in vec4 units
	fields:  []DXBC_CB_Field,
}

DXBC_CB_Field :: struct {
	name:       string,
	vec4_index: u32, // which vec4 slot
	comp_offset: u32, // component offset within vec4 (0-3)
	comp_count: int,  // number of components
	is_matrix:  bool,
	mat_cols:   int,
	mat_rows:   int,
}

DXBC_Indexable_Info :: struct {
	xreg:       u32, // indexable temp register number (x0, x1, ...)
	elem_count: u32, // number of elements
	elem_comps: int, // components per element
}

DXBC_Sig_Element :: struct {
	semantic_name:  string,
	semantic_index: u32,
	system_value:   u32,
	component_type: u32, // 1=uint, 2=int, 3=float
	register_num:   u32,
	mask:           u8,
	rw_mask:        u8,
}

DXBC_Builder :: struct {
	module:     ^IR_Module,
	current_fn: ^IR_Function,

	// Bytecode sections
	shdr_decls: [dynamic]u32,
	shdr_body:  [dynamic]u32,

	// Register allocation
	next_temp:          u32,
	var_map:            map[IR_Var_Id]DXBC_Reg,
	binding_map:        map[string]DXBC_Binding_Info,
	input_map:          map[string]DXBC_Reg,
	output_map:         map[string]DXBC_Reg,
	builtin_input_map:  map[string]DXBC_Reg,
	builtin_output_map: map[string]DXBC_Reg,

	// Signatures
	input_sigs:  [dynamic]DXBC_Sig_Element,
	output_sigs: [dynamic]DXBC_Sig_Element,

	// Resource tracking
	num_cb:       u32,
	num_textures: u32,
	num_samplers: u32,
	num_tgsm:     u32,
	shared_var_map:     map[string]u32,
	// Indexable temp tracking (for dynamic array indexing)
	num_indexable_temps: u32,
	indexable_temp_map:  map[IR_Var_Id]DXBC_Indexable_Info,

	// Output register counter
	next_output_reg: u32,

	// Stats
	stat_inst_count: u32,

	diagnostics: [dynamic]Diagnostic,
}

emit_dxbc :: proc(module: ^IR_Module, allocator := context.allocator) -> ([]u8, []Diagnostic) {
	b := DXBC_Builder{
		module       = module,
		shdr_decls   = make([dynamic]u32, allocator),
		shdr_body    = make([dynamic]u32, allocator),
		var_map      = make(map[IR_Var_Id]DXBC_Reg, allocator = allocator),
		binding_map  = make(map[string]DXBC_Binding_Info, allocator = allocator),
		input_map    = make(map[string]DXBC_Reg, allocator = allocator),
		output_map   = make(map[string]DXBC_Reg, allocator = allocator),
		builtin_input_map  = make(map[string]DXBC_Reg, allocator = allocator),
		builtin_output_map = make(map[string]DXBC_Reg, allocator = allocator),
		input_sigs      = make([dynamic]DXBC_Sig_Element, allocator),
		output_sigs     = make([dynamic]DXBC_Sig_Element, allocator),
		shared_var_map      = make(map[string]u32, allocator = allocator),
		indexable_temp_map  = make(map[IR_Var_Id]DXBC_Indexable_Info, allocator = allocator),
		diagnostics         = make([dynamic]Diagnostic, allocator),
	}

	// Find the entry point function
	entry_fn: ^IR_Function
	for &fn in module.functions {
		if fn.is_entry {
			entry_fn = &fn
			break
		}
	}

	if entry_fn == nil {
		append(&b.diagnostics, Diagnostic{level = .Error, message = "no entry point found for DXBC emission"})
		return nil, b.diagnostics[:]
	}

	// Setup bindings
	dxbc_setup_bindings(&b)

	// Emit the entry point
	b.current_fn = entry_fn
	dxbc_emit_entry_point(&b, entry_fn)
	b.current_fn = nil

	// Patch dcl_temps at the start of declarations
	if b.next_temp > 0 {
		temps_inst := [2]u32{dxbc_opcode_token(DXBC_OP_DCL_TEMPS, 2), b.next_temp}
		// Prepend to shdr_decls
		old_decls := b.shdr_decls[:]
		new_decls := make([dynamic]u32, allocator)
		append(&new_decls, ..temps_inst[:])
		append(&new_decls, ..old_decls)
		b.shdr_decls = new_decls
	}

	// Assemble container
	result := dxbc_assemble(&b, entry_fn, allocator)
	return result, b.diagnostics[:]
}

// ---- Register allocation ----

dxbc_alloc_temp :: proc(b: ^DXBC_Builder, components: int = 4) -> DXBC_Reg {
	idx := b.next_temp
	b.next_temp += 1
	return DXBC_Reg{
		kind       = DXBC_OPERAND_TEMP,
		index      = idx,
		components = components,
		mask       = dxbc_mask_from_count(components),
	}
}

dxbc_type_components :: proc(t: ^Resolved_Type) -> int {
	if t == nil do return 1
	#partial switch v in t^ {
	case Type_Scalar: return 1
	case Type_Vector: return v.size
	case Type_Matrix: return 4 // matrices use multiple vec4 registers
	case Type_Void:   return 0
	case: return 4
	}
}

// ---- Binding setup ----

dxbc_setup_bindings :: proc(b: ^DXBC_Builder) {
	for &binding in b.module.bindings {
		switch binding.kind {
		case .Uniform, .Push_Constant:
			info := DXBC_Binding_Info{
				kind    = binding.kind,
				reg_num = b.num_cb,
				space   = u32(binding.group),
			}
			// Compute cbuffer layout
			if binding.struct_ref != nil {
				info.fields, info.size = dxbc_compute_cb_layout(binding.struct_ref)
			} else {
				info.size = 1
			}
			b.binding_map[binding.name] = info
			b.num_cb += 1

		case .Texture:
			b.binding_map[binding.name] = DXBC_Binding_Info{
				kind    = .Texture,
				reg_num = b.num_textures,
			}
			b.num_textures += 1

		case .Sampler:
			b.binding_map[binding.name] = DXBC_Binding_Info{
				kind    = .Sampler,
				reg_num = b.num_samplers,
			}
			b.num_samplers += 1

		case .Buffer:
			// TODO: structured buffers
		}
	}
}

dxbc_compute_cb_layout :: proc(s: ^Type_Struct_Resolved) -> ([]DXBC_CB_Field, u32) {
	fields := make([dynamic]DXBC_CB_Field)
	current_vec4: u32 = 0
	current_comp: u32 = 0

	for f in s.fields {
		if f.type == nil do continue
		#partial switch v in f.type^ {
		case Type_Scalar:
			// Check if we need to advance to next vec4
			if current_comp >= 4 {
				current_vec4 += 1
				current_comp = 0
			}
			append(&fields, DXBC_CB_Field{
				name        = f.name,
				vec4_index  = current_vec4,
				comp_offset = current_comp,
				comp_count  = 1,
			})
			current_comp += 1

		case Type_Vector:
			// Vectors don't cross vec4 boundaries
			if current_comp + u32(v.size) > 4 {
				current_vec4 += 1
				current_comp = 0
			}
			append(&fields, DXBC_CB_Field{
				name        = f.name,
				vec4_index  = current_vec4,
				comp_offset = current_comp,
				comp_count  = v.size,
			})
			current_comp += u32(v.size)

		case Type_Matrix:
			// Each column starts on a new vec4
			if current_comp > 0 {
				current_vec4 += 1
				current_comp = 0
			}
			append(&fields, DXBC_CB_Field{
				name        = f.name,
				vec4_index  = current_vec4,
				comp_offset = 0,
				comp_count  = v.rows,
				is_matrix   = true,
				mat_cols    = v.cols,
				mat_rows    = v.rows,
			})
			current_vec4 += u32(v.cols)
			current_comp = 0

		case Type_Array_Resolved:
			if current_comp > 0 {
				current_vec4 += 1
				current_comp = 0
			}
			elem_size := u32(1) // TODO: proper element sizing
			append(&fields, DXBC_CB_Field{
				name        = f.name,
				vec4_index  = current_vec4,
				comp_offset = 0,
				comp_count  = 4,
			})
			current_vec4 += u32(v.size) * elem_size
			current_comp = 0

		case: // struct, sampler, void — skip
		}
	}

	total := current_vec4
	if current_comp > 0 do total += 1
	return fields[:], total
}

// ---- Entry point emission ----

dxbc_emit_entry_point :: proc(b: ^DXBC_Builder, fn: ^IR_Function) {
	// Emit global flags (bit 11 = refactoring allowed)
	append(&b.shdr_decls, dxbc_opcode_token_ex(DXBC_OP_DCL_GLOBAL_FLAGS, 1, 1))

	// Emit binding declarations
	for name, info in b.binding_map {
		switch info.kind {
		case .Uniform, .Push_Constant:
			dxbc_emit_dcl_cb(b, info.reg_num, info.size)
		case .Texture:
			dxbc_emit_dcl_resource(b, info.reg_num)
		case .Sampler:
			dxbc_emit_dcl_sampler(b, info.reg_num, false)
		case .Buffer:
			// TODO
		}
	}

	// Setup input registers and declarations
	for io, i in fn.inputs {
		comp := dxbc_type_components(io.type)
		mask := dxbc_mask_from_count(comp)

		if io.builtin != "" {
			reg := dxbc_alloc_input_builtin(b, io, fn.stage)
			b.builtin_input_map[io.builtin] = reg
		} else {
			reg := DXBC_Reg{
				kind       = DXBC_OPERAND_INPUT,
				index      = u32(io.location),
				components = comp,
				mask       = mask,
			}
			b.input_map[io.name] = reg
			dxbc_emit_dcl_input(b, fn.stage, u32(io.location), mask)
			append(&b.input_sigs, DXBC_Sig_Element{
				semantic_name  = fmt.aprintf("TEXCOORD"),
				semantic_index = u32(io.location),
				system_value   = DXBC_SV_UNDEFINED,
				component_type = 3, // float
				register_num   = u32(io.location),
				mask           = u8(mask),
				rw_mask        = u8(mask),
			})
		}
	}

	// Setup output registers and declarations
	for io, i in fn.outputs {
		comp := dxbc_type_components(io.type)
		mask := dxbc_mask_from_count(comp)

		if io.builtin != "" {
			reg := dxbc_alloc_output_builtin(b, io, fn.stage)
			b.builtin_output_map[io.builtin] = reg
		} else {
			out_reg := b.next_output_reg
			b.next_output_reg += 1
			reg := DXBC_Reg{
				kind       = DXBC_OPERAND_OUTPUT,
				index      = out_reg,
				components = comp,
				mask       = mask,
			}
			b.output_map[io.name] = reg
			dxbc_emit_dcl_output(b, fn.stage, out_reg, mask)
			if fn.stage == .Fragment {
				append(&b.output_sigs, DXBC_Sig_Element{
					semantic_name  = fmt.aprintf("SV_Target"),
					semantic_index = u32(io.location),
					system_value   = DXBC_SV_UNDEFINED,
					component_type = 3,
					register_num   = out_reg,
					mask           = u8(mask),
					rw_mask        = 0,
				})
			} else {
				append(&b.output_sigs, DXBC_Sig_Element{
					semantic_name  = fmt.aprintf("TEXCOORD"),
					semantic_index = u32(io.location),
					system_value   = DXBC_SV_UNDEFINED,
					component_type = 3,
					register_num   = out_reg,
					mask           = u8(mask),
					rw_mask        = 0,
				})
			}
		}
	}

	// Compute shader: emit thread group declaration and TGSM
	if fn.stage == .Compute {
		x := u32(fn.workgroup_size[0])
		y := u32(fn.workgroup_size[1])
		z := u32(fn.workgroup_size[2])
		if x == 0 do x = 1
		if y == 0 do y = 1
		if z == 0 do z = 1
		append(&b.shdr_decls, dxbc_opcode_token(DXBC_OP_DCL_THREAD_GROUP, 4))
		append(&b.shdr_decls, x)
		append(&b.shdr_decls, y)
		append(&b.shdr_decls, z)

		// Declare shared (TGSM) variables
		for sv in b.module.shared_vars {
			tgsm_idx := b.num_tgsm
			b.shared_var_map[sv.name] = tgsm_idx
			b.num_tgsm += 1
			stride, count := dxbc_tgsm_stride_count(sv.type)
			// dcl_tgsm_structured gN, stride_bytes, count
			// Instruction: opcode + TGSM operand (2 DWORDs) + stride + count = 5 DWORDs
			append(&b.shdr_decls, dxbc_opcode_token(DXBC_OP_DCL_TGSM_STRUCTURED, 5))
			append(&b.shdr_decls, dxbc_mask_operand(DXBC_OPERAND_THREAD_GROUP_SHARED_MEMORY, DXBC_INDEX_1D, DXBC_WRITEMASK_ALL) | dxbc_index_rep(DXBC_INDEX_IMM32))
			append(&b.shdr_decls, tgsm_idx)
			append(&b.shdr_decls, stride)
			append(&b.shdr_decls, count)
		}
	}

	// Emit body
	dxbc_emit_stmts(b, fn.body[:])

	// Emit RET
	append(&b.shdr_body, dxbc_opcode_token(DXBC_OP_RET, 1))
	b.stat_inst_count += 1
}

dxbc_alloc_input_builtin :: proc(b: ^DXBC_Builder, io: IR_IO_Var, stage: Shader_Stage) -> DXBC_Reg {
	comp := dxbc_type_components(io.type)
	mask := dxbc_mask_from_count(comp)
	sv: u32
	semantic_name: string
	reg_idx: u32

	switch io.builtin {
	case "vertex_id":
		sv = DXBC_SV_VERTEX_ID
		semantic_name = "SV_VertexID"
		reg_idx = u32(len(b.input_sigs))
		// dcl_input_sgv
		dxbc_emit_dcl_input_sgv(b, reg_idx, sv)
	case "instance_id":
		sv = DXBC_SV_INSTANCE_ID
		semantic_name = "SV_InstanceID"
		reg_idx = u32(len(b.input_sigs))
		dxbc_emit_dcl_input_sgv(b, reg_idx, sv)
	case "frag_coord":
		sv = DXBC_SV_POSITION
		semantic_name = "SV_Position"
		reg_idx = u32(len(b.input_sigs))
		dxbc_emit_dcl_input_ps_siv(b, reg_idx, mask, sv)
	case "front_facing":
		sv = DXBC_SV_IS_FRONT_FACE
		semantic_name = "SV_IsFrontFace"
		reg_idx = u32(len(b.input_sigs))
		dxbc_emit_dcl_input_sgv(b, reg_idx, sv)
	case "local_invocation_id":
		reg := DXBC_Reg{kind = DXBC_OPERAND_INPUT_THREAD_ID_IN_GROUP, index = 0, components = 3, mask = 0x7}
		append(&b.shdr_decls, dxbc_opcode_token(DXBC_OP_DCL_INPUT, 2))
		append(&b.shdr_decls, dxbc_mask_operand(DXBC_OPERAND_INPUT_THREAD_ID_IN_GROUP, DXBC_INDEX_0D, 0x7))
		return reg
	case "global_invocation_id":
		reg := DXBC_Reg{kind = DXBC_OPERAND_INPUT_THREAD_ID, index = 0, components = 3, mask = 0x7}
		append(&b.shdr_decls, dxbc_opcode_token(DXBC_OP_DCL_INPUT, 2))
		append(&b.shdr_decls, dxbc_mask_operand(DXBC_OPERAND_INPUT_THREAD_ID, DXBC_INDEX_0D, 0x7))
		return reg
	case "workgroup_id":
		reg := DXBC_Reg{kind = DXBC_OPERAND_INPUT_THREAD_GROUP_ID, index = 0, components = 3, mask = 0x7}
		append(&b.shdr_decls, dxbc_opcode_token(DXBC_OP_DCL_INPUT, 2))
		append(&b.shdr_decls, dxbc_mask_operand(DXBC_OPERAND_INPUT_THREAD_GROUP_ID, DXBC_INDEX_0D, 0x7))
		return reg
	case:
		sv = DXBC_SV_UNDEFINED
		semantic_name = io.builtin
		reg_idx = u32(len(b.input_sigs))
	}

	reg := DXBC_Reg{
		kind       = DXBC_OPERAND_INPUT,
		index      = reg_idx,
		components = comp,
		mask       = mask,
	}

	append(&b.input_sigs, DXBC_Sig_Element{
		semantic_name  = semantic_name,
		semantic_index = 0,
		system_value   = sv,
		component_type = 3,
		register_num   = reg_idx,
		mask           = u8(mask),
		rw_mask        = u8(mask),
	})

	return reg
}

dxbc_alloc_output_builtin :: proc(b: ^DXBC_Builder, io: IR_IO_Var, stage: Shader_Stage) -> DXBC_Reg {
	comp := dxbc_type_components(io.type)
	mask := dxbc_mask_from_count(comp)
	sv: u32
	semantic_name: string
	reg_idx: u32

	switch io.builtin {
	case "position":
		sv = DXBC_SV_POSITION
		semantic_name = "SV_Position"
		reg_idx = b.next_output_reg
		b.next_output_reg += 1
		dxbc_emit_dcl_output_siv(b, reg_idx, mask, sv)
	case "frag_depth":
		reg := DXBC_Reg{kind = DXBC_OPERAND_OUTPUT_DEPTH, index = 0, components = 1, mask = DXBC_WRITEMASK_X}
		append(&b.output_sigs, DXBC_Sig_Element{
			semantic_name  = "SV_Depth",
			semantic_index = 0,
			system_value   = DXBC_SV_DEPTH,
			component_type = 3,
			register_num   = 0xFFFF,
			mask           = 1,
			rw_mask        = 0,
		})
		return reg
	case:
		sv = DXBC_SV_UNDEFINED
		semantic_name = io.builtin
		reg_idx = b.next_output_reg
		b.next_output_reg += 1
		dxbc_emit_dcl_output(b, stage, reg_idx, mask)
	}

	reg := DXBC_Reg{
		kind       = DXBC_OPERAND_OUTPUT,
		index      = reg_idx,
		components = comp,
		mask       = mask,
	}

	append(&b.output_sigs, DXBC_Sig_Element{
		semantic_name  = semantic_name,
		semantic_index = 0,
		system_value   = sv,
		component_type = 3,
		register_num   = reg_idx,
		mask           = u8(mask),
		rw_mask        = 0,
	})

	return reg
}

// ---- Declaration emission ----

dxbc_emit_dcl_cb :: proc(b: ^DXBC_Builder, reg: u32, size: u32) {
	// dcl_constantbuffer CB0[size], immediateIndexed
	append(&b.shdr_decls, dxbc_opcode_token(DXBC_OP_DCL_CONSTANT_BUFFER, 4))
	append(&b.shdr_decls, dxbc_mask_operand(DXBC_OPERAND_CONSTANT_BUFFER, DXBC_INDEX_2D, DXBC_WRITEMASK_ALL) | dxbc_index_rep(DXBC_INDEX_IMM32, DXBC_INDEX_IMM32))
	append(&b.shdr_decls, reg)
	append(&b.shdr_decls, size)
}

dxbc_emit_dcl_resource :: proc(b: ^DXBC_Builder, reg: u32) {
	// dcl_resource_texture2d (float,float,float,float) t0
	append(&b.shdr_decls, dxbc_opcode_token_ex(DXBC_OP_DCL_RESOURCE, 4, (DXBC_RESOURCE_DIM_TEXTURE2D << 0)))
	append(&b.shdr_decls, dxbc_operand_token(DXBC_COMP_4, DXBC_COMP_MASK, DXBC_OPERAND_RESOURCE, DXBC_INDEX_1D) | (DXBC_WRITEMASK_ALL << 4) | dxbc_index_rep(DXBC_INDEX_IMM32))
	append(&b.shdr_decls, reg)
	// Return type token: float for all 4 components
	ret_type := u32(DXBC_RETURN_TYPE_FLOAT) | (u32(DXBC_RETURN_TYPE_FLOAT) << 4) | (u32(DXBC_RETURN_TYPE_FLOAT) << 8) | (u32(DXBC_RETURN_TYPE_FLOAT) << 12)
	append(&b.shdr_decls, ret_type)
}

dxbc_emit_dcl_sampler :: proc(b: ^DXBC_Builder, reg: u32, comparison: bool) {
	mode := comparison ? u32(DXBC_SAMPLER_COMPARISON) : u32(DXBC_SAMPLER_DEFAULT)
	append(&b.shdr_decls, dxbc_opcode_token_ex(DXBC_OP_DCL_SAMPLER, 3, mode))
	append(&b.shdr_decls, dxbc_operand_token(DXBC_COMP_0, 0, DXBC_OPERAND_SAMPLER, DXBC_INDEX_1D) | dxbc_index_rep(DXBC_INDEX_IMM32))
	append(&b.shdr_decls, reg)
}

dxbc_emit_dcl_input :: proc(b: ^DXBC_Builder, stage: Shader_Stage, reg: u32, mask: u32) {
	if stage == .Fragment {
		// dcl_input_ps linear vN.mask
		append(&b.shdr_decls, dxbc_opcode_token_ex(DXBC_OP_DCL_INPUT_PS, 3, DXBC_INTERPOLATION_LINEAR))
		append(&b.shdr_decls, dxbc_mask_operand(DXBC_OPERAND_INPUT, DXBC_INDEX_1D, mask) | dxbc_index_rep(DXBC_INDEX_IMM32))
		append(&b.shdr_decls, reg)
	} else {
		// dcl_input vN.mask
		append(&b.shdr_decls, dxbc_opcode_token(DXBC_OP_DCL_INPUT, 3))
		append(&b.shdr_decls, dxbc_mask_operand(DXBC_OPERAND_INPUT, DXBC_INDEX_1D, mask) | dxbc_index_rep(DXBC_INDEX_IMM32))
		append(&b.shdr_decls, reg)
	}
}

dxbc_emit_dcl_input_sgv :: proc(b: ^DXBC_Builder, reg: u32, sv: u32) {
	opcode: u32
	if b.current_fn != nil && b.current_fn.stage == .Fragment {
		opcode = DXBC_OP_DCL_INPUT_PS_SGV
	} else {
		opcode = DXBC_OP_DCL_INPUT_SGV
	}
	append(&b.shdr_decls, dxbc_opcode_token(opcode, 4))
	append(&b.shdr_decls, dxbc_mask_operand(DXBC_OPERAND_INPUT, DXBC_INDEX_1D, DXBC_WRITEMASK_X) | dxbc_index_rep(DXBC_INDEX_IMM32))
	append(&b.shdr_decls, reg)
	append(&b.shdr_decls, sv)
}

dxbc_emit_dcl_input_ps_siv :: proc(b: ^DXBC_Builder, reg: u32, mask: u32, sv: u32) {
	append(&b.shdr_decls, dxbc_opcode_token_ex(DXBC_OP_DCL_INPUT_PS_SIV, 4, DXBC_INTERPOLATION_LINEAR))
	append(&b.shdr_decls, dxbc_mask_operand(DXBC_OPERAND_INPUT, DXBC_INDEX_1D, mask) | dxbc_index_rep(DXBC_INDEX_IMM32))
	append(&b.shdr_decls, reg)
	append(&b.shdr_decls, sv)
}

dxbc_emit_dcl_output :: proc(b: ^DXBC_Builder, stage: Shader_Stage, reg: u32, mask: u32) {
	append(&b.shdr_decls, dxbc_opcode_token(DXBC_OP_DCL_OUTPUT, 3))
	append(&b.shdr_decls, dxbc_mask_operand(DXBC_OPERAND_OUTPUT, DXBC_INDEX_1D, mask) | dxbc_index_rep(DXBC_INDEX_IMM32))
	append(&b.shdr_decls, reg)
}

dxbc_emit_dcl_output_siv :: proc(b: ^DXBC_Builder, reg: u32, mask: u32, sv: u32) {
	append(&b.shdr_decls, dxbc_opcode_token(DXBC_OP_DCL_OUTPUT_SIV, 4))
	append(&b.shdr_decls, dxbc_mask_operand(DXBC_OPERAND_OUTPUT, DXBC_INDEX_1D, mask) | dxbc_index_rep(DXBC_INDEX_IMM32))
	append(&b.shdr_decls, reg)
	append(&b.shdr_decls, sv)
}

// ---- Operand encoding helpers ----

// Encode a destination operand (register with write mask) into the buffer
dxbc_encode_dest :: proc(buf: ^[dynamic]u32, reg: DXBC_Reg) {
	if reg.kind == DXBC_OPERAND_NULL {
		// NULL operand: 0D, no following index
		append(buf, dxbc_mask_operand(DXBC_OPERAND_NULL, DXBC_INDEX_0D, DXBC_WRITEMASK_ALL))
		return
	}
	if reg.kind == DXBC_OPERAND_OUTPUT_DEPTH {
		// Special scalar output
		append(buf, dxbc_scalar_operand(DXBC_OPERAND_OUTPUT_DEPTH, DXBC_INDEX_0D))
		return
	}
	if reg.kind == DXBC_OPERAND_INDEXABLE_TEMP {
		// xN[M]: 2D index (register, element)
		mask := reg.mask
		if mask == 0 do mask = DXBC_WRITEMASK_ALL
		append(buf, dxbc_mask_operand(DXBC_OPERAND_INDEXABLE_TEMP, DXBC_INDEX_2D, mask) | dxbc_index_rep(DXBC_INDEX_IMM32, DXBC_INDEX_IMM32))
		append(buf, reg.index)
		append(buf, reg.index2)
		return
	}
	mask := reg.mask
	if mask == 0 do mask = DXBC_WRITEMASK_ALL
	append(buf, dxbc_mask_operand(reg.kind, DXBC_INDEX_1D, mask) | dxbc_index_rep(DXBC_INDEX_IMM32))
	append(buf, reg.index)
}

// Encode a source operand (register with swizzle) into the buffer
dxbc_encode_src :: proc(buf: ^[dynamic]u32, reg: DXBC_Reg) {
	if reg.kind == DXBC_OPERAND_IMMEDIATE32 {
		dxbc_encode_imm(buf, reg)
		return
	}

	// Build source modifier extended token if needed
	has_modifier := reg.negate || reg.abs
	extended: u32

	if reg.kind == DXBC_OPERAND_INPUT_THREAD_ID ||
	   reg.kind == DXBC_OPERAND_INPUT_THREAD_GROUP_ID ||
	   reg.kind == DXBC_OPERAND_INPUT_THREAD_ID_IN_GROUP {
		// Compute builtins: 0D index
		swz := reg.has_swizzle ? reg.swizzle : dxbc_swizzle_from_mask(reg.mask)
		tok := dxbc_swizzle_operand(reg.kind, DXBC_INDEX_0D, swz)
		if has_modifier do tok |= (1 << 31) // extended bit
		append(buf, tok)
		if has_modifier {
			mod: u32 = 1 // 1 = extended operand modifier
			if reg.negate do mod |= (1 << 6)
			if reg.abs do mod |= (1 << 7)
			append(buf, mod)
		}
		return
	}

	if reg.kind == DXBC_OPERAND_INDEXABLE_TEMP {
		// xN[M]: 2D index (register, element) — static element index
		swz := reg.has_swizzle ? reg.swizzle : dxbc_swizzle_from_mask(reg.mask)
		tok := dxbc_swizzle_operand(DXBC_OPERAND_INDEXABLE_TEMP, DXBC_INDEX_2D, swz) | dxbc_index_rep(DXBC_INDEX_IMM32, DXBC_INDEX_IMM32)
		append(buf, tok)
		append(buf, reg.index)
		append(buf, reg.index2)
		return
	}

	if reg.kind == DXBC_OPERAND_CONSTANT_BUFFER {
		// 2D index: cb[reg.index][reg.index2]
		swz := reg.has_swizzle ? reg.swizzle : dxbc_swizzle_from_mask(reg.mask)
		tok := dxbc_swizzle_operand(DXBC_OPERAND_CONSTANT_BUFFER, DXBC_INDEX_2D, swz) | dxbc_index_rep(DXBC_INDEX_IMM32, DXBC_INDEX_IMM32)
		if has_modifier do tok |= (1 << 31)
		append(buf, tok)
		if has_modifier {
			mod: u32 = 1
			if reg.negate do mod |= (1 << 6)
			if reg.abs do mod |= (1 << 7)
			append(buf, mod)
		}
		append(buf, reg.index)
		append(buf, reg.index2)
		return
	}

	if reg.kind == DXBC_OPERAND_SAMPLER || reg.kind == DXBC_OPERAND_RESOURCE {
		// sampler/resource: 4-component swizzle, 1D index
		swz := DXBC_SWIZZLE_IDENTITY
		if reg.kind == DXBC_OPERAND_RESOURCE {
			swz = reg.has_swizzle ? reg.swizzle : dxbc_swizzle_from_mask(reg.mask)
		}
		tok := dxbc_swizzle_operand(reg.kind, DXBC_INDEX_1D, swz) | dxbc_index_rep(DXBC_INDEX_IMM32)
		append(buf, tok)
		append(buf, reg.index)
		return
	}

	// Standard 1D register with swizzle
	swz := reg.has_swizzle ? reg.swizzle : dxbc_swizzle_from_mask(reg.mask)
	tok := dxbc_swizzle_operand(reg.kind, DXBC_INDEX_1D, swz) | dxbc_index_rep(DXBC_INDEX_IMM32)
	if has_modifier do tok |= (1 << 31)
	append(buf, tok)
	if has_modifier {
		mod: u32 = 1
		if reg.negate do mod |= (1 << 6)
		if reg.abs do mod |= (1 << 7)
		append(buf, mod)
	}
	append(buf, reg.index)
}

dxbc_encode_imm :: proc(buf: ^[dynamic]u32, reg: DXBC_Reg) {
	if reg.components <= 1 {
		// Scalar immediate
		append(buf, dxbc_operand_token(DXBC_COMP_1, 0, DXBC_OPERAND_IMMEDIATE32, DXBC_INDEX_0D))
		append(buf, reg.index) // the value is stored in index for scalar imm
	} else {
		// 4-component immediate
		append(buf, dxbc_operand_token(DXBC_COMP_4, DXBC_COMP_SWIZZLE, DXBC_OPERAND_IMMEDIATE32, DXBC_INDEX_0D))
		append(buf, reg.index)  // x
		append(buf, reg.index2) // y (repurposed)
		append(buf, reg.mask)   // z (repurposed)
		append(buf, 0)          // w
	}
}

dxbc_swizzle_from_mask :: proc(mask: u32) -> u32 {
	// Convert a write mask to a source swizzle
	// e.g., mask 0x7 (xyz) -> swizzle x,y,z,z
	comps: [4]u32
	idx: u32
	for i in u32(0)..<4 {
		if mask & (1 << i) != 0 {
			comps[i] = i
			idx = i
		} else {
			comps[i] = idx // repeat last valid component
		}
	}
	return dxbc_swizzle(comps[0], comps[1], comps[2], comps[3])
}

// Set an explicit source swizzle on a register copy
dxbc_with_swizzle :: proc(reg: DXBC_Reg, swz: u32) -> DXBC_Reg {
	r := reg
	r.swizzle = swz
	r.has_swizzle = true
	return r
}

dxbc_make_imm_f32 :: proc(val: f32) -> DXBC_Reg {
	return DXBC_Reg{
		kind       = DXBC_OPERAND_IMMEDIATE32,
		index      = transmute(u32)val,
		components = 1,
	}
}

dxbc_make_imm_u32 :: proc(val: u32) -> DXBC_Reg {
	return DXBC_Reg{
		kind       = DXBC_OPERAND_IMMEDIATE32,
		index      = val,
		components = 1,
	}
}

dxbc_make_imm_i32 :: proc(val: i32) -> DXBC_Reg {
	return DXBC_Reg{
		kind       = DXBC_OPERAND_IMMEDIATE32,
		index      = transmute(u32)val,
		components = 1,
	}
}

// Emit an instruction: write opcode+operands to shdr_body, patch length
dxbc_emit_op :: proc(b: ^DXBC_Builder, opcode: u32, operands: ..DXBC_Reg) {
	start := len(b.shdr_body)
	append(&b.shdr_body, u32(0)) // placeholder opcode token

	is_dest := true
	for op in operands {
		if is_dest {
			dxbc_encode_dest(&b.shdr_body, op)
			is_dest = false
		} else {
			dxbc_encode_src(&b.shdr_body, op)
		}
	}

	length := u32(len(b.shdr_body) - start)
	b.shdr_body[start] = dxbc_opcode_token(opcode, length)
	b.stat_inst_count += 1
}

// Emit instruction with custom control bits (e.g., IF_NZ)
dxbc_emit_op_ex :: proc(b: ^DXBC_Builder, opcode: u32, control: u32, operands: ..DXBC_Reg) {
	start := len(b.shdr_body)
	append(&b.shdr_body, u32(0)) // placeholder

	is_dest := true
	for op in operands {
		if is_dest && opcode != DXBC_OP_IF && opcode != DXBC_OP_BREAKC && opcode != DXBC_OP_DISCARD {
			dxbc_encode_dest(&b.shdr_body, op)
			is_dest = false
		} else {
			dxbc_encode_src(&b.shdr_body, op)
			is_dest = false
		}
	}

	length := u32(len(b.shdr_body) - start)
	b.shdr_body[start] = dxbc_opcode_token_ex(opcode, length, control)
	b.stat_inst_count += 1
}

// Emit a source-only instruction (no dest), like SAMPLE which needs custom encoding
dxbc_emit_op_raw :: proc(b: ^DXBC_Builder, opcode: u32, tokens: ..u32) {
	start := len(b.shdr_body)
	append(&b.shdr_body, u32(0)) // placeholder
	for t in tokens {
		append(&b.shdr_body, t)
	}
	length := u32(len(b.shdr_body) - start)
	b.shdr_body[start] = dxbc_opcode_token(opcode, length)
	b.stat_inst_count += 1
}

// ---- Statement emission ----

dxbc_emit_stmts :: proc(b: ^DXBC_Builder, stmts: []IR_Stmt) {
	for stmt in stmts {
		dxbc_emit_stmt(b, stmt)
	}
}

dxbc_emit_stmt :: proc(b: ^DXBC_Builder, stmt: IR_Stmt) {
	switch s in stmt {
	case ^IR_Let:
		if arr, ok := s.type^.(Type_Array_Resolved); ok {
			// Array variables need indexable temps for dynamic indexing
			elem_comps := dxbc_type_components(arr.elem)
			elem_count := u32(arr.size)
			if elem_count == 0 do elem_count = 1
			xreg := b.num_indexable_temps
			b.num_indexable_temps += 1
			info := DXBC_Indexable_Info{xreg = xreg, elem_count = elem_count, elem_comps = elem_comps}
			b.indexable_temp_map[s.id] = info
			// Emit dcl_indexable_temp x{xreg}[elem_count], 4
			// Format: opcode + operand(INDEXABLE_TEMP, 1D, xreg) + elem_count + components_per_reg = 5 DWORDs
			append(&b.shdr_decls, dxbc_opcode_token(DXBC_OP_DCL_INDEXABLE_TEMP, 5))
			append(&b.shdr_decls, dxbc_mask_operand(DXBC_OPERAND_INDEXABLE_TEMP, DXBC_INDEX_1D, DXBC_WRITEMASK_ALL) | dxbc_index_rep(DXBC_INDEX_IMM32))
			append(&b.shdr_decls, xreg)
			append(&b.shdr_decls, elem_count)
			append(&b.shdr_decls, u32(elem_comps))
			// Emit array constructor if value is a Construct
			if con, ok2 := s.value.derived.(^IR_Construct); ok2 {
				for arg_expr, i in con.args {
					arg := dxbc_emit_expr(b, arg_expr)
					itmp := DXBC_Reg{
						kind       = DXBC_OPERAND_INDEXABLE_TEMP,
						index      = xreg,
						index2     = u32(i),
						components = elem_comps,
						mask       = dxbc_mask_from_count(elem_comps),
					}
					dxbc_emit_mov_indexable(b, itmp, arg)
				}
			}
			// Register a regular temp as a placeholder (for var_map lookups)
			placeholder := dxbc_alloc_temp(b, 1)
			b.var_map[s.id] = placeholder
			break
		}
		result := dxbc_emit_expr(b, s.value)
		dest := dxbc_alloc_temp(b, dxbc_type_components(s.type))
		dxbc_emit_mov(b, dest, result)
		b.var_map[s.id] = dest

	case ^IR_Assign:
		value := dxbc_emit_expr(b, s.value)
		// TGSM writes need store_structured, not MOV
		if sr, ok := s.target.derived.(^IR_Shared_Ref); ok {
			if tgsm_idx, ok2 := b.shared_var_map[sr.name]; ok2 {
				dxbc_emit_store_structured(b, tgsm_idx, value)
				break
			}
		}
		target := dxbc_emit_lvalue(b, s.target)
		dxbc_emit_mov(b, target, value)

	case ^IR_Return:
		// Return is handled by falling through to RET at the end
		// If there's a return value in a helper function, we'd need different handling
		// For entry points, outputs are written via Store_Output

	case ^IR_Store_Output:
		value := dxbc_emit_expr(b, s.value)
		fn := b.current_fn
		if fn != nil && s.io_index >= 0 && s.io_index < len(fn.outputs) {
			io := fn.outputs[s.io_index]
			out_reg: DXBC_Reg
			if io.builtin != "" {
				if r, ok := b.builtin_output_map[io.builtin]; ok {
					out_reg = r
				}
			} else {
				if r, ok := b.output_map[io.name]; ok {
					out_reg = r
				}
			}
			dxbc_emit_mov(b, out_reg, value)
		}

	case ^IR_If:
		cond := dxbc_emit_expr(b, s.condition)
		// Ensure condition is scalar
		cond_scalar := dxbc_ensure_scalar(b, cond)
		// IF_NZ cond
		dxbc_emit_op_ex(b, DXBC_OP_IF, 1, cond_scalar) // 1 = if_nz
		dxbc_emit_stmts(b, s.then_body[:])

		for ei in s.elseif_clauses {
			append(&b.shdr_body, dxbc_opcode_token(DXBC_OP_ELSE, 1))
			eicond := dxbc_emit_expr(b, ei.condition)
			eicond_scalar := dxbc_ensure_scalar(b, eicond)
			dxbc_emit_op_ex(b, DXBC_OP_IF, 1, eicond_scalar)
			dxbc_emit_stmts(b, ei.body[:])
			append(&b.shdr_body, dxbc_opcode_token(DXBC_OP_ENDIF, 1))
		}

		if len(s.else_body) > 0 {
			append(&b.shdr_body, dxbc_opcode_token(DXBC_OP_ELSE, 1))
			dxbc_emit_stmts(b, s.else_body[:])
		}

		append(&b.shdr_body, dxbc_opcode_token(DXBC_OP_ENDIF, 1))

	case ^IR_For:
		// Initialize loop variable
		start_val := dxbc_emit_expr(b, s.start)
		loop_var := dxbc_alloc_temp(b, 1)
		dxbc_emit_mov(b, loop_var, start_val)
		b.var_map[s.var_id] = loop_var

		stop_val := dxbc_emit_expr(b, s.stop)

		append(&b.shdr_body, dxbc_opcode_token(DXBC_OP_LOOP, 1))

		// Condition: loop_var > stop -> break (using IGE for int)
		cond_tmp := dxbc_alloc_temp(b, 1)
		// ILT: is loop_var <= stop? If NOT, break.
		// Actually use: if loop_var > stop => break
		// IGE loop_var, stop => true when loop_var >= stop+1
		// Simpler: use ILT stop, loop_var => true when stop < loop_var => break
		dxbc_emit_op(b, DXBC_OP_ILT, cond_tmp, stop_val, loop_var)
		dxbc_emit_op_ex(b, DXBC_OP_BREAKC, 1, cond_tmp) // breakc_nz

		dxbc_emit_stmts(b, s.body[:])

		// Increment
		if s.step != nil {
			step_val := dxbc_emit_expr(b, s.step)
			dxbc_emit_op(b, DXBC_OP_IADD, loop_var, loop_var, step_val)
		} else {
			one := dxbc_make_imm_i32(1)
			dxbc_emit_op(b, DXBC_OP_IADD, loop_var, loop_var, one)
		}

		append(&b.shdr_body, dxbc_opcode_token(DXBC_OP_ENDLOOP, 1))

	case ^IR_While:
		append(&b.shdr_body, dxbc_opcode_token(DXBC_OP_LOOP, 1))

		cond := dxbc_emit_expr(b, s.condition)
		cond_scalar := dxbc_ensure_scalar(b, cond)
		// Break if condition is false (breakc_z)
		not_cond := dxbc_alloc_temp(b, 1)
		dxbc_emit_op(b, DXBC_OP_NOT, not_cond, cond_scalar)
		dxbc_emit_op_ex(b, DXBC_OP_BREAKC, 1, not_cond)

		dxbc_emit_stmts(b, s.body[:])
		append(&b.shdr_body, dxbc_opcode_token(DXBC_OP_ENDLOOP, 1))

	case ^IR_Expr_Stmt:
		dxbc_emit_expr(b, s.expr)

	case ^IR_Barrier:
		flags := u32(DXBC_SYNC_THREADS_IN_GROUP | DXBC_SYNC_THREAD_GROUP_SHARED)
		append(&b.shdr_body, dxbc_opcode_token_ex(DXBC_OP_SYNC, 1, flags >> 11))
		b.stat_inst_count += 1

	case ^IR_Discard:
		// discard_nz with all-ones condition
		all_ones := dxbc_make_imm_u32(0xFFFFFFFF)
		dxbc_emit_op_ex(b, DXBC_OP_DISCARD, 1, all_ones) // 1 = discard_nz

	case ^IR_Break:
		append(&b.shdr_body, dxbc_opcode_token(DXBC_OP_BREAK, 1))
		b.stat_inst_count += 1

	case ^IR_Continue:
		append(&b.shdr_body, dxbc_opcode_token(DXBC_OP_CONTINUE, 1))
		b.stat_inst_count += 1
	}
}

dxbc_ensure_scalar :: proc(b: ^DXBC_Builder, reg: DXBC_Reg) -> DXBC_Reg {
	if reg.components == 1 do return reg
	// Extract .x component
	r := reg
	r.mask = DXBC_WRITEMASK_X
	r.components = 1
	return r
}

dxbc_emit_lvalue :: proc(b: ^DXBC_Builder, expr: ^IR_Expr) -> DXBC_Reg {
	if expr == nil do return {}
	#partial switch e in expr.derived {
	case ^IR_Var_Ref:
		if reg, ok := b.var_map[e.id]; ok {
			return reg
		}
	case ^IR_Shared_Ref:
		if tgsm_idx, ok := b.shared_var_map[e.name]; ok {
			return DXBC_Reg{kind = DXBC_OPERAND_THREAD_GROUP_SHARED_MEMORY, index = tgsm_idx, components = 4, mask = DXBC_WRITEMASK_ALL}
		}
	case ^IR_Swizzle:
		base := dxbc_emit_lvalue(b, e.object)
		base.mask = dxbc_swizzle_string_to_mask(e.components)
		base.components = len(e.components)
		return base
	case ^IR_Index:
		if vr, ok := e.object.derived.(^IR_Var_Ref); ok {
			if info, ok2 := b.indexable_temp_map[vr.id]; ok2 {
				if lit, ok3 := e.index.derived.(^IR_Literal); ok3 {
					elem_idx: u32
					if iv, ok4 := lit.value.(i64); ok4 { elem_idx = u32(iv) }
					return DXBC_Reg{
						kind       = DXBC_OPERAND_INDEXABLE_TEMP,
						index      = info.xreg,
						index2     = elem_idx,
						components = info.elem_comps,
						mask       = dxbc_mask_from_count(info.elem_comps),
					}
				}
			}
		}
		// Fallback: return base object
		return dxbc_emit_lvalue(b, e.object)
	case: // other lvalues
	}
	return {}
}

dxbc_swizzle_string_to_mask :: proc(s: string) -> u32 {
	mask: u32
	for c in s {
		switch c {
		case 'x', 'r': mask |= DXBC_WRITEMASK_X
		case 'y', 'g': mask |= DXBC_WRITEMASK_Y
		case 'z', 'b': mask |= DXBC_WRITEMASK_Z
		case 'w', 'a': mask |= DXBC_WRITEMASK_W
		}
	}
	return mask
}

// ---- Expression emission ----

dxbc_emit_expr :: proc(b: ^DXBC_Builder, expr: ^IR_Expr) -> DXBC_Reg {
	if expr == nil do return dxbc_make_imm_f32(0)

	#partial switch e in expr.derived {
	case ^IR_Literal:
		return dxbc_emit_literal(b, e, expr.type)

	case ^IR_Var_Ref:
		if reg, ok := b.var_map[e.id]; ok {
			return reg
		}
		return dxbc_make_imm_f32(0)

	case ^IR_Binary:
		return dxbc_emit_binary(b, e, expr.type)

	case ^IR_Unary:
		return dxbc_emit_unary(b, e, expr.type)

	case ^IR_Call:
		return dxbc_emit_call(b, e, expr.type)

	case ^IR_Field_Access:
		return dxbc_emit_field_access(b, e, expr.type)

	case ^IR_Swizzle:
		return dxbc_emit_swizzle(b, e, expr.type)

	case ^IR_Index:
		// Check if the object is an array variable with an indexable temp
		if vr, ok := e.object.derived.(^IR_Var_Ref); ok {
			if info, ok2 := b.indexable_temp_map[vr.id]; ok2 {
				comp := info.elem_comps
				dest := dxbc_alloc_temp(b, comp)
				// Check if index is a literal (static) or dynamic
				if lit, ok3 := e.index.derived.(^IR_Literal); ok3 {
					elem_idx: u32
					if iv, ok4 := lit.value.(i64); ok4 {
						elem_idx = u32(iv)
					}
					itmp := DXBC_Reg{
						kind       = DXBC_OPERAND_INDEXABLE_TEMP,
						index      = info.xreg,
						index2     = elem_idx,
						components = comp,
						mask       = dxbc_mask_from_count(comp),
					}
					dxbc_emit_mov_indexable_src(b, dest, itmp)
				} else {
					idx := dxbc_emit_expr(b, e.index)
					dxbc_emit_indexed_load(b, dest, info.xreg, idx, comp)
				}
				return dest
			}
		}
		// Vector component indexing
		obj := dxbc_emit_expr(b, e.object)
		comp := dxbc_type_components(expr.type)
		if lit, ok := e.index.derived.(^IR_Literal); ok {
			// Static index: extract component via swizzle
			comp_idx: u32
			if iv, ok2 := lit.value.(i64); ok2 {
				comp_idx = u32(iv)
			}
			result := dxbc_with_swizzle(obj, dxbc_swizzle_replicate(comp_idx))
			result.components = comp
			return result
		}
		// Dynamic vector index: use MOVC chain
		idx := dxbc_emit_expr(b, e.index)
		return dxbc_emit_dynamic_vec_index(b, obj, idx, comp)

	case ^IR_Construct:
		return dxbc_emit_construct(b, e, expr.type)

	case ^IR_Type_Cast:
		return dxbc_emit_type_cast(b, e, expr.type)

	case ^IR_Load_Binding:
		return dxbc_emit_load_binding(b, e)

	case ^IR_Input_Field:
		return dxbc_emit_input_field(b, e)

	case ^IR_Builtin_Var:
		if e.is_input {
			if reg, ok := b.builtin_input_map[e.name]; ok {
				return reg
			}
		} else {
			if reg, ok := b.builtin_output_map[e.name]; ok {
				return reg
			}
		}
		return dxbc_make_imm_f32(0)

	case ^IR_Composite_Extract:
		return dxbc_emit_composite_extract(b, e, expr.type)

	case ^IR_Vector_Shuffle:
		return dxbc_emit_vector_shuffle(b, e, expr.type)

	case ^IR_Shared_Ref:
		if tgsm_idx, ok := b.shared_var_map[e.name]; ok {
			comp := dxbc_type_components(expr.type)
			dest := dxbc_alloc_temp(b, comp)
			dxbc_emit_load_structured(b, dest, tgsm_idx, comp)
			return dest
		}
		return dxbc_make_imm_f32(0)

	case ^IR_Select:
		cond := dxbc_emit_expr(b, e.condition)
		true_val := dxbc_emit_expr(b, e.true_val)
		false_val := dxbc_emit_expr(b, e.false_val)
		comp := dxbc_type_components(expr.type)
		dest := dxbc_alloc_temp(b, comp)
		// MOVC dest, cond, true, false
		// Need to broadcast condition if scalar
		cond_bc := cond
		if cond.components == 1 && comp > 1 {
			cond_bc.mask = dxbc_mask_from_count(comp)
			cond_bc.components = comp
		}
		dxbc_emit_op(b, DXBC_OP_MOVC, dest, cond_bc, true_val, false_val)
		return dest
	}

	return dxbc_make_imm_f32(0)
}

dxbc_emit_literal :: proc(b: ^DXBC_Builder, lit: ^IR_Literal, t: ^Resolved_Type) -> DXBC_Reg {
	switch v in lit.value {
	case f64:
		return dxbc_make_imm_f32(f32(v))
	case i64:
		return dxbc_make_imm_i32(i32(v))
	case bool:
		return dxbc_make_imm_u32(v ? 0xFFFFFFFF : 0)
	}
	return dxbc_make_imm_f32(0)
}

dxbc_emit_binary :: proc(b: ^DXBC_Builder, bin: ^IR_Binary, t: ^Resolved_Type) -> DXBC_Reg {
	left := dxbc_emit_expr(b, bin.left)
	right := dxbc_emit_expr(b, bin.right)
	comp := dxbc_type_components(t)
	is_float := is_float_type(t) || is_float_type(bin.left.type)
	is_uint := is_uint_type(t)

	dest := dxbc_alloc_temp(b, comp)

	switch bin.op {
	case .Add:
		dxbc_emit_op(b, is_float ? DXBC_OP_ADD : DXBC_OP_IADD, dest, left, right)
	case .Sub:
		if is_float {
			neg_right := right
			neg_right.negate = !neg_right.negate
			dxbc_emit_op(b, DXBC_OP_ADD, dest, left, neg_right)
		} else {
			neg_right := dxbc_alloc_temp(b, comp)
			dxbc_emit_op(b, DXBC_OP_INEG, neg_right, right)
			dxbc_emit_op(b, DXBC_OP_IADD, dest, left, neg_right)
		}
	case .Mul:
		if is_float {
			// Check for matrix multiply
			if is_matrix(bin.left.type) || is_matrix(bin.right.type) {
				return dxbc_emit_matrix_mul(b, bin.left, bin.right, t)
			}
			dxbc_emit_op(b, DXBC_OP_MUL, dest, left, right)
		} else {
			dxbc_emit_op(b, DXBC_OP_IMUL, dest, dest, left, right)
		}
	case .Div:
		if is_float {
			dxbc_emit_op(b, DXBC_OP_DIV, dest, left, right)
		} else {
			dxbc_emit_op(b, DXBC_OP_UDIV, dest, DXBC_Reg{kind = DXBC_OPERAND_NULL}, left, right)
		}
	case .Mod:
		if is_float {
			// fmod: a - floor(a/b) * b
			tmp := dxbc_alloc_temp(b, comp)
			dxbc_emit_op(b, DXBC_OP_DIV, tmp, left, right)
			dxbc_emit_op(b, DXBC_OP_ROUND_NI, tmp, tmp)
			dxbc_emit_op(b, DXBC_OP_MUL, tmp, tmp, right)
			neg_tmp := tmp
			neg_tmp.negate = true
			dxbc_emit_op(b, DXBC_OP_ADD, dest, left, neg_tmp)
		} else {
			// UDIV writes quotient to dest1 and remainder to dest2
			dxbc_emit_op(b, DXBC_OP_UDIV, DXBC_Reg{kind = DXBC_OPERAND_NULL}, dest, left, right)
		}
	case .Eq:
		dxbc_emit_op(b, is_float ? DXBC_OP_EQ : DXBC_OP_IEQ, dest, left, right)
	case .Neq:
		dxbc_emit_op(b, is_float ? DXBC_OP_NE : DXBC_OP_INE, dest, left, right)
	case .Lt:
		dxbc_emit_op(b, is_float ? DXBC_OP_LT : (is_uint ? DXBC_OP_ULT : DXBC_OP_ILT), dest, left, right)
	case .Gt:
		dxbc_emit_op(b, is_float ? DXBC_OP_LT : (is_uint ? DXBC_OP_ULT : DXBC_OP_ILT), dest, right, left)
	case .Lte:
		dxbc_emit_op(b, is_float ? DXBC_OP_GE : (is_uint ? DXBC_OP_UGE : DXBC_OP_IGE), dest, right, left)
	case .Gte:
		dxbc_emit_op(b, is_float ? DXBC_OP_GE : (is_uint ? DXBC_OP_UGE : DXBC_OP_IGE), dest, left, right)
	case .And:
		dxbc_emit_op(b, DXBC_OP_AND, dest, left, right)
	case .Or:
		dxbc_emit_op(b, DXBC_OP_OR, dest, left, right)
	case .Neg, .Not:
		// These are unary ops, shouldn't appear in binary
	}

	return dest
}

is_uint_type :: proc(t: ^Resolved_Type) -> bool {
	if t == nil do return false
	#partial switch v in t^ {
	case Type_Scalar: return v.kind == .Uint
	case Type_Vector: return v.elem == .Uint
	case: return false
	}
}

dxbc_emit_unary :: proc(b: ^DXBC_Builder, un: ^IR_Unary, t: ^Resolved_Type) -> DXBC_Reg {
	operand := dxbc_emit_expr(b, un.operand)
	comp := dxbc_type_components(t)
	dest := dxbc_alloc_temp(b, comp)

	#partial switch un.op {
	case .Neg:
		if is_float_type(t) {
			neg_op := operand
			neg_op.negate = !neg_op.negate
			dxbc_emit_mov(b, dest, neg_op)
		} else {
			dxbc_emit_op(b, DXBC_OP_INEG, dest, operand)
		}
	case .Not:
		dxbc_emit_op(b, DXBC_OP_NOT, dest, operand)
	case:
		dxbc_emit_mov(b, dest, operand)
	}

	return dest
}

dxbc_emit_mov :: proc(b: ^DXBC_Builder, dest: DXBC_Reg, src: DXBC_Reg) {
	dxbc_emit_op(b, DXBC_OP_MOV, dest, src)
}

// SINCOS has two destination operands: sincos dst_sin, dst_cos, src
// Use NULL_REG for an unused destination.
dxbc_emit_sincos :: proc(b: ^DXBC_Builder, sin_dest: DXBC_Reg, cos_dest: DXBC_Reg, src: DXBC_Reg) {
	start := len(b.shdr_body)
	append(&b.shdr_body, u32(0))
	dxbc_encode_dest(&b.shdr_body, sin_dest)
	dxbc_encode_dest(&b.shdr_body, cos_dest)
	dxbc_encode_src(&b.shdr_body, src)
	length := u32(len(b.shdr_body) - start)
	b.shdr_body[start] = dxbc_opcode_token(DXBC_OP_SINCOS, length)
	b.stat_inst_count += 1
}

dxbc_emit_call :: proc(b: ^DXBC_Builder, call: ^IR_Call, t: ^Resolved_Type) -> DXBC_Reg {
	comp := dxbc_type_components(t)
	if !call.is_builtin {
		return dxbc_emit_inline_call(b, call, t)
	}

	switch call.name {
	case "sample":
		return dxbc_emit_sample(b, call, t)
	case "sample_shadow":
		return dxbc_emit_sample_shadow(b, call, t)
	case "normalize":
		return dxbc_emit_normalize(b, call, t)
	case "dot":
		return dxbc_emit_dot(b, call, t)
	case "length":
		return dxbc_emit_length(b, call, t)
	case "cross":
		return dxbc_emit_cross(b, call, t)
	case "min":
		a := dxbc_emit_expr(b, call.args[0])
		c := dxbc_emit_expr(b, call.args[1])
		dest := dxbc_alloc_temp(b, comp)
		dxbc_emit_op(b, is_float_type(t) ? DXBC_OP_MIN : (is_uint_type(t) ? DXBC_OP_UMIN : DXBC_OP_IMIN), dest, a, c)
		return dest
	case "max":
		a := dxbc_emit_expr(b, call.args[0])
		c := dxbc_emit_expr(b, call.args[1])
		dest := dxbc_alloc_temp(b, comp)
		dxbc_emit_op(b, is_float_type(t) ? DXBC_OP_MAX : (is_uint_type(t) ? DXBC_OP_UMAX : DXBC_OP_IMAX), dest, a, c)
		return dest
	case "clamp":
		x := dxbc_emit_expr(b, call.args[0])
		lo := dxbc_emit_expr(b, call.args[1])
		hi := dxbc_emit_expr(b, call.args[2])
		tmp := dxbc_alloc_temp(b, comp)
		dest := dxbc_alloc_temp(b, comp)
		dxbc_emit_op(b, DXBC_OP_MAX, tmp, x, lo)
		dxbc_emit_op(b, DXBC_OP_MIN, dest, tmp, hi)
		return dest
	case "mix", "lerp":
		a := dxbc_emit_expr(b, call.args[0])
		c := dxbc_emit_expr(b, call.args[1])
		t_val := dxbc_emit_expr(b, call.args[2])
		// mix(a, b, t) = a + t*(b-a) = mad(t, b-a, a)
		tmp := dxbc_alloc_temp(b, comp)
		dest := dxbc_alloc_temp(b, comp)
		neg_a := a
		neg_a.negate = !neg_a.negate
		dxbc_emit_op(b, DXBC_OP_ADD, tmp, c, neg_a) // tmp = b - a
		dxbc_emit_op(b, DXBC_OP_MAD, dest, t_val, tmp, a) // dest = t*(b-a)+a
		return dest
	case "abs":
		a := dxbc_emit_expr(b, call.args[0])
		dest := dxbc_alloc_temp(b, comp)
		abs_a := a
		abs_a.abs = true
		dxbc_emit_mov(b, dest, abs_a)
		return dest
	case "floor":
		a := dxbc_emit_expr(b, call.args[0])
		dest := dxbc_alloc_temp(b, comp)
		dxbc_emit_op(b, DXBC_OP_ROUND_NI, dest, a)
		return dest
	case "ceil":
		a := dxbc_emit_expr(b, call.args[0])
		dest := dxbc_alloc_temp(b, comp)
		dxbc_emit_op(b, DXBC_OP_ROUND_PI, dest, a)
		return dest
	case "fract", "frac":
		a := dxbc_emit_expr(b, call.args[0])
		dest := dxbc_alloc_temp(b, comp)
		dxbc_emit_op(b, DXBC_OP_FRC, dest, a)
		return dest
	case "sqrt":
		a := dxbc_emit_expr(b, call.args[0])
		dest := dxbc_alloc_temp(b, comp)
		dxbc_emit_op(b, DXBC_OP_SQRT, dest, a)
		return dest
	case "inversesqrt", "rsqrt":
		a := dxbc_emit_expr(b, call.args[0])
		dest := dxbc_alloc_temp(b, comp)
		dxbc_emit_op(b, DXBC_OP_RSQ, dest, a)
		return dest
	case "exp", "exp2":
		a := dxbc_emit_expr(b, call.args[0])
		dest := dxbc_alloc_temp(b, comp)
		dxbc_emit_op(b, DXBC_OP_EXP, dest, a)
		return dest
	case "log", "log2":
		a := dxbc_emit_expr(b, call.args[0])
		dest := dxbc_alloc_temp(b, comp)
		dxbc_emit_op(b, DXBC_OP_LOG, dest, a)
		return dest
	case "sin":
		a := dxbc_emit_expr(b, call.args[0])
		dest := dxbc_alloc_temp(b, comp)
		null_reg := DXBC_Reg{kind = DXBC_OPERAND_NULL}
		dxbc_emit_sincos(b, dest, null_reg, a)
		return dest
	case "cos":
		a := dxbc_emit_expr(b, call.args[0])
		dest := dxbc_alloc_temp(b, comp)
		null_reg := DXBC_Reg{kind = DXBC_OPERAND_NULL}
		dxbc_emit_sincos(b, null_reg, dest, a)
		return dest
	case "pow":
		// pow(a,b) = exp2(b * log2(a))
		a := dxbc_emit_expr(b, call.args[0])
		p := dxbc_emit_expr(b, call.args[1])
		tmp := dxbc_alloc_temp(b, comp)
		dest := dxbc_alloc_temp(b, comp)
		dxbc_emit_op(b, DXBC_OP_LOG, tmp, a)
		dxbc_emit_op(b, DXBC_OP_MUL, tmp, tmp, p)
		dxbc_emit_op(b, DXBC_OP_EXP, dest, tmp)
		return dest
	case "step":
		// step(edge, x) = x >= edge ? 1 : 0
		edge := dxbc_emit_expr(b, call.args[0])
		x := dxbc_emit_expr(b, call.args[1])
		cmp := dxbc_alloc_temp(b, comp)
		dest := dxbc_alloc_temp(b, comp)
		dxbc_emit_op(b, DXBC_OP_GE, cmp, x, edge)
		dxbc_emit_op(b, DXBC_OP_AND, dest, cmp, dxbc_make_imm_f32(1.0))
		return dest
	case "smoothstep":
		// smoothstep(edge0, edge1, x) = t*t*(3-2*t) where t = clamp((x-edge0)/(edge1-edge0), 0, 1)
		e0 := dxbc_emit_expr(b, call.args[0])
		e1 := dxbc_emit_expr(b, call.args[1])
		x := dxbc_emit_expr(b, call.args[2])
		diff := dxbc_alloc_temp(b, comp)
		t := dxbc_alloc_temp(b, comp)
		dest := dxbc_alloc_temp(b, comp)
		neg_e0 := e0
		neg_e0.negate = true
		dxbc_emit_op(b, DXBC_OP_ADD, diff, e1, neg_e0) // edge1 - edge0
		dxbc_emit_op(b, DXBC_OP_ADD, t, x, neg_e0) // x - edge0
		dxbc_emit_op(b, DXBC_OP_DIV, t, t, diff) // (x-e0)/(e1-e0)
		dxbc_emit_op(b, DXBC_OP_MAX, t, t, dxbc_make_imm_f32(0))
		dxbc_emit_op(b, DXBC_OP_MIN, t, t, dxbc_make_imm_f32(1))
		// t*t*(3-2*t)
		tmp := dxbc_alloc_temp(b, comp)
		dxbc_emit_op(b, DXBC_OP_MUL, tmp, t, t) // t*t
		two_t := dxbc_alloc_temp(b, comp)
		dxbc_emit_op(b, DXBC_OP_ADD, two_t, t, t) // 2*t
		neg_2t := two_t
		neg_2t.negate = true
		dxbc_emit_op(b, DXBC_OP_ADD, dest, dxbc_make_imm_f32(3), neg_2t) // 3-2*t
		dxbc_emit_op(b, DXBC_OP_MUL, dest, tmp, dest)
		return dest
	case "reflect":
		// reflect(I, N) = I - 2*dot(N,I)*N
		i_val := dxbc_emit_expr(b, call.args[0])
		n_val := dxbc_emit_expr(b, call.args[1])
		d := dxbc_alloc_temp(b, 1)
		in_comp := dxbc_type_components(call.args[0].type)
		if in_comp == 3 {
			dxbc_emit_op(b, DXBC_OP_DP3, d, n_val, i_val)
		} else {
			dxbc_emit_op(b, DXBC_OP_DP4, d, n_val, i_val)
		}
		two_d := dxbc_alloc_temp(b, 1)
		dxbc_emit_op(b, DXBC_OP_ADD, two_d, d, d)
		// Broadcast and mul by N
		tmp := dxbc_alloc_temp(b, in_comp)
		two_d.mask = dxbc_mask_from_count(in_comp)
		two_d.components = in_comp
		dxbc_emit_op(b, DXBC_OP_MUL, tmp, two_d, n_val)
		dest := dxbc_alloc_temp(b, in_comp)
		neg_tmp := tmp
		neg_tmp.negate = true
		dxbc_emit_op(b, DXBC_OP_ADD, dest, i_val, neg_tmp)
		return dest
	case "refract":
		return dxbc_emit_refract(b, call, t)
	case "tan":
		// tan(x) = sin(x) / cos(x) via SINCOS
		a := dxbc_emit_expr(b, call.args[0])
		sin_dest := dxbc_alloc_temp(b, comp)
		cos_dest := dxbc_alloc_temp(b, comp)
		dxbc_emit_sincos(b, sin_dest, cos_dest, a)
		dest := dxbc_alloc_temp(b, comp)
		dxbc_emit_op(b, DXBC_OP_DIV, dest, sin_dest, cos_dest)
		return dest
	case "distance":
		// distance(a, b) = length(a - b)
		a := dxbc_emit_expr(b, call.args[0])
		c_val := dxbc_emit_expr(b, call.args[1])
		in_comp := dxbc_type_components(call.args[0].type)
		diff := dxbc_alloc_temp(b, in_comp)
		neg_c := c_val
		neg_c.negate = true
		dxbc_emit_op(b, DXBC_OP_ADD, diff, a, neg_c)
		d := dxbc_alloc_temp(b, 1)
		switch in_comp {
		case 2: dxbc_emit_op(b, DXBC_OP_DP2, d, diff, diff)
		case 3: dxbc_emit_op(b, DXBC_OP_DP3, d, diff, diff)
		case:   dxbc_emit_op(b, DXBC_OP_DP4, d, diff, diff)
		}
		dest := dxbc_alloc_temp(b, 1)
		dxbc_emit_op(b, DXBC_OP_SQRT, dest, d)
		return dest
	case "sign":
		// sign(x): x>0 ? 1 : (x<0 ? -1 : 0)
		a := dxbc_emit_expr(b, call.args[0])
		lt_op := u32(is_float_type(call.args[0].type) ? DXBC_OP_LT : DXBC_OP_ILT)
		tmp_neg := dxbc_alloc_temp(b, comp) // x < 0?
		tmp_pos := dxbc_alloc_temp(b, comp) // 0 < x?
		dxbc_emit_op(b, lt_op, tmp_neg, a, dxbc_make_imm_f32(0))
		dxbc_emit_op(b, lt_op, tmp_pos, dxbc_make_imm_f32(0), a)
		dest := dxbc_alloc_temp(b, comp)
		dxbc_emit_op(b, DXBC_OP_MOVC, dest, tmp_pos, dxbc_make_imm_f32(1), dxbc_make_imm_f32(0))
		dxbc_emit_op(b, DXBC_OP_MOVC, dest, tmp_neg, dxbc_make_imm_f32(-1), dest)
		return dest
	case "dfdx":
		a := dxbc_emit_expr(b, call.args[0])
		dest := dxbc_alloc_temp(b, comp)
		dxbc_emit_op(b, DXBC_OP_DERIV_RTX, dest, a)
		return dest
	case "dfdy":
		a := dxbc_emit_expr(b, call.args[0])
		dest := dxbc_alloc_temp(b, comp)
		dxbc_emit_op(b, DXBC_OP_DERIV_RTY, dest, a)
		return dest
	case "sample_level":
		return dxbc_emit_sample_level(b, call, t)
	case "transpose":
		return dxbc_emit_transpose(b, call, t)
	case "asin", "acos", "atan", "atan2":
		append(&b.diagnostics, Diagnostic{level = .Error, message = fmt.aprintf("DXBC: '%s' has no native SM5.0 equivalent; use a polynomial approximation", call.name)})
		return dxbc_alloc_temp(b, comp)
	case "determinant", "inverse":
		append(&b.diagnostics, Diagnostic{level = .Error, message = fmt.aprintf("DXBC: '%s' is not supported in the DXBC backend", call.name)})
		return dxbc_alloc_temp(b, comp)
	case:
		// Unknown builtin — emit as error
		append(&b.diagnostics, Diagnostic{level = .Warning, message = fmt.aprintf("DXBC: unsupported builtin '%s'", call.name)})
		return dxbc_alloc_temp(b, comp)
	}
}

dxbc_emit_inline_call :: proc(b: ^DXBC_Builder, call: ^IR_Call, t: ^Resolved_Type) -> DXBC_Reg {
	comp := dxbc_type_components(t)

	// Find the function in the module
	target_fn: ^IR_Function
	for &fn in b.module.functions {
		if fn.name == call.name {
			target_fn = &fn
			break
		}
	}
	if target_fn == nil {
		append(&b.diagnostics, Diagnostic{level = .Error, message = fmt.aprintf("DXBC: function not found for inlining: %s", call.name)})
		return dxbc_alloc_temp(b, comp)
	}

	// Save current var_map state for restoration after inlining
	saved_vars := make(map[IR_Var_Id]DXBC_Reg)
	for k, v in b.var_map {
		saved_vars[k] = v
	}

	// Emit argument expressions and map to parameter var IDs
	for i := 0; i < len(call.args) && i < len(target_fn.params); i += 1 {
		arg_reg := dxbc_emit_expr(b, call.args[i])
		param := target_fn.params[i]
		param_comp := dxbc_type_components(param.type)
		dest := dxbc_alloc_temp(b, param_comp)
		dxbc_emit_mov(b, dest, arg_reg)
		b.var_map[param.id] = dest
	}

	// Emit the function body — capture return value
	return_reg := dxbc_alloc_temp(b, comp)
	for stmt in target_fn.body {
		#partial switch s in stmt {
		case ^IR_Return:
			if s.value != nil {
				val := dxbc_emit_expr(b, s.value)
				dxbc_emit_mov(b, return_reg, val)
			}
		case:
			dxbc_emit_stmt(b, stmt)
		}
	}

	// Restore var_map
	for k, _ in b.var_map {
		if !(k in saved_vars) {
			delete_key(&b.var_map, k)
		}
	}
	for k, v in saved_vars {
		b.var_map[k] = v
	}

	return return_reg
}

dxbc_emit_normalize :: proc(b: ^DXBC_Builder, call: ^IR_Call, t: ^Resolved_Type) -> DXBC_Reg {
	v := dxbc_emit_expr(b, call.args[0])
	comp := dxbc_type_components(t)
	d := dxbc_alloc_temp(b, 1)
	if comp <= 3 {
		dxbc_emit_op(b, DXBC_OP_DP3, d, v, v)
	} else {
		dxbc_emit_op(b, DXBC_OP_DP4, d, v, v)
	}
	dxbc_emit_op(b, DXBC_OP_RSQ, d, d)
	// Broadcast d and multiply
	dest := dxbc_alloc_temp(b, comp)
	d_bc := d
	d_bc.mask = dxbc_mask_from_count(comp)
	d_bc.components = comp
	dxbc_emit_op(b, DXBC_OP_MUL, dest, v, d_bc)
	return dest
}

dxbc_emit_dot :: proc(b: ^DXBC_Builder, call: ^IR_Call, t: ^Resolved_Type) -> DXBC_Reg {
	a := dxbc_emit_expr(b, call.args[0])
	c := dxbc_emit_expr(b, call.args[1])
	comp := dxbc_type_components(call.args[0].type)
	dest := dxbc_alloc_temp(b, 1)
	switch comp {
	case 2: dxbc_emit_op(b, DXBC_OP_DP2, dest, a, c)
	case 3: dxbc_emit_op(b, DXBC_OP_DP3, dest, a, c)
	case:   dxbc_emit_op(b, DXBC_OP_DP4, dest, a, c)
	}
	return dest
}

dxbc_emit_length :: proc(b: ^DXBC_Builder, call: ^IR_Call, t: ^Resolved_Type) -> DXBC_Reg {
	v := dxbc_emit_expr(b, call.args[0])
	comp := dxbc_type_components(call.args[0].type)
	d := dxbc_alloc_temp(b, 1)
	switch comp {
	case 2: dxbc_emit_op(b, DXBC_OP_DP2, d, v, v)
	case 3: dxbc_emit_op(b, DXBC_OP_DP3, d, v, v)
	case:   dxbc_emit_op(b, DXBC_OP_DP4, d, v, v)
	}
	dest := dxbc_alloc_temp(b, 1)
	dxbc_emit_op(b, DXBC_OP_SQRT, dest, d)
	return dest
}

dxbc_emit_cross :: proc(b: ^DXBC_Builder, call: ^IR_Call, t: ^Resolved_Type) -> DXBC_Reg {
	a := dxbc_emit_expr(b, call.args[0])
	c := dxbc_emit_expr(b, call.args[1])
	// cross(a,b) = a.yzx*b.zxy - a.zxy*b.yzx
	t1 := dxbc_alloc_temp(b, 3)
	t2 := dxbc_alloc_temp(b, 3)
	a_yzx := dxbc_with_swizzle(a, dxbc_swizzle(1,2,0,0))
	b_zxy := dxbc_with_swizzle(c, dxbc_swizzle(2,0,1,0))
	a_zxy := dxbc_with_swizzle(a, dxbc_swizzle(2,0,1,0))
	b_yzx := dxbc_with_swizzle(c, dxbc_swizzle(1,2,0,0))
	dxbc_emit_op(b, DXBC_OP_MUL, t1, a_yzx, b_zxy)
	dxbc_emit_op(b, DXBC_OP_MUL, t2, a_zxy, b_yzx)
	dest := dxbc_alloc_temp(b, 3)
	neg_t2 := t2
	neg_t2.negate = true
	dxbc_emit_op(b, DXBC_OP_ADD, dest, t1, neg_t2)
	return dest
}

dxbc_emit_sample :: proc(b: ^DXBC_Builder, call: ^IR_Call, t: ^Resolved_Type) -> DXBC_Reg {
	// sample(combined_sampler, coord) -> SAMPLE dest, coord, texture, sampler
	if len(call.args) < 2 do return dxbc_alloc_temp(b, 4)

	coord := dxbc_emit_expr(b, call.args[1])
	dest := dxbc_alloc_temp(b, 4)

	// Get texture and sampler registers from the binding
	tex_reg := DXBC_Reg{kind = DXBC_OPERAND_RESOURCE, index = 0, components = 4, mask = DXBC_WRITEMASK_ALL}
	samp_reg := DXBC_Reg{kind = DXBC_OPERAND_SAMPLER, index = 0, components = 4}

	if lb, ok := call.args[0].derived.(^IR_Load_Binding); ok {
		tex_name, samp_name := find_split_bindings(b.module, lb.name)
		if info, tok := b.binding_map[tex_name]; tok {
			tex_reg.index = info.reg_num
		}
		if info, sok := b.binding_map[samp_name]; sok {
			samp_reg.index = info.reg_num
		}
	}

	// Ensure coord is at least 4-component for DXBC
	if coord.components < 4 {
		padded := dxbc_alloc_temp(b, 4)
		dxbc_emit_mov(b, padded, coord)
		coord = padded
	}

	// SAMPLE uses custom encoding: dest, coord, resource, sampler
	// Build instruction manually
	start := len(b.shdr_body)
	append(&b.shdr_body, u32(0)) // placeholder
	dxbc_encode_dest(&b.shdr_body, dest)
	dxbc_encode_src(&b.shdr_body, coord)
	dxbc_encode_src(&b.shdr_body, tex_reg)
	dxbc_encode_src(&b.shdr_body, samp_reg)
	length := u32(len(b.shdr_body) - start)
	b.shdr_body[start] = dxbc_opcode_token(DXBC_OP_SAMPLE, length)
	b.stat_inst_count += 1

	return dest
}

dxbc_emit_sample_shadow :: proc(b: ^DXBC_Builder, call: ^IR_Call, t: ^Resolved_Type) -> DXBC_Reg {
	// sample_shadow(sampler, coord, dref) -> SAMPLE_C dest, coord, resource, sampler, dref
	if len(call.args) < 3 do return dxbc_alloc_temp(b, 4)

	coord := dxbc_emit_expr(b, call.args[1])
	dref := dxbc_emit_expr(b, call.args[2])
	dest := dxbc_alloc_temp(b, 4)

	tex_reg := DXBC_Reg{kind = DXBC_OPERAND_RESOURCE, index = 0, components = 4, mask = DXBC_WRITEMASK_ALL}
	samp_reg := DXBC_Reg{kind = DXBC_OPERAND_SAMPLER, index = 0, components = 4}

	if lb, ok := call.args[0].derived.(^IR_Load_Binding); ok {
		tex_name, samp_name := find_split_bindings(b.module, lb.name)
		if info, tok := b.binding_map[tex_name]; tok {
			tex_reg.index = info.reg_num
		}
		if info, sok := b.binding_map[samp_name]; sok {
			samp_reg.index = info.reg_num
		}
	}

	if coord.components < 4 {
		padded := dxbc_alloc_temp(b, 4)
		dxbc_emit_mov(b, padded, coord)
		coord = padded
	}

	start := len(b.shdr_body)
	append(&b.shdr_body, u32(0))
	dxbc_encode_dest(&b.shdr_body, dest)
	dxbc_encode_src(&b.shdr_body, coord)
	dxbc_encode_src(&b.shdr_body, tex_reg)
	dxbc_encode_src(&b.shdr_body, samp_reg)
	dxbc_encode_src(&b.shdr_body, dref)
	length := u32(len(b.shdr_body) - start)
	b.shdr_body[start] = dxbc_opcode_token(DXBC_OP_SAMPLE_C, length)
	b.stat_inst_count += 1

	return dest
}

dxbc_emit_sample_level :: proc(b: ^DXBC_Builder, call: ^IR_Call, t: ^Resolved_Type) -> DXBC_Reg {
	// sample_level(combined_sampler, coord, lod) -> SAMPLE_L dest, coord, resource, sampler, lod
	if len(call.args) < 3 do return dxbc_alloc_temp(b, 4)

	coord := dxbc_emit_expr(b, call.args[1])
	lod   := dxbc_emit_expr(b, call.args[2])
	dest  := dxbc_alloc_temp(b, 4)

	tex_reg  := DXBC_Reg{kind = DXBC_OPERAND_RESOURCE, index = 0, components = 4, mask = DXBC_WRITEMASK_ALL}
	samp_reg := DXBC_Reg{kind = DXBC_OPERAND_SAMPLER,  index = 0, components = 4}

	if lb, ok := call.args[0].derived.(^IR_Load_Binding); ok {
		tex_name, samp_name := find_split_bindings(b.module, lb.name)
		if info, tok := b.binding_map[tex_name]; tok {
			tex_reg.index = info.reg_num
		}
		if info, sok := b.binding_map[samp_name]; sok {
			samp_reg.index = info.reg_num
		}
	}

	if coord.components < 4 {
		padded := dxbc_alloc_temp(b, 4)
		dxbc_emit_mov(b, padded, coord)
		coord = padded
	}

	start := len(b.shdr_body)
	append(&b.shdr_body, u32(0))
	dxbc_encode_dest(&b.shdr_body, dest)
	dxbc_encode_src(&b.shdr_body, coord)
	dxbc_encode_src(&b.shdr_body, tex_reg)
	dxbc_encode_src(&b.shdr_body, samp_reg)
	dxbc_encode_src(&b.shdr_body, lod)
	length := u32(len(b.shdr_body) - start)
	b.shdr_body[start] = dxbc_opcode_token(DXBC_OP_SAMPLE_L, length)
	b.stat_inst_count += 1

	return dest
}

dxbc_emit_refract :: proc(b: ^DXBC_Builder, call: ^IR_Call, t: ^Resolved_Type) -> DXBC_Reg {
	// refract(I, N, eta):
	//   k = 1 - eta^2 * (1 - dot(N,I)^2)
	//   if k < 0: result = 0
	//   else:     result = eta*I - (eta*dot(N,I) + sqrt(k))*N
	i_val    := dxbc_emit_expr(b, call.args[0])
	n_val    := dxbc_emit_expr(b, call.args[1])
	eta      := dxbc_emit_expr(b, call.args[2])
	in_comp  := dxbc_type_components(call.args[0].type)

	dp_op := u32(in_comp == 3 ? DXBC_OP_DP3 : DXBC_OP_DP4)

	// dot_ni = dot(N, I)
	dot_ni := dxbc_alloc_temp(b, 1)
	dxbc_emit_op(b, dp_op, dot_ni, n_val, i_val)

	// eta_sq = eta * eta
	eta_sq := dxbc_alloc_temp(b, 1)
	eta_s := dxbc_with_swizzle(eta, dxbc_swizzle_replicate(0))
	eta_s.components = 1
	dxbc_emit_op(b, DXBC_OP_MUL, eta_sq, eta_s, eta_s)

	// ni_sq = dot_ni * dot_ni
	ni_sq := dxbc_alloc_temp(b, 1)
	dxbc_emit_op(b, DXBC_OP_MUL, ni_sq, dot_ni, dot_ni)

	// k = 1 - eta_sq * (1 - ni_sq)
	//   = 1 - eta_sq + eta_sq * ni_sq
	one_minus_ni_sq := dxbc_alloc_temp(b, 1)
	neg_ni_sq := ni_sq
	neg_ni_sq.negate = true
	dxbc_emit_op(b, DXBC_OP_ADD, one_minus_ni_sq, dxbc_make_imm_f32(1), neg_ni_sq)

	k := dxbc_alloc_temp(b, 1)
	neg_eta_sq_term := dxbc_alloc_temp(b, 1)
	dxbc_emit_op(b, DXBC_OP_MUL, neg_eta_sq_term, eta_sq, one_minus_ni_sq)
	neg_eta_sq_term_neg := neg_eta_sq_term
	neg_eta_sq_term_neg.negate = true
	dxbc_emit_op(b, DXBC_OP_ADD, k, dxbc_make_imm_f32(1), neg_eta_sq_term_neg)

	// sqrt_k = sqrt(k)
	sqrt_k := dxbc_alloc_temp(b, 1)
	dxbc_emit_op(b, DXBC_OP_SQRT, sqrt_k, k)

	// coeff = eta * dot_ni + sqrt_k
	eta_dot := dxbc_alloc_temp(b, 1)
	dxbc_emit_op(b, DXBC_OP_MUL, eta_dot, eta_s, dot_ni)
	coeff := dxbc_alloc_temp(b, 1)
	dxbc_emit_op(b, DXBC_OP_ADD, coeff, eta_dot, sqrt_k)

	// result = eta * I - coeff * N
	coeff_bc := coeff
	coeff_bc.mask = dxbc_mask_from_count(in_comp)
	coeff_bc.components = in_comp
	eta_bc := eta_s
	eta_bc.mask = dxbc_mask_from_count(in_comp)
	eta_bc.components = in_comp

	eta_i := dxbc_alloc_temp(b, in_comp)
	dxbc_emit_op(b, DXBC_OP_MUL, eta_i, eta_bc, i_val)
	coeff_n := dxbc_alloc_temp(b, in_comp)
	dxbc_emit_op(b, DXBC_OP_MUL, coeff_n, coeff_bc, n_val)

	result := dxbc_alloc_temp(b, in_comp)
	neg_coeff_n := coeff_n
	neg_coeff_n.negate = true
	dxbc_emit_op(b, DXBC_OP_ADD, result, eta_i, neg_coeff_n)

	// If k < 0, result should be 0; use MOVC with (k >= 0) condition
	k_ge_zero := dxbc_alloc_temp(b, 1)
	dxbc_emit_op(b, DXBC_OP_GE, k_ge_zero, k, dxbc_make_imm_f32(0))
	k_ge_zero_bc := k_ge_zero
	k_ge_zero_bc.mask = dxbc_mask_from_count(in_comp)
	k_ge_zero_bc.components = in_comp
	dest := dxbc_alloc_temp(b, in_comp)
	dxbc_emit_op(b, DXBC_OP_MOVC, dest, k_ge_zero_bc, result, dxbc_make_imm_f32(0))
	return dest
}

dxbc_emit_transpose :: proc(b: ^DXBC_Builder, call: ^IR_Call, t: ^Resolved_Type) -> DXBC_Reg {
	// transpose(m): swap rows and columns
	// Only 4x4 is fully supported; other sizes emit a warning
	m := dxbc_emit_expr(b, call.args[0])

	mat_type, is_mat := call.args[0].type^.(Type_Matrix)
	if !is_mat {
		return m
	}

	cols := mat_type.cols
	rows := mat_type.rows

	if cols != 4 || rows != 4 {
		append(&b.diagnostics, Diagnostic{level = .Warning, message = fmt.aprintf("DXBC: transpose only fully supported for mat4; got mat%dx%d", cols, rows)})
	}

	if m.kind != DXBC_OPERAND_CONSTANT_BUFFER {
		append(&b.diagnostics, Diagnostic{level = .Warning, message = "DXBC: transpose only supported for cbuffer matrices; result may be incorrect"})
		return m
	}

	// For a cbuffer matrix, columns are at m.index2, m.index2+1, ...
	// transpose swaps: result column j = [col0[j], col1[j], col2[j], col3[j]]
	// i.e., result_col[j].x = original_col[0][j], etc.
	cb  := m.index
	base := m.index2

	// Allocate result columns as temp registers
	result_col := [4]DXBC_Reg{}
	for j in 0..<cols {
		result_col[j] = dxbc_alloc_temp(b, rows)
	}

	// For each output column j, fill component c from original column c, component j
	for j in 0..<cols {
		for c in 0..<rows {
			src_col := DXBC_Reg{
				kind       = DXBC_OPERAND_CONSTANT_BUFFER,
				index      = cb,
				index2     = base + u32(c),
				components = 1,
				swizzle    = dxbc_swizzle_replicate(u32(j)),
				has_swizzle = true,
			}
			dst := result_col[j]
			dst.mask = 1 << u32(c)
			dst.components = 1
			dxbc_emit_mov(b, dst, src_col)
		}
	}

	// Return the first column register; the matrix multiply code
	// will access subsequent columns by incrementing the index
	return result_col[0]
}

// Compute (stride_in_bytes, element_count) for a TGSM dcl_tgsm_structured
dxbc_tgsm_stride_count :: proc(t: ^Resolved_Type) -> (stride: u32, count: u32) {
	if t == nil do return 4, 1
	#partial switch v in t^ {
	case Type_Array_Resolved:
		elem_stride, _ := dxbc_tgsm_stride_count(v.elem)
		cnt := u32(v.size)
		if cnt == 0 do cnt = 1
		return elem_stride, cnt
	case Type_Scalar:
		return 4, 1
	case Type_Vector:
		return u32(v.size) * 4, 1
	case Type_Matrix:
		return u32(v.cols) * 16, 1
	case:
		return 16, 1
	}
}

// MOV to an indexable temp destination (static index)
dxbc_emit_mov_indexable :: proc(b: ^DXBC_Builder, dest: DXBC_Reg, src: DXBC_Reg) {
	// dest must be DXBC_OPERAND_INDEXABLE_TEMP with static index2
	start := len(b.shdr_body)
	append(&b.shdr_body, u32(0))
	dxbc_encode_dest(&b.shdr_body, dest)
	dxbc_encode_src(&b.shdr_body, src)
	length := u32(len(b.shdr_body) - start)
	b.shdr_body[start] = dxbc_opcode_token(DXBC_OP_MOV, length)
	b.stat_inst_count += 1
}

// MOV from an indexable temp source (static index) to dest
dxbc_emit_mov_indexable_src :: proc(b: ^DXBC_Builder, dest: DXBC_Reg, src: DXBC_Reg) {
	start := len(b.shdr_body)
	append(&b.shdr_body, u32(0))
	dxbc_encode_dest(&b.shdr_body, dest)
	dxbc_encode_src(&b.shdr_body, src)
	length := u32(len(b.shdr_body) - start)
	b.shdr_body[start] = dxbc_opcode_token(DXBC_OP_MOV, length)
	b.stat_inst_count += 1
}

// Emit a dynamically-indexed load from indexable temp xN[dynamic_idx]
// Uses xN[r_idx.x] encoding (DXBC_INDEX_RELATIVE)
dxbc_emit_indexed_load :: proc(b: ^DXBC_Builder, dest: DXBC_Reg, xreg: u32, idx_reg: DXBC_Reg, comp: int) {
	swz := dxbc_swizzle_from_mask(dxbc_mask_from_count(comp))
	// Operand: INDEXABLE_TEMP, 2D, dim0=IMM32 (xreg), dim1=RELATIVE (element idx)
	// token: COMP_4 | COMP_SWIZZLE | swz<<4 | INDEXABLE_TEMP<<12 | INDEX_2D<<20 | IMM32<<22 | RELATIVE<<25
	tok := dxbc_swizzle_operand(DXBC_OPERAND_INDEXABLE_TEMP, DXBC_INDEX_2D, swz) | dxbc_index_rep(DXBC_INDEX_IMM32, DXBC_INDEX_RELATIVE)
	start := len(b.shdr_body)
	append(&b.shdr_body, u32(0))
	dxbc_encode_dest(&b.shdr_body, dest)
	// Source operand for xN[idx_reg.x]
	append(&b.shdr_body, tok)
	append(&b.shdr_body, xreg)
	// Relative index: TEMP scalar
	idx_scalar := idx_reg
	idx_scalar.components = 1
	idx_scalar.mask = DXBC_WRITEMASK_X
	dxbc_encode_src(&b.shdr_body, idx_scalar)
	length := u32(len(b.shdr_body) - start)
	b.shdr_body[start] = dxbc_opcode_token(DXBC_OP_MOV, length)
	b.stat_inst_count += 1
}

// Dynamic vector component selection via MOVC chain (for vec2..vec4)
dxbc_emit_dynamic_vec_index :: proc(b: ^DXBC_Builder, vec: DXBC_Reg, idx: DXBC_Reg, comp: int) -> DXBC_Reg {
	vec_comp := vec.components
	if vec_comp <= 0 do vec_comp = 4
	dest := dxbc_alloc_temp(b, comp)
	// Start with component 0 as default
	comp0 := dxbc_with_swizzle(vec, dxbc_swizzle_replicate(0))
	comp0.components = comp
	dxbc_emit_mov(b, dest, comp0)
	// For each component i > 0, MOVC dest, (idx == i), vec.comp_i, dest
	for i in 1..<vec_comp {
		cmp := dxbc_alloc_temp(b, comp)
		dxbc_emit_op(b, DXBC_OP_IEQ, cmp, idx, dxbc_make_imm_i32(i32(i)))
		compi := dxbc_with_swizzle(vec, dxbc_swizzle_replicate(u32(i)))
		compi.components = comp
		dxbc_emit_op(b, DXBC_OP_MOVC, dest, cmp, compi, dest)
	}
	return dest
}

// Emit: ld_structured dest, 0, 0, gN
dxbc_emit_load_structured :: proc(b: ^DXBC_Builder, dest: DXBC_Reg, tgsm_idx: u32, comp: int) {
	start := len(b.shdr_body)
	append(&b.shdr_body, u32(0))
	dxbc_encode_dest(&b.shdr_body, dest)
	dxbc_encode_src(&b.shdr_body, dxbc_make_imm_u32(0)) // element address
	dxbc_encode_src(&b.shdr_body, dxbc_make_imm_u32(0)) // byte offset
	// Encode TGSM source operand
	swz := dxbc_swizzle_from_mask(dxbc_mask_from_count(comp))
	append(&b.shdr_body, dxbc_swizzle_operand(DXBC_OPERAND_THREAD_GROUP_SHARED_MEMORY, DXBC_INDEX_1D, swz) | dxbc_index_rep(DXBC_INDEX_IMM32))
	append(&b.shdr_body, tgsm_idx)
	length := u32(len(b.shdr_body) - start)
	b.shdr_body[start] = dxbc_opcode_token(DXBC_OP_LD_STRUCTURED, length)
	b.stat_inst_count += 1
}

// Emit: store_structured gN, 0, 0, src
dxbc_emit_store_structured :: proc(b: ^DXBC_Builder, tgsm_idx: u32, src: DXBC_Reg) {
	start := len(b.shdr_body)
	append(&b.shdr_body, u32(0))
	// TGSM destination operand
	append(&b.shdr_body, dxbc_mask_operand(DXBC_OPERAND_THREAD_GROUP_SHARED_MEMORY, DXBC_INDEX_1D, DXBC_WRITEMASK_ALL) | dxbc_index_rep(DXBC_INDEX_IMM32))
	append(&b.shdr_body, tgsm_idx)
	dxbc_encode_src(&b.shdr_body, dxbc_make_imm_u32(0)) // element address
	dxbc_encode_src(&b.shdr_body, dxbc_make_imm_u32(0)) // byte offset
	dxbc_encode_src(&b.shdr_body, src)
	length := u32(len(b.shdr_body) - start)
	b.shdr_body[start] = dxbc_opcode_token(DXBC_OP_STORE_STRUCTURED, length)
	b.stat_inst_count += 1
}

dxbc_emit_field_access :: proc(b: ^DXBC_Builder, fa: ^IR_Field_Access, t: ^Resolved_Type) -> DXBC_Reg {
	// Check if this is a cbuffer field access
	if lb, ok := fa.object.derived.(^IR_Load_Binding); ok {
		return dxbc_emit_cb_field(b, lb.name, fa.field_name, t)
	}
	// Otherwise, emit the object and extract field
	obj := dxbc_emit_expr(b, fa.object)
	// Need struct field index — look up from type
	if fa.object.type != nil {
		if s, ok := fa.object.type^.(Type_Struct_Resolved); ok {
			offset: u32 = 0
			for f in s.fields {
				fc := dxbc_type_components(f.type)
				if f.name == fa.field_name {
					comp := dxbc_type_components(t)
					dest := dxbc_alloc_temp(b, comp)
					// Swizzle to extract the right components
					src := dxbc_with_swizzle(obj, dxbc_mask_range(offset, u32(comp)))
					dxbc_emit_mov(b, dest, src)
					return dest
				}
				offset += u32(fc)
			}
		}
	}
	return obj
}

dxbc_emit_cb_field :: proc(b: ^DXBC_Builder, binding_name: string, field_name: string, t: ^Resolved_Type) -> DXBC_Reg {
	info, ok := b.binding_map[binding_name]
	if !ok do return dxbc_make_imm_f32(0)

	comp := dxbc_type_components(t)

	// Find the field in the layout
	for f in info.fields {
		if f.name == field_name {
			if f.is_matrix {
				return dxbc_load_cb_matrix(b, info.reg_num, f.vec4_index, f.mat_cols, f.mat_rows, t)
			}
			reg := DXBC_Reg{
				kind        = DXBC_OPERAND_CONSTANT_BUFFER,
				index       = info.reg_num,
				index2      = f.vec4_index,
				components  = comp,
				swizzle     = dxbc_mask_range(f.comp_offset, u32(comp)),
				has_swizzle = true,
			}
			return reg
		}
	}

	// Fallback: return first vec4 of the cbuffer
	return DXBC_Reg{
		kind       = DXBC_OPERAND_CONSTANT_BUFFER,
		index      = info.reg_num,
		index2     = 0,
		components = comp,
		mask       = dxbc_mask_from_count(comp),
	}
}

dxbc_load_cb_matrix :: proc(b: ^DXBC_Builder, cb_reg: u32, start_vec4: u32, cols: int, rows: int, t: ^Resolved_Type) -> DXBC_Reg {
	// Load a matrix from cbuffer as multiple vec4 registers
	// Store the start info in a temp and return it — the matrix multiply code
	// will reference the individual columns directly
	// For now, we store just the "base" reference
	return DXBC_Reg{
		kind       = DXBC_OPERAND_CONSTANT_BUFFER,
		index      = cb_reg,
		index2     = start_vec4,
		components = 4,
		mask       = DXBC_WRITEMASK_ALL,
	}
}

dxbc_emit_matrix_mul :: proc(b: ^DXBC_Builder, left_expr: ^IR_Expr, right_expr: ^IR_Expr, t: ^Resolved_Type) -> DXBC_Reg {
	comp := dxbc_type_components(t)
	dest := dxbc_alloc_temp(b, comp)

	left_is_mat := is_matrix(left_expr.type)
	right_is_mat := is_matrix(right_expr.type)

	if left_is_mat && !right_is_mat {
		// mat * vec: column-major multiplication
		// result = col0*v.x + col1*v.y + col2*v.z + col3*v.w
		vec := dxbc_emit_expr(b, right_expr)
		mat_reg := dxbc_emit_expr(b, left_expr)

		mat_type := left_expr.type^.(Type_Matrix)
		cols := mat_type.cols

		if mat_reg.kind == DXBC_OPERAND_CONSTANT_BUFFER {
			// Matrix comes from cbuffer — load columns directly
			cb := mat_reg.index
			base := mat_reg.index2

			// col0 * v.xxxx
			col0 := DXBC_Reg{kind = DXBC_OPERAND_CONSTANT_BUFFER, index = cb, index2 = base, components = 4, mask = DXBC_WRITEMASK_ALL}
			vec_x := dxbc_with_swizzle(vec, dxbc_swizzle_replicate(0))
			dxbc_emit_op(b, DXBC_OP_MUL, dest, col0, vec_x)

			for c in 1..<cols {
				col := DXBC_Reg{kind = DXBC_OPERAND_CONSTANT_BUFFER, index = cb, index2 = base + u32(c), components = 4, mask = DXBC_WRITEMASK_ALL}
				vec_c := dxbc_with_swizzle(vec, dxbc_swizzle_replicate(u32(c)))
				dxbc_emit_op(b, DXBC_OP_MAD, dest, col, vec_c, dest)
			}
		} else {
			// Matrix in temp registers — less common, simplified handling
			dxbc_emit_mov(b, dest, vec)
		}

		return dest
	}

	if !left_is_mat && right_is_mat {
		// vec * mat: row vector multiply
		// result.i = dot(vec, col_i)
		vec := dxbc_emit_expr(b, left_expr)
		mat_reg := dxbc_emit_expr(b, right_expr)

		if mat_reg.kind == DXBC_OPERAND_CONSTANT_BUFFER {
			mat_type := right_expr.type^.(Type_Matrix)
			cols := mat_type.cols
			cb := mat_reg.index
			base := mat_reg.index2

			for c in 0..<cols {
				col := DXBC_Reg{kind = DXBC_OPERAND_CONSTANT_BUFFER, index = cb, index2 = base + u32(c), components = 4, mask = DXBC_WRITEMASK_ALL}
				dest_c := dest
				dest_c.mask = 1 << u32(c)
				dxbc_emit_op(b, DXBC_OP_DP4, dest_c, vec, col)
			}
		}
		return dest
	}

	// mat * mat: result_col[j] = sum_k( left_col[k] * right_col[j][k] )
	// Both matrices must come from cbuffers for full support.
	left_mat  := left_expr.type^.(Type_Matrix)
	right_mat := right_expr.type^.(Type_Matrix)
	// left  is left_mat.cols x left_mat.rows  (cols columns, each rows-tall)
	// right is right_mat.cols x right_mat.rows
	// result is right_mat.cols x left_mat.rows
	left_reg  := dxbc_emit_expr(b, left_expr)
	right_reg := dxbc_emit_expr(b, right_expr)

	res_cols := right_mat.cols
	res_rows := left_mat.rows
	k_count  := left_mat.cols // == right_mat.rows

	// Allocate result columns (one temp per output column)
	result_first := b.next_temp
	for _ in 0..<res_cols {
		_ = dxbc_alloc_temp(b, res_rows)
	}

	// Helper to get column register
	get_col :: proc(reg: DXBC_Reg, col_idx: int, rows: int) -> DXBC_Reg {
		if reg.kind == DXBC_OPERAND_CONSTANT_BUFFER {
			return DXBC_Reg{
				kind       = DXBC_OPERAND_CONSTANT_BUFFER,
				index      = reg.index,
				index2     = reg.index2 + u32(col_idx),
				components = rows,
				mask       = dxbc_mask_from_count(rows),
			}
		}
		// Temp matrix: columns are consecutive temp registers
		return DXBC_Reg{
			kind       = DXBC_OPERAND_TEMP,
			index      = reg.index + u32(col_idx),
			components = rows,
			mask       = dxbc_mask_from_count(rows),
		}
	}

	for j in 0..<res_cols {
		res_reg := DXBC_Reg{
			kind       = DXBC_OPERAND_TEMP,
			index      = result_first + u32(j),
			components = res_rows,
			mask       = dxbc_mask_from_count(res_rows),
		}
		// First term: left_col[0] * right_col[j].xxxx
		left_col0  := get_col(left_reg,  0, res_rows)
		right_col_j := get_col(right_reg, j, k_count)
		right_x := dxbc_with_swizzle(right_col_j, dxbc_swizzle_replicate(0))
		right_x.components = res_rows
		dxbc_emit_op(b, DXBC_OP_MUL, res_reg, left_col0, right_x)

		for k in 1..<k_count {
			left_colk  := get_col(left_reg, k, res_rows)
			right_k    := dxbc_with_swizzle(right_col_j, dxbc_swizzle_replicate(u32(k)))
			right_k.components = res_rows
			dxbc_emit_op(b, DXBC_OP_MAD, res_reg, left_colk, right_k, res_reg)
		}
	}

	// Return first result column register
	return DXBC_Reg{
		kind       = DXBC_OPERAND_TEMP,
		index      = result_first,
		components = res_rows,
		mask       = dxbc_mask_from_count(res_rows),
	}
}

dxbc_emit_swizzle :: proc(b: ^DXBC_Builder, swz: ^IR_Swizzle, t: ^Resolved_Type) -> DXBC_Reg {
	obj := dxbc_emit_expr(b, swz.object)
	comp := len(swz.components)
	// Build swizzle from component string
	comps: [4]u32
	for c, i in swz.components {
		if i >= 4 do break
		switch c {
		case 'x', 'r': comps[i] = 0
		case 'y', 'g': comps[i] = 1
		case 'z', 'b': comps[i] = 2
		case 'w', 'a': comps[i] = 3
		}
	}
	// Fill remaining with last valid
	for i := comp; i < 4; i += 1 {
		comps[i] = comps[comp - 1]
	}
	result := dxbc_with_swizzle(obj, dxbc_swizzle(comps[0], comps[1], comps[2], comps[3]))
	result.components = comp
	return result
}

dxbc_emit_construct :: proc(b: ^DXBC_Builder, con: ^IR_Construct, t: ^Resolved_Type) -> DXBC_Reg {
	comp := dxbc_type_components(t)

	if len(con.args) == 1 {
		// Splat: vec4(x) -> mov dest.xyzw, x.xxxx
		arg := dxbc_emit_expr(b, con.args[0])
		if arg.components == 1 && comp > 1 {
			dest := dxbc_alloc_temp(b, comp)
			splat := dxbc_with_swizzle(arg, dxbc_swizzle_replicate(0))
			splat.components = comp
			dxbc_emit_mov(b, dest, splat)
			return dest
		}
		if arg.components == comp {
			return arg
		}
		dest := dxbc_alloc_temp(b, comp)
		dxbc_emit_mov(b, dest, arg)
		return dest
	}

	// Multi-arg constructor: mov each component
	dest := dxbc_alloc_temp(b, comp)
	dest_idx: u32 = 0
	for arg_expr in con.args {
		arg := dxbc_emit_expr(b, arg_expr)
		arg_comp := dxbc_type_components(arg_expr.type)
		for c in 0..<arg_comp {
			if dest_idx >= u32(comp) do break
			// MOV dest.{component}, arg.{component}
			d := dest
			d.mask = 1 << dest_idx
			s := arg
			if arg.kind == DXBC_OPERAND_IMMEDIATE32 {
				// For scalar immediates, just use as-is
			} else if arg_comp > 1 {
				s = dxbc_with_swizzle(s, dxbc_swizzle_replicate(u32(c)))
			}
			dxbc_emit_mov(b, d, s)
			dest_idx += 1
		}
	}

	return dest
}

dxbc_emit_type_cast :: proc(b: ^DXBC_Builder, tc: ^IR_Type_Cast, t: ^Resolved_Type) -> DXBC_Reg {
	src := dxbc_emit_expr(b, tc.value)
	comp := dxbc_type_components(t)
	dest := dxbc_alloc_temp(b, comp)

	src_float := is_float_type(tc.value.type)
	dst_float := is_float_type(t)
	dst_uint := is_uint_type(t)
	src_uint := is_uint_type(tc.value.type)

	if src_float && !dst_float {
		dxbc_emit_op(b, dst_uint ? DXBC_OP_FTOU : DXBC_OP_FTOI, dest, src)
	} else if !src_float && dst_float {
		dxbc_emit_op(b, src_uint ? DXBC_OP_UTOF : DXBC_OP_ITOF, dest, src)
	} else {
		dxbc_emit_mov(b, dest, src)
	}

	return dest
}

dxbc_emit_load_binding :: proc(b: ^DXBC_Builder, lb: ^IR_Load_Binding) -> DXBC_Reg {
	if info, ok := b.binding_map[lb.name]; ok {
		if info.kind == .Uniform || info.kind == .Push_Constant {
			return DXBC_Reg{
				kind       = DXBC_OPERAND_CONSTANT_BUFFER,
				index      = info.reg_num,
				index2     = 0,
				components = 4,
				mask       = DXBC_WRITEMASK_ALL,
			}
		}
	}
	return dxbc_make_imm_f32(0)
}

dxbc_emit_input_field :: proc(b: ^DXBC_Builder, inf: ^IR_Input_Field) -> DXBC_Reg {
	if reg, ok := b.input_map[inf.field_name]; ok {
		return reg
	}
	// Check builtins
	if reg, ok := b.builtin_input_map[inf.field_name]; ok {
		return reg
	}
	return dxbc_make_imm_f32(0)
}

dxbc_emit_composite_extract :: proc(b: ^DXBC_Builder, ce: ^IR_Composite_Extract, t: ^Resolved_Type) -> DXBC_Reg {
	obj := dxbc_emit_expr(b, ce.object)
	comp := dxbc_type_components(t)
	// Extract field by index — use swizzle
	result := obj
	if ce.index < 4 {
		start := u32(ce.index)
		result = dxbc_with_swizzle(result, dxbc_mask_range(start, u32(comp)))
		result.components = comp
	}
	return result
}

dxbc_emit_vector_shuffle :: proc(b: ^DXBC_Builder, vs: ^IR_Vector_Shuffle, t: ^Resolved_Type) -> DXBC_Reg {
	obj := dxbc_emit_expr(b, vs.object)
	comps: [4]u32
	for idx, i in vs.components {
		if i >= 4 do break
		comps[i] = u32(idx)
	}
	for i := len(vs.components); i < 4; i += 1 {
		comps[i] = comps[len(vs.components) - 1]
	}
	result := dxbc_with_swizzle(obj, dxbc_swizzle(comps[0], comps[1], comps[2], comps[3]))
	result.components = len(vs.components)
	return result
}

dxbc_mask_range :: proc(start: u32, count: u32) -> u32 {
	// For source operands, this creates a swizzle starting at component 'start'
	comps: [4]u32
	for i in u32(0)..<4 {
		if i < count {
			comps[i] = start + i
		} else {
			comps[i] = start + count - 1
		}
	}
	return dxbc_swizzle(comps[0], comps[1], comps[2], comps[3])
}

// ---- Chunk emission ----

dxbc_emit_isgn_chunk :: proc(b: ^DXBC_Builder, allocator := context.allocator) -> [dynamic]u8 {
	return dxbc_emit_signature_chunk(b.input_sigs[:], allocator)
}

dxbc_emit_osgn_chunk :: proc(b: ^DXBC_Builder, allocator := context.allocator) -> [dynamic]u8 {
	return dxbc_emit_signature_chunk(b.output_sigs[:], allocator)
}

dxbc_emit_signature_chunk :: proc(sigs: []DXBC_Sig_Element, allocator := context.allocator) -> [dynamic]u8 {
	buf := make([dynamic]u8, allocator)

	// Header: element count + magic 8
	dxbc_write_u32(&buf, u32(len(sigs)))
	dxbc_write_u32(&buf, 8) // constant

	// Calculate string table offset: after header (8 bytes) + elements (24 bytes each)
	str_offset_base := 8 + u32(len(sigs)) * 24

	// Build string table and element data
	str_table := make([dynamic]u8, allocator)
	str_offsets := make([dynamic]u32, allocator)

	for sig in sigs {
		offset := str_offset_base + u32(len(str_table))
		append(&str_offsets, offset)
		for c in sig.semantic_name {
			append(&str_table, u8(c))
		}
		append(&str_table, 0) // null terminator
		// Pad to 4-byte alignment (DXBC uses 0xAB for padding)
		for len(str_table) % 4 != 0 {
			append(&str_table, 0xAB)
		}
	}

	// Write elements
	for sig, i in sigs {
		dxbc_write_u32(&buf, str_offsets[i])
		dxbc_write_u32(&buf, sig.semantic_index)
		dxbc_write_u32(&buf, sig.system_value)
		dxbc_write_u32(&buf, sig.component_type)
		dxbc_write_u32(&buf, sig.register_num)
		// mask byte + rw_mask byte packed into u32
		dxbc_write_u32(&buf, u32(sig.mask) | (u32(sig.rw_mask) << 8))
	}

	// Append string table
	append(&buf, ..str_table[:])

	return buf
}

dxbc_emit_rdef_chunk :: proc(b: ^DXBC_Builder, fn: ^IR_Function, allocator := context.allocator) -> [dynamic]u8 {
	buf := make([dynamic]u8, allocator)

	// RDEF for SM5.0 with RD11 sub-header
	num_cb := b.num_cb
	num_resources := num_cb + b.num_textures + b.num_samplers

	// Collect string data (names) to emit after all structured data
	strings_buf := make([dynamic]u8, allocator)

	creator_str := "Luma Shader Compiler"

	// Compute layout offsets relative to start of RDEF data
	// Header: 28 bytes (7 DWORDs)
	// RD11 sub-header: 32 bytes (8 DWORDs)
	header_size :: 28
	rd11_size :: 32
	fixed_header := u32(header_size + rd11_size)

	// Cbuffer descriptors: 24 bytes each (name_offset, var_count, var_offset, size, flags, type)
	cb_desc_offset := fixed_header
	cb_desc_total := num_cb * 24

	// Resource binding descriptors: 32 bytes each (name, type, rettype, dim, samples, bind, count, flags)
	resource_desc_offset := cb_desc_offset + cb_desc_total
	resource_desc_total := num_resources * 32

	// Strings start after resource descriptors
	strings_offset := resource_desc_offset + resource_desc_total

	// Helper to add string and return its offset from RDEF start
	add_string :: proc(strings_buf: ^[dynamic]u8, base_offset: u32, s: string) -> u32 {
		off := base_offset + u32(len(strings_buf))
		for c in s {
			append(strings_buf, u8(c))
		}
		append(strings_buf, 0) // null terminator
		// Pad to 4-byte alignment (DXBC uses 0xAB for padding)
		for len(strings_buf) % 4 != 0 {
			append(strings_buf, 0xAB)
		}
		return off
	}

	// Pre-collect all string offsets
	creator_offset := add_string(&strings_buf, strings_offset, creator_str)

	// Collect cbuffer and resource names
	CB_Info :: struct { name: string, info: DXBC_Binding_Info, name_offset: u32 }
	Res_Info :: struct { name: string, info: DXBC_Binding_Info, name_offset: u32 }
	cb_list := make([dynamic]CB_Info, allocator)
	res_list := make([dynamic]Res_Info, allocator)

	for name, info in b.binding_map {
		if info.kind == .Uniform || info.kind == .Push_Constant {
			name_off := add_string(&strings_buf, strings_offset, name)
			append(&cb_list, CB_Info{name = name, info = info, name_offset = name_off})
		}
	}
	for name, info in b.binding_map {
		name_off := add_string(&strings_buf, strings_offset, name)
		append(&res_list, Res_Info{name = name, info = info, name_offset = name_off})
	}

	// Shader type for version field
	shader_type: u16
	#partial switch fn.stage {
	case .Vertex:   shader_type = 0xFFFE
	case .Fragment:  shader_type = 0xFFFF
	case .Compute:  shader_type = 0x4353
	case:           shader_type = 0xFFFE
	}

	// ---- Write RDEF header (28 bytes) ----
	dxbc_write_u32(&buf, num_cb)                           // cbuffer count
	dxbc_write_u32(&buf, num_cb > 0 ? cb_desc_offset : 0) // cbuffer desc offset
	dxbc_write_u32(&buf, num_resources)          // bound resource count
	dxbc_write_u32(&buf, resource_desc_offset)   // bound resource desc offset
	dxbc_write_u32(&buf, u32(shader_type) << 16 | u32(5) << 8) // version: shader_type, major=5, minor=0
	dxbc_write_u32(&buf, 0x00000100)             // compile flags (refactoring allowed = 0x100)
	dxbc_write_u32(&buf, creator_offset)         // creator string offset

	// ---- Write RD11 sub-header (32 bytes) ----
	dxbc_write_u32(&buf, 0x31314452) // "RD11" magic
	dxbc_write_u32(&buf, 60)  // header size (total RDEF header = 60 with RD11)
	dxbc_write_u32(&buf, 24)  // cbuffer desc stride
	dxbc_write_u32(&buf, 32)  // resource binding desc stride
	dxbc_write_u32(&buf, 40)  // shader variable desc stride
	dxbc_write_u32(&buf, 36)  // shader type desc stride
	dxbc_write_u32(&buf, 12)  // member desc stride
	dxbc_write_u32(&buf, 0)   // interface slot count

	// ---- Write cbuffer descriptors (24 bytes each) ----
	for cb in cb_list {
		dxbc_write_u32(&buf, cb.name_offset)      // name offset
		dxbc_write_u32(&buf, 0)                     // variable count (simplified)
		dxbc_write_u32(&buf, 0)                     // variable desc offset
		dxbc_write_u32(&buf, cb.info.size * 16)     // size in bytes
		dxbc_write_u32(&buf, 0)                     // flags
		dxbc_write_u32(&buf, 0)                     // type (0 = cbuffer)
	}

	// ---- Write resource binding descriptors (32 bytes each) ----
	for res in res_list {
		dxbc_write_u32(&buf, res.name_offset) // name offset
		switch res.info.kind {
		case .Uniform, .Push_Constant:
			dxbc_write_u32(&buf, 0)   // type = cbuffer
			dxbc_write_u32(&buf, 0)   // return type
			dxbc_write_u32(&buf, 0)   // dimension
			dxbc_write_u32(&buf, 0)   // num samples
			dxbc_write_u32(&buf, res.info.reg_num) // bind point
			dxbc_write_u32(&buf, 1)   // bind count
			dxbc_write_u32(&buf, 0)   // flags
		case .Texture:
			dxbc_write_u32(&buf, 2)   // type = texture
			dxbc_write_u32(&buf, 5)   // return type = float
			dxbc_write_u32(&buf, 3)   // dimension = texture2d
			dxbc_write_u32(&buf, 0xFFFFFFFF) // num samples
			dxbc_write_u32(&buf, res.info.reg_num)
			dxbc_write_u32(&buf, 1)
			dxbc_write_u32(&buf, 12)  // flags = D3D_SIF_TEXTURE_COMPONENT_1
		case .Sampler:
			dxbc_write_u32(&buf, 3)   // type = sampler
			dxbc_write_u32(&buf, 0)
			dxbc_write_u32(&buf, 0)
			dxbc_write_u32(&buf, 0)
			dxbc_write_u32(&buf, res.info.reg_num)
			dxbc_write_u32(&buf, 1)
			dxbc_write_u32(&buf, 0)
		case .Buffer:
			dxbc_write_u32(&buf, 7)   // type = structured buffer
			dxbc_write_u32(&buf, 0)
			dxbc_write_u32(&buf, 0)
			dxbc_write_u32(&buf, 0)
			dxbc_write_u32(&buf, res.info.reg_num)
			dxbc_write_u32(&buf, 1)
			dxbc_write_u32(&buf, 0)
		}
	}

	// ---- Append strings ----
	append(&buf, ..strings_buf[:])

	return buf
}

dxbc_emit_stat_chunk :: proc(b: ^DXBC_Builder, allocator := context.allocator) -> [dynamic]u8 {
	buf := make([dynamic]u8, allocator)
	// STAT chunk is a fixed-size structure. SM5.0 uses ~37 DWORDs.
	// We fill in the basics and zero the rest.
	dxbc_write_u32(&buf, b.stat_inst_count) // instruction count
	dxbc_write_u32(&buf, b.next_temp)       // temp register count
	// Remaining ~35 fields: zeroed
	for _ in 0..<35 {
		dxbc_write_u32(&buf, 0)
	}
	return buf
}

dxbc_emit_shex_chunk :: proc(b: ^DXBC_Builder, fn: ^IR_Function, allocator := context.allocator) -> [dynamic]u8 {
	buf := make([dynamic]u8, allocator)

	// Version token
	prog_type: u32
	#partial switch fn.stage {
	case .Vertex:   prog_type = DXBC_PROG_VERTEX
	case .Fragment:  prog_type = DXBC_PROG_PIXEL
	case .Compute:  prog_type = DXBC_PROG_COMPUTE
	case:           prog_type = DXBC_PROG_VERTEX
	}
	dxbc_write_u32(&buf, dxbc_version_token(prog_type))

	// Length placeholder (will be patched)
	length_offset := len(buf)
	dxbc_write_u32(&buf, 0)

	// Write declarations
	for word in b.shdr_decls {
		dxbc_write_u32(&buf, word)
	}

	// Write body
	for word in b.shdr_body {
		dxbc_write_u32(&buf, word)
	}

	// Patch length (in DWORDs, including version token and length field)
	total_dwords := u32(len(buf) / 4)
	buf[length_offset + 0] = u8(total_dwords)
	buf[length_offset + 1] = u8(total_dwords >> 8)
	buf[length_offset + 2] = u8(total_dwords >> 16)
	buf[length_offset + 3] = u8(total_dwords >> 24)

	return buf
}

// ---- Container assembly ----

dxbc_assemble :: proc(b: ^DXBC_Builder, fn: ^IR_Function, allocator := context.allocator) -> []u8 {
	// Build chunks
	rdef_data := dxbc_emit_rdef_chunk(b, fn, allocator)
	isgn_data := dxbc_emit_isgn_chunk(b, allocator)
	osgn_data := dxbc_emit_osgn_chunk(b, allocator)
	shex_data := dxbc_emit_shex_chunk(b, fn, allocator)
	stat_data := dxbc_emit_stat_chunk(b, allocator)

	Chunk :: struct {
		magic: u32,
		data:  [dynamic]u8,
	}

	chunks := []Chunk{
		{DXBC_CHUNK_RDEF, rdef_data},
		{DXBC_CHUNK_ISGN, isgn_data},
		{DXBC_CHUNK_OSGN, osgn_data},
		{DXBC_CHUNK_SHEX, shex_data},
		{DXBC_CHUNK_STAT, stat_data},
	}
	num_chunks := u32(len(chunks))

	// Header size: 4 (magic) + 16 (hash) + 4 (version) + 4 (total size) + 4 (chunk count) + num_chunks*4 (offsets)
	header_size := u32(32 + num_chunks * 4)

	// Calculate chunk offsets and total size
	offset := header_size
	chunk_offsets := make([dynamic]u32, allocator)
	for chunk in chunks {
		append(&chunk_offsets, offset)
		offset += 8 + u32(len(chunk.data)) // 8 = magic + size fields
	}
	total_size := offset

	// Build the container
	out := make([dynamic]u8, allocator)

	// Magic
	dxbc_write_u32(&out, DXBC_MAGIC)

	// MD5 hash placeholder (16 bytes)
	hash_offset := len(out)
	for _ in 0..<16 {
		append(&out, 0)
	}

	// Version
	dxbc_write_u32(&out, 1)

	// Total size
	dxbc_write_u32(&out, total_size)

	// Chunk count
	dxbc_write_u32(&out, num_chunks)

	// Chunk offsets
	for off in chunk_offsets {
		dxbc_write_u32(&out, off)
	}

	// Chunks
	for chunk in chunks {
		dxbc_write_u32(&out, chunk.magic)
		dxbc_write_u32(&out, u32(len(chunk.data)))
		append(&out, ..chunk.data[:])
	}

	// Compute MD5 hash of bytes 20..end
	hash := dxbc_hash(out[20:])
	for i in 0..<16 {
		out[hash_offset + i] = hash[i]
	}

	return out[:]
}

// ---- Utility ----

dxbc_write_u32 :: proc(buf: ^[dynamic]u8, val: u32) {
	append(buf, u8(val))
	append(buf, u8(val >> 8))
	append(buf, u8(val >> 16))
	append(buf, u8(val >> 24))
}

// ---- MD5 implementation (RFC 1321) ----

// DXBC uses a modified MD5 hash with custom padding (not standard RFC 1321).
// Based on the reverse-engineered algorithm from vkd3d-proton/GPUOpen.
dxbc_hash :: proc(data: []u8) -> [16]u8 {
	// Per-round shift amounts
	S := [64]u32{
		7,12,17,22, 7,12,17,22, 7,12,17,22, 7,12,17,22,
		5, 9,14,20, 5, 9,14,20, 5, 9,14,20, 5, 9,14,20,
		4,11,16,23, 4,11,16,23, 4,11,16,23, 4,11,16,23,
		6,10,15,21, 6,10,15,21, 6,10,15,21, 6,10,15,21,
	}

	// Pre-computed T constants: floor(2^32 * abs(sin(i+1)))
	K := [64]u32{
		0xd76aa478, 0xe8c7b756, 0x242070db, 0xc1bdceee,
		0xf57c0faf, 0x4787c62a, 0xa8304613, 0xfd469501,
		0x698098d8, 0x8b44f7af, 0xffff5bb1, 0x895cd7be,
		0x6b901122, 0xfd987193, 0xa679438e, 0x49b40821,
		0xf61e2562, 0xc040b340, 0x265e5a51, 0xe9b6c7aa,
		0xd62f105d, 0x02441453, 0xd8a1e681, 0xe7d3fbc8,
		0x21e1cde6, 0xc33707d6, 0xf4d50d87, 0x455a14ed,
		0xa9e3e905, 0xfcefa3f8, 0x676f02d9, 0x8d2a4c8a,
		0xfffa3942, 0x8771f681, 0x6d9d6122, 0xfde5380c,
		0xa4beea44, 0x4bdecfa9, 0xf6bb4b60, 0xbebfbc70,
		0x289b7ec6, 0xeaa127fa, 0xd4ef3085, 0x04881d05,
		0xd9d4d039, 0xe6db99e5, 0x1fa27cf8, 0xc4ac5665,
		0xf4292244, 0x432aff97, 0xab9423a7, 0xfc93a039,
		0x655b59c3, 0x8f0ccc92, 0xffeff47d, 0x85845dd1,
		0x6fa87e4f, 0xfe2ce6e0, 0xa3014314, 0x4e0811a1,
		0xf7537e82, 0xbd3af235, 0x2ad7d2bb, 0xeb86d391,
	}

	// MD5 transform: process a 64-byte block
	md5_transform :: proc(a0, b0, c0, d0: ^u32, block: [16]u32, S_: [64]u32, K_: [64]u32) {
		A, B, C, D := a0^, b0^, c0^, d0^
		for i in u32(0)..<64 {
			F, g: u32
			if i < 16 {
				F = (B & C) | (~B & D)
				g = i
			} else if i < 32 {
				F = (D & B) | (~D & C)
				g = (5 * i + 1) % 16
			} else if i < 48 {
				F = B ~ C ~ D
				g = (3 * i + 5) % 16
			} else {
				F = C ~ (B | ~D)
				g = (7 * i) % 16
			}
			F = F + A + K_[i] + block[g]
			A = D
			D = C
			C = B
			rot := S_[i]
			B = B + ((F << rot) | (F >> (32 - rot)))
		}
		a0^ += A
		b0^ += B
		c0^ += C
		d0^ += D
	}

	// Read a little-endian u32 from a byte slice
	read_le32 :: proc(d: []u8, off: int) -> u32 {
		return u32(d[off]) | (u32(d[off+1]) << 8) | (u32(d[off+2]) << 16) | (u32(d[off+3]) << 24)
	}

	// Initialize MD5 state
	a0: u32 = 0x67452301
	b0: u32 = 0xefcdab89
	c0: u32 = 0x98badcfe
	d0: u32 = 0x10325476

	length := len(data)
	num_bits := u32(length) * 8
	num_bits2 := (num_bits >> 2) | 1

	leftover_length := length % 64
	full_blocks := length - leftover_length

	// Process full 64-byte blocks
	for block_start := 0; block_start < full_blocks; block_start += 64 {
		M: [16]u32
		for i in 0..<16 {
			j := block_start + i * 4
			M[i] = read_le32(data, j)
		}
		md5_transform(&a0, &b0, &c0, &d0, M, S, K)
	}

	// Custom DXBC finalization (differs from standard MD5)
	leftover := data[full_blocks:]

	if leftover_length >= 56 {
		// Process the leftover data
		M: [16]u32
		for i in 0..<leftover_length/4 {
			M[i] = read_le32(leftover, i * 4)
		}
		// Handle partial last word
		partial := leftover_length % 4
		if partial > 0 {
			word_idx := leftover_length / 4
			val: u32
			for j in 0..<partial {
				val |= u32(leftover[word_idx * 4 + j]) << u32(j * 8)
			}
			M[word_idx] = val
		}
		md5_transform(&a0, &b0, &c0, &d0, M, S, K)

		// Pad block: 0x80 at start, num_bits2 at word[15]
		P: [16]u32
		P[0] = 0x80
		P[15] = num_bits2
		md5_transform(&a0, &b0, &c0, &d0, P, S, K)

		// Additional block with num_bits
		Q: [16]u32
		Q[0] = num_bits
		Q[15] = num_bits2
		md5_transform(&a0, &b0, &c0, &d0, Q, S, K)
	} else {
		// Build a single padded block:
		// [num_bits (4 bytes)] [leftover data] [0x80] [zeros...] [num_bits2 at end]
		// All fed as one or two MD5_Update calls
		block: [16]u32

		// Start building a byte buffer for the final block
		buf: [64]u8

		// First 4 bytes: num_bits as LE u32
		buf[0] = u8(num_bits)
		buf[1] = u8(num_bits >> 8)
		buf[2] = u8(num_bits >> 16)
		buf[3] = u8(num_bits >> 24)

		// Copy leftover data
		for i in 0..<leftover_length {
			buf[4 + i] = leftover[i]
		}

		// 0x80 byte after data
		buf[4 + leftover_length] = 0x80

		// The rest is zeros (already zeroed)

		// Total used: 4 + leftover_length + 1 + padding
		// Total size must be 64 bytes
		// num_bits2 goes at the last 4 bytes (offset 60)
		buf[60] = u8(num_bits2)
		buf[61] = u8(num_bits2 >> 8)
		buf[62] = u8(num_bits2 >> 16)
		buf[63] = u8(num_bits2 >> 24)

		// Convert to u32 block
		for i in 0..<16 {
			block[i] = u32(buf[i*4]) | (u32(buf[i*4+1]) << 8) | (u32(buf[i*4+2]) << 16) | (u32(buf[i*4+3]) << 24)
		}
		md5_transform(&a0, &b0, &c0, &d0, block, S, K)
	}

	// Output raw state (no standard MD5 finalization)
	result: [16]u8
	result[0]  = u8(a0);       result[1]  = u8(a0 >> 8);  result[2]  = u8(a0 >> 16); result[3]  = u8(a0 >> 24)
	result[4]  = u8(b0);       result[5]  = u8(b0 >> 8);  result[6]  = u8(b0 >> 16); result[7]  = u8(b0 >> 24)
	result[8]  = u8(c0);       result[9]  = u8(c0 >> 8);  result[10] = u8(c0 >> 16); result[11] = u8(c0 >> 24)
	result[12] = u8(d0);       result[13] = u8(d0 >> 8);  result[14] = u8(d0 >> 16); result[15] = u8(d0 >> 24)
	return result
}