use goblin::Object;
fn classify_jmp_reg(bytes: &[u8]) -> Option<&'static str> {
if bytes.is_empty() {
return None;
}
if bytes.len() >= 2 && bytes[0] == 0xff {
return match bytes[1] {
0xe0 => Some("RAX"),
0xe1 => Some("RCX"),
0xe2 => Some("RDX"),
0xe3 => Some("RBX"),
0xe4 => Some("RSP"),
0xe5 => Some("RBP"),
0xe6 => Some("RSI"),
0xe7 => Some("RDI"),
_ => None,
};
}
if bytes.len() >= 3 && bytes[0] == 0x41 && bytes[1] == 0xff {
return match bytes[2] {
0xe0 => Some("R8"),
0xe1 => Some("R9"),
0xe2 => Some("R10"),
0xe3 => Some("R11"),
0xe4 => Some("R12"),
0xe5 => Some("R13"),
0xe6 => Some("R14"),
0xe7 => Some("R15"),
_ => None,
};
}
None
}
#[derive(Debug, Clone)]
pub struct Trampoline {
pub addr: u64,
pub reg: &'static str,
}
pub fn scan_region(code: &[u8], base_va: u64) -> Vec<Trampoline> {
let mut hits = Vec::new();
let mut i = 0;
while i + 2 < code.len() {
if let Some(reg) = classify_jmp_reg(&code[i..i + 2]) {
let after = code.get(i + 2).copied().unwrap_or(0);
if after == 0xCC || after == 0x90 {
hits.push(Trampoline {
addr: base_va + i as u64,
reg,
});
i += 3;
continue;
}
}
if i + 3 < code.len() {
if let Some(reg) = classify_jmp_reg(&code[i..i + 3]) {
let after = code.get(i + 3).copied().unwrap_or(0);
if after == 0xCC || after == 0x90 {
hits.push(Trampoline {
addr: base_va + i as u64,
reg,
});
i += 4;
continue;
}
}
}
i += 1;
}
hits
}
pub fn scan(obj: &Object<'_>, data: &[u8]) -> Vec<Trampoline> {
match obj {
Object::PE(pe) => {
let mut hits = Vec::new();
const IMAGE_SCN_MEM_EXECUTE: u32 = 0x2000_0000;
for sec in &pe.sections {
if sec.characteristics & IMAGE_SCN_MEM_EXECUTE == 0 {
continue;
}
let raddr = sec.pointer_to_raw_data as usize;
let rsize = sec.size_of_raw_data as usize;
if raddr + rsize > data.len() {
continue;
}
let base_va = pe.image_base as u64 + sec.virtual_address as u64;
let region_hits = scan_region(&data[raddr..raddr + rsize], base_va);
hits.extend(region_hits);
}
hits
}
Object::Elf(elf) => {
let mut hits = Vec::new();
for sh in &elf.section_headers {
if sh.sh_flags & 0x4 == 0 {
continue;
}
let raddr = sh.sh_offset as usize;
let rsize = sh.sh_size as usize;
if raddr + rsize > data.len() {
continue;
}
let base_va = sh.sh_addr;
let region_hits = scan_region(&data[raddr..raddr + rsize], base_va);
hits.extend(region_hits);
}
hits
}
_ => Vec::new(),
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn detects_jmp_rax() {
let code = b"\xff\xe0\xcc\xcc\xcc\xcc";
let hits = scan_region(code, 0x1000);
assert_eq!(hits.len(), 1);
assert_eq!(hits[0].reg, "RAX");
assert_eq!(hits[0].addr, 0x1000);
}
#[test]
fn detects_jmp_r8() {
let code = b"\x41\xff\xe0\xcc\xcc";
let hits = scan_region(code, 0x2000);
assert_eq!(hits.len(), 1);
assert_eq!(hits[0].reg, "R8");
}
#[test]
fn rejects_no_padding() {
let code = b"\xff\xe0\x48\x89\xc1";
let hits = scan_region(code, 0x3000);
assert!(hits.is_empty());
}
#[test]
fn finds_multiple_in_region() {
let code = b"\xff\xe0\xcc\xcc\x41\xff\xe2\xcc\xcc";
let hits = scan_region(code, 0x4000);
assert_eq!(hits.len(), 2);
assert_eq!(hits[0].reg, "RAX");
assert_eq!(hits[0].addr, 0x4000);
assert_eq!(hits[1].reg, "R10");
assert_eq!(hits[1].addr, 0x4004);
}
}