use crate::breakpoints::{Breakpoint, BreakpointMap, ModuleBaseCache};
use crate::defs::{HookAction, HookContext, HookHandler};
use crate::utils::{
get_thread_context, normalize_name, read_remote_or, read_single_byte, resolve_export,
set_thread_context, write_process_memory, write_single_byte,
};
use cradle_shared::{CradleError, CradleResult};
use windows_sys::Win32::Foundation::HANDLE;
pub enum DispatchResult {
Handled,
Skip,
}
pub struct HookEngine {
breakpoints: BreakpointMap,
module_cache: ModuleBaseCache,
process: HANDLE,
pending_restore: Option<PendingRestore>,
}
struct PendingRestore {
addr: usize,
thread: HANDLE,
}
impl HookEngine {
pub fn empty() -> HookEngine {
Self {
breakpoints: BreakpointMap::new(),
module_cache: ModuleBaseCache::new(),
process: HANDLE::default(),
pending_restore: None,
}
}
}
impl HookEngine {
pub fn new(process: HANDLE) -> Self {
Self {
breakpoints: BreakpointMap::new(),
module_cache: ModuleBaseCache::new(),
process,
pending_restore: None,
}
}
pub fn register_module(&mut self, name: &str, base: usize) {
self.module_cache.insert(normalize_name(name), base);
}
pub fn has_breakpoint(&self, addr: usize) -> bool {
self.breakpoints.contains_key(&addr)
}
pub unsafe fn dispatch(&mut self, thread: HANDLE) -> CradleResult<DispatchResult> {
unsafe {
let mut ctx = get_thread_context(thread).map_err(|e| {
CradleError::InvalidValue(format!("get thread context failed: {e}"))
})?;
let bp_addr = (ctx.Rip - 1) as usize;
let bp = match self.breakpoints.get(&bp_addr) {
Some(bp) => bp,
None => return Ok(DispatchResult::Skip),
};
let mut hook_ctx = HookContext {
registers: &mut ctx,
process: self.process,
ret_value: 0,
export_name: &bp.export_name,
module_name: &bp.module_name,
#[cfg(feature = "unstable")]
return_hook: None,
};
let mut action = HookAction::Continue;
for handler in &bp.handlers {
match handler(&mut hook_ctx) {
Ok(HookAction::Continue) => {}
Ok(HookAction::Modify) => action = HookAction::Modify,
Ok(block @ HookAction::Block(_)) => {
action = block;
break;
}
Err(_) => {}
}
}
let ret = hook_ctx.ret_value;
match action {
HookAction::Continue | HookAction::Modify => {
write_process_memory(self.process, bp_addr, &[bp.original_byte]).map_err(
|e| {
CradleError::InvalidValue(format!(
"{}!{} restore byte at {bp_addr:#x} failed: {e}",
bp.module_name, bp.export_name
))
},
)?;
ctx.Rip = bp_addr as u64;
ctx.EFlags |= 0x100;
set_thread_context(thread, &ctx).map_err(|e| {
CradleError::InvalidValue(format!(
"{}!{} set context failed at {bp_addr:#x}: {e}",
bp.module_name, bp.export_name
))
})?;
self.pending_restore = Some(PendingRestore {
addr: bp_addr,
thread,
});
}
HookAction::Block(u) => {
let ret_addr: u64 = read_remote_or(self.process, ctx.Rsp as usize, u);
ctx.Rip = ret_addr;
ctx.Rsp += 8;
ctx.Rax = ret;
set_thread_context(thread, &ctx).map_err(|e| {
CradleError::InvalidValue(format!(
"{}!{} block set context failed at {bp_addr:#x}: {e}",
bp.module_name, bp.export_name
))
})?;
}
}
Ok(DispatchResult::Handled)
}
}
pub unsafe fn dispatch_single_step(&mut self, thread: HANDLE) -> CradleResult<DispatchResult> {
match self.pending_restore.take() {
Some(restore) if restore.thread == thread => {
if self.breakpoints.contains_key(&restore.addr) {
unsafe {
write_process_memory(self.process, restore.addr, &[0xCC])?;
}
}
Ok(DispatchResult::Handled)
}
Some(restore) => {
self.pending_restore = Some(restore);
Ok(DispatchResult::Skip)
}
None => Ok(DispatchResult::Skip),
}
}
pub unsafe fn hook_export(
&mut self,
module: &str,
export: &str,
handler: HookHandler,
) -> CradleResult {
if self.module_cache.is_empty() {
return Err(CradleError::ModuleNotFound(
"module cache is empty".to_owned(),
));
}
let module_lower = normalize_name(module);
let base = self
.module_cache
.get(&module_lower)
.copied()
.ok_or_else(|| CradleError::ModuleNotFound(module.to_string()))?;
let addr = resolve_export(self.process, base, export)?;
if let Some(bp) = self.breakpoints.get_mut(&addr) {
bp.push_handler(handler);
return Ok(());
}
let mut orig = [0u8; 1];
unsafe {
read_single_byte(self.process, addr, &mut orig)?;
write_single_byte(self.process, addr, &[0xCC])?;
}
self.breakpoints.insert(
addr,
Breakpoint::new(orig[0], handler, module_lower, export.to_string()),
);
Ok(())
}
pub unsafe fn unhook(&mut self, addr: usize) -> CradleResult {
if let Some(bp) = self.breakpoints.remove(&addr) {
unsafe {
write_process_memory(self.process, addr, &[bp.original_byte])?;
}
}
Ok(())
}
pub unsafe fn unhook_all(&mut self) -> CradleResult {
let addrs: Vec<usize> = self.breakpoints.keys().copied().collect();
for addr in addrs {
unsafe {
self.unhook(addr)?;
}
}
Ok(())
}
pub fn hooked_count(&self) -> usize {
self.breakpoints.len()
}
}
unsafe impl Send for HookEngine {}