use core::ffi::c_void;
use super::loader::{cuda_library, resolve_sym};
pub unsafe fn create_temp_context() -> *mut c_void {
let lib = unsafe { cuda_library() };
if lib.is_null() {
return core::ptr::null_mut();
}
let device_get = unsafe { resolve_sym(lib, c"cuDeviceGet") };
let ctx_create = unsafe { resolve_sym(lib, c"cuCtxCreate_v2") };
if device_get.is_null() || ctx_create.is_null() {
return core::ptr::null_mut();
}
type CuDeviceGetFn = unsafe extern "system" fn(*mut i32, i32) -> i32;
type CuCtxCreateFn = unsafe extern "system" fn(*mut *mut c_void, u32, i32) -> i32;
let cu_device_get: CuDeviceGetFn = unsafe { core::mem::transmute(device_get) };
let cu_ctx_create: CuCtxCreateFn = unsafe { core::mem::transmute(ctx_create) };
let mut dev: i32 = 0;
unsafe {
if cu_device_get(&mut dev, 0) == 0 {
let mut ctx: *mut c_void = core::ptr::null_mut();
if cu_ctx_create(&mut ctx, 0, dev) == 0 {
return ctx;
}
}
}
core::ptr::null_mut()
}
pub unsafe fn destroy_temp_context(ctx: *mut c_void) {
if ctx.is_null() {
return;
}
let lib = unsafe { cuda_library() };
if lib.is_null() {
return;
}
let ctx_destroy = unsafe { resolve_sym(lib, c"cuCtxDestroy_v2") };
if !ctx_destroy.is_null() {
type CuCtxDestroyFn = unsafe extern "system" fn(*mut c_void) -> i32;
let cu_ctx_destroy: CuCtxDestroyFn = unsafe { core::mem::transmute(ctx_destroy) };
let _destroy_status = unsafe { cu_ctx_destroy(ctx) };
}
}