use std::collections::HashMap;
use crate::ir::{Expr, SsaCfg, Stmt, SsaTerminator, VarDef, VarId};
use pcode_ir::AddressSpaceId;
#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash, PartialOrd, Ord)]
pub struct Region(pub u32);
#[derive(Debug, Clone, PartialEq, Eq, Hash)]
pub enum AllocSite {
Param(u8),
StackFrame,
Global(u64),
Heap(u64),
Const(u64),
Unknown(u32),
}
#[derive(Debug, Clone, PartialEq, Eq, Hash)]
pub enum OffsetClass {
ConstOffset(i64),
Symbolic,
}
#[derive(Debug, Clone, Default)]
pub struct RegionMap {
by_var: Vec<Region>,
sites: Vec<AllocSite>,
intern: HashMap<AllocSite, Region>,
}
impl RegionMap {
pub fn region_of(&self, var: VarId) -> Region {
*self
.by_var
.get(var.0 as usize)
.unwrap_or(&Region(u32::MAX))
}
pub fn site_of(&self, r: Region) -> Option<&AllocSite> {
self.sites.get(r.0 as usize)
}
fn intern_site(&mut self, site: AllocSite) -> Region {
if let Some(r) = self.intern.get(&site) {
return *r;
}
let r = Region(self.sites.len() as u32);
self.sites.push(site.clone());
self.intern.insert(site, r);
r
}
}
pub fn infer_regions(ssa: &SsaCfg) -> RegionMap {
let mut map = RegionMap::default();
map.by_var = vec![Region(u32::MAX); ssa.vars.len()];
let heap_returns = collect_heap_returns(ssa);
let spill_map = build_spill_map(ssa);
for _iter in 0..4 {
let mut changed = false;
for i in 0..ssa.vars.len() {
let id = VarId(i as u32);
let cur = map.by_var[i];
let site_opt = classify(&ssa.vars[i], &map, &heap_returns, id, ssa, &spill_map);
let new = match site_opt {
Some(s) => map.intern_site(s),
None => Region(u32::MAX),
};
if new != cur {
map.by_var[i] = new;
changed = true;
}
}
if !changed {
break;
}
}
for (i, slot) in map.by_var.iter_mut().enumerate() {
if slot.0 == u32::MAX {
let r = Region(map.sites.len() as u32);
map.sites.push(AllocSite::Unknown(i as u32));
map.intern
.insert(AllocSite::Unknown(i as u32), r);
*slot = r;
}
}
map
}
fn classify(
def: &VarDef,
map: &RegionMap,
heap_returns: &HashMap<VarId, u64>,
id: VarId,
ssa: &SsaCfg,
spill_map: &SpillMap,
) -> Option<AllocSite> {
let _ = id;
if let Some(name) = &def.param_name {
if let Some(rest) = name.strip_prefix("param_") {
if let Ok(n) = rest.parse::<u8>() {
return Some(AllocSite::Param(n));
}
}
}
if def.call_return {
if let Some(call_site) = heap_returns.get(&id) {
return Some(AllocSite::Heap(*call_site));
}
return None;
}
match &def.expr {
Expr::Const(c, _) => {
if def.varnode.space == AddressSpaceId::Ram && *c != 0 {
Some(AllocSite::Global(*c as u64))
} else {
Some(AllocSite::Const(*c as u64))
}
}
Expr::Var(inner) => site_of_var(*inner, map),
Expr::BinOp(_, a, b) => site_of_var(*a, map).or_else(|| site_of_var(*b, map)),
Expr::UnaryOp(_, a) => site_of_var(*a, map),
Expr::FieldAccess(base, off) => {
let key = format!(
"BAdd({},C{}.8)",
addr_canon_local(*base, &ssa.vars).unwrap_or_else(|| "?".to_string()),
off
);
if let Some(stored) = spill_map.by_canon.get(&key) {
if let Some(site) = site_of_var(*stored, map) {
return Some(site);
}
}
site_of_var(*base, map)
}
Expr::Load(addr) => {
if let Some(stored) = spill_map.lookup(*addr, &ssa.vars) {
if let Some(site) = site_of_var(stored, map) {
return Some(site);
}
}
site_of_var(*addr, map)
}
Expr::Phi(args) => args.iter().find_map(|a| site_of_var(*a, map)),
Expr::Unknown if def.varnode.space == AddressSpaceId::Register => {
Some(AllocSite::StackFrame)
}
_ => None,
}
}
fn site_of_var(v: VarId, map: &RegionMap) -> Option<AllocSite> {
let r = *map.by_var.get(v.0 as usize)?;
if r.0 == u32::MAX {
return None;
}
let site = map.sites.get(r.0 as usize)?;
if matches!(site, AllocSite::Unknown(_)) {
return None;
}
Some(site.clone())
}
fn collect_heap_returns(ssa: &SsaCfg) -> HashMap<VarId, u64> {
let mut m = HashMap::new();
for block in &ssa.blocks {
for stmt in &block.stmts {
if let Stmt::Call {
target,
out: Some(o),
..
} = stmt
{
if let crate::ir::CallTarget::Direct(addr) = target {
m.insert(*o, *addr);
}
}
}
if let SsaTerminator::Call {
target,
out: Some(o),
..
} = &block.terminator
{
if let crate::ir::CallTarget::Direct(addr) = target {
m.insert(*o, *addr);
}
}
}
m
}
struct SpillMap {
by_canon: HashMap<String, VarId>,
by_varnode: HashMap<pcode_ir::Varnode, VarId>,
}
fn build_spill_map(ssa: &SsaCfg) -> SpillMap {
use crate::ir::Stmt;
let mut by_canon: HashMap<String, VarId> = HashMap::new();
let mut by_varnode: HashMap<pcode_ir::Varnode, VarId> = HashMap::new();
for block in &ssa.blocks {
for stmt in &block.stmts {
if let Stmt::Store { addr, val } = stmt {
if let Some(key) = addr_canon_local(*addr, &ssa.vars) {
by_canon.insert(key, *val);
}
if let Some(addr_def) = ssa.vars.get(addr.0 as usize) {
by_varnode.insert(addr_def.varnode, *val);
}
}
}
}
SpillMap { by_canon, by_varnode }
}
impl SpillMap {
fn lookup(&self, addr: VarId, vars: &[VarDef]) -> Option<VarId> {
if let Some(key) = addr_canon_local(addr, vars) {
if let Some(v) = self.by_canon.get(&key) {
return Some(*v);
}
}
if let Some(d) = vars.get(addr.0 as usize) {
if let Some(v) = self.by_varnode.get(&d.varnode) {
return Some(*v);
}
}
None
}
}
fn addr_canon_local(var: VarId, vars: &[VarDef]) -> Option<String> {
fn rec(var: VarId, vars: &[VarDef], depth: u32) -> Option<String> {
if depth > 16 {
return None;
}
let def = vars.get(var.0 as usize)?;
Some(match &def.expr {
Expr::Var(inner) => rec(*inner, vars, depth + 1)?,
Expr::Const(c, sz) => format!("C{}.{}", c, sz),
Expr::BinOp(op, a, b) => {
let ka = rec(*a, vars, depth + 1).unwrap_or_else(|| "?".to_string());
let kb = rec(*b, vars, depth + 1).unwrap_or_else(|| "?".to_string());
format!("B{:?}({},{})", op, ka, kb)
}
Expr::UnaryOp(op, a) => {
let ka = rec(*a, vars, depth + 1).unwrap_or_else(|| "?".to_string());
format!("U{:?}({})", op, ka)
}
_ => format!(
"V{:?}/{}/{}",
def.varnode.space, def.varnode.offset, def.varnode.size
),
})
}
rec(var, vars, 0)
}
#[cfg(test)]
mod tests {
use super::*;
use crate::ir::{
BlockId, Diagnostic, Expr, InferredType, SsaBlock, SsaCfg, SsaTerminator, VarDef,
};
use pcode_ir::Varnode;
fn mk_var(id: u32, expr: Expr, vn: Varnode, param: Option<&str>) -> VarDef {
VarDef {
id: VarId(id),
varnode: vn,
expr,
size: 8,
use_count: 1,
param_name: param.map(String::from),
call_return: false,
inferred_type: InferredType::Unknown,
display_type: None,
}
}
fn empty_ssa(vars: Vec<VarDef>) -> SsaCfg {
SsaCfg {
blocks: vec![SsaBlock {
id: BlockId(0),
addr: 0,
stmts: vec![],
terminator: SsaTerminator::Return(None),
}],
vars,
entry: BlockId(0),
diagnostics: Vec::<Diagnostic>::new(),
}
}
#[test]
fn param_var_classified_as_param() {
let vars = vec![mk_var(
0,
Expr::Unknown,
Varnode::register(0, 8),
Some("param_2"),
)];
let map = infer_regions(&empty_ssa(vars));
let r = map.region_of(VarId(0));
match map.site_of(r) {
Some(AllocSite::Param(2)) => {}
other => panic!("expected Param(2), got {other:?}"),
}
}
#[test]
fn pointer_arith_inherits_base_region() {
let vars = vec![
mk_var(0, Expr::Unknown, Varnode::register(0, 8), Some("param_0")),
mk_var(1, Expr::Const(8, 8), Varnode::constant(8, 8), None),
mk_var(
2,
Expr::BinOp(crate::ir::BinOpKind::Add, VarId(0), VarId(1)),
Varnode::register(64, 8),
None,
),
];
let map = infer_regions(&empty_ssa(vars));
let r0 = map.region_of(VarId(0));
let r2 = map.region_of(VarId(2));
assert_eq!(r0, r2, "BinOp(Add, ptr, const) must inherit ptr's region");
}
#[test]
fn distinct_params_get_distinct_regions() {
let vars = vec![
mk_var(0, Expr::Unknown, Varnode::register(0, 8), Some("param_0")),
mk_var(1, Expr::Unknown, Varnode::register(8, 8), Some("param_1")),
];
let map = infer_regions(&empty_ssa(vars));
assert_ne!(map.region_of(VarId(0)), map.region_of(VarId(1)));
}
#[test]
fn const_global_classified_as_global() {
let vars = vec![mk_var(
0,
Expr::Const(0x602080, 8),
Varnode {
space: pcode_ir::AddressSpaceId::Ram,
offset: 0x602080,
size: 8,
},
None,
)];
let map = infer_regions(&empty_ssa(vars));
let r = map.region_of(VarId(0));
match map.site_of(r) {
Some(AllocSite::Global(0x602080)) => {}
other => panic!("expected Global(0x602080), got {other:?}"),
}
}
}