use std::collections::HashSet;
use rucc_base::Interner;
use rucc_ir::{Extra, Func, FuncId, Inst, InstData, Module, Opcode, Value};
use crate::origin;
const ENDS: &[(&str, usize)] = &[("free", 0), ("realloc", 0), ("reallocf", 0)];
pub fn checks(module: &mut Module, names: &Interner) -> usize {
let defined: HashSet<&str> = module
.funcs()
.filter(|&id| !module[id].is_declaration())
.map(|id| names.resolve(module[id].name))
.collect();
let ids: Vec<FuncId> = module.funcs().collect();
let mut done = 0;
for id in ids {
if module[id].is_declaration() {
continue;
}
done += one(&mut module[id], names, &defined);
}
done
}
fn one(func: &mut Func, names: &Interner, defined: &HashSet<&str>) -> usize {
let ending: Vec<(Inst, Value)> = func
.blocks()
.flat_map(|block| func.insts(block).collect::<Vec<Inst>>())
.filter_map(|inst| handed(func, names, defined, inst).map(|value| (inst, value)))
.collect();
if ending.is_empty() {
return 0;
}
let mut origins = origin::Origins::new();
for (pointer, cap) in origin::existing(func) {
origins.seed(pointer, cap);
}
let mut done = 0;
for (inst, pointer) in ending {
let span = func.span(inst);
let capability = origins.of(func, pointer, inst);
let args = func.push_values(&[capability, pointer]);
let data = InstData { args, ..InstData::new(Opcode::CheckFree) };
let made = func.create_inst(data, &[], span);
func.insert_before(made, inst);
done += 1;
}
done
}
fn handed(func: &Func, names: &Interner, defined: &HashSet<&str>, inst: Inst) -> Option<Value> {
if func[inst].opcode != Opcode::Call {
return None;
}
let Extra::Call(at) = func[inst].extra else { return None };
let name = names.resolve(func[at].callee?);
if defined.contains(name) {
return None;
}
let (_, which) = ENDS.iter().find(|&&(each, _)| each == name)?;
let &value = func[func[inst].args].get(*which)?;
func[value].ty.is_ptr().then_some(value)
}
#[cfg(test)]
mod tests {
use rucc_base::Interner;
use rucc_ir::{
Builder, CallInfo, Extra, Func, FuncId, InstData, Module, Opcode, Signature, Type,
};
use rucc_target::{Arch, Env, Os, TargetInfo, Triple};
use super::{ENDS, checks};
fn unit(names: &mut Interner) -> Module {
let target = TargetInfo::new(Triple::new(Arch::X86_64, Os::Linux, Env::Gnu));
Module::new(names.intern("a.c"), &target)
}
fn declares(names: &mut Interner, module: &mut Module, name: &str, returns: bool) -> FuncId {
let mut signature = Signature::new().with_params(&[Type::PTR]);
if returns {
signature = signature.with_returns(&[Type::PTR]);
}
module.add_func(Func::new(names.intern(name), signature))
}
fn passes_it_to(names: &mut Interner, module: &mut Module, callee: &str, returns: bool) {
declares(names, module, callee, returns);
calls(names, module, callee, returns);
}
fn calls(names: &mut Interner, module: &mut Module, callee: &str, returns: bool) {
let mut signature = Signature::new().with_params(&[Type::PTR]);
signature = signature.with_returns(&[]);
let mut func = Func::new(names.intern("caller"), signature);
let entry = func.create_block();
let p = func.append_param(entry, Type::PTR);
let mut called = Signature::new().with_params(&[Type::PTR]);
if returns {
called = called.with_returns(&[Type::PTR]);
}
let sig = func.add_signature(called);
let varargs = func.push_abis(&[]);
let info =
func.add_call(CallInfo { callee: Some(names.intern(callee)), signature: sig, varargs });
let mut b = Builder::new(&mut func, entry);
let args = b.func().push_values(&[p]);
let data = InstData { args, extra: Extra::Call(info), ..InstData::new(Opcode::Call) };
if returns {
b.value(data, Type::PTR);
} else {
b.inst(data, &[]);
}
b.ret(&[]);
module.add_func(func);
}
fn counted(module: &Module) -> usize {
module
.funcs()
.filter(|&id| !module[id].is_declaration())
.map(|id| &module[id])
.flat_map(|func| {
func.blocks()
.flat_map(|block| func.insts(block).collect::<Vec<_>>())
.filter(|&inst| func[inst].opcode == Opcode::CheckFree)
.collect::<Vec<_>>()
})
.count()
}
#[test]
fn the_table_is_sorted_and_says_each_name_once() {
for pair in ENDS.windows(2) {
assert!(pair[0].0 < pair[1].0, "{} then {}", pair[0].0, pair[1].0);
}
}
#[test]
fn a_free_gets_a_check_in_front_of_it() {
let mut names = Interner::new();
let mut module = unit(&mut names);
passes_it_to(&mut names, &mut module, "free", false);
assert_eq!(checks(&mut module, &names), 1);
assert_eq!(counted(&module), 1);
}
#[test]
fn a_realloc_gets_one_too_because_the_pointer_going_in_is_over() {
let mut names = Interner::new();
let mut module = unit(&mut names);
passes_it_to(&mut names, &mut module, "realloc", true);
assert_eq!(checks(&mut module, &names), 1);
}
#[test]
fn a_call_of_something_else_gets_nothing() {
let mut names = Interner::new();
let mut module = unit(&mut names);
passes_it_to(&mut names, &mut module, "puts", false);
assert_eq!(checks(&mut module, &names), 0);
}
#[test]
fn a_module_that_writes_its_own_free_keeps_it() {
let mut names = Interner::new();
let mut module = unit(&mut names);
let mut own = Func::new(names.intern("free"), Signature::new().with_params(&[Type::PTR]));
let entry = own.create_block();
own.append_param(entry, Type::PTR);
let mut b = Builder::new(&mut own, entry);
b.ret(&[]);
module.add_func(own);
calls(&mut names, &mut module, "free", false);
assert_eq!(checks(&mut module, &names), 0);
}
#[test]
fn the_check_reads_the_capability_the_pointer_already_had() {
let mut names = Interner::new();
let mut module = unit(&mut names);
passes_it_to(&mut names, &mut module, "free", false);
assert_eq!(checks(&mut module, &names), 1);
let id = module
.funcs()
.find(|&id| !module[id].is_declaration())
.expect("the module defines one function");
let func = &module[id];
let made: Vec<_> = func
.blocks()
.flat_map(|block| func.insts(block).collect::<Vec<_>>())
.filter(|&inst| func[inst].opcode == Opcode::CapOf)
.collect();
assert_eq!(made.len(), 1, "one producer for the one pointer");
let check = func
.blocks()
.flat_map(|block| func.insts(block).collect::<Vec<_>>())
.find(|&inst| func[inst].opcode == Opcode::CheckFree)
.expect("the check went in");
let cap = func[made[0]].results().next().expect("a cap_of produces one value");
assert_eq!(func[func[check].args].first(), Some(&cap));
}
}