use std::collections::{HashMap, HashSet};
use rucc_ir::{Block, BlockCall, Def, Func, Inst, InstData, Opcode, Type, Value};
#[derive(Debug, Default)]
pub(crate) struct Origins {
held: HashMap<Value, Value>,
joins: Vec<(Block, usize, Value)>,
}
impl Origins {
pub(crate) fn new() -> Self {
Self::default()
}
pub(crate) fn of(&mut self, func: &mut Func, pointer: Value, at: Inst) -> Value {
if let Some(&held) = self.held.get(&pointer) {
return held;
}
let base = root(func, pointer);
let cap = match self.held.get(&base) {
Some(&held) => held,
None => {
let made = match joinable(func, base) {
Some((block, index)) => {
let made = func.append_param(block, Type::CAP);
self.joins.push((block, index, made));
made
}
None => cap_of(func, base, at),
};
self.held.insert(base, made);
made
}
};
let mut each = pointer;
while each != base {
self.held.insert(each, cap);
let Def::Result { inst, .. } = func[each].def else { break };
let Some(&next) = func[func[inst].args].first() else { break };
each = next;
}
cap
}
pub(crate) fn seed(&mut self, pointer: Value, cap: Value) {
self.held.insert(pointer, cap);
}
pub(crate) fn join(&mut self, func: &mut Func) {
let mut next = 0;
while let Some(&(block, index, _)) = self.joins.get(next) {
next += 1;
for pred in func.blocks().collect::<Vec<Block>>() {
let Some(term) = func.terminator(pred) else { continue };
for at in func.target_list(term).iter() {
let call = func[at];
if call.block != block {
continue;
}
let Some(&passed) = func[call.args].get(index) else { continue };
let cap = self.of(func, passed, term);
let call = func[at];
let args = func.append_arg(call.args, cap);
func.set_block_call(at, BlockCall { args, ..call });
}
}
}
loop {
let mut again = false;
for &(block, _, cap) in &self.joins {
let Some(at) = func[block].params.iter().position(|¶m| param == cap) else {
continue;
};
let arriving: HashSet<Value> =
edges(func, block).map(|args| func[args][at]).filter(|&v| v != cap).collect();
let Some(&only) = arriving.iter().next() else { continue };
if arriving.len() != 1 {
continue;
}
renamed(func, cap, only);
unjoin(func, block, at);
again = true;
}
if !again {
return;
}
}
}
}
fn edges(func: &Func, block: Block) -> impl Iterator<Item = rucc_ir::ValueList> + use<'_> {
func.blocks()
.filter_map(|pred| func.terminator(pred))
.flat_map(|term| func.successors(term))
.filter(move |call| call.block == block)
.map(|call| call.args)
}
fn renamed(func: &mut Func, from: Value, to: Value) {
let with = |value: Value| if value == from { to } else { value };
for block in func.blocks().collect::<Vec<Block>>() {
for inst in func.insts(block).collect::<Vec<Inst>>() {
func.rewrite(func[inst].args, with);
for call in func.successors(inst).collect::<Vec<_>>() {
func.rewrite(call.args, with);
}
}
}
}
fn unjoin(func: &mut Func, block: Block, at: usize) {
for pred in func.blocks().collect::<Vec<Block>>() {
let Some(term) = func.terminator(pred) else { continue };
for place in func.target_list(term).iter() {
let call = func[place];
if call.block != block {
continue;
}
let mut kept = func[call.args].to_vec();
kept.remove(at);
let args = func.push_values(&kept);
func.set_block_call(place, BlockCall { args, ..call });
}
}
let gone = func[block].params[at];
func.retain_params(block, |param| param != gone);
}
fn joinable(func: &Func, pointer: Value) -> Option<(Block, usize)> {
let Def::Param { block, index } = func[pointer].def else { return None };
if func.entry() == Some(block) {
return None;
}
let index = index as usize;
let mut reached = false;
for pred in func.blocks() {
let Some(term) = func.terminator(pred) else { continue };
for call in func.successors(term) {
if call.block != block {
continue;
}
if func[term].opcode == Opcode::IndirectBr {
return None;
}
let &passed = func[call.args].get(index)?;
let def = func[root(func, passed)].def;
if matches!(def, Def::Result { inst, .. } if func.is_terminator(inst)) {
return None;
}
reached = true;
}
}
reached.then_some((block, index))
}
pub(crate) fn existing(func: &Func) -> HashMap<Value, Value> {
let alive = kept(func);
let mut held = HashMap::new();
for block in func.blocks() {
for inst in func.insts(block) {
let Some(at) = func[inst].opcode.capability_names() else { continue };
let Some(&pointer) = func[func[inst].args].get(at) else { continue };
let Some(cap) = func[inst].results().next() else { continue };
if !alive.contains(&cap) {
continue;
}
held.entry(root(func, pointer)).or_insert(cap);
}
}
joins(func, &alive, &mut held);
held
}
fn joins(func: &Func, alive: &HashSet<Value>, held: &mut HashMap<Value, Value>) {
loop {
let mut again = false;
for block in func.blocks() {
let params = &func[block].params;
for (at, &cap) in params.iter().enumerate() {
if !func[cap].ty.is_cap() || !alive.contains(&cap) {
continue;
}
for (index, &pointer) in params.iter().enumerate() {
if !func[pointer].ty.is_ptr() || held.contains_key(&pointer) {
continue;
}
if paired(func, held, block, (index, pointer), (at, cap)) {
held.insert(pointer, cap);
again = true;
break;
}
}
}
}
if !again {
return;
}
}
}
fn paired(
func: &Func,
held: &HashMap<Value, Value>,
block: Block,
(index, pointer): (usize, Value),
(at, cap): (usize, Value),
) -> bool {
let mut reached = false;
for pred in func.blocks() {
let Some(term) = func.terminator(pred) else { continue };
for call in func.successors(term) {
if call.block != block {
continue;
}
let args = &func[call.args];
let (Some(&passed), Some(&carried)) = (args.get(index), args.get(at)) else {
return false;
};
let round = passed == pointer && carried == cap;
if !round && held.get(&root(func, passed)) != Some(&carried) {
return false;
}
reached = true;
}
}
reached
}
fn kept(func: &Func) -> HashSet<Value> {
let mut alive: HashSet<Value> = HashSet::new();
for block in func.blocks() {
for inst in func.insts(block) {
if !crate::lower::keeps(func[inst].opcode) {
continue;
}
let read = func[func[inst].args].iter().copied();
alive.extend(read.filter(|&value| func[value].ty.is_cap()));
}
}
loop {
let mut again = false;
for block in func.blocks() {
if let Some(term) = func.terminator(block) {
for call in func.successors(term) {
for (&value, ¶m) in func[call.args].iter().zip(&func[call.block].params) {
if func[param].ty.is_cap() && alive.contains(¶m) && alive.insert(value)
{
again = true;
}
}
}
}
for inst in func.insts(block) {
if !func[inst].opcode.makes_capability() {
continue;
}
if !func[inst].results().any(|value| alive.contains(&value)) {
continue;
}
for &value in func[func[inst].args].iter() {
if func[value].ty.is_cap() && alive.insert(value) {
again = true;
}
}
}
}
if !again {
return alive;
}
}
}
pub(crate) fn already(func: &Func, held: &HashMap<Value, Value>, pointer: Value) -> Option<Value> {
held.get(&root(func, pointer)).copied()
}
fn root(func: &Func, pointer: Value) -> Value {
let mut at = pointer;
loop {
let Def::Result { inst, .. } = func[at].def else { return at };
if func[inst].opcode != Opcode::PtrAdd {
return at;
}
let Some(&base) = func[func[inst].args].first() else { return at };
if !func[base].ty.is_ptr() {
return at;
}
at = base;
}
}
fn cap_of(func: &mut Func, pointer: Value, at: Inst) -> Value {
let (anchor, behind) = place(func, pointer, at);
let span = func.span(at);
let args = func.push_values(&[pointer]);
let data = InstData { args, ..InstData::new(Opcode::CapOf) };
let cap = func.create_inst(data, &[Type::CAP], span);
if behind {
func.insert_after(cap, anchor);
} else {
func.insert_before(cap, anchor);
}
func[cap].results().next().expect("cap_of produces one value")
}
fn place(func: &Func, pointer: Value, at: Inst) -> (Inst, bool) {
match func[pointer].def {
Def::Result { inst, .. } if !func.is_terminator(inst) => (inst, true),
Def::Param { block, .. } => {
let first = func.insts(block).find(|&inst| func[inst].opcode != Opcode::Alloca);
(first.unwrap_or(at), false)
}
Def::Result { .. } => (at, false),
}
}