use std::env;
use std::path::{Path, PathBuf};
const TOOLKIT_ENV_VARS: &[&str] = &["CUDA_TOOLKIT_PATH", "CUDA_HOME"];
const TOOLKIT_TARGET_DIR_ENV: &str = "CUDA_TOOLKIT_TARGET_DIR";
#[cfg(windows)]
const DEFAULT_TOOLKIT_DIRS: &[&str] = &[
r"C:\Program Files\NVIDIA GPU Computing Toolkit\CUDA\v13.3",
r"C:\Program Files\NVIDIA GPU Computing Toolkit\CUDA\v13.2",
];
#[cfg(not(windows))]
const DEFAULT_TOOLKIT_DIRS: &[&str] = &[
"/usr/local/cuda-13.3",
"/usr/local/cuda-13.2",
"/usr/local/cuda-13",
"/usr/local/cuda",
];
const MIN_CUDA_VERSION: u32 = 13000;
fn main() {
println!("cargo::rustc-check-cfg=cfg(cuda_has_multicast)");
for var in TOOLKIT_ENV_VARS {
println!("cargo:rerun-if-env-changed={var}");
}
println!("cargo:rerun-if-env-changed={TOOLKIT_TARGET_DIR_ENV}");
let Some(cuda_h) = find_cuda_header() else {
return;
};
println!("cargo:rerun-if-changed={}", cuda_h.display());
if std::fs::read_to_string(&cuda_h).is_ok_and(|header| header.contains("cuMulticastCreate")) {
println!("cargo:rustc-cfg=cuda_has_multicast");
}
}
fn toolkit_target_dirs() -> Vec<String> {
let arch = env::var("CARGO_CFG_TARGET_ARCH").unwrap_or_default();
let os = env::var("CARGO_CFG_TARGET_OS").unwrap_or_default();
if os != "linux" {
return vec![];
}
match arch.as_str() {
"x86_64" => vec!["x86_64-linux".to_string()],
"aarch64" => vec!["sbsa-linux".to_string(), "aarch64-linux".to_string()],
_ => vec![],
}
}
fn include_candidates(toolkit: &Path) -> Vec<PathBuf> {
if let Some(dir) = env::var(TOOLKIT_TARGET_DIR_ENV)
.ok()
.filter(|dir| !dir.trim().is_empty())
{
return vec![toolkit.join("targets").join(dir).join("include")];
}
let mut candidates = vec![toolkit.join("include")];
for target_dir in toolkit_target_dirs() {
candidates.push(toolkit.join("targets").join(target_dir).join("include"));
}
candidates
}
fn find_cuda_header_in(toolkit: &Path) -> Option<PathBuf> {
include_candidates(toolkit)
.into_iter()
.map(|dir| dir.join("cuda.h"))
.find(|header| header.is_file())
}
fn cuda_version_from_header(cuda_h: &Path) -> Option<u32> {
let source = std::fs::read_to_string(cuda_h).ok()?;
source.lines().find_map(|line| {
let mut parts = line.split_whitespace();
match (parts.next(), parts.next(), parts.next()) {
(Some("#define"), Some("CUDA_VERSION"), Some(version)) => version.parse().ok(),
_ => None,
}
})
}
fn find_cuda_header() -> Option<PathBuf> {
for var in TOOLKIT_ENV_VARS {
if let Ok(toolkit) = env::var(var) {
return find_cuda_header_in(Path::new(&toolkit));
}
}
DEFAULT_TOOLKIT_DIRS.iter().find_map(|toolkit| {
let header = find_cuda_header_in(Path::new(toolkit))?;
(cuda_version_from_header(&header)? >= MIN_CUDA_VERSION).then_some(header)
})
}