use std::env;
use std::fs::{self, File};
use std::io::{BufReader, copy};
use std::path::{Path, PathBuf};
use serde_json::Value;
const DEFAULT_REPO: &str = "launcher-rs/audio-cpp-rs";
pub fn asset_name(target: &str, use_shared_libs: bool) -> Option<String> {
let os = platform_os(target)?;
let backend = backend_suffix()?;
let modelset = modelset_suffix()?;
let library_type = if use_shared_libs { "dynamic" } else { "static" };
let crt = crt_suffix();
Some(asset_name_of(
&os,
target,
&backend,
&modelset,
&library_type,
crt.as_deref(),
))
}
fn asset_name_of(
os: &str,
target: &str,
backend: &str,
modelset: &str,
library_type: &str,
crt: Option<&str>,
) -> String {
match crt {
Some(crt) => format!(
"audio-cpp-prebuilt-{os}-{target}-{backend}-{crt}-{modelset}-{library_type}.tar.gz"
),
None => {
format!("audio-cpp-prebuilt-{os}-{target}-{backend}-{modelset}-{library_type}.tar.gz")
}
}
}
fn fetch_prebuilt(
target: &str,
use_shared_libs: bool,
modelset_override: Option<&str>,
) -> Option<PathBuf> {
if is_disabled() {
return None;
}
if env::var("AUDIOCPP_PREBUILT_DIR").is_ok() {
return None;
}
let modelset = modelset_override
.map(|s| s.to_string())
.or_else(modelset_suffix)?;
let os = platform_os(target)?;
let backend = backend_suffix()?;
let library_type = if use_shared_libs { "dynamic" } else { "static" };
let asset = asset_name_of(
&os,
target,
&backend,
&modelset,
&library_type,
crt_suffix().as_deref(),
);
let tag = release_tag();
let cache_root = cache_root()?;
let extract_dir = cache_root
.join(tag.trim_start_matches('v'))
.join(asset.strip_suffix(".tar.gz").unwrap_or(&asset));
if is_valid_prebuilt_root(&extract_dir) {
println!(
"cargo:warning=使用缓存中的 audio.cpp 预编译库:{}",
extract_dir.display()
);
return Some(extract_dir);
}
let url = download_url(&tag, &asset);
println!("cargo:warning=下载 audio.cpp 预编译库:{url}");
match download_and_extract(&url, &extract_dir) {
Ok(()) if is_valid_prebuilt_root(&extract_dir) => {
if !identity_matches(&extract_dir) {
println!(
"cargo:warning=预编译归档的 audio.cpp commit 与本地 submodule 不符,回落到源码构建"
);
let _ = fs::remove_dir_all(&extract_dir);
return None;
}
println!(
"cargo:warning=audio.cpp 预编译库就绪:{}",
extract_dir.display()
);
Some(extract_dir)
}
Ok(()) => {
println!("cargo:warning=预编译归档已解压但未找到库文件,回落到源码构建");
let _ = fs::remove_dir_all(&extract_dir);
None
}
Err(err) => {
println!("cargo:warning=预编译库下载失败({err}),该资产不可用");
let _ = fs::remove_dir_all(&extract_dir);
None
}
}
}
pub fn ensure_prebuilt(target: &str, use_shared_libs: bool) -> Option<PathBuf> {
let exact = modelset_suffix();
match exact.as_deref() {
Some("full") => fetch_prebuilt(target, use_shared_libs, Some("full")),
Some(mset) => fetch_prebuilt(target, use_shared_libs, Some(mset)).or_else(|| {
println!("cargo:warning=custom/core 资产不可用,回退到 full 全模型资产");
fetch_prebuilt(target, use_shared_libs, Some("full"))
}),
None => fetch_prebuilt(target, use_shared_libs, None),
}
}
fn is_disabled() -> bool {
matches!(
env::var("AUDIOCPP_PREBUILT_OFF").as_deref(),
Ok("1") | Ok("true") | Ok("TRUE") | Ok("on") | Ok("ON")
)
}
fn release_tag() -> String {
env::var("AUDIOCPP_PREBUILT_TAG").unwrap_or_else(|_| {
format!(
"v{}",
env::var("CARGO_PKG_VERSION").unwrap_or_else(|_| "0.1.0".into())
)
})
}
fn github_repo() -> String {
env::var("AUDIOCPP_PREBUILT_REPO").unwrap_or_else(|_| DEFAULT_REPO.to_string())
}
fn download_url(tag: &str, asset: &str) -> String {
if let Ok(template) = env::var("AUDIOCPP_PREBUILT_URL") {
return template.replace("{tag}", tag).replace("{asset}", asset);
}
format!(
"https://github.com/{}/releases/download/{}/{}",
github_repo(),
tag,
asset
)
}
fn cache_root() -> Option<PathBuf> {
let out_dir = env::var("OUT_DIR").ok()?;
let profile = env::var("PROFILE").ok()?;
let mut target_dir = None;
let mut sub_path = Path::new(&out_dir);
while let Some(parent) = sub_path.parent() {
if parent.ends_with(&profile) {
target_dir = Some(parent);
break;
}
sub_path = parent;
}
Some(target_dir?.join("audio-cpp-prebuilt-cache"))
}
fn platform_os(target: &str) -> Option<&'static str> {
if target.contains("linux") {
Some("linux")
} else if target.contains("windows") {
Some("windows")
} else if target.contains("apple") {
Some("macos")
} else {
None
}
}
fn backend_suffix() -> Option<String> {
if cfg!(feature = "cuda") || cfg!(feature = "hip") {
return None;
}
if cfg!(feature = "metal") {
return Some("metal".to_string());
}
if cfg!(feature = "vulkan") {
return Some("vulkan".to_string());
}
Some("cpu".to_string())
}
fn crt_suffix() -> Option<String> {
if std::env::consts::OS != "windows" {
return None;
}
let static_crt = env::var("CARGO_CFG_TARGET_FEATURE")
.map(|f| f.split(',').any(|s| s.trim() == "crt-static"))
.unwrap_or(false);
Some(if static_crt {
"mt".to_string()
} else {
"md".to_string()
})
}
fn modelset_suffix() -> Option<String> {
if cfg!(feature = "full-models") {
return Some("full".to_string());
}
if cfg!(feature = "custom-models") {
let mut families: Vec<String> = Vec::new();
for (key, _) in env::vars() {
if let Some(suffix) = key.strip_prefix("CARGO_FEATURE_MODEL_") {
families.push(suffix.to_lowercase());
}
}
if let Ok(env_models) = env::var("AUDIOCPP_MODELS") {
for m in env_models
.split(',')
.map(str::trim)
.filter(|s| !s.is_empty())
{
if !families.iter().any(|s| s == m) {
families.push(m.to_string());
}
}
}
families.sort();
if families.is_empty() {
return None;
}
return Some(format!("custom-{}", families.join("-")));
}
Some("core".to_string())
}
fn identity_matches(root: &Path) -> bool {
let local_commit = local_audio_commit();
let metadata_path = root.join("metadata.json");
if local_commit.is_none() || !metadata_path.is_file() {
return true;
}
let Ok(text) = fs::read_to_string(&metadata_path) else {
return true;
};
let Ok(json) = serde_json::from_str::<Value>(&text) else {
return true;
};
match json.get("audio_commit").and_then(|v| v.as_str()) {
Some(recorded) => {
if recorded != local_commit.as_deref().unwrap_or("") {
return false;
}
}
None => {}
}
if let Some(recorded_msvc) = json.get("msvc_ver").and_then(|v| v.as_i64()) {
if recorded_msvc > 0 {
if let Some(local_msvc) = local_msvc_ver() {
if local_msvc < recorded_msvc {
println!(
"cargo:warning=预编译归档由 MSVC {} 构建,本地为 MSVC {}(偏低),\
可能导致链接失败,回落到源码构建",
recorded_msvc, local_msvc
);
return false;
}
}
}
}
true
}
fn local_msvc_ver() -> Option<i64> {
if std::env::consts::OS != "windows" {
return None;
}
let cc = cc::Build::new();
let compiler = cc.try_get_compiler().ok()?;
if !compiler.is_like_msvc() {
return None;
}
let cl = compiler.path();
let banner = std::process::Command::new(cl).arg("/Bv").output().ok()?;
let text = String::from_utf8_lossy(&banner.stdout);
let text = format!("{}{}", text, String::from_utf8_lossy(&banner.stderr));
let nums: Vec<i64> = text
.split(|c: char| !c.is_ascii_digit())
.filter_map(|tok| tok.parse().ok())
.collect();
for w in nums.windows(2) {
if w[0] == 19 && (40..=99).contains(&w[1]) {
return Some(1900 + w[1]);
}
}
None
}
fn local_audio_commit() -> Option<String> {
let manifest_dir = env::var("CARGO_MANIFEST_DIR").ok()?;
let sub = Path::new(&manifest_dir).join("audio.cpp");
if !sub.join("CMakeLists.txt").exists() {
return None;
}
let output = std::process::Command::new("git")
.args(["rev-parse", "--short=8", "HEAD"])
.current_dir(&sub)
.output()
.ok()?;
if !output.status.success() {
return None;
}
let s = String::from_utf8(output.stdout).ok()?;
let s = s.trim().to_string();
if s.is_empty() { None } else { Some(s) }
}
fn is_valid_prebuilt_root(root: &Path) -> bool {
if !root.is_dir() {
return false;
}
for dir in [
root.to_path_buf(),
root.join("lib"),
root.join("lib64"),
root.join("bin"),
] {
if !dir.is_dir() {
continue;
}
let Ok(entries) = fs::read_dir(&dir) else {
continue;
};
for entry in entries.flatten() {
let Ok(file_type) = entry.file_type() else {
continue;
};
if !file_type.is_file() && !file_type.is_symlink() {
continue;
}
let name = entry.file_name();
let Some(name) = name.to_str() else {
continue;
};
if is_audio_lib_name(name) {
return true;
}
}
}
false
}
fn is_audio_lib_name(name: &str) -> bool {
let base = name
.strip_prefix("lib")
.unwrap_or(name)
.split('.')
.next()
.unwrap_or(name);
matches!(
base,
"engine_runtime"
| "ggml"
| "ggml-base"
| "ggml-cpu"
| "sentencepiece"
| "cjson_vendor"
| "yaml_vendor"
)
}
fn download_and_extract(url: &str, extract_dir: &Path) -> Result<(), String> {
if extract_dir.exists() {
fs::remove_dir_all(extract_dir).map_err(|e| e.to_string())?;
}
fs::create_dir_all(extract_dir).map_err(|e| e.to_string())?;
let archive_path = extract_dir.with_extension("tar.gz");
download_file(url, &archive_path)?;
let file = File::open(&archive_path).map_err(|e| e.to_string())?;
let reader = BufReader::new(file);
let decoder = flate2::read::GzDecoder::new(reader);
let mut archive = tar::Archive::new(decoder);
archive.unpack(extract_dir).map_err(|e| e.to_string())?;
let _ = fs::remove_file(&archive_path);
Ok(())
}
fn download_file(url: &str, dest: &Path) -> Result<(), String> {
if let Some(parent) = dest.parent() {
fs::create_dir_all(parent).map_err(|e| e.to_string())?;
}
if let Some(local) = url.strip_prefix("file://") {
let src = PathBuf::from(local);
if !src.is_file() {
return Err(format!("本地归档不存在:{}", src.display()));
}
fs::copy(&src, dest).map_err(|e| e.to_string())?;
return Ok(());
}
let partial = dest.with_extension("partial");
download_with_retry(url, &partial)?;
fs::rename(&partial, dest).map_err(|e| e.to_string())?;
Ok(())
}
fn download_with_retry(url: &str, partial: &Path) -> Result<(), String> {
const MAX_ATTEMPTS: usize = 3;
let mut last_err = format!("下载 {url} 失败");
for attempt in 1..=MAX_ATTEMPTS {
match try_download(url, partial) {
Ok(()) => return Ok(()),
Err(err) => {
if is_non_retryable(&err) {
println!("cargo:warning=预编译下载失败({err})");
return Err(err);
}
println!("cargo:warning=预编译下载第 {attempt}/{MAX_ATTEMPTS} 次失败:{err}");
last_err = err;
std::thread::sleep(std::time::Duration::from_secs(attempt as u64));
}
}
}
Err(last_err)
}
fn is_non_retryable(err: &str) -> bool {
if let Some(pos) = err.rfind("status code ") {
if let Some(code) = err[pos + "status code ".len()..].split_whitespace().next() {
if let Ok(n) = code.parse::<u16>() {
return (400..500).contains(&n) && n != 429;
}
}
}
false
}
fn try_download(url: &str, partial: &Path) -> Result<(), String> {
let response = ureq::get(url)
.call()
.map_err(|e| format!("HTTP GET {url}: {e}"))?;
if !(200..300).contains(&response.status()) {
return Err(format!("HTTP {} for {url}", response.status()));
}
let mut reader = response.into_reader();
let mut file = File::create(partial).map_err(|e| e.to_string())?;
copy(&mut reader, &mut file).map_err(|e| e.to_string())?;
Ok(())
}