use std::ffi::{c_int, c_void};
use std::marker::PhantomData;
use std::ops::Range;
use std::ptr;
use idakit_sys as sys;
use crate::Database;
use crate::address::Address;
use crate::error::{Error, PatternRejection, Result};
use crate::ffi::{cstr, with_cstr};
impl Database {
#[must_use]
#[doc(alias("bin_search"))]
pub fn search<'p, 'db>(&self, pattern: &'p Pattern<'db>) -> Matches<'p, 'db> {
match self.address_range() {
Some(range) => self.search_in(range, pattern),
None => Matches::empty(pattern),
}
}
#[must_use]
#[doc(alias("bin_search"))]
pub fn search_in<'p, 'db>(
&self,
range: Range<Address>,
pattern: &'p Pattern<'db>,
) -> Matches<'p, 'db> {
Matches {
pat: pattern,
cur: Some(range.start),
end: range.end,
}
}
}
#[doc(alias("compiled_binpat_t"))]
pub struct Pattern<'db> {
handle: *mut c_void,
flags: c_int,
_db: PhantomData<&'db Database>,
}
impl<'db> Pattern<'db> {
pub fn hex(db: &'db Database, pattern: impl AsRef<str>) -> Result<Self> {
let pattern = pattern.as_ref();
let (bytes, mask) = parse_hex(pattern).map_err(|kind| Error::PatternRejected {
pattern: pattern.to_owned(),
kind,
})?;
Self::from_parts(db, pattern.to_owned(), &bytes, Some(&mask))
}
pub fn code_mask(db: &'db Database, code: &[u8], mask: impl AsRef<str>) -> Result<Self> {
let mask = mask.as_ref();
let repr = render(code);
let mask_bytes = parse_mask(mask).map_err(|kind| Error::PatternRejected {
pattern: repr.clone(),
kind,
})?;
if mask_bytes.len() != code.len() {
return Err(Error::PatternRejected {
pattern: repr,
kind: PatternRejection::MaskMismatch {
bytes: code.len(),
mask: mask_bytes.len(),
},
});
}
Self::from_parts(db, repr, code, Some(&mask_bytes))
}
fn from_parts(
db: &'db Database,
repr: String,
bytes: &[u8],
mask: Option<&[u8]>,
) -> Result<Self> {
let _ = db; let anchors = match mask {
Some(m) => m.iter().filter(|&&b| b != 0).count(),
None => bytes.len(),
};
if anchors == 0 {
return Err(Error::PatternRejected {
pattern: repr,
kind: PatternRejection::NoAnchor { total: bytes.len() },
});
}
let mask_ptr = mask.map_or(ptr::null(), <[u8]>::as_ptr);
let handle =
unsafe { sys::idakit_binpat_from_bytes(bytes.as_ptr(), mask_ptr, bytes.len()) };
Ok(Self {
handle,
flags: sys::BIN_SEARCH_BITMASK,
_db: PhantomData,
})
}
}
#[bon::bon]
impl<'db> Pattern<'db> {
#[builder]
pub fn bytes<'a>(
#[builder(start_fn)] db: &'db Database,
#[builder(start_fn)] data: &'a [u8],
mask: Option<&'a [u8]>,
) -> Result<Self> {
if let Some(m) = mask
&& m.len() != data.len()
{
return Err(Error::PatternRejected {
pattern: render(data),
kind: PatternRejection::MaskMismatch {
bytes: data.len(),
mask: m.len(),
},
});
}
Self::from_parts(db, render(data), data, mask)
}
#[builder]
#[doc(alias("parse_binpat_str"))]
pub fn ida(
#[builder(start_fn)] db: &'db Database,
#[builder(start_fn)] pattern: impl AsRef<str>,
#[builder(default = 16)] radix: u32,
#[builder(default)] case_sensitive: bool,
) -> Result<Self> {
let pattern = pattern.as_ref();
let flags = if case_sensitive {
sys::BIN_SEARCH_CASE
} else {
0
};
let ctx = db.min_ea();
let mut err = [0u8; 256];
let handle = with_cstr(pattern, "pattern", |p| unsafe {
sys::idakit_binpat_compile(ctx, p, radix as c_int, err.as_mut_ptr().cast(), err.len())
})?;
if handle.is_null() {
let detail = unsafe { cstr(err.as_ptr().cast()) };
return Err(Error::PatternRejected {
pattern: pattern.to_owned(),
kind: PatternRejection::Unparseable {
detail: (!detail.is_empty()).then_some(detail),
},
});
}
let (mut total, mut anchors) = (0usize, 0usize);
unsafe { sys::idakit_binpat_stats(handle, &mut total, &mut anchors) };
if anchors == 0 {
unsafe { sys::idakit_binpat_free(handle) };
return Err(Error::PatternRejected {
pattern: pattern.to_owned(),
kind: PatternRejection::NoAnchor { total },
});
}
Ok(Self {
handle,
flags,
_db: PhantomData,
})
}
}
impl Drop for Pattern<'_> {
#[inline]
fn drop(&mut self) {
unsafe { sys::idakit_binpat_free(self.handle) };
}
}
impl std::fmt::Debug for Pattern<'_> {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("Pattern")
.field("flags", &self.flags)
.finish()
}
}
fn parse_hex(s: &str) -> std::result::Result<(Vec<u8>, Vec<u8>), PatternRejection> {
let mut bytes = Vec::new();
let mut mask = Vec::new();
for (index, tok) in s.split_whitespace().enumerate() {
let (b, m) = parse_hex_token(tok).ok_or_else(|| PatternRejection::BadToken {
token: tok.to_owned(),
index,
})?;
bytes.push(b);
mask.push(m);
}
Ok((bytes, mask))
}
fn parse_hex_token(tok: &str) -> Option<(u8, u8)> {
let mut chars = tok.chars();
let a = chars.next()?;
match chars.next() {
None => match nibble(a)? {
(_, 0) => Some((0, 0)), (v, _) => Some((v, 0xFF)), },
Some(b) => {
if chars.next().is_some() {
return None; }
let (hv, hm) = nibble(a)?;
let (lv, lm) = nibble(b)?;
Some(((hv << 4) | lv, (hm << 4) | lm))
}
}
}
fn nibble(c: char) -> Option<(u8, u8)> {
if c == '?' {
Some((0, 0))
} else {
c.to_digit(16).map(|d| (d as u8, 0xF))
}
}
fn parse_mask(mask: &str) -> std::result::Result<Vec<u8>, PatternRejection> {
mask.chars()
.enumerate()
.map(|(index, ch)| match ch {
'x' | 'X' => Ok(0xFF),
'?' | '.' => Ok(0x00),
_ => Err(PatternRejection::BadMaskChar { ch, index }),
})
.collect()
}
fn render(bytes: &[u8]) -> String {
bytes
.iter()
.map(|b| format!("{b:02X}"))
.collect::<Vec<_>>()
.join(" ")
}
#[doc(alias("bin_search"))]
pub struct Matches<'p, 'db> {
pat: &'p Pattern<'db>,
cur: Option<Address>,
end: Address,
}
impl<'p, 'db> Matches<'p, 'db> {
#[inline]
fn empty(pat: &'p Pattern<'db>) -> Self {
Self {
pat,
cur: None,
end: Address::new_const(0),
}
}
}
impl Iterator for Matches<'_, '_> {
type Item = Address;
fn next(&mut self) -> Option<Address> {
let start = self.cur?;
if start >= self.end {
self.cur = None;
return None;
}
let hit = unsafe {
sys::idakit_bin_search(start.get(), self.end.get(), self.pat.handle, self.pat.flags)
};
match Address::try_new(hit) {
Some(address) => {
self.cur = Some(address + 1);
Some(address)
}
None => {
self.cur = None;
None
}
}
}
}
#[cfg(test)]
mod tests {
use assert2::assert;
use rstest::rstest;
use super::*;
#[rstest]
#[case("48 8B 90", &[0x48, 0x8B, 0x90], &[0xFF, 0xFF, 0xFF])]
#[case("48 ? 90", &[0x48, 0x00, 0x90], &[0xFF, 0x00, 0xFF])]
#[case("48 ?? 90", &[0x48, 0x00, 0x90], &[0xFF, 0x00, 0xFF])]
#[case("4? ?B", &[0x40, 0x0B], &[0xF0, 0x0F])]
#[case("?? ??", &[0x00, 0x00], &[0x00, 0x00])]
#[case("3", &[0x03], &[0xFF])]
#[case(" 0A\t0b ", &[0x0A, 0x0B], &[0xFF, 0xFF])]
fn hex_parses(#[case] input: &str, #[case] bytes: &[u8], #[case] mask: &[u8]) {
assert!(let Ok((b, m)) = parse_hex(input));
assert!(b == bytes);
assert!(m == mask);
}
#[rstest]
#[case("48 GG 90", "GG", 1)]
#[case("xyz", "xyz", 0)]
#[case("48 8BB 90", "8BB", 1)] #[case("Z", "Z", 0)]
#[case("4G", "4G", 0)] fn hex_rejects_bad_token(#[case] input: &str, #[case] token: &str, #[case] index: usize) {
assert!(let Err(PatternRejection::BadToken { token: t, index: i }) = parse_hex(input));
assert!(t == token);
assert!(i == index);
}
#[rstest]
#[case("xx?x", &[0xFF, 0xFF, 0x00, 0xFF])]
#[case("X.X", &[0xFF, 0x00, 0xFF])]
fn mask_parses(#[case] input: &str, #[case] expect: &[u8]) {
assert!(let Ok(m) = parse_mask(input));
assert!(m == expect);
}
#[rstest]
#[case("xx_x", '_', 2)]
#[case("?y", 'y', 1)]
fn mask_rejects_bad_char(#[case] input: &str, #[case] ch: char, #[case] index: usize) {
assert!(let Err(PatternRejection::BadMaskChar { ch: c, index: i }) = parse_mask(input));
assert!(c == ch);
assert!(i == index);
}
}