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    fn new(
103        evidence: CudaAbiEvidence,
104        driver: Library,
105        runtime: Library,
106        cublas: Library,
107        cublaslt: Library,
108    ) -> Self {
109        Self {
110            evidence,
111            driver,
112            runtime,
113            cublas,
114            cublaslt,
115        }
116    }
117
118    /// Returns checked ABI evidence.
119    pub fn evidence(&self) -> &CudaAbiEvidence {
120        &self.evidence
121    }
122
123    /// Returns loaded library handles to keep symbols alive.
124    pub fn handles(&self) -> (&Library, &Library, &Library) {
125        (&self.driver, &self.cublas, &self.cublaslt)
126    }
127
128    pub(crate) fn execution_handles(&self) -> (&Library, &Library) {
129        (&self.runtime, &self.cublas)
130    }
131}
132
133impl fmt::Debug for CudaLibrarySet {
134    fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
135        formatter
136            .debug_struct("CudaLibrarySet")
137            .field("evidence", &self.evidence)
138            .finish_non_exhaustive()
139    }
140}
141
142/// Result of CUDA runtime discovery.
143#[derive(Clone, Debug)]
144pub struct CudaRuntimeProbe {
145    /// Validated loaded runtime, when discovery succeeded.
146    pub runtime: Option<Arc<CudaLibrarySet>>,
147    /// ABI evidence from the successful runtime or the best failed probe.
148    pub evidence: Option<CudaAbiEvidence>,
149    /// Diagnostics collected while searching dynamic libraries.
150    pub diagnostics: Vec<String>,
151}
152
153impl CudaRuntimeProbe {
154    /// Builds a successful probe from validated evidence without library
155    /// handles. This is intended for deterministic fake-loader tests.
156    pub fn fake_present(evidence: CudaAbiEvidence) -> Self {
157        Self {
158            runtime: None,
159            evidence: Some(evidence),
160            diagnostics: Vec::new(),
161        }
162    }
163
164    /// Returns true when discovery validated a usable CUDA provider.
165    pub fn is_available(&self) -> bool {
166        self.evidence
167            .as_ref()
168            .is_some_and(CudaAbiEvidence::is_complete)
169    }
170}
171
172/// CUDA dynamic-loading failure.
173#[derive(Clone, Debug, PartialEq, Eq)]
174pub struct CudaLoadError {
175    /// Human-readable failure message.
176    pub message: String,
177}
178
179impl fmt::Display for CudaLoadError {
180    fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
181        formatter.write_str(&self.message)
182    }
183}
184
185impl std::error::Error for CudaLoadError {}
186
187/// Loader abstraction used by real and fake CUDA discovery.
188pub trait DynamicCudaLoader {
189    /// Performs CUDA runtime discovery.
190    fn discover(&self) -> Result<CudaRuntimeProbe, CudaLoadError>;
191}
192
193/// Real dynamic loader using platform CUDA shared libraries.
194#[derive(Clone, Debug)]
195pub struct CudaRuntimeLoader {
196    search_dirs: Vec<PathBuf>,
197    search_system: bool,
198}
199
200impl Default for CudaRuntimeLoader {
201    fn default() -> Self {
202        Self {
203            search_dirs: Vec::new(),
204            search_system: true,
205        }
206    }
207}
208
209impl CudaRuntimeLoader {
210    /// Builds a loader that searches platform library paths.
211    pub fn new() -> Self {
212        Self::default()
213    }
214
215    /// Builds a loader that first searches explicit directories.
216    pub fn with_search_dirs(search_dirs: Vec<PathBuf>) -> Self {
217        Self {
218            search_dirs,
219            search_system: true,
220        }
221    }
222
223    /// Builds a loader restricted to explicit directories. This is used to
224    /// prove fail-closed behavior when vendor libraries are unavailable.
225    pub fn with_search_dirs_only(search_dirs: Vec<PathBuf>) -> Self {
226        Self {
227            search_dirs,
228            search_system: false,
229        }
230    }
231}
232
233impl DynamicCudaLoader for CudaRuntimeLoader {
234    fn discover(&self) -> Result<CudaRuntimeProbe, CudaLoadError> {
235        let mut diagnostics = Vec::new();
236        let (driver_name, driver) = self.open_first(DRIVER_NAMES, &mut diagnostics)?;
237        let (runtime_name, runtime) = self.open_first(RUNTIME_NAMES, &mut diagnostics)?;
238        let (cublas_name, cublas) = self.open_first(CUBLAS_NAMES, &mut diagnostics)?;
239        let (cublaslt_name, cublaslt) = self.open_first(CUBLASLT_NAMES, &mut diagnostics)?;
240
241        let driver_symbols = symbol_evidence(&driver, DRIVER_SYMBOLS);
242        let runtime_symbols = symbol_evidence(&runtime, RUNTIME_SYMBOLS);
243        let cublas_symbols = symbol_evidence(&cublas, CUBLAS_SYMBOLS);
244        let cublaslt_symbols = symbol_evidence(&cublaslt, CUBLASLT_SYMBOLS);
245        let driver_version = driver_version(&driver).ok();
246        let evidence = CudaAbiEvidence {
247            driver_library: driver_name,
248            runtime_library: runtime_name,
249            cublas_library: cublas_name,
250            cublaslt_library: cublaslt_name,
251            driver_version,
252            driver_symbols,
253            runtime_symbols,
254            cublas_symbols,
255            cublaslt_symbols,
256        };
257        if !evidence.is_complete() {
258            return Ok(CudaRuntimeProbe {
259                runtime: None,
260                evidence: Some(evidence),
261                diagnostics,
262            });
263        }
264        let runtime = Arc::new(CudaLibrarySet::new(
265            evidence.clone(),
266            driver,
267            runtime,
268            cublas,
269            cublaslt,
270        ));
271        Ok(CudaRuntimeProbe {
272            runtime: Some(runtime),
273            evidence: Some(evidence),
274            diagnostics,
275        })
276    }
277}
278
279impl CudaRuntimeLoader {
280    fn open_first(
281        &self,
282        names: &[&str],
283        diagnostics: &mut Vec<String>,
284    ) -> Result<(String, Library), CudaLoadError> {
285        for name in candidate_paths(&self.search_dirs, names, self.search_system) {
286            match open_library(&name) {
287                Ok(library) => return Ok((name.display().to_string(), library)),
288                Err(error) => diagnostics.push(format!("{}: {error}", name.display())),
289            }
290        }
291        Err(CudaLoadError {
292            message: format!("CUDA library was not found; tried {}", names.join(", ")),
293        })
294    }
295}
296
297/// Fake loader for deterministic tests.
298#[derive(Clone, Debug)]
299pub struct FakeCudaLoader {
300    probe: Result<CudaRuntimeProbe, CudaLoadError>,
301}
302
303impl FakeCudaLoader {
304    /// Builds a fake loader that returns validated CUDA evidence.
305    pub fn available() -> Self {
306        Self {
307            probe: Ok(CudaRuntimeProbe::fake_present(complete_fake_evidence())),
308        }
309    }
310
311    /// Builds a fake loader with incomplete ABI evidence.
312    pub fn incomplete() -> Self {
313        let mut evidence = complete_fake_evidence();
314        if let Some(symbol) = evidence
315            .cublaslt_symbols
316            .iter_mut()
317            .find(|symbol| symbol.name == "cublasLtMatmul")
318        {
319            symbol.present = false;
320        }
321        Self {
322            probe: Ok(CudaRuntimeProbe {
323                runtime: None,
324                evidence: Some(evidence),
325                diagnostics: vec!["missing cublasLtMatmul".to_owned()],
326            }),
327        }
328    }
329
330    /// Builds a fake loader that reports CUDA as absent.
331    pub fn absent() -> Self {
332        Self {
333            probe: Err(CudaLoadError {
334                message: "CUDA runtime absent".to_owned(),
335            }),
336        }
337    }
338}
339
340impl DynamicCudaLoader for FakeCudaLoader {
341    fn discover(&self) -> Result<CudaRuntimeProbe, CudaLoadError> {
342        self.probe.clone()
343    }
344}
345
346/// Discovers CUDA using the real platform dynamic loader.
347pub fn discover_cuda_runtime() -> Result<CudaRuntimeProbe, CudaLoadError> {
348    CudaRuntimeLoader::new().discover()
349}
350
351fn complete_fake_evidence() -> CudaAbiEvidence {
352    CudaAbiEvidence {
353        driver_library: "fake-libcuda".to_owned(),
354        runtime_library: "fake-libcudart".to_owned(),
355        cublas_library: "fake-libcublas".to_owned(),
356        cublaslt_library: "fake-libcublasLt".to_owned(),
357        driver_version: Some(12_000),
358        driver_symbols: DRIVER_SYMBOLS
359            .iter()
360            .map(|name| CudaSymbolEvidence {
361                name: (*name).to_owned(),
362                present: true,
363            })
364            .collect(),
365        runtime_symbols: RUNTIME_SYMBOLS
366            .iter()
367            .map(|name| CudaSymbolEvidence {
368                name: (*name).to_owned(),
369                present: true,
370            })
371            .collect(),
372        cublas_symbols: CUBLAS_SYMBOLS
373            .iter()
374            .map(|name| CudaSymbolEvidence {
375                name: (*name).to_owned(),
376                present: true,
377            })
378            .collect(),
379        cublaslt_symbols: CUBLASLT_SYMBOLS
380            .iter()
381            .map(|name| CudaSymbolEvidence {
382                name: (*name).to_owned(),
383                present: true,
384            })
385            .collect(),
386    }
387}
388
389fn candidate_paths(search_dirs: &[PathBuf], names: &[&str], search_system: bool) -> Vec<PathBuf> {
390    let mut candidates = Vec::new();
391    for directory in search_dirs {
392        for name in names {
393            candidates.push(directory.join(name));
394        }
395    }
396    if search_system {
397        candidates.extend(names.iter().map(PathBuf::from));
398    }
399    candidates
400}
401
402fn symbol_evidence(library: &Library, names: &[&str]) -> Vec<CudaSymbolEvidence> {
403    names
404        .iter()
405        .map(|name| CudaSymbolEvidence {
406            name: (*name).to_owned(),
407            present: symbol_present(library, name),
408        })
409        .collect()
410}
411
412fn open_library(path: &Path) -> Result<Library, libloading::Error> {
413    // SAFETY: Loading a CUDA shared library is the intended boundary of this
414    // crate. The handle is stored in CudaLibrarySet for at least as long as any
415    // validated symbol evidence derived from it is used.
416    unsafe { Library::new(path) }
417}
418
419fn symbol_present(library: &Library, name: &str) -> bool {
420    let mut bytes = name.as_bytes().to_vec();
421    bytes.push(0);
422    // SAFETY: The lookup only checks whether the library exports the named
423    // symbol as an opaque address. The address is not called or dereferenced.
424    unsafe { library.get::<*mut c_void>(&bytes).is_ok() }
425}
426
427fn driver_version(library: &Library) -> Result<i32, CudaLoadError> {
428    type CuInit = unsafe extern "C" fn(u32) -> i32;
429    type CuDriverGetVersion = unsafe extern "C" fn(*mut i32) -> i32;
430    // SAFETY: Symbols were loaded from the CUDA driver library by their official
431    // C ABI names. The calls use the documented signatures for cuInit and
432    // cuDriverGetVersion and pass initialized pointers.
433    unsafe {
434        let cu_init = library
435            .get::<CuInit>(b"cuInit\0")
436            .map_err(|error| CudaLoadError {
437                message: error.to_string(),
438            })?;
439        let get_version = library
440            .get::<CuDriverGetVersion>(b"cuDriverGetVersion\0")
441            .map_err(|error| CudaLoadError {
442                message: error.to_string(),
443            })?;
444        let init_status = cu_init(0);
445        if init_status != 0 {
446            return Err(CudaLoadError {
447                message: format!("cuInit failed with status {init_status}"),
448            });
449        }
450        let mut version = 0;
451        let version_status = get_version(&mut version);
452        if version_status != 0 {
453            return Err(CudaLoadError {
454                message: format!("cuDriverGetVersion failed with status {version_status}"),
455            });
456        }
457        Ok(version)
458    }
459}