use goblin::Object;
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum BranchDir {
Take,
Skip,
}
#[derive(Debug, Clone)]
pub struct TagCase {
pub cmp_va: u64,
pub tag: u8,
pub target_va: u64,
pub dir: BranchDir,
}
pub fn scan_region(code: &[u8], base_va: u64) -> Vec<TagCase> {
let mut cases = Vec::new();
let mut k = 0;
while k + 4 <= code.len() {
let (tag, cmp_len) =
if code[k] == 0x80 && (code[k + 1] & 0xf8) == 0xf8 && k + 3 <= code.len() {
(code[k + 2], 3)
} else if code[k] == 0x3c {
(code[k + 1], 2)
} else if code[k] == 0x41
&& code[k + 1] == 0x80
&& (code[k + 2] & 0xf8) == 0xf8
&& k + 4 <= code.len()
{
(code[k + 3], 4)
} else {
k += 1;
continue;
};
let jz_off = k + cmp_len;
if jz_off >= code.len() {
break;
}
if jz_off + 2 <= code.len() && (code[jz_off] == 0x74 || code[jz_off] == 0x75) {
let rel = code[jz_off + 1] as i8 as i64;
let target = (base_va + (jz_off + 2) as u64).wrapping_add(rel as u64);
let dir = if code[jz_off] == 0x74 {
BranchDir::Take
} else {
BranchDir::Skip
};
cases.push(TagCase {
cmp_va: base_va + k as u64,
tag,
target_va: target,
dir,
});
k = jz_off + 2;
continue;
}
if jz_off + 6 <= code.len()
&& code[jz_off] == 0x0f
&& (code[jz_off + 1] == 0x84 || code[jz_off + 1] == 0x85)
{
let rel = i32::from_le_bytes([
code[jz_off + 2],
code[jz_off + 3],
code[jz_off + 4],
code[jz_off + 5],
]) as i64;
let target = (base_va + (jz_off + 6) as u64).wrapping_add(rel as u64);
let dir = if code[jz_off + 1] == 0x84 {
BranchDir::Take
} else {
BranchDir::Skip
};
cases.push(TagCase {
cmp_va: base_va + k as u64,
tag,
target_va: target,
dir,
});
k = jz_off + 6;
continue;
}
k += 1;
}
cases
}
pub fn scan_function(obj: &Object<'_>, data: &[u8], func_va: u64) -> Vec<TagCase> {
if let Object::PE(pe) = obj {
for sec in &pe.sections {
let svaddr = pe.image_base as u64 + sec.virtual_address as u64;
let sv = sec.virtual_size as u64;
if func_va >= svaddr && func_va < svaddr + sv {
let raddr = sec.pointer_to_raw_data as usize;
let rsize = sec.size_of_raw_data as usize;
let off_in_section = (func_va - svaddr) as usize;
if off_in_section < rsize {
let scan_len = (0x600).min(rsize - off_in_section);
return scan_region(
&data[raddr + off_in_section..raddr + off_in_section + scan_len],
func_va,
);
}
}
}
}
Vec::new()
}
pub fn render(cases: &[TagCase]) -> Vec<String> {
cases
.iter()
.map(|c| {
let mnem = match c.dir {
BranchDir::Take => "JZ",
BranchDir::Skip => "JNZ",
};
format!(
"{:#x}: CMP r8, {:#04x} → {} {:#x}",
c.cmp_va, c.tag, mnem, c.target_va
)
})
.collect()
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn detects_cmp_dl_jz_short() {
let code = b"\x80\xfa\x2f\x74\x05";
let cases = scan_region(code, 0x1000);
assert_eq!(cases.len(), 1);
assert_eq!(cases[0].tag, 0x2f);
assert_eq!(cases[0].dir, BranchDir::Take);
assert_eq!(cases[0].target_va, 0x100a);
}
#[test]
fn detects_cmp_dl_jnz_skip() {
let code = b"\x80\xfa\x2f\x75\x05";
let cases = scan_region(code, 0x1000);
assert_eq!(cases.len(), 1);
assert_eq!(cases[0].tag, 0x2f);
assert_eq!(cases[0].dir, BranchDir::Skip);
}
#[test]
fn detects_cmp_al_jz_long() {
let mut code = vec![0x3c, 0xc8, 0x0f, 0x84];
code.extend_from_slice(&0x100_i32.to_le_bytes());
let cases = scan_region(&code, 0x2000);
assert_eq!(cases.len(), 1);
assert_eq!(cases[0].tag, 0xc8);
assert_eq!(cases[0].target_va, 0x2108);
}
#[test]
fn extracts_chain() {
let code = b"\x80\xfa\x2f\x74\x00\x80\xfa\x6e\x74\x00";
let cases = scan_region(code, 0x3000);
assert_eq!(cases.len(), 2);
assert_eq!(cases[0].tag, 0x2f);
assert_eq!(cases[1].tag, 0x6e);
}
#[test]
fn ignores_cmp_without_jz() {
let code = b"\x80\xfa\x2f\x90\x90";
let cases = scan_region(code, 0x4000);
assert!(cases.is_empty());
}
}