extern crate bindgen;
use std::env;
use std::path::{Path, PathBuf};
#[derive(Clone, Copy, Eq, PartialEq)]
enum OS {
Linux,
#[allow(clippy::enum_variant_names)]
MacOS,
Windows,
}
impl OS {
fn get() -> Self {
let os = env::var("CARGO_CFG_TARGET_OS").expect("Unable to get TARGET_OS");
match os.as_str() {
"linux" => Self::Linux,
"macos" => Self::MacOS,
"windows" => Self::Windows,
os => panic!("Unsupported system {os}"),
}
}
}
fn make_shared_lib<P: AsRef<Path>>(os: OS, xla_dir: P) {
println!("cargo:rerun-if-changed=xla_rs/xla_rs.cc");
println!("cargo:rerun-if-changed=xla_rs/xla_rs.h");
match os {
OS::Linux | OS::MacOS => {
cc::Build::new()
.cpp(true)
.pic(true)
.warnings(false)
.include(xla_dir.as_ref().join("include"))
.flag("-std=c++17")
.flag("-Wno-deprecated-declarations")
.flag("-DLLVM_ON_UNIX=1")
.flag("-DLLVM_VERSION_STRING=")
.flag("-DNDEBUG")
.file("xla_rs/xla_rs.cc")
.compile("xla_rs");
}
OS::Windows => {
cc::Build::new()
.cpp(true)
.pic(true)
.warnings(false)
.include(xla_dir.as_ref().join("include"))
.flag("/DNDEBUG")
.file("xla_rs/xla_rs.cc")
.compile("xla_rs");
}
};
}
fn env_var_rerun(name: &str) -> Option<String> {
println!("cargo:rerun-if-env-changed={name}");
env::var(name).ok()
}
fn main() {
let os = OS::get();
let xla_dir = env_var_rerun("XLA_EXTENSION_DIR")
.map_or_else(|| env::current_dir().unwrap().join("xla_extension"), PathBuf::from);
println!("cargo:rerun-if-changed=xla_rs/xla_rs.h");
println!("cargo:rerun-if-changed=xla_rs/xla_rs.cc");
let bindings = bindgen::Builder::default()
.header("xla_rs/xla_rs.h")
.parse_callbacks(Box::new(bindgen::CargoCallbacks::new()))
.generate()
.expect("Unable to generate bindings");
let out_path = PathBuf::from(env::var("OUT_DIR").unwrap());
bindings.write_to_file(out_path.join("c_xla.rs")).expect("Couldn't write bindings!");
if std::env::var("DOCS_RS").is_ok() {
return;
}
make_shared_lib(os, &xla_dir);
if os == OS::Linux {
println!("cargo:rustc-link-arg=-Wl,-lstdc++");
}
println!("cargo:rustc-link-search=native={}", xla_dir.join("lib").display());
println!("cargo:rustc-link-lib=static=xla_rs");
if os == OS::MacOS {
println!("cargo:rustc-link-arg=-Wl,-rpath,{}", xla_dir.join("lib").display());
} else {
println!("cargo:rustc-link-arg=-Wl,-rpath={}", xla_dir.join("lib").display());
}
println!("cargo:rustc-link-lib=xla_extension");
}