use std::collections::HashMap;
use crate::{
analysis::{
cfg::SsaCfg,
memory::{MemoryDefSite, MemoryLocation, MemoryOp, MemorySsa, MemoryVersion},
},
events::{Event, EventKind, EventListener},
ir::{
function::{SsaEditOptions, SsaFunction, SsaRollbackPolicy},
variable::SsaVarId,
},
pointer::PointerSize,
target::Target,
};
enum Rewrite {
Forward {
redundant: SsaVarId,
available: SsaVarId,
block: usize,
instr: usize,
},
DeadStore {
block: usize,
instr: usize,
},
}
pub fn run<T, L>(
ssa: &mut SsaFunction<T>,
method: &T::MethodRef,
events: &L,
ptr_size: PointerSize,
) -> bool
where
T: Target,
L: EventListener<T> + ?Sized,
{
let rewrites = {
let cfg = SsaCfg::from_ssa(ssa);
let mem_ssa = MemorySsa::build(ssa, &cfg, ptr_size);
plan_rewrites(ssa, &mem_ssa)
};
if rewrites.is_empty() {
return false;
}
apply_rewrites(ssa, method, events, rewrites)
}
fn plan_rewrites<T: Target>(ssa: &SsaFunction<T>, mem_ssa: &MemorySsa<T>) -> Vec<Rewrite> {
let mut by_position: HashMap<(usize, usize), &MemoryOp<T>> = HashMap::new();
for op in mem_ssa.operations() {
by_position.insert((op.block(), op.instr()), op);
}
let mut rewrites = Vec::new();
for block_idx in 0..ssa.block_count() {
plan_block(ssa, mem_ssa, &by_position, block_idx, &mut rewrites);
}
rewrites
}
fn plan_block<T: Target>(
ssa: &SsaFunction<T>,
mem_ssa: &MemorySsa<T>,
by_position: &HashMap<(usize, usize), &MemoryOp<T>>,
block_idx: usize,
rewrites: &mut Vec<Rewrite>,
) {
let Some(block) = ssa.block(block_idx) else {
return;
};
let mut available: Vec<(MemoryLocation<T>, SsaVarId)> = Vec::new();
let mut pending_stores: Vec<(MemoryLocation<T>, usize)> = Vec::new();
seed_from_dominator(mem_ssa, ssa, by_position, block_idx, &mut available);
for (instr_idx, _) in block.instructions().iter().enumerate() {
let Some(op) = by_position.get(&(block_idx, instr_idx)) else {
continue;
};
match op {
MemoryOp::Load { location, dest, .. } => {
pending_stores.retain(|(pending, _)| !pending.may_alias(location));
if let Some(value) = lookup_available(&available, location) {
rewrites.push(Rewrite::Forward {
redundant: *dest,
available: value,
block: block_idx,
instr: instr_idx,
});
} else {
set_available(&mut available, location.clone(), *dest);
}
}
MemoryOp::Store {
location, value, ..
} => {
if let Some(position) = pending_stores
.iter()
.position(|(pending, _)| pending.must_alias(location))
{
let (_, dead_instr) = pending_stores.swap_remove(position);
rewrites.push(Rewrite::DeadStore {
block: block_idx,
instr: dead_instr,
});
}
available.retain(|(known, _)| !known.may_alias(location));
pending_stores.retain(|(pending, _)| !pending.may_alias(location));
set_available(&mut available, location.clone(), *value);
pending_stores.push((location.clone(), instr_idx));
}
MemoryOp::ReadWrite { location, .. } => {
pending_stores.retain(|(pending, _)| !pending.may_alias(location));
available.retain(|(known, _)| !known.may_alias(location));
}
MemoryOp::Barrier { .. } => {
available.clear();
pending_stores.clear();
}
}
}
}
fn seed_from_dominator<T: Target>(
mem_ssa: &MemorySsa<T>,
ssa: &SsaFunction<T>,
by_position: &HashMap<(usize, usize), &MemoryOp<T>>,
block_idx: usize,
available: &mut Vec<(MemoryLocation<T>, SsaVarId)>,
) {
for location in mem_ssa.locations() {
if mem_ssa
.memory_phis(block_idx)
.iter()
.any(|phi| phi.location == *location)
{
continue;
}
let Some(version) = mem_ssa.version_at_entry(location, block_idx) else {
continue;
};
let Some(MemoryDefSite::Store {
block: store_block,
instr: store_instr,
}) = mem_ssa.definition(&MemoryVersion::new(location.clone(), version))
else {
continue;
};
if store_block == block_idx {
continue;
}
let Some(MemoryOp::Store {
location: stored_location,
value,
..
}) = by_position.get(&(store_block, store_instr)).copied()
else {
continue;
};
if !stored_location.must_alias(location) {
continue;
}
if ssa.get_definition(*value).is_none() {
continue;
}
set_available(available, location.clone(), *value);
}
}
fn lookup_available<T: Target>(
available: &[(MemoryLocation<T>, SsaVarId)],
location: &MemoryLocation<T>,
) -> Option<SsaVarId> {
available
.iter()
.find(|(known, _)| known.must_alias(location))
.map(|(_, value)| *value)
}
fn set_available<T: Target>(
available: &mut Vec<(MemoryLocation<T>, SsaVarId)>,
location: MemoryLocation<T>,
value: SsaVarId,
) {
if let Some(slot) = available
.iter_mut()
.find(|(known, _)| *known == location)
.map(|(_, existing)| existing)
{
*slot = value;
return;
}
available.push((location, value));
}
fn apply_rewrites<T, L>(
ssa: &mut SsaFunction<T>,
method: &T::MethodRef,
events: &L,
rewrites: Vec<Rewrite>,
) -> bool
where
T: Target,
L: EventListener<T> + ?Sized,
{
let mut forwarded = 0usize;
let mut dead_stores = 0usize;
let report = ssa.edit(
SsaEditOptions::new().with_rollback(SsaRollbackPolicy::OnFailure),
|editor| {
for rewrite in rewrites {
match rewrite {
Rewrite::Forward {
redundant,
available,
block,
instr,
..
} => {
let result = editor.replace_uses_checked(redundant, available);
if result.is_complete() {
editor.nop_instruction(block, instr)?;
forwarded = forwarded.saturating_add(1);
}
}
Rewrite::DeadStore { block, instr, .. } => {
editor.nop_instruction(block, instr)?;
dead_stores = dead_stores.saturating_add(1);
}
}
}
Ok(())
},
);
if report.is_err() {
return false;
}
if forwarded > 0 {
events.push(Event {
kind: EventKind::LoadForwarded,
method: Some(method.clone()),
location: None,
message: format!("forwarded {forwarded} load(s) from available memory"),
pass: None,
});
}
if dead_stores > 0 {
events.push(Event {
kind: EventKind::DeadStoreRemoved,
method: Some(method.clone()),
location: None,
message: format!("removed {dead_stores} dead store(s)"),
pass: None,
});
}
forwarded > 0 || dead_stores > 0
}