use core::{
cell::{Ref, RefCell, RefMut},
ops::{Deref, DerefMut},
};
use alloc::{
boxed::Box,
string::{String, ToString},
vec::Vec,
};
use downcast_rs::{Downcast, impl_downcast};
use thiserror::Error;
use crate::{
arg_error_noloc,
context::{Context, Ptr},
identifier::Identifier,
irbuild::IRStatus,
op::{Op, OpInterfaceMarker, op_impls},
operation::{OpDbg, Operation, verify_operation},
printable::Printable,
result::Result,
std_deps::{
self,
fs::{create_dir_all, write},
path::PathBuf,
},
utils::{
table::{HMap, HSet, IMap},
timer::Timer,
},
};
#[derive(Default)]
pub struct PassResult {
pub ir_changed: IRStatus,
preserved_analyses: HSet<core::any::TypeId>,
}
impl PassResult {
pub fn set_preserved<A: Analysis + 'static>(&mut self) {
self.preserved_analyses.insert(core::any::TypeId::of::<A>());
}
}
pub trait Pass {
fn name(&self) -> &str;
fn run(
&mut self,
op: Ptr<Operation>,
ctx: &mut Context,
analyses: &mut AnalysisManager,
) -> Result<PassResult>;
fn as_pass_manager(&mut self) -> Option<&mut dyn PassManager> {
None
}
}
#[derive(Default)]
pub struct Passes {
passes: Vec<Box<dyn Pass>>,
}
impl Pass for Passes {
fn name(&self) -> &str {
"passes"
}
fn run(
&mut self,
op: Ptr<Operation>,
ctx: &mut Context,
analyses: &mut AnalysisManager,
) -> Result<PassResult> {
let mut pass_res = PassResult::default();
for pass in &mut self.passes {
let res = <Self as PassManager>::run_pass(&mut **pass, op, ctx, analyses)?;
pass_res.ir_changed |= res.ir_changed;
analyses.retain_preserved(&res);
}
let preserved_analyses = analyses.list_analyses();
pass_res.preserved_analyses = preserved_analyses;
Ok(pass_res)
}
fn as_pass_manager(&mut self) -> Option<&mut dyn PassManager> {
Some(self)
}
}
impl Passes {
pub fn add_pass(&mut self, pass: impl Pass + 'static) {
self.passes.push(Box::new(pass));
}
}
impl PassManager for Passes {}
pub struct NestedOpsPass {
pass: Box<dyn Pass>,
}
impl Pass for NestedOpsPass {
fn name(&self) -> &str {
"nested_ops_pass"
}
fn run(
&mut self,
op: Ptr<Operation>,
ctx: &mut Context,
analyses: &mut AnalysisManager,
) -> Result<PassResult> {
use crate::linked_list::ContainsLinkedList;
let mut pass_res = PassResult::default();
let regions = op.deref(ctx).regions().collect::<Vec<_>>();
for region in regions {
let blocks = region.deref(ctx).iter(ctx).collect::<Vec<_>>();
for block in blocks {
let ops = block.deref(ctx).iter(ctx).collect::<Vec<_>>();
for nested_op in ops {
let res =
<Self as PassManager>::run_pass(&mut *self.pass, nested_op, ctx, analyses)?;
pass_res.ir_changed |= res.ir_changed;
analyses.retain_preserved(&res);
}
}
}
let preserved_analyses = analyses.list_analyses();
pass_res.preserved_analyses = preserved_analyses;
Ok(pass_res)
}
fn as_pass_manager(&mut self) -> Option<&mut dyn PassManager> {
Some(self)
}
}
impl NestedOpsPass {
pub fn new(pass: impl Pass + 'static) -> Self {
Self {
pass: Box::new(pass),
}
}
}
impl PassManager for NestedOpsPass {}
pub trait Guard {
fn is_allowed(&self, op: Ptr<Operation>, ctx: &Context) -> bool;
}
pub struct OpGuard<T: Op> {
_marker: core::marker::PhantomData<T>,
}
impl<T: Op> Default for OpGuard<T> {
fn default() -> Self {
Self {
_marker: core::marker::PhantomData,
}
}
}
impl<T: Op> Guard for OpGuard<T> {
fn is_allowed(&self, op: Ptr<Operation>, ctx: &Context) -> bool {
Operation::is_op::<T>(op, ctx)
}
}
pub struct OpInterfaceGuard<T: ?Sized + OpInterfaceMarker + 'static> {
_marker: core::marker::PhantomData<T>,
}
impl<T: ?Sized + OpInterfaceMarker + 'static> Default for OpInterfaceGuard<T> {
fn default() -> Self {
Self {
_marker: core::marker::PhantomData,
}
}
}
impl<T: ?Sized + OpInterfaceMarker + 'static> Guard for OpInterfaceGuard<T> {
fn is_allowed(&self, op: Ptr<Operation>, ctx: &Context) -> bool {
let op = Operation::get_op_dyn(op, ctx);
op_impls::<T>(&*op)
}
}
#[derive(Default)]
pub struct GuardedPass<G: Guard, P: Pass> {
guard: G,
pass: P,
}
impl<G: Guard, P: Pass> GuardedPass<G, P> {
pub fn new(guard: G, pass: P) -> Self {
Self { guard, pass }
}
}
impl<G: Guard, P: Pass> PassManager for GuardedPass<G, P> {}
impl<G: Guard, P: Pass> Pass for GuardedPass<G, P> {
fn name(&self) -> &str {
"guarded_pass"
}
fn run(
&mut self,
op: Ptr<Operation>,
ctx: &mut Context,
analyses: &mut AnalysisManager,
) -> Result<PassResult> {
if self.guard.is_allowed(op, ctx) {
<Self as PassManager>::run_pass(&mut self.pass, op, ctx, analyses)
} else {
Ok(PassResult::default())
}
}
fn as_pass_manager(&mut self) -> Option<&mut dyn PassManager> {
Some(self)
}
}
impl<G: Guard, P: Pass> Deref for GuardedPass<G, P> {
type Target = P;
fn deref(&self) -> &Self::Target {
&self.pass
}
}
impl<G: Guard, P: Pass> DerefMut for GuardedPass<G, P> {
fn deref_mut(&mut self) -> &mut Self::Target {
&mut self.pass
}
}
pub type OpPass<T, P> = GuardedPass<OpGuard<T>, P>;
pub type OpInterfacePass<T, P> = GuardedPass<OpInterfaceGuard<T>, P>;
pub trait PassManager {
fn run_pass(
pass: &mut dyn Pass,
op: Ptr<Operation>,
ctx: &mut Context,
analyses: &mut AnalysisManager,
) -> Result<PassResult>
where
Self: Sized,
{
let is_pass_manager = pass.as_pass_manager().is_some();
let config = analyses.pm_data().config();
let pass_run_count = analyses.pm_data().state().pass_run_count;
let skip_pass = !is_pass_manager && config.skip_passes.contains(pass.name());
let pre_print_pass = !is_pass_manager
&& (config.print_before_all || config.print_before.contains(pass.name()));
let post_print_pass = !is_pass_manager
&& (config.print_after_all || config.print_after.contains(pass.name()));
let pre_verify_pass = !is_pass_manager
&& (config.verify_before_all || config.verify_before.contains(pass.name()));
let post_verify_pass = !is_pass_manager
&& (config.verify_after_all || config.verify_after.contains(pass.name()));
let should_time = !is_pass_manager
&& (config.time_all_passes || config.time_passes.contains(pass.name()));
let ir_printing_dir = config.ir_printing_dir.clone();
if skip_pass {
log::debug!("Skipping pass {} on {}", pass.name(), OpDbg { op, ctx });
return Ok(PassResult::default());
}
if !is_pass_manager {
log::debug!("Running pass {} on {}", pass.name(), OpDbg { op, ctx });
}
if pre_print_pass {
log::info!("IR before pass {}:\n{}", pass.name(), op.disp(ctx));
if let Some(dir) = &ir_printing_dir {
let filename = alloc::format!("{}-before-{}.plir", pass_run_count, pass.name());
print_op_to_file(ctx, dir, filename, op)?;
}
}
if pre_verify_pass {
verify_operation(op, ctx).inspect_err(|e| {
log::error!(
"Verification failed before pass {} on {}:\n{}",
pass.name(),
OpDbg { op, ctx },
e.disp(ctx)
);
})?;
}
let timer = Timer::start();
let result = pass.run(op, ctx, analyses);
if should_time {
let elapsed = timer.elapsed();
log::info!(
"Pass {} on {} completed in {:?}",
pass.name(),
OpDbg { op, ctx },
elapsed
);
}
if post_print_pass {
log::info!("IR after pass {}:\n{}", pass.name(), op.disp(ctx));
if let Some(dir) = &ir_printing_dir {
let filename = alloc::format!("{}-after-{}.plir", pass_run_count, pass.name());
print_op_to_file(ctx, dir, filename, op)?;
}
}
if post_verify_pass {
verify_operation(op, ctx).inspect_err(|e| {
log::error!(
"Verification failed after pass {} on {}:\n{}",
pass.name(),
OpDbg { op, ctx },
e.disp(ctx)
);
})?;
}
if !is_pass_manager {
analyses.pm_data_mut().state_mut().pass_run_count += 1;
}
result
}
}
#[derive(Default)]
pub struct PMConfig {
pub print_before_all: bool,
pub print_after_all: bool,
pub ir_printing_dir: Option<PathBuf>,
pub print_before: HSet<String>,
pub print_after: HSet<String>,
pub verify_before_all: bool,
pub verify_after_all: bool,
pub verify_before: HSet<String>,
pub verify_after: HSet<String>,
pub time_all_passes: bool,
pub time_passes: HSet<String>,
pub skip_passes: HSet<String>,
pub custom_config: HMap<Identifier, Box<dyn core::any::Any>>,
}
#[derive(Default)]
pub struct PMState {
pub stats: IMap<&'static str, Box<dyn Printable>>,
pub custom_state: HMap<Identifier, Box<dyn core::any::Any>>,
pub pass_run_count: usize,
}
#[derive(Default)]
pub struct PMData {
config: PMConfig,
state: PMState,
}
impl PMData {
pub fn config(&self) -> &PMConfig {
&self.config
}
pub fn set_config(&mut self, config: PMConfig) {
self.config = config;
}
pub fn state(&self) -> &PMState {
&self.state
}
pub fn state_mut(&mut self) -> &mut PMState {
&mut self.state
}
}
pub trait Analysis: Downcast {
fn name(&self) -> &str;
fn compute(op: Ptr<Operation>, ctx: &Context, analyses: &mut AnalysisManager) -> Result<Self>
where
Self: Sized;
}
impl_downcast!(Analysis);
type AnalysisManagerKey = (core::any::TypeId, Ptr<Operation>);
#[derive(Default)]
pub struct AnalysisManager {
pub pm_data: PMData,
analyses: IMap<AnalysisManagerKey, Box<RefCell<dyn Analysis>>>,
}
impl AnalysisManager {
pub fn compute_analysis<A: Analysis + 'static>(
&mut self,
op: Ptr<Operation>,
ctx: &Context,
) -> Result<()> {
let key = (core::any::TypeId::of::<A>(), op);
if !self.analyses.contains_key(&key) {
let analysis = A::compute(op, ctx, self)?;
self.analyses.insert(key, Box::new(RefCell::new(analysis)));
}
Ok(())
}
pub fn get_analysis_mut<'a, A: Analysis + 'static>(
&'a mut self,
op: Ptr<Operation>,
ctx: &Context,
) -> Result<RefMut<'a, A>> {
self.compute_analysis::<A>(op, ctx)?;
let key = (core::any::TypeId::of::<A>(), op);
let analysis = self.analyses.get(&key).unwrap();
Ok(RefMut::map(analysis.borrow_mut(), |a| {
a.downcast_mut::<A>().unwrap()
}))
}
pub fn get_analysis<'a, A: Analysis + 'static>(
&'a mut self,
op: Ptr<Operation>,
ctx: &Context,
) -> Result<Ref<'a, A>> {
self.compute_analysis::<A>(op, ctx)?;
let key = (core::any::TypeId::of::<A>(), op);
let analysis = self.analyses.get(&key).unwrap();
Ok(Ref::map(analysis.borrow(), |a| {
a.downcast_ref::<A>().unwrap()
}))
}
pub fn try_get_analysis<'a, A: Analysis + 'static>(
&'a self,
op: Ptr<Operation>,
) -> Option<Ref<'a, A>> {
let key = (core::any::TypeId::of::<A>(), op);
self.analyses
.get(&key)
.map(|analysis| Ref::map(analysis.borrow(), |a| a.downcast_ref::<A>().unwrap()))
}
pub fn try_get_analysis_mut<'a, A: Analysis + 'static>(
&'a self,
op: Ptr<Operation>,
) -> Option<RefMut<'a, A>> {
let key = (core::any::TypeId::of::<A>(), op);
self.analyses
.get(&key)
.map(|analysis| RefMut::map(analysis.borrow_mut(), |a| a.downcast_mut::<A>().unwrap()))
}
pub fn retain_preserved(&mut self, pass_res: &PassResult) {
if pass_res.ir_changed == IRStatus::Unchanged {
return;
}
self.analyses
.retain(|(type_id, _), _| pass_res.preserved_analyses.contains(type_id));
}
fn list_analyses(&self) -> HSet<core::any::TypeId> {
self.analyses.keys().map(|(type_id, _)| *type_id).collect()
}
pub fn set_config(&mut self, config: PMConfig) {
self.pm_data.set_config(config);
}
pub fn pm_data(&self) -> &PMData {
&self.pm_data
}
pub fn pm_data_mut(&mut self) -> &mut PMData {
&mut self.pm_data
}
}
#[derive(Debug, Error)]
pub enum PrintOpToFileErr {
#[error("Failed to write to file {}: {}", .0.display(), .1)]
FileWriteError(std_deps::path::PathBuf, std_deps::io::Error),
#[error("Failed to create directory {}: {}", .0.display(), .1)]
DirCreateError(std_deps::path::PathBuf, std_deps::io::Error),
}
fn print_op_to_file(
ctx: &Context,
dir: &PathBuf,
file_name: String,
op: Ptr<Operation>,
) -> Result<()> {
create_dir_all(dir)
.map_err(|err| arg_error_noloc!(PrintOpToFileErr::DirCreateError(dir.clone(), err)))?;
let path = dir.join(file_name);
write(&path, op.disp(ctx).to_string().as_bytes())
.map_err(|err| arg_error_noloc!(PrintOpToFileErr::FileWriteError(path, err)))
}