use alloc::vec::Vec;
use crate::{
basic_block::BasicBlock,
context::{Context, Ptr},
irbuild::{
inserter::{BlockInsertionPoint, Inserter, OpInsertionPoint},
rewriter::{Rewriter, ScopedRewriter},
},
linked_list::ContainsLinkedList,
location::Located,
operation::Operation,
region::Region,
r#type::Typed,
utils::table::HMap,
value::Value,
};
#[derive(Debug, Default)]
pub struct IrMapping {
values: HMap<Value, Value>,
blocks: HMap<Ptr<BasicBlock>, Ptr<BasicBlock>>,
ops: HMap<Ptr<Operation>, Ptr<Operation>>,
}
impl IrMapping {
pub fn new() -> Self {
Self::default()
}
pub fn map_value(&mut self, from: Value, to: Value) {
self.values.insert(from, to);
}
pub fn map_block(&mut self, from: Ptr<BasicBlock>, to: Ptr<BasicBlock>) {
self.blocks.insert(from, to);
}
pub fn map_op(&mut self, from: Ptr<Operation>, to: Ptr<Operation>) {
self.ops.insert(from, to);
}
pub fn lookup_value(&self, from: Value) -> Option<Value> {
self.values.get(&from).copied()
}
pub fn lookup_block(&self, from: Ptr<BasicBlock>) -> Option<Ptr<BasicBlock>> {
self.blocks.get(&from).copied()
}
pub fn lookup_op(&self, from: Ptr<Operation>) -> Option<Ptr<Operation>> {
self.ops.get(&from).copied()
}
pub fn lookup_value_or_default(&self, from: Value) -> Value {
self.lookup_value(from).unwrap_or(from)
}
pub fn lookup_block_or_default(&self, from: Ptr<BasicBlock>) -> Ptr<BasicBlock> {
self.lookup_block(from).unwrap_or(from)
}
}
pub fn clone_operation(
op: Ptr<Operation>,
ctx: &mut Context,
rewriter: &mut dyn Rewriter,
mapper: &mut IrMapping,
) -> Ptr<Operation> {
let new_op = clone_op_shell(op, ctx, mapper);
fill_operation(op, ctx, rewriter, mapper);
new_op
}
fn clone_op_shell(op: Ptr<Operation>, ctx: &mut Context, mapper: &mut IrMapping) -> Ptr<Operation> {
let (concrete_op, result_types, successors, num_operands, num_regions, attributes, loc) = {
let op_ref = op.deref(ctx);
let successors: Vec<Ptr<BasicBlock>> = op_ref
.successors()
.map(|b| mapper.lookup_block_or_default(b))
.collect();
(
op_ref.concrete_op_info(),
op_ref.result_types().collect::<Vec<_>>(),
successors,
op_ref.get_num_operands(),
op_ref.num_regions(),
op_ref.attributes.clone(),
op_ref.loc(),
)
};
let new_op = Operation::new(
ctx,
concrete_op,
result_types,
Vec::with_capacity(num_operands),
successors,
num_regions,
);
{
let mut new_ref = new_op.deref_mut(ctx);
new_ref.attributes = attributes;
new_ref.set_loc(loc);
}
mapper.map_op(op, new_op);
let old_ref = op.deref(ctx);
let new_ref = new_op.deref(ctx);
for (old, new) in old_ref.results().zip(new_ref.results()) {
mapper.map_value(old, new);
}
new_op
}
fn fill_operation(
op: Ptr<Operation>,
ctx: &mut Context,
rewriter: &mut dyn Rewriter,
mapper: &mut IrMapping,
) {
let new_op = mapper
.lookup_op(op)
.expect("op shell must be created before it is filled");
let operands: Vec<Value> = op
.deref(ctx)
.operands()
.map(|v| mapper.lookup_value_or_default(v))
.collect();
for operand in operands {
Operation::push_operand(new_op, ctx, operand);
}
let num_regions = op.deref(ctx).num_regions();
for region_idx in 0..num_regions {
let src_region = op.deref(ctx).get_region(region_idx);
let dest_region = new_op.deref(ctx).get_region(region_idx);
clone_region_into(src_region, dest_region, ctx, rewriter, mapper);
}
}
pub fn clone_region_into(
src_region: Ptr<Region>,
dest_region: Ptr<Region>,
ctx: &mut Context,
rewriter: &mut dyn Rewriter,
mapper: &mut IrMapping,
) {
let blocks: Vec<Ptr<BasicBlock>> = src_region.deref(ctx).iter(ctx).collect();
clone_blocks_into(&blocks, dest_region, ctx, rewriter, mapper);
}
pub fn clone_blocks_into(
blocks: &[Ptr<BasicBlock>],
dest_region: Ptr<Region>,
ctx: &mut Context,
rewriter: &mut dyn Rewriter,
mapper: &mut IrMapping,
) {
let mut rewriter = ScopedRewriter::new(rewriter, OpInsertionPoint::Unset);
for &src_block in blocks {
let (arg_types, label, attrs, loc) = {
let block_ref = src_block.deref(ctx);
let arg_types: Vec<_> = block_ref.arguments().map(|arg| arg.get_type(ctx)).collect();
(
arg_types,
block_ref.label.clone(),
block_ref.attributes.clone(),
block_ref.loc(),
)
};
let new_block = rewriter.create_block(
ctx,
BlockInsertionPoint::AtRegionEnd(dest_region),
label,
arg_types,
);
{
let mut new_ref = new_block.deref_mut(ctx);
new_ref.attributes = attrs;
new_ref.set_loc(loc);
}
let src_ref = src_block.deref(ctx);
let new_ref = new_block.deref(ctx);
for (old, new) in src_ref.arguments().zip(new_ref.arguments()) {
mapper.map_value(old, new);
}
mapper.map_block(src_block, new_block);
}
for &src_block in blocks {
let new_block = mapper
.lookup_block(src_block)
.expect("block was mapped in phase one");
rewriter.set_insertion_point(OpInsertionPoint::AtBlockEnd(new_block));
let ops: Vec<Ptr<Operation>> = src_block.deref(ctx).iter(ctx).collect();
for src_op in ops {
let shell = clone_op_shell(src_op, ctx, mapper);
rewriter.append_operation(ctx, shell);
}
}
for &src_block in blocks {
let ops: Vec<Ptr<Operation>> = src_block.deref(ctx).iter(ctx).collect();
for src_op in ops {
fill_operation(src_op, ctx, &mut rewriter, mapper);
}
}
}