use std::ffi::CStr;
use std::mem::MaybeUninit;
use std::{
error,
fmt::{self, Display, Formatter},
};
#[derive(Clone, Copy, PartialEq, Eq)]
pub struct DriverError(pub cuda_bindings::CUresult);
impl DriverError {
pub fn is_unsupported_ptx_version(&self) -> bool {
self.0 == cuda_bindings::cudaError_enum_CUDA_ERROR_UNSUPPORTED_PTX_VERSION
}
fn _fmt(&self, formatter: &mut Formatter) -> fmt::Result {
self.fmt_with_loader_error(formatter, cuda_bindings::cuda_driver_load_error())
}
fn fmt_with_loader_error(
&self,
formatter: &mut Formatter,
load_error: Option<&cuda_bindings::DynLoadError>,
) -> fmt::Result {
if let Some(load_error) = load_error {
if self.0 == cuda_bindings::cudaError_enum_CUDA_ERROR_NOT_INITIALIZED
|| self.0 == cuda_bindings::cudaError_enum_CUDA_ERROR_SHARED_OBJECT_INIT_FAILED
{
return formatter
.debug_tuple("DriverError")
.field(&self.0)
.field(&format!("CUDA driver library unavailable: {load_error}"))
.finish();
}
}
let help = "the CUDA driver cannot JIT PTX from the selected toolkit; upgrade the driver \
or select a compatible toolkit with CUDA_TOOLKIT_PATH or CUDA_HOME";
let mut output = formatter.debug_tuple("DriverError");
output.field(&self.0);
match self.error_string() {
Ok(err_str) => {
output.field(&err_str);
}
Err(_) => {
output.field(&"<Failure when calling cuGetErrorString()>");
}
}
if self.is_unsupported_ptx_version() {
output.field(&help);
}
output.finish()
}
}
impl Display for DriverError {
fn fmt(&self, formatter: &mut Formatter) -> fmt::Result {
self._fmt(formatter)
}
}
impl std::fmt::Debug for DriverError {
fn fmt(&self, formatter: &mut std::fmt::Formatter) -> fmt::Result {
self._fmt(formatter)
}
}
impl error::Error for DriverError {}
pub trait IntoResult<T> {
fn result(self) -> Result<T, DriverError>
where
Self: Sized;
}
impl IntoResult<()> for cuda_bindings::CUresult {
fn result(self) -> Result<(), DriverError> {
match self {
cuda_bindings::cudaError_enum_CUDA_SUCCESS => Ok(()),
_ => Err(DriverError(self)),
}
}
}
impl<T> IntoResult<T> for (cuda_bindings::CUresult, T) {
fn result(self) -> Result<T, DriverError> {
match self.0 {
cuda_bindings::cudaError_enum_CUDA_SUCCESS => Ok(self.1),
_ => Err(DriverError(self.0)),
}
}
}
impl<T> IntoResult<T> for (cuda_bindings::CUresult, MaybeUninit<T>) {
fn result(self) -> Result<T, DriverError> {
match self.0 {
cuda_bindings::cudaError_enum_CUDA_SUCCESS => Ok(unsafe { self.1.assume_init() }),
_ => Err(DriverError(self.0)),
}
}
}
impl DriverError {
pub fn error_name(&self) -> Result<&CStr, DriverError> {
let mut err_str = MaybeUninit::uninit();
unsafe {
cuda_bindings::cuGetErrorName(self.0, err_str.as_mut_ptr()).result()?;
Ok(CStr::from_ptr(err_str.assume_init()))
}
}
pub fn error_string(&self) -> Result<&CStr, DriverError> {
let mut err_str = MaybeUninit::uninit();
unsafe {
cuda_bindings::cuGetErrorString(self.0, err_str.as_mut_ptr()).result()?;
Ok(CStr::from_ptr(err_str.assume_init()))
}
}
}
#[cfg(test)]
mod tests {
use super::DriverError;
#[test]
fn identifies_unsupported_ptx_version() {
let unsupported =
DriverError(cuda_bindings::cudaError_enum_CUDA_ERROR_UNSUPPORTED_PTX_VERSION);
let unrelated = DriverError(cuda_bindings::cudaError_enum_CUDA_ERROR_INVALID_VALUE);
assert!(unsupported.is_unsupported_ptx_version());
assert!(!unrelated.is_unsupported_ptx_version());
}
const NOT_INITIALIZED: cuda_bindings::CUresult =
cuda_bindings::cudaError_enum_CUDA_ERROR_NOT_INITIALIZED;
const SHARED_OBJECT_INIT_FAILED: cuda_bindings::CUresult =
cuda_bindings::cudaError_enum_CUDA_ERROR_SHARED_OBJECT_INIT_FAILED;
fn format_with(error: DriverError, load_error: Option<&cuda_bindings::DynLoadError>) -> String {
struct WithLoader<'a>(DriverError, Option<&'a cuda_bindings::DynLoadError>);
impl std::fmt::Display for WithLoader<'_> {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
self.0.fmt_with_loader_error(f, self.1)
}
}
WithLoader(error, load_error).to_string()
}
#[test]
fn loader_hint_is_attached_to_both_unavailable_driver_codes() {
let load_error = cuda_bindings::DynLoadError::RuntimeTooOld {
compile_version: 13030,
runtime_version: 12040,
};
for code in [NOT_INITIALIZED, SHARED_OBJECT_INIT_FAILED] {
let formatted = format_with(DriverError(code), Some(&load_error));
assert!(
formatted.contains("CUDA driver library unavailable"),
"code {code}: expected the loader hint, got: {formatted}"
);
assert!(
formatted.contains("CUDA driver too old"),
"code {code}: expected the loader's own message, got: {formatted}"
);
}
let unrelated = format_with(
DriverError(cuda_bindings::cudaError_enum_CUDA_ERROR_INVALID_VALUE),
Some(&load_error),
);
assert!(
!unrelated.contains("CUDA driver library unavailable"),
"an unrelated code must not carry the loader hint, got: {unrelated}"
);
for code in [NOT_INITIALIZED, SHARED_OBJECT_INIT_FAILED] {
let formatted = format_with(DriverError(code), None);
assert!(
!formatted.contains("CUDA driver library unavailable"),
"code {code}: no loader failure, no hint; got: {formatted}"
);
}
}
#[test]
fn display_surfaces_the_real_loader_error_when_driver_is_missing() {
let Some(load_error) = cuda_bindings::cuda_driver_load_error() else {
return;
};
let expected_detail = match load_error {
cuda_bindings::DynLoadError::LoadFailed { .. } => "failed to load any of",
cuda_bindings::DynLoadError::RuntimeTooOld { .. } => "CUDA driver too old",
};
for code in [NOT_INITIALIZED, SHARED_OBJECT_INIT_FAILED] {
let formatted = DriverError(code).to_string();
assert!(
formatted.contains("CUDA driver library unavailable"),
"code {code}: expected a human-readable loader hint, got: {formatted}"
);
assert!(
formatted.contains(expected_detail),
"code {code}: expected the cached loader failure context, got: {formatted}"
);
}
}
}