use alloc::collections::VecDeque;
use crate::{
context::{Context, Ptr},
graph::walkers::{IRNode, WalkConfig, uninterruptible::immutable::walk_op},
irbuild::{
IRStatus,
inserter::{Inserter, OpInsertionPoint},
listener::{Recorder, RecorderEvent},
rewriter::{IRRewriter, Rewriter},
},
operation::Operation,
pass::{AnalysisManager, Pass, PassResult},
result::Result,
utils::table::HSet,
};
pub type MatchRewriter = IRRewriter<Recorder>;
pub trait MatchRewrite {
fn r#match(&mut self, ctx: &Context, op: Ptr<Operation>) -> bool;
fn rewrite(
&mut self,
ctx: &mut Context,
rewriter: &mut MatchRewriter,
op: Ptr<Operation>,
) -> Result<()>;
}
#[derive(Clone, Copy, Debug, Default)]
pub enum EnqueueOrder {
EnqueFront,
#[default]
EnqueBack,
}
#[derive(Clone, Debug, Default)]
pub struct RewriterOrder {
pub collect: WalkConfig,
pub enque: EnqueueOrder,
}
pub fn apply_match_rewrite<M: MatchRewrite>(
ctx: &mut Context,
match_rewrite: &mut M,
order: RewriterOrder,
op: Ptr<Operation>,
) -> Result<IRStatus> {
let mut to_rewrite = VecDeque::new();
struct WalkerState<'a, M> {
match_rewrite: &'a mut M,
to_rewrite: &'a mut VecDeque<Ptr<Operation>>,
}
let mut state = WalkerState {
match_rewrite,
to_rewrite: &mut to_rewrite,
};
fn walker_callback<M: MatchRewrite>(ctx: &Context, state: &mut WalkerState<M>, node: IRNode) {
if let IRNode::Operation(op) = node
&& state.match_rewrite.r#match(ctx, op)
{
state.to_rewrite.push_back(op);
}
}
walk_op(ctx, &mut state, &order.collect, op, walker_callback);
let mut erased = HSet::<Ptr<Operation>>::default();
let mut rewriter = MatchRewriter::default();
rewriter.set_listener(Recorder::default());
while !to_rewrite.is_empty() {
let op = to_rewrite.pop_front().unwrap();
if erased.contains(&op) {
continue;
}
rewriter.set_insertion_point(OpInsertionPoint::BeforeOperation(op));
match_rewrite.rewrite(ctx, &mut rewriter, op)?;
let listener = rewriter.get_listener_mut();
for event in &listener.events {
if let RecorderEvent::ErasedOperation(erased_op) = event {
erased.insert(*erased_op);
}
}
for event in &listener.events {
match event {
RecorderEvent::ErasedOperation(_) => {
}
RecorderEvent::InsertedOperation(new_op) => {
if !erased.contains(new_op) && match_rewrite.r#match(ctx, *new_op) {
match order.enque {
EnqueueOrder::EnqueFront => to_rewrite.push_front(*new_op),
EnqueueOrder::EnqueBack => to_rewrite.push_back(*new_op),
}
}
}
RecorderEvent::ReplacedValueUses { .. } => {
}
RecorderEvent::InsertedBlock(_) => {
}
RecorderEvent::ErasedBlock(_) => {
}
RecorderEvent::ErasedRegion(_) => {
}
RecorderEvent::ValueTypeChanged { .. } => {
}
RecorderEvent::UnlinkedOperation(_op, _prev_position) => {
}
RecorderEvent::UnlinkedBlock(_block, _prev_position) => {
}
}
}
listener.clear();
}
Ok(rewriter.is_modified().into())
}
pub struct PassWrapper<M: MatchRewrite> {
name: &'static str,
match_rewrite: M,
rewrite_order: RewriterOrder,
}
impl<M: MatchRewrite> PassWrapper<M> {
pub fn new(name: &'static str, match_rewrite: M) -> Self {
Self {
name,
match_rewrite,
rewrite_order: RewriterOrder::default(),
}
}
pub fn set_rewrite_order(mut self, order: RewriterOrder) -> Self {
self.rewrite_order = order;
self
}
}
impl<M: MatchRewrite> Pass for PassWrapper<M> {
fn run(
&mut self,
op: Ptr<Operation>,
ctx: &mut Context,
_analyses: &mut AnalysisManager,
) -> Result<PassResult> {
let mut pass_result = PassResult::default();
pass_result.ir_changed |=
apply_match_rewrite(ctx, &mut self.match_rewrite, self.rewrite_order.clone(), op)?;
Ok(pass_result)
}
fn name(&self) -> &str {
self.name
}
}