Skip to main content

sim_lib_compute_cuda/
loader.rs

1//! Runtime CUDA/cuBLAS symbol discovery.
2
3use std::{
4    ffi::c_void,
5    fmt,
6    path::{Path, PathBuf},
7    sync::Arc,
8};
9
10use libloading::Library;
11
12const DRIVER_NAMES: &[&str] = &["libcuda.so.1", "libcuda.so", "nvcuda.dll"];
13const RUNTIME_NAMES: &[&str] = &[
14    "libcudart.so.13",
15    "libcudart.so.12",
16    "libcudart.so.11.0",
17    "libcudart.so",
18    "cudart64_130.dll",
19    "cudart64_12.dll",
20];
21const CUBLAS_NAMES: &[&str] = &["libcublas.so.13", "libcublas.so.12", "libcublas.so"];
22const CUBLASLT_NAMES: &[&str] = &["libcublasLt.so.13", "libcublasLt.so.12", "libcublasLt.so"];
23
24const DRIVER_SYMBOLS: &[&str] = &["cuInit", "cuDriverGetVersion"];
25const RUNTIME_SYMBOLS: &[&str] = &[
26    "cudaMalloc",
27    "cudaFree",
28    "cudaMemcpy",
29    "cudaDeviceSynchronize",
30];
31const CUBLAS_SYMBOLS: &[&str] = &[
32    "cublasCreate_v2",
33    "cublasDestroy_v2",
34    "cublasSgemm_v2",
35    "cublasGemmEx",
36];
37const CUBLASLT_SYMBOLS: &[&str] = &["cublasLtCreate", "cublasLtDestroy", "cublasLtMatmul"];
38
39/// One validated runtime symbol.
40#[derive(Clone, Debug, PartialEq, Eq)]
41pub struct CudaSymbolEvidence {
42    /// Symbol name.
43    pub name: String,
44    /// Whether the dynamic library exported the symbol.
45    pub present: bool,
46}
47
48/// Dynamic-library ABI evidence required by the CUDA provider.
49#[derive(Clone, Debug, PartialEq, Eq)]
50pub struct CudaAbiEvidence {
51    /// Loaded CUDA driver library path or platform name.
52    pub driver_library: String,
53    /// Loaded CUDA runtime library path or platform name.
54    pub runtime_library: String,
55    /// Loaded cuBLAS library path or platform name.
56    pub cublas_library: String,
57    /// Loaded cuBLASLt library path or platform name.
58    pub cublaslt_library: String,
59    /// CUDA driver version when the runtime can report it.
60    pub driver_version: Option<i32>,
61    /// Checked driver symbols.
62    pub driver_symbols: Vec<CudaSymbolEvidence>,
63    /// Checked CUDA runtime symbols.
64    pub runtime_symbols: Vec<CudaSymbolEvidence>,
65    /// Checked cuBLAS symbols.
66    pub cublas_symbols: Vec<CudaSymbolEvidence>,
67    /// Checked cuBLASLt symbols.
68    pub cublaslt_symbols: Vec<CudaSymbolEvidence>,
69}
70
71impl CudaAbiEvidence {
72    /// Returns true when all required driver/cuBLAS/cuBLASLt symbols exist.
73    pub fn is_complete(&self) -> bool {
74        self.driver_symbols.iter().all(|symbol| symbol.present)
75            && self.runtime_symbols.iter().all(|symbol| symbol.present)
76            && self.cublas_symbols.iter().all(|symbol| symbol.present)
77            && self.cublaslt_symbols.iter().all(|symbol| symbol.present)
78    }
79
80    /// Returns true when half-family matmul may use the validated cuBLASLt path.
81    pub fn supports_half_matmul(&self) -> bool {
82        self.cublas_symbols
83            .iter()
84            .any(|symbol| symbol.name == "cublasGemmEx" && symbol.present)
85            && self
86                .cublaslt_symbols
87                .iter()
88                .any(|symbol| symbol.name == "cublasLtMatmul" && symbol.present)
89    }
90}
91
92/// Loaded CUDA runtime libraries kept alive for function-pointer validity.
93pub struct CudaLibrarySet {
94    evidence: CudaAbiEvidence,
95    driver: Library,
96    runtime: Library,
97    cublas: Library,
98    cublaslt: Library,
99}
100
101impl CudaLibrarySet {
102    /// Joins capsule-loaded libraries to their validated ABI evidence.
103    pub fn new(
104        evidence: CudaAbiEvidence,
105        driver: Library,
106        runtime: Library,
107        cublas: Library,
108        cublaslt: Library,
109    ) -> Self {
110        Self {
111            evidence,
112            driver,
113            runtime,
114            cublas,
115            cublaslt,
116        }
117    }
118
119    /// Returns checked ABI evidence.
120    pub fn evidence(&self) -> &CudaAbiEvidence {
121        &self.evidence
122    }
123
124    /// Returns loaded library handles to keep symbols alive.
125    pub fn handles(&self) -> (&Library, &Library, &Library) {
126        (&self.driver, &self.cublas, &self.cublaslt)
127    }
128
129    pub(crate) fn execution_handles(&self) -> (&Library, &Library) {
130        (&self.runtime, &self.cublas)
131    }
132}
133
134impl fmt::Debug for CudaLibrarySet {
135    fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
136        formatter
137            .debug_struct("CudaLibrarySet")
138            .field("evidence", &self.evidence)
139            .finish_non_exhaustive()
140    }
141}
142
143/// Result of CUDA runtime discovery.
144#[derive(Clone, Debug)]
145pub struct CudaRuntimeProbe {
146    /// Validated loaded runtime, when discovery succeeded.
147    pub runtime: Option<Arc<CudaLibrarySet>>,
148    /// ABI evidence from the successful runtime or the best failed probe.
149    pub evidence: Option<CudaAbiEvidence>,
150    /// Diagnostics collected while searching dynamic libraries.
151    pub diagnostics: Vec<String>,
152}
153
154impl CudaRuntimeProbe {
155    /// Builds a successful probe from validated evidence without library
156    /// handles. This is intended for deterministic fake-loader tests.
157    pub fn fake_present(evidence: CudaAbiEvidence) -> Self {
158        Self {
159            runtime: None,
160            evidence: Some(evidence),
161            diagnostics: Vec::new(),
162        }
163    }
164
165    /// Returns true when discovery validated a usable CUDA provider.
166    pub fn is_available(&self) -> bool {
167        self.evidence
168            .as_ref()
169            .is_some_and(CudaAbiEvidence::is_complete)
170    }
171}
172
173/// CUDA dynamic-loading failure.
174#[derive(Clone, Debug, PartialEq, Eq)]
175pub struct CudaLoadError {
176    /// Human-readable failure message.
177    pub message: String,
178}
179
180impl fmt::Display for CudaLoadError {
181    fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
182        formatter.write_str(&self.message)
183    }
184}
185
186impl std::error::Error for CudaLoadError {}
187
188/// Loader abstraction used by real and fake CUDA discovery.
189pub trait DynamicCudaLoader {
190    /// Performs CUDA runtime discovery.
191    fn discover(&self) -> Result<CudaRuntimeProbe, CudaLoadError>;
192}
193
194/// Capsule membrane used by the provider to receive an explicit probe.
195pub trait CudaProbePort {
196    /// Returns a capsule-owned CUDA probe without ambient rediscovery.
197    fn probe_cuda(&self) -> Result<CudaRuntimeProbe, CudaLoadError>;
198}
199
200impl<T: DynamicCudaLoader + ?Sized> CudaProbePort for T {
201    fn probe_cuda(&self) -> Result<CudaRuntimeProbe, CudaLoadError> {
202        self.discover()
203    }
204}
205
206/// Real dynamic loader using platform CUDA shared libraries.
207#[derive(Clone, Debug)]
208pub struct CudaRuntimeLoader {
209    search_dirs: Vec<PathBuf>,
210    search_system: bool,
211}
212
213impl Default for CudaRuntimeLoader {
214    fn default() -> Self {
215        Self {
216            search_dirs: Vec::new(),
217            search_system: true,
218        }
219    }
220}
221
222impl CudaRuntimeLoader {
223    /// Builds a loader that searches platform library paths.
224    pub fn new() -> Self {
225        Self::default()
226    }
227
228    /// Builds a loader that first searches explicit directories.
229    pub fn with_search_dirs(search_dirs: Vec<PathBuf>) -> Self {
230        Self {
231            search_dirs,
232            search_system: true,
233        }
234    }
235
236    /// Builds a loader restricted to explicit directories. This is used to
237    /// prove fail-closed behavior when vendor libraries are unavailable.
238    pub fn with_search_dirs_only(search_dirs: Vec<PathBuf>) -> Self {
239        Self {
240            search_dirs,
241            search_system: false,
242        }
243    }
244}
245
246impl DynamicCudaLoader for CudaRuntimeLoader {
247    fn discover(&self) -> Result<CudaRuntimeProbe, CudaLoadError> {
248        let mut diagnostics = Vec::new();
249        let (driver_name, driver) = self.open_first(DRIVER_NAMES, &mut diagnostics)?;
250        let (runtime_name, runtime) = self.open_first(RUNTIME_NAMES, &mut diagnostics)?;
251        let (cublas_name, cublas) = self.open_first(CUBLAS_NAMES, &mut diagnostics)?;
252        let (cublaslt_name, cublaslt) = self.open_first(CUBLASLT_NAMES, &mut diagnostics)?;
253
254        let driver_symbols = symbol_evidence(&driver, DRIVER_SYMBOLS);
255        let runtime_symbols = symbol_evidence(&runtime, RUNTIME_SYMBOLS);
256        let cublas_symbols = symbol_evidence(&cublas, CUBLAS_SYMBOLS);
257        let cublaslt_symbols = symbol_evidence(&cublaslt, CUBLASLT_SYMBOLS);
258        let driver_version = driver_version(&driver).ok();
259        let evidence = CudaAbiEvidence {
260            driver_library: driver_name,
261            runtime_library: runtime_name,
262            cublas_library: cublas_name,
263            cublaslt_library: cublaslt_name,
264            driver_version,
265            driver_symbols,
266            runtime_symbols,
267            cublas_symbols,
268            cublaslt_symbols,
269        };
270        if !evidence.is_complete() {
271            return Ok(CudaRuntimeProbe {
272                runtime: None,
273                evidence: Some(evidence),
274                diagnostics,
275            });
276        }
277        let runtime = Arc::new(CudaLibrarySet::new(
278            evidence.clone(),
279            driver,
280            runtime,
281            cublas,
282            cublaslt,
283        ));
284        Ok(CudaRuntimeProbe {
285            runtime: Some(runtime),
286            evidence: Some(evidence),
287            diagnostics,
288        })
289    }
290}
291
292impl CudaRuntimeLoader {
293    fn open_first(
294        &self,
295        names: &[&str],
296        diagnostics: &mut Vec<String>,
297    ) -> Result<(String, Library), CudaLoadError> {
298        for name in candidate_paths(&self.search_dirs, names, self.search_system) {
299            match open_library(&name) {
300                Ok(library) => return Ok((name.display().to_string(), library)),
301                Err(error) => diagnostics.push(format!("{}: {error}", name.display())),
302            }
303        }
304        Err(CudaLoadError {
305            message: format!("CUDA library was not found; tried {}", names.join(", ")),
306        })
307    }
308}
309
310/// Fake loader for deterministic tests.
311#[derive(Clone, Debug)]
312pub struct FakeCudaLoader {
313    probe: Result<CudaRuntimeProbe, CudaLoadError>,
314}
315
316impl FakeCudaLoader {
317    /// Builds a fake loader that returns validated CUDA evidence.
318    pub fn available() -> Self {
319        Self {
320            probe: Ok(CudaRuntimeProbe::fake_present(complete_fake_evidence())),
321        }
322    }
323
324    /// Builds a fake loader with incomplete ABI evidence.
325    pub fn incomplete() -> Self {
326        let mut evidence = complete_fake_evidence();
327        if let Some(symbol) = evidence
328            .cublaslt_symbols
329            .iter_mut()
330            .find(|symbol| symbol.name == "cublasLtMatmul")
331        {
332            symbol.present = false;
333        }
334        Self {
335            probe: Ok(CudaRuntimeProbe {
336                runtime: None,
337                evidence: Some(evidence),
338                diagnostics: vec!["missing cublasLtMatmul".to_owned()],
339            }),
340        }
341    }
342
343    /// Builds a fake loader that reports CUDA as absent.
344    pub fn absent() -> Self {
345        Self {
346            probe: Err(CudaLoadError {
347                message: "CUDA runtime absent".to_owned(),
348            }),
349        }
350    }
351}
352
353impl DynamicCudaLoader for FakeCudaLoader {
354    fn discover(&self) -> Result<CudaRuntimeProbe, CudaLoadError> {
355        self.probe.clone()
356    }
357}
358
359/// Discovers CUDA using the real platform dynamic loader.
360pub fn discover_cuda_runtime() -> Result<CudaRuntimeProbe, CudaLoadError> {
361    CudaRuntimeLoader::new().discover()
362}
363
364fn complete_fake_evidence() -> CudaAbiEvidence {
365    CudaAbiEvidence {
366        driver_library: "fake-libcuda".to_owned(),
367        runtime_library: "fake-libcudart".to_owned(),
368        cublas_library: "fake-libcublas".to_owned(),
369        cublaslt_library: "fake-libcublasLt".to_owned(),
370        driver_version: Some(12_000),
371        driver_symbols: DRIVER_SYMBOLS
372            .iter()
373            .map(|name| CudaSymbolEvidence {
374                name: (*name).to_owned(),
375                present: true,
376            })
377            .collect(),
378        runtime_symbols: RUNTIME_SYMBOLS
379            .iter()
380            .map(|name| CudaSymbolEvidence {
381                name: (*name).to_owned(),
382                present: true,
383            })
384            .collect(),
385        cublas_symbols: CUBLAS_SYMBOLS
386            .iter()
387            .map(|name| CudaSymbolEvidence {
388                name: (*name).to_owned(),
389                present: true,
390            })
391            .collect(),
392        cublaslt_symbols: CUBLASLT_SYMBOLS
393            .iter()
394            .map(|name| CudaSymbolEvidence {
395                name: (*name).to_owned(),
396                present: true,
397            })
398            .collect(),
399    }
400}
401
402fn candidate_paths(search_dirs: &[PathBuf], names: &[&str], search_system: bool) -> Vec<PathBuf> {
403    let mut candidates = Vec::new();
404    for directory in search_dirs {
405        for name in names {
406            candidates.push(directory.join(name));
407        }
408    }
409    if search_system {
410        candidates.extend(names.iter().map(PathBuf::from));
411    }
412    candidates
413}
414
415fn symbol_evidence(library: &Library, names: &[&str]) -> Vec<CudaSymbolEvidence> {
416    names
417        .iter()
418        .map(|name| CudaSymbolEvidence {
419            name: (*name).to_owned(),
420            present: symbol_present(library, name),
421        })
422        .collect()
423}
424
425fn open_library(path: &Path) -> Result<Library, libloading::Error> {
426    // SAFETY: Loading a CUDA shared library is the intended boundary of this
427    // crate. The handle is stored in CudaLibrarySet for at least as long as any
428    // validated symbol evidence derived from it is used.
429    unsafe { Library::new(path) }
430}
431
432fn symbol_present(library: &Library, name: &str) -> bool {
433    let mut bytes = name.as_bytes().to_vec();
434    bytes.push(0);
435    // SAFETY: The lookup only checks whether the library exports the named
436    // symbol as an opaque address. The address is not called or dereferenced.
437    unsafe { library.get::<*mut c_void>(&bytes).is_ok() }
438}
439
440fn driver_version(library: &Library) -> Result<i32, CudaLoadError> {
441    type CuInit = unsafe extern "C" fn(u32) -> i32;
442    type CuDriverGetVersion = unsafe extern "C" fn(*mut i32) -> i32;
443    // SAFETY: Symbols were loaded from the CUDA driver library by their official
444    // C ABI names. The calls use the documented signatures for cuInit and
445    // cuDriverGetVersion and pass initialized pointers.
446    unsafe {
447        let cu_init = library
448            .get::<CuInit>(b"cuInit\0")
449            .map_err(|error| CudaLoadError {
450                message: error.to_string(),
451            })?;
452        let get_version = library
453            .get::<CuDriverGetVersion>(b"cuDriverGetVersion\0")
454            .map_err(|error| CudaLoadError {
455                message: error.to_string(),
456            })?;
457        let init_status = cu_init(0);
458        if init_status != 0 {
459            return Err(CudaLoadError {
460                message: format!("cuInit failed with status {init_status}"),
461            });
462        }
463        let mut version = 0;
464        let version_status = get_version(&mut version);
465        if version_status != 0 {
466            return Err(CudaLoadError {
467                message: format!("cuDriverGetVersion failed with status {version_status}"),
468            });
469        }
470        Ok(version)
471    }
472}