use super::{BranchResult, CodeGen, CodeGenError, TailPosition};
use crate::ast::{MatchArm, Pattern, Statement};
use std::fmt::Write as _;
pub(super) struct ArmInfo {
variant_name: String,
block: String,
field_count: usize,
field_names: Vec<String>,
}
impl CodeGen {
fn emit_dup(&mut self, stack_var: &str) -> Result<String, CodeGenError> {
let dup_stack = self.fresh_temp();
writeln!(
&mut self.output,
" %{} = call ptr @patch_seq_dup(ptr %{})",
dup_stack, stack_var
)?;
Ok(dup_stack)
}
pub(super) fn codegen_if_statement(
&mut self,
stack_var: &str,
then_branch: &[Statement],
else_branch: Option<&Vec<Statement>>,
position: TailPosition,
) -> Result<String, CodeGenError> {
let stack_var = self.spill_virtual_stack(stack_var)?;
let stack_var = stack_var.as_str();
let top_ptr = self.emit_stack_gep(stack_var, -1)?;
let cond_val = self.emit_load_int_payload(&top_ptr)?;
let popped_stack = top_ptr.clone();
let cmp_temp = self.fresh_temp();
writeln!(
&mut self.output,
" %{} = icmp ne i64 %{}, 0",
cmp_temp, cond_val
)?;
let then_block = self.fresh_block("if_then");
let else_block = self.fresh_block("if_else");
let merge_block = self.fresh_block("if_merge");
writeln!(
&mut self.output,
" br i1 %{}, label %{}, label %{}",
cmp_temp, then_block, else_block
)?;
writeln!(&mut self.output, "{}:", then_block)?;
let then_result = self.codegen_branch(
then_branch,
&popped_stack,
position,
&merge_block,
"if_then",
)?;
writeln!(&mut self.output, "{}:", else_block)?;
let else_result = if let Some(eb) = else_branch {
self.codegen_branch(eb, &popped_stack, position, &merge_block, "if_else")?
} else {
let else_pred = self.fresh_block("if_else_end");
writeln!(&mut self.output, " br label %{}", else_pred)?;
writeln!(&mut self.output, "{}:", else_pred)?;
writeln!(&mut self.output, " br label %{}", merge_block)?;
BranchResult {
stack_var: popped_stack.clone(),
emitted_tail_call: false,
predecessor: else_pred,
}
};
if then_result.emitted_tail_call && else_result.emitted_tail_call {
return Ok(then_result.stack_var);
}
writeln!(&mut self.output, "{}:", merge_block)?;
let result_var = self.fresh_temp();
if then_result.emitted_tail_call {
writeln!(
&mut self.output,
" %{} = phi ptr [ %{}, %{} ]",
result_var, else_result.stack_var, else_result.predecessor
)?;
} else if else_result.emitted_tail_call {
writeln!(
&mut self.output,
" %{} = phi ptr [ %{}, %{} ]",
result_var, then_result.stack_var, then_result.predecessor
)?;
} else {
writeln!(
&mut self.output,
" %{} = phi ptr [ %{}, %{} ], [ %{}, %{} ]",
result_var,
then_result.stack_var,
then_result.predecessor,
else_result.stack_var,
else_result.predecessor
)?;
}
Ok(result_var)
}
pub(super) fn codegen_match_statement(
&mut self,
stack_var: &str,
arms: &[MatchArm],
position: TailPosition,
) -> Result<String, CodeGenError> {
let stack_var = self.spill_virtual_stack(stack_var)?;
let stack_var = stack_var.as_str();
if let Err((union_name, missing)) = self.check_match_exhaustiveness(arms) {
return Err(CodeGenError::Logic(format!(
"Non-exhaustive match on union '{}'. Missing variants: {}",
union_name,
missing.join(", ")
)));
}
let dup_stack = self.emit_dup(stack_var)?;
let tagged_stack = self.fresh_temp();
writeln!(
&mut self.output,
" %{} = call ptr @patch_seq_variant_tag(ptr %{})",
tagged_stack, dup_stack
)?;
let default_block = self.fresh_block("match_unreachable");
let merge_block = self.fresh_block("match_merge");
let mut arm_info: Vec<ArmInfo> = Vec::new();
for (i, arm) in arms.iter().enumerate() {
let block = self.fresh_block(&format!("match_arm_{}", i));
let variant_name = match &arm.pattern {
Pattern::Variant(name) => name.clone(),
Pattern::VariantWithBindings { name, .. } => name.clone(),
};
let (_tag, field_count, field_names) = self.find_variant_info(&variant_name)?;
arm_info.push(ArmInfo {
variant_name,
block,
field_count,
field_names,
});
}
self.codegen_match_dispatch(&tagged_stack, &arm_info, &default_block)?;
let arm_results =
self.codegen_match_arms(stack_var, arms, &arm_info, position, &merge_block)?;
self.codegen_match_merge(&arm_results, &merge_block)
}
pub(super) fn codegen_extract_variant_bindings(
&mut self,
stack_var: &str,
bindings: &[String],
field_names: &[String],
) -> Result<String, CodeGenError> {
if bindings.is_empty() {
let drop_stack = self.fresh_temp();
writeln!(
&mut self.output,
" %{} = call ptr @patch_seq_drop_op(ptr %{})",
drop_stack, stack_var
)?;
return Ok(drop_stack);
}
let mut current_stack = stack_var.to_string();
let last_idx = bindings.len() - 1;
for (bind_idx, binding) in bindings.iter().enumerate() {
let field_idx = field_names
.iter()
.position(|f| f == binding)
.expect("binding validation should have caught unknown field");
current_stack = if bind_idx != last_idx {
self.codegen_extract_field_middle(¤t_stack, field_idx)?
} else {
self.codegen_extract_field_last(¤t_stack, field_idx)?
};
}
Ok(current_stack)
}
pub(super) fn codegen_extract_field_middle(
&mut self,
stack_var: &str,
field_idx: usize,
) -> Result<String, CodeGenError> {
let dup_stack = self.emit_dup(stack_var)?;
let idx_stack = self.fresh_temp();
writeln!(
&mut self.output,
" %{} = call ptr @patch_seq_push_int(ptr %{}, i64 {})",
idx_stack, dup_stack, field_idx
)?;
let field_stack = self.fresh_temp();
writeln!(
&mut self.output,
" %{} = call ptr @patch_seq_variant_field_at(ptr %{})",
field_stack, idx_stack
)?;
let swap_stack = self.fresh_temp();
writeln!(
&mut self.output,
" %{} = call ptr @patch_seq_swap(ptr %{})",
swap_stack, field_stack
)?;
Ok(swap_stack)
}
pub(super) fn codegen_extract_field_last(
&mut self,
stack_var: &str,
field_idx: usize,
) -> Result<String, CodeGenError> {
let idx_stack = self.fresh_temp();
writeln!(
&mut self.output,
" %{} = call ptr @patch_seq_push_int(ptr %{}, i64 {})",
idx_stack, stack_var, field_idx
)?;
let field_stack = self.fresh_temp();
writeln!(
&mut self.output,
" %{} = call ptr @patch_seq_variant_field_at(ptr %{})",
field_stack, idx_stack
)?;
Ok(field_stack)
}
pub(super) fn codegen_match_dispatch(
&mut self,
tagged_stack: &str,
arm_info: &[ArmInfo],
default_block: &str,
) -> Result<(), CodeGenError> {
let mut current_tag_stack = tagged_stack.to_string();
for (i, arm) in arm_info.iter().enumerate() {
let is_last = i == arm_info.len() - 1;
let next_check = if is_last {
default_block.to_string()
} else {
self.fresh_block(&format!("match_check_{}", i + 1))
};
let compare_stack = if !is_last {
self.emit_dup(¤t_tag_stack)?
} else {
current_tag_stack.clone()
};
let str_const = self.get_string_global(arm.variant_name.as_bytes())?;
let cmp_stack = self.fresh_temp();
writeln!(
&mut self.output,
" %{} = call ptr @patch_seq_symbol_eq_cstr(ptr %{}, ptr {})",
cmp_stack, compare_stack, str_const
)?;
let cmp_val = self.fresh_temp();
writeln!(
&mut self.output,
" %{} = call i1 @patch_seq_peek_bool_value(ptr %{})",
cmp_val, cmp_stack
)?;
let popped = self.fresh_temp();
writeln!(
&mut self.output,
" %{} = call ptr @patch_seq_pop_stack(ptr %{})",
popped, cmp_stack
)?;
writeln!(
&mut self.output,
" br i1 %{}, label %{}, label %{}",
cmp_val, arm.block, next_check
)?;
if !is_last {
writeln!(&mut self.output, "{}:", next_check)?;
current_tag_stack = popped;
}
}
writeln!(&mut self.output, "{}:", default_block)?;
writeln!(&mut self.output, " unreachable")?;
Ok(())
}
pub(super) fn codegen_match_arms(
&mut self,
stack_var: &str,
arms: &[MatchArm],
arm_info: &[ArmInfo],
position: TailPosition,
merge_block: &str,
) -> Result<Vec<BranchResult>, CodeGenError> {
let mut arm_results: Vec<BranchResult> = Vec::new();
for (i, (arm, info)) in arms.iter().zip(arm_info.iter()).enumerate() {
writeln!(&mut self.output, "{}:", info.block)?;
let unpacked_stack = match &arm.pattern {
Pattern::Variant(_) => {
let result = self.fresh_temp();
writeln!(
&mut self.output,
" %{} = call ptr @patch_seq_unpack_variant(ptr %{}, i64 {})",
result, stack_var, info.field_count
)?;
result
}
Pattern::VariantWithBindings { bindings, .. } => {
self.codegen_extract_variant_bindings(stack_var, bindings, &info.field_names)?
}
};
let result = self.codegen_branch(
&arm.body,
&unpacked_stack,
position,
merge_block,
&format!("match_arm_{}", i),
)?;
arm_results.push(result);
}
Ok(arm_results)
}
pub(super) fn codegen_match_merge(
&mut self,
arm_results: &[BranchResult],
merge_block: &str,
) -> Result<String, CodeGenError> {
let all_tail_calls = arm_results.iter().all(|r| r.emitted_tail_call);
if all_tail_calls {
return Ok(arm_results[0].stack_var.clone());
}
writeln!(&mut self.output, "{}:", merge_block)?;
let result_var = self.fresh_temp();
let phi_entries: Vec<_> = arm_results
.iter()
.filter(|r| !r.emitted_tail_call)
.map(|r| format!("[ %{}, %{} ]", r.stack_var, r.predecessor))
.collect();
if phi_entries.is_empty() {
return Err(CodeGenError::Logic(
"Match codegen: unexpected empty phi".to_string(),
));
}
writeln!(
&mut self.output,
" %{} = phi ptr {}",
result_var,
phi_entries.join(", ")
)?;
Ok(result_var)
}
}