use alloc::vec::Vec;
use pliron::{
builtin::{
op_interfaces::{OneRegionInterface, SymbolOpInterface},
ops::{ConstantOp as BuiltinConstantOp, FuncOp as BuiltinFuncOp, ModuleOp},
type_interfaces::FunctionTypeInterface,
types::{FunctionType as BuiltinFunctionType, UnitType},
},
common_traits::Verify,
context::{Context, Ptr},
derive::{op_interface_impl, type_interface_impl},
input_err_noloc, input_error_noloc,
irbuild::{
dialect_conversion::{self, DialectConversion, DialectConversionRewriter, OperandsInfo},
inserter::Inserter,
rewriter::Rewriter,
},
op::{Op, op_impls},
operation::Operation,
pass::{GuardedPass, OpGuard, OpPass, Pass, PassResult},
region::Region,
result::{Error, ErrorKind, Result},
r#type::{TypeHandle, TypedHandle, type_cast},
};
use crate::{
ToLLVMDialect, ToLLVMType,
ops::{ConstantOp as LLVMConstantOp, FuncOp as LLVMFuncOp},
types::{FuncType as LLVMFuncType, VoidType},
};
#[derive(thiserror::Error, Debug)]
pub enum BuiltinToLLVMConversionError {
#[error("Invalid function type, cannot be converted to LLVM function type")]
InvalidFunctionType,
}
#[type_interface_impl]
impl ToLLVMType for BuiltinFunctionType {
fn convert(&self, ctx: &Context) -> Result<TypeHandle> {
let arg_types = self.arg_types();
let res_types = self.res_types();
let convert_type_to_llvm = |ty: TypeHandle| {
type_cast::<dyn ToLLVMType>(&*ty.deref(ctx))
.map(|ty| ty.convert(ctx))
.unwrap_or(Ok(ty))
};
let arg_types = arg_types
.into_iter()
.map(convert_type_to_llvm)
.collect::<Result<Vec<_>>>()?;
let res_types = res_types
.into_iter()
.map(convert_type_to_llvm)
.collect::<Result<Vec<_>>>()?;
if res_types.is_empty() || res_types.len() > 1 {
return input_err_noloc!(BuiltinToLLVMConversionError::InvalidFunctionType);
}
let result_type = res_types[0];
let llvm_func_type = LLVMFuncType::get(ctx, result_type, arg_types, false);
Ok(llvm_func_type.into())
}
}
#[type_interface_impl]
impl ToLLVMType for UnitType {
fn convert(&self, ctx: &Context) -> Result<TypeHandle> {
Ok(VoidType::get(ctx).to_handle())
}
}
#[op_interface_impl]
impl ToLLVMDialect for BuiltinConstantOp {
fn rewrite(
&self,
ctx: &mut Context,
rewriter: &mut DialectConversionRewriter,
_operands_info: &OperandsInfo,
) -> Result<()> {
let const_value = self.get_value(ctx);
let llvm_const = LLVMConstantOp::new(ctx, const_value);
if let Err(e @ Error { .. }) = llvm_const.verify(ctx) {
return Err(Error {
kind: ErrorKind::InvalidInput,
backtrace: pliron::std_deps::backtrace::Backtrace::capture(),
..e
});
}
rewriter.insert_operation(ctx, llvm_const.get_operation());
let old_op = self.get_operation();
rewriter.replace_operation(ctx, old_op, llvm_const.get_operation());
Ok(())
}
}
#[op_interface_impl]
impl ToLLVMDialect for BuiltinFuncOp {
fn rewrite(
&self,
ctx: &mut Context,
rewriter: &mut DialectConversionRewriter,
_operands_info: &OperandsInfo,
) -> Result<()> {
let func_name = self.get_symbol_name(ctx);
let builtin_func_type = self.get_type(ctx);
let llvm_func_type = type_cast::<dyn ToLLVMType>(&*builtin_func_type.deref(ctx))
.ok_or_else(|| {
input_error_noloc!("builtin.func type does not implement ToLLVMType interface")
})?
.convert(ctx)?;
let llvm_func_type = TypedHandle::from_handle(llvm_func_type, ctx)?;
let llvm_func = LLVMFuncOp::new(ctx, func_name, llvm_func_type);
Region::move_to_op(self.get_region(ctx), llvm_func.get_operation(), ctx);
let old_op = self.get_operation();
rewriter.insert_operation(ctx, llvm_func.get_operation());
rewriter.replace_operation(ctx, old_op, llvm_func.get_operation());
Ok(())
}
}
#[derive(Default)]
pub struct BuiltinToLLVMConversion;
impl DialectConversion for BuiltinToLLVMConversion {
fn can_convert_op(&self, ctx: &Context, op: Ptr<Operation>) -> bool {
let op_dyn = Operation::get_op_dyn(op, ctx);
let op_ref = op_dyn.op_ref();
op_impls::<dyn ToLLVMDialect>(op_ref)
}
fn rewrite(
&mut self,
ctx: &mut Context,
rewriter: &mut DialectConversionRewriter,
op: Ptr<Operation>,
operands_info: &OperandsInfo,
) -> Result<()> {
let op_dyn = Operation::get_op_dyn(op, ctx);
let op_ref = op_dyn.op_ref();
if let Some(to_llvm) = pliron::op::op_cast::<dyn ToLLVMDialect>(op_ref) {
to_llvm.rewrite(ctx, rewriter, operands_info)?;
}
Ok(())
}
}
pub fn convert_builtin_to_llvm(ctx: &mut Context, module: ModuleOp) -> Result<PassResult> {
builtin_to_llvm_pass().run(
module.get_operation(),
ctx,
&mut pliron::pass::AnalysisManager::default(),
)
}
pub fn builtin_to_llvm_pass()
-> OpPass<ModuleOp, dialect_conversion::PassWrapper<BuiltinToLLVMConversion>> {
let pass = dialect_conversion::PassWrapper::<BuiltinToLLVMConversion>::new(
"builtin_to_llvm",
BuiltinToLLVMConversion,
);
GuardedPass::new(OpGuard::<ModuleOp>::default(), pass)
}