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 machine = pe.header.coff_header.machine;
if machine != 0x014c && machine != 0x8664 {
return Vec::new();
}
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) => {
if elf.header.e_machine != 3 && elf.header.e_machine != 62 {
return Vec::new();
}
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 skips_non_x86_elf() {
use goblin::Object;
#[rustfmt::skip]
let elf: Vec<u8> = {
let mut v: Vec<u8> = vec![
0x7f, b'E', b'L', b'F', 1, 1, 1, 0, 0, 0, 0, 0, 0, 0, 0, 0,
2, 0, 8, 0, 1, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0x34, 0, 0, 0, 0, 0, 0, 0, 0x34, 0, 0, 0, 0, 0, 0x28, 0, 3, 0, 2, 0, ];
v.extend(std::iter::repeat(0).take(40));
let text_off: u32 = 0xC0;
let text_size: u32 = 4;
v.extend_from_slice(&1u32.to_le_bytes());
v.extend_from_slice(&1u32.to_le_bytes());
v.extend_from_slice(&6u32.to_le_bytes());
v.extend_from_slice(&0x1000u32.to_le_bytes());
v.extend_from_slice(&text_off.to_le_bytes());
v.extend_from_slice(&text_size.to_le_bytes());
v.extend_from_slice(&0u32.to_le_bytes()); v.extend_from_slice(&0u32.to_le_bytes()); v.extend_from_slice(&4u32.to_le_bytes()); v.extend_from_slice(&0u32.to_le_bytes()); let str_off: u32 = 0xAC;
let str_size: u32 = 0x11;
v.extend_from_slice(&7u32.to_le_bytes());
v.extend_from_slice(&3u32.to_le_bytes());
v.extend_from_slice(&0u32.to_le_bytes());
v.extend_from_slice(&0u32.to_le_bytes());
v.extend_from_slice(&str_off.to_le_bytes());
v.extend_from_slice(&str_size.to_le_bytes());
v.extend_from_slice(&0u32.to_le_bytes());
v.extend_from_slice(&0u32.to_le_bytes());
v.extend_from_slice(&1u32.to_le_bytes());
v.extend_from_slice(&0u32.to_le_bytes());
assert_eq!(v.len(), 0xAC);
v.extend_from_slice(b"\0.text\0.shstrtab\0");
v.extend_from_slice(&[0u8; 3]);
assert_eq!(v.len(), 0xC0);
v.extend_from_slice(b"\xff\xe7\xcc\x90");
v
};
let obj = Object::parse(&elf).expect("parseable MIPS ELF");
let hits = scan(&obj, &elf);
assert!(
hits.is_empty(),
"scan emitted {} bogus trampoline(s) on MIPS ELF",
hits.len()
);
}
#[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);
}
}