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