use core::ffi::c_void;
use core::sync::atomic::{AtomicBool, AtomicI32, AtomicPtr, Ordering};
#[repr(C)]
struct EXCEPTION_RECORD {
exception_code: u32,
exception_flags: u32,
exception_record: *mut EXCEPTION_RECORD,
exception_address: *mut c_void,
number_parameters: u32,
exception_information: [usize; 15],
}
#[repr(C)]
struct EXCEPTION_POINTERS {
exception_record: *mut EXCEPTION_RECORD,
context_record: *mut c_void,
}
unsafe extern "system" {
fn GetModuleHandleA(lpModuleName: *const u8) -> *mut c_void;
fn CreateThread(
lpThreadAttributes: *mut c_void,
dwStackSize: usize,
lpStartAddress: unsafe extern "system" fn(*mut c_void) -> u32,
lpParameter: *mut c_void,
dwCreationFlags: u32,
lpThreadId: *mut u32,
) -> *mut c_void;
fn WaitForSingleObject(hHandle: *mut c_void, dwMilliseconds: u32) -> u32;
fn GetExitCodeThread(hThread: *mut c_void, lpExitCode: *mut u32) -> i32;
fn CloseHandle(hObject: *mut c_void) -> i32;
fn AddVectoredExceptionHandler(
FirstHandler: u32,
VectoredHandler: unsafe extern "system" fn(*mut EXCEPTION_POINTERS) -> i32,
) -> *mut c_void;
fn RemoveVectoredExceptionHandler(Handler: *mut c_void) -> u32;
#[cfg(target_arch = "x86_64")]
fn ExitThread(dwExitCode: u32) -> !;
}
const STATUS_ACCESS_VIOLATION: u32 = 0xC000_0005;
const STATUS_IN_PAGE_ERROR: u32 = 0xC000_0006;
const EXCEPTION_CONTINUE_SEARCH: i32 = 0;
#[cfg(target_arch = "x86_64")]
const EXCEPTION_CONTINUE_EXECUTION: i32 = -1;
const PROBE_TIMEOUT_MS: u32 = 5000;
const CU_INIT_FAILED: i32 = -1;
#[cfg(target_arch = "x86_64")]
const CONTEXT_RSP_OFFSET: usize = 0x98;
#[cfg(target_arch = "x86_64")]
const CONTEXT_RIP_OFFSET: usize = 0xF8;
static PROBE_ACTIVE: AtomicBool = AtomicBool::new(false);
static CU_INIT_PTR: AtomicPtr<c_void> = AtomicPtr::new(core::ptr::null_mut());
static CU_INIT_RESULT: AtomicI32 = AtomicI32::new(CU_INIT_FAILED);
unsafe extern "system" fn worker_thread_fn(_param: *mut c_void) -> u32 {
type CuInitFn = unsafe extern "system" fn(u32) -> i32;
let cu_init: CuInitFn = unsafe { core::mem::transmute(CU_INIT_PTR.load(Ordering::Acquire)) };
let res = unsafe { cu_init(0) };
CU_INIT_RESULT.store(res, Ordering::Release);
0
}
unsafe fn is_address_in_nvcuda(addr: *mut c_void) -> bool {
let h_module = unsafe { GetModuleHandleA(c"nvcuda.dll".as_ptr() as *const u8) };
if h_module.is_null() {
return false;
}
let base = h_module as usize;
let addr_val = addr as usize;
if addr_val < base {
return false;
}
unsafe {
let dos_header = base as *const u8;
let e_lfanew = *(dos_header.add(0x3c) as *const i32) as usize;
let nt_headers = base + e_lfanew;
let size_of_image = *((nt_headers + 24 + 56) as *const u32) as usize;
addr_val < base + size_of_image
}
}
#[cfg(target_arch = "x86_64")]
unsafe extern "system" fn probe_thread_exit() -> ! {
unsafe { ExitThread(1) }
}
unsafe extern "system" fn probe_veh_handler(exception_info: *mut EXCEPTION_POINTERS) -> i32 {
if !PROBE_ACTIVE.load(Ordering::Acquire) {
return EXCEPTION_CONTINUE_SEARCH;
}
let (code, addr) = unsafe {
let record = (*exception_info).exception_record;
((*record).exception_code, (*record).exception_address)
};
let is_probe_fault = code == STATUS_IN_PAGE_ERROR
|| (code == STATUS_ACCESS_VIOLATION
&& unsafe { is_address_in_nvcuda(addr) });
if !is_probe_fault {
return EXCEPTION_CONTINUE_SEARCH;
}
CU_INIT_RESULT.store(CU_INIT_FAILED, Ordering::Release);
#[cfg(target_arch = "x86_64")]
{
unsafe {
let context = (*exception_info).context_record as *mut u8;
let rsp = context.add(CONTEXT_RSP_OFFSET) as *mut u64;
let rip = context.add(CONTEXT_RIP_OFFSET) as *mut u64;
*rsp = (*rsp & !0xF).wrapping_sub(8);
*rip = probe_thread_exit as *const () as usize as u64;
}
EXCEPTION_CONTINUE_EXECUTION
}
#[cfg(not(target_arch = "x86_64"))]
{
EXCEPTION_CONTINUE_SEARCH
}
}
pub(super) unsafe fn run_cu_init_isolated(init_sym: *mut c_void) -> i32 {
CU_INIT_PTR.store(init_sym, Ordering::Release);
CU_INIT_RESULT.store(CU_INIT_FAILED, Ordering::Release);
let handler = unsafe { AddVectoredExceptionHandler(1, probe_veh_handler) };
if handler.is_null() {
return CU_INIT_FAILED;
}
PROBE_ACTIVE.store(true, Ordering::Release);
let mut thread_id = 0_u32;
let thread_handle = unsafe {
CreateThread(
core::ptr::null_mut(),
0,
worker_thread_fn,
core::ptr::null_mut(),
0,
&mut thread_id,
)
};
if thread_handle.is_null() {
PROBE_ACTIVE.store(false, Ordering::Release);
unsafe { RemoveVectoredExceptionHandler(handler) };
return CU_INIT_FAILED;
}
let mut exit_code = 0_u32;
unsafe {
WaitForSingleObject(thread_handle, PROBE_TIMEOUT_MS);
GetExitCodeThread(thread_handle, &mut exit_code);
CloseHandle(thread_handle);
}
PROBE_ACTIVE.store(false, Ordering::Release);
unsafe { RemoveVectoredExceptionHandler(handler) };
if exit_code == 0 {
CU_INIT_RESULT.load(Ordering::Acquire)
} else {
CU_INIT_FAILED
}
}