use super::CodeGen;
use super::context::RegisterContext;
use super::types::{RegisterType, SpecSignature};
use crate::ast::{Statement, WordDef};
use crate::codegen::CodeGenError;
use crate::codegen::mangle_name;
use std::fmt::Write as _;
use std::ops::{Deref, DerefMut};
pub(super) struct SpecializedEmitter<'a> {
codegen: &'a mut CodeGen,
word_name: &'a str,
sig: &'a SpecSignature,
}
impl<'a> SpecializedEmitter<'a> {
pub(super) fn new(
codegen: &'a mut CodeGen,
word_name: &'a str,
sig: &'a SpecSignature,
) -> Self {
Self {
codegen,
word_name,
sig,
}
}
pub(super) fn word_name(&self) -> &str {
self.word_name
}
pub(super) fn sig(&self) -> &SpecSignature {
self.sig
}
}
impl Deref for SpecializedEmitter<'_> {
type Target = CodeGen;
fn deref(&self) -> &CodeGen {
self.codegen
}
}
impl DerefMut for SpecializedEmitter<'_> {
fn deref_mut(&mut self) -> &mut CodeGen {
self.codegen
}
}
impl CodeGen {
pub fn codegen_specialized_word(
&mut self,
word: &WordDef,
sig: &SpecSignature,
) -> Result<(), CodeGenError> {
SpecializedEmitter::new(self, &word.name, sig).emit_word(word)
}
}
impl SpecializedEmitter<'_> {
fn emit_word(&mut self, word: &WordDef) -> Result<(), CodeGenError> {
let base_name = format!("seq_{}", mangle_name(self.word_name));
let spec_name = format!("{}{}", base_name, self.sig.suffix());
let return_type = if self.sig.outputs.len() == 1 {
self.sig.outputs[0].llvm_type().to_string()
} else {
let types: Vec<_> = self.sig.outputs.iter().map(|t| t.llvm_type()).collect();
format!("{{ {} }}", types.join(", "))
};
let params: Vec<String> = self
.sig
.inputs
.iter()
.enumerate()
.map(|(i, ty)| format!("{} %arg{}", ty.llvm_type(), i))
.collect();
writeln!(
&mut self.output,
"define {} @{}({}) {{",
return_type,
spec_name,
params.join(", ")
)?;
writeln!(&mut self.output, "entry:")?;
let initial_params: Vec<(String, RegisterType)> = self
.sig
.inputs
.iter()
.enumerate()
.map(|(i, ty)| (format!("arg{}", i), *ty))
.collect();
let mut ctx = RegisterContext::from_params(&initial_params);
let body_len = word.body.len();
let mut prev_int_literal: Option<i64> = None;
for (i, stmt) in word.body.iter().enumerate() {
let is_last = i == body_len - 1;
self.emit_statement(&mut ctx, stmt, is_last, &mut prev_int_literal)?;
}
writeln!(&mut self.output, "}}")?;
writeln!(&mut self.output)?;
let key = self.word_name.to_string();
let sig = self.sig.clone();
self.specialized_words.insert(key, sig);
Ok(())
}
pub(super) fn emit_statement(
&mut self,
ctx: &mut RegisterContext,
stmt: &Statement,
is_last: bool,
prev_int_literal: &mut Option<i64>,
) -> Result<(), CodeGenError> {
let prev_int = *prev_int_literal;
*prev_int_literal = None;
match stmt {
Statement::IntLiteral(n) => {
let var = self.fresh_temp();
writeln!(&mut self.output, " %{} = add i64 0, {}", var, n)?;
ctx.push(var, RegisterType::I64);
*prev_int_literal = Some(*n); }
Statement::FloatLiteral(f) => {
let var = self.fresh_temp();
let bits = f.to_bits();
writeln!(
&mut self.output,
" %{} = bitcast i64 {} to double",
var, bits
)?;
ctx.push(var, RegisterType::Double);
}
Statement::BoolLiteral(b) => {
let var = self.fresh_temp();
let val = if *b { 1 } else { 0 };
writeln!(&mut self.output, " %{} = add i64 0, {}", var, val)?;
ctx.push(var, RegisterType::I64);
}
Statement::WordCall { name, .. } => {
self.emit_word_call(ctx, name, is_last, prev_int)?;
}
Statement::If {
then_branch,
else_branch,
span: _,
} => {
self.emit_if(ctx, then_branch, else_branch.as_deref(), is_last)?;
}
Statement::StringLiteral(_)
| Statement::Symbol(_)
| Statement::Quotation { .. }
| Statement::Match { .. } => {
return Err(CodeGenError::Logic(format!(
"Non-specializable statement in specialized word: {:?}",
stmt
)));
}
}
let already_returns = match stmt {
Statement::If { .. } => true,
Statement::WordCall { name, .. } if name == self.word_name => true,
_ => false,
};
if is_last && !already_returns {
self.emit_return(ctx)?;
}
Ok(())
}
pub(super) fn emit_return(&mut self, ctx: &RegisterContext) -> Result<(), CodeGenError> {
let output_count = self.sig.outputs.len();
if output_count == 0 {
writeln!(&mut self.output, " ret void")?;
} else if output_count == 1 {
let (var, ty) = ctx
.values
.last()
.ok_or_else(|| CodeGenError::Logic("Empty context at return".to_string()))?;
writeln!(&mut self.output, " ret {} %{}", ty.llvm_type(), var)?;
} else {
if ctx.values.len() < output_count {
return Err(CodeGenError::Logic(format!(
"Not enough values for multi-output return: need {}, have {}",
output_count,
ctx.values.len()
)));
}
let start_idx = ctx.values.len() - output_count;
let return_values: Vec<_> = ctx.values[start_idx..].to_vec();
let struct_type = self.sig.llvm_return_type();
let mut current_struct = "undef".to_string();
for (i, (var, ty)) in return_values.iter().enumerate() {
let new_struct = self.fresh_temp();
writeln!(
&mut self.output,
" %{} = insertvalue {} {}, {} %{}, {}",
new_struct,
struct_type,
current_struct,
ty.llvm_type(),
var,
i
)?;
current_struct = format!("%{}", new_struct);
}
writeln!(&mut self.output, " ret {} {}", struct_type, current_struct)?;
}
Ok(())
}
fn emit_branch(
&mut self,
parent_ctx: &RegisterContext,
branch: &[Statement],
branch_label: &str,
merge_label: &str,
is_last: bool,
) -> Result<(RegisterContext, Option<String>), CodeGenError> {
writeln!(&mut self.output, "{}:", branch_label)?;
let mut branch_ctx = parent_ctx.clone();
let mut branch_prev_int: Option<i64> = None;
for (i, stmt) in branch.iter().enumerate() {
let is_stmt_last = i == branch.len() - 1 && is_last;
self.emit_statement(&mut branch_ctx, stmt, is_stmt_last, &mut branch_prev_int)?;
}
if is_last && branch.is_empty() {
self.emit_return(&branch_ctx)?;
}
let predecessor = if is_last {
None
} else {
writeln!(&mut self.output, " br label %{}", merge_label)?;
Some(branch_label.to_string())
};
Ok((branch_ctx, predecessor))
}
pub(super) fn emit_if(
&mut self,
ctx: &mut RegisterContext,
then_branch: &[Statement],
else_branch: Option<&[Statement]>,
is_last: bool,
) -> Result<(), CodeGenError> {
let (cond_var, _) = ctx
.pop()
.ok_or_else(|| CodeGenError::Logic("Empty context at if condition".to_string()))?;
let cmp_result = self.fresh_temp();
writeln!(
&mut self.output,
" %{} = icmp ne i64 %{}, 0",
cmp_result, cond_var
)?;
let then_label = self.fresh_block("if_then");
let else_label = self.fresh_block("if_else");
let merge_label = self.fresh_block("if_merge");
writeln!(
&mut self.output,
" br i1 %{}, label %{}, label %{}",
cmp_result, then_label, else_label
)?;
let (then_ctx, then_pred) =
self.emit_branch(ctx, then_branch, &then_label, &merge_label, is_last)?;
let else_slice: &[Statement] = else_branch.unwrap_or(&[]);
let (else_ctx, else_pred) =
self.emit_branch(ctx, else_slice, &else_label, &merge_label, is_last)?;
if then_pred.is_some() || else_pred.is_some() {
writeln!(&mut self.output, "{}:", merge_label)?;
if let (Some(then_p), Some(else_p)) = (&then_pred, &else_pred) {
if then_ctx.values.len() != else_ctx.values.len() {
return Err(CodeGenError::Logic(format!(
"Stack depth mismatch in if branches: then has {}, else has {}",
then_ctx.values.len(),
else_ctx.values.len()
)));
}
ctx.values.clear();
for i in 0..then_ctx.values.len() {
let (then_var, then_ty) = &then_ctx.values[i];
let (else_var, else_ty) = &else_ctx.values[i];
if then_ty != else_ty {
return Err(CodeGenError::Logic(format!(
"Type mismatch at position {} in if branches: {:?} vs {:?}",
i, then_ty, else_ty
)));
}
if then_var == else_var {
ctx.push(then_var.clone(), *then_ty);
} else {
let phi_result = self.fresh_temp();
writeln!(
&mut self.output,
" %{} = phi {} [ %{}, %{} ], [ %{}, %{} ]",
phi_result,
then_ty.llvm_type(),
then_var,
then_p,
else_var,
else_p
)?;
ctx.push(phi_result, *then_ty);
}
}
} else if then_pred.is_some() {
*ctx = then_ctx;
} else {
*ctx = else_ctx;
}
if is_last && (then_pred.is_some() || else_pred.is_some()) {
self.emit_return(ctx)?;
}
}
Ok(())
}
}