use std::io::{Read, Write};
use std::path::{Path, PathBuf};
use std::time::Duration;
use base64::Engine;
use sha2::Digest;
pub(crate) fn main() {
println!("cargo:rerun-if-env-changed=DOCS_RS");
println!("cargo:rerun-if-env-changed=COPILOT_SKIP_CLI_DOWNLOAD");
println!("cargo:rerun-if-env-changed=COPILOT_CLI_EXTRACT_DIR");
println!("cargo:rerun-if-env-changed=BUNDLED_CLI_CACHE_DIR");
println!("cargo::rustc-check-cfg=cfg(has_bundled_cli)");
println!("cargo::rustc-check-cfg=cfg(has_extracted_cli)");
println!("cargo:rerun-if-changed=cli-version-in-process.txt");
let manifest_dir = std::env::var("CARGO_MANIFEST_DIR").expect("CARGO_MANIFEST_DIR is set");
let lockfile = Path::new(&manifest_dir)
.join("..")
.join("nodejs")
.join("package-lock.json");
if lockfile.is_file() {
println!("cargo:rerun-if-changed={}", lockfile.display());
}
if std::env::var_os("COPILOT_SKIP_CLI_DOWNLOAD").is_some() {
println!(
"cargo:warning=COPILOT_SKIP_CLI_DOWNLOAD is set — skipping CLI download/bundle/cache"
);
return;
}
if std::env::var_os("DOCS_RS").is_some() {
println!("cargo:warning=DOCS_RS is set — skipping CLI download/bundle/cache");
return;
}
let Some(platform) = target_platform() else {
println!("cargo:warning=Unsupported target platform for Copilot CLI bundling — skipping");
return;
};
let out_dir = std::env::var("OUT_DIR").expect("OUT_DIR is always set by cargo");
let out = Path::new(&out_dir);
let (version, expected_integrity) = resolve_version_and_integrity(platform.package_name);
println!("cargo:rustc-env=COPILOT_SDK_CLI_VERSION={version}");
let archive_name = format!("{}-{version}.tgz", platform.package_name);
let download_url = format!(
"https://registry.npmjs.org/@github/{}/-/{}",
platform.package_name, archive_name
);
let cache_dir = std::env::var("BUNDLED_CLI_CACHE_DIR")
.ok()
.map(std::path::PathBuf::from);
let cache_key = format!("v{version}-{archive_name}");
let include_runtime = std::env::var_os("CARGO_FEATURE_BUNDLED_IN_PROCESS").is_some();
if std::env::var_os("CARGO_FEATURE_BUNDLED_CLI").is_some() {
let archive = cached_download(&download_url, &cache_key, &expected_integrity, &cache_dir);
verify_binary_present_in_archive(&archive, platform.binary_name, &archive_name);
emit_embedded(out, &archive, platform, include_runtime);
println!("cargo:rustc-cfg=has_bundled_cli");
} else {
let install_dir = extracted_install_dir(&version);
let final_path = install_dir.join(platform.binary_name);
println!("cargo:rerun-if-changed={}", final_path.display());
if !final_path.is_file() {
let archive =
cached_download(&download_url, &cache_key, &expected_integrity, &cache_dir);
verify_binary_present_in_archive(&archive, platform.binary_name, &archive_name);
extract_to_cache(&archive, &install_dir, platform);
}
if final_path.is_file() {
println!("cargo:rustc-cfg=has_extracted_cli");
}
}
}
fn extracted_install_dir(version: &str) -> PathBuf {
if let Some(custom) = std::env::var_os("COPILOT_CLI_EXTRACT_DIR") {
PathBuf::from(custom)
} else {
let cache = dirs::cache_dir().unwrap_or_else(std::env::temp_dir);
cache
.join("github-copilot-sdk")
.join("cli")
.join(sanitize_version(version))
}
}
fn emit_embedded(out: &Path, package: &[u8], platform: Platform, include_runtime: bool) {
let archive = build_embedded_archive(package, platform, include_runtime);
std::fs::write(out.join("copilot_cli.archive"), archive)
.expect("failed to write copilot_cli.archive");
let generated = r#"// Auto-generated by github-copilot-sdk build.rs. Do not edit.
pub(super) static CLI_ARCHIVE: &[u8] = include_bytes!("copilot_cli.archive");
"#;
std::fs::write(out.join("bundled_cli.rs"), generated).expect("failed to write bundled_cli.rs");
}
fn build_embedded_archive(package: &[u8], platform: Platform, include_runtime: bool) -> Vec<u8> {
let encoder = flate2::GzBuilder::new()
.mtime(0)
.write(Vec::new(), flate2::Compression::default());
let mut archive = tar::Builder::new(encoder);
append_archive_file(
&mut archive,
platform.binary_name,
&extract_binary_bytes(package, platform),
0o755,
);
if include_runtime {
let runtime = extract_runtime_library_bytes(package).unwrap_or_else(|| {
panic!(
"package `{}` does not contain the native runtime library required by the `bundled-in-process` feature",
platform.package_name
)
});
append_archive_file(
&mut archive,
platform.runtime_library_name(),
&runtime,
0o644,
);
}
let encoder = archive
.into_inner()
.expect("failed to finish minimal embedded CLI archive");
encoder
.finish()
.expect("failed to compress minimal embedded CLI archive")
}
fn append_archive_file<W: Write>(
archive: &mut tar::Builder<W>,
path: &str,
bytes: &[u8],
mode: u32,
) {
let mut header = tar::Header::new_gnu();
header.set_size(bytes.len() as u64);
header.set_mode(mode);
header.set_uid(0);
header.set_gid(0);
header.set_mtime(0);
header.set_cksum();
archive
.append_data(&mut header, path, bytes)
.unwrap_or_else(|e| panic!("failed to add `{path}` to embedded CLI archive: {e}"));
}
fn resolve_version_and_integrity(package_name: &str) -> (String, String) {
let manifest_dir = std::env::var("CARGO_MANIFEST_DIR").expect("CARGO_MANIFEST_DIR is set");
let snapshot = Path::new(&manifest_dir).join("cli-version-in-process.txt");
if snapshot.is_file() {
let contents = std::fs::read_to_string(&snapshot)
.unwrap_or_else(|e| panic!("failed to read {}: {e}", snapshot.display()));
return parse_snapshot(&contents, package_name)
.unwrap_or_else(|e| panic!("invalid {}: {e}", snapshot.display()));
}
let lockfile = Path::new(&manifest_dir)
.join("..")
.join("nodejs")
.join("package-lock.json");
if lockfile.is_file() {
return read_version_and_integrity_from_package_lock(&lockfile, package_name);
}
panic!(
"Could not resolve the Copilot CLI version.\n\
Tried:\n\
- {} (missing)\n\
- {} (missing)\n\
In a published crate or vendored slot, `cli-version-in-process.txt` should be present.\n\
Inside the github/copilot-sdk repo, `../nodejs/package-lock.json` is the source.",
snapshot.display(),
lockfile.display(),
);
}
fn parse_snapshot(contents: &str, package_name: &str) -> Result<(String, String), String> {
let mut version: Option<String> = None;
let mut integrity: Option<String> = None;
for (line_no, raw) in contents.lines().enumerate() {
let line = raw.trim();
if line.is_empty() || line.starts_with('#') {
continue;
}
let Some((key, value)) = line.split_once('=') else {
return Err(format!(
"line {}: expected `key=value`, got `{raw}`",
line_no + 1
));
};
match key.trim() {
"version" => version = Some(value.trim().to_string()),
k if k == package_name => integrity = Some(value.trim().to_string()),
_ => {}
}
}
let version = version.ok_or("missing `version=` line")?;
let integrity =
integrity.ok_or_else(|| format!("missing integrity for package `{package_name}`"))?;
Ok((version, integrity))
}
fn read_version_and_integrity_from_package_lock(
path: &Path,
package_name: &str,
) -> (String, String) {
let contents = std::fs::read_to_string(path)
.unwrap_or_else(|e| panic!("failed to read {}: {e}", path.display()));
let lock: serde_json::Value = serde_json::from_str(&contents)
.unwrap_or_else(|e| panic!("failed to parse {}: {e}", path.display()));
let cli_key = "node_modules/@github/copilot";
let version = lock["packages"][cli_key]["version"]
.as_str()
.unwrap_or_else(|| panic!("{cli_key} has no version in {}", path.display()));
let platform_key = format!("node_modules/@github/{package_name}");
let integrity = lock["packages"][&platform_key]["integrity"]
.as_str()
.unwrap_or_else(|| panic!("{platform_key} has no integrity in {}", path.display()));
(version.to_string(), integrity.to_string())
}
#[derive(Clone, Copy)]
struct Platform {
package_name: &'static str,
binary_name: &'static str,
}
impl Platform {
fn runtime_library_name(&self) -> &'static str {
if self.package_name.contains("win32") {
"copilot_runtime.dll"
} else if self.package_name.contains("darwin") {
"libcopilot_runtime.dylib"
} else {
"libcopilot_runtime.so"
}
}
}
fn target_platform() -> Option<Platform> {
let os = std::env::var("CARGO_CFG_TARGET_OS").ok()?;
let arch = std::env::var("CARGO_CFG_TARGET_ARCH").ok()?;
let target_env = std::env::var("CARGO_CFG_TARGET_ENV").unwrap_or_default();
match (os.as_str(), arch.as_str(), target_env.as_str()) {
("macos", "aarch64", _) => Some(Platform {
package_name: "copilot-darwin-arm64",
binary_name: "copilot",
}),
("macos", "x86_64", _) => Some(Platform {
package_name: "copilot-darwin-x64",
binary_name: "copilot",
}),
("linux", "x86_64", "musl") => Some(Platform {
package_name: "copilot-linuxmusl-x64",
binary_name: "copilot",
}),
("linux", "aarch64", "musl") => Some(Platform {
package_name: "copilot-linuxmusl-arm64",
binary_name: "copilot",
}),
("linux", "x86_64", _) => Some(Platform {
package_name: "copilot-linux-x64",
binary_name: "copilot",
}),
("linux", "aarch64", _) => Some(Platform {
package_name: "copilot-linux-arm64",
binary_name: "copilot",
}),
("windows", "x86_64", _) => Some(Platform {
package_name: "copilot-win32-x64",
binary_name: "copilot.exe",
}),
("windows", "aarch64", _) => Some(Platform {
package_name: "copilot-win32-arm64",
binary_name: "copilot.exe",
}),
_ => None,
}
}
fn extract_to_cache(archive: &[u8], install_dir: &Path, platform: Platform) -> PathBuf {
let final_path = install_dir.join(platform.binary_name);
if final_path.is_file() {
return final_path;
}
std::fs::create_dir_all(install_dir).unwrap_or_else(|e| {
panic!(
"failed to create install dir {}: {e}",
install_dir.display()
)
});
let bytes = extract_binary_bytes(archive, platform);
let nanos = std::time::SystemTime::now()
.duration_since(std::time::UNIX_EPOCH)
.map(|d| d.as_nanos())
.unwrap_or(0);
let staging_path = install_dir.join(format!(
".{}.staging-{}-{nanos}",
platform.binary_name,
std::process::id(),
));
{
let mut f = std::fs::File::create(&staging_path).unwrap_or_else(|e| {
let _ = std::fs::remove_file(&staging_path);
panic!(
"failed to create staging file {}: {e}",
staging_path.display()
);
});
if let Err(e) = f.write_all(&bytes) {
let _ = std::fs::remove_file(&staging_path);
panic!(
"failed to write staging file {}: {e}",
staging_path.display()
);
}
#[cfg(unix)]
{
use std::os::unix::fs::PermissionsExt;
if let Err(e) = f.set_permissions(std::fs::Permissions::from_mode(0o755)) {
let _ = std::fs::remove_file(&staging_path);
panic!("failed to chmod {}: {e}", staging_path.display());
}
}
if let Err(e) = f.set_modified(std::time::SystemTime::UNIX_EPOCH) {
println!(
"cargo:warning=Could not backdate {} (a redundant rebuild may occur): {e}",
staging_path.display()
);
}
}
if let Err(e) = std::fs::rename(&staging_path, &final_path) {
let _ = std::fs::remove_file(&staging_path);
panic!(
"failed to rename {} -> {}: {e}",
staging_path.display(),
final_path.display()
);
}
println!(
"cargo:warning=Extracted Copilot CLI to {}",
final_path.display()
);
final_path
}
fn extract_runtime_library_bytes(archive: &[u8]) -> Option<Vec<u8>> {
let gz = flate2::read::GzDecoder::new(archive);
let mut tar = tar::Archive::new(gz);
for entry in tar.entries().ok()? {
let mut entry = entry.ok()?;
let name = entry.path().ok()?.to_string_lossy().into_owned();
if name == "runtime.node" || name.ends_with("/runtime.node") {
let mut bytes = Vec::with_capacity(entry.size() as usize);
entry.read_to_end(&mut bytes).ok()?;
return Some(bytes);
}
}
None
}
fn sanitize_version(version: &str) -> String {
version
.chars()
.map(|c| match c {
'a'..='z' | 'A'..='Z' | '0'..='9' | '.' | '-' | '_' => c,
_ => '_',
})
.collect()
}
fn extract_binary_bytes(archive: &[u8], platform: Platform) -> Vec<u8> {
let gz = flate2::read::GzDecoder::new(archive);
let mut tar = tar::Archive::new(gz);
for entry in tar
.entries()
.unwrap_or_else(|e| panic!("failed to read tar entries: {e}"))
{
let mut entry = entry.unwrap_or_else(|e| panic!("failed to read tar entry: {e}"));
let path = entry
.path()
.unwrap_or_else(|e| panic!("failed to read tar entry path: {e}"));
let name = path.to_string_lossy().into_owned();
if name == platform.binary_name || name.ends_with(&format!("/{}", platform.binary_name)) {
let mut bytes = Vec::with_capacity(entry.size() as usize);
entry
.read_to_end(&mut bytes)
.unwrap_or_else(|e| panic!("failed to read tar entry bytes: {e}"));
return bytes;
}
}
panic!(
"binary `{}` not found in package `{}`",
platform.binary_name, platform.package_name
);
}
fn cached_download(
url: &str,
cache_key: &str,
expected_integrity: &str,
cache_dir: &Option<std::path::PathBuf>,
) -> Vec<u8> {
if let Some(dir) = cache_dir {
let cached_path = dir.join(cache_key);
if cached_path.is_file() {
match std::fs::read(&cached_path) {
Ok(data) if verify_integrity(&data, expected_integrity) => {
return data;
}
Ok(_) => {
println!("cargo:warning=Cached archive hash mismatch, re-downloading");
let _ = std::fs::remove_file(&cached_path);
}
Err(e) => {
println!(
"cargo:warning=Failed to read cache {}, re-downloading: {e}",
cached_path.display()
);
}
}
}
}
println!("cargo:warning=Downloading {url}");
let data = download_with_retry(url);
if !verify_integrity(&data, expected_integrity) {
panic!(
"Archive integrity check failed for {url}!\n expected: {expected_integrity}\n \
This could indicate a corrupted download or a supply-chain attack."
);
}
if let Some(dir) = cache_dir {
if let Err(e) = std::fs::create_dir_all(dir) {
println!(
"cargo:warning=Failed to create cache directory {}: {e}",
dir.display()
);
} else {
let cached_path = dir.join(cache_key);
println!("cargo:warning=Caching archive at {}", cached_path.display());
if let Err(e) = std::fs::write(&cached_path, &data) {
println!(
"cargo:warning=Failed to write cache file {}: {e}",
cached_path.display()
);
}
}
}
data
}
const MAX_RETRIES: u32 = 3;
fn download_with_retry(url: &str) -> Vec<u8> {
let mut attempt = 0u32;
loop {
attempt += 1;
match try_download(url) {
Ok(bytes) => return bytes,
Err(err) if err.transient && attempt <= MAX_RETRIES => {
let backoff = Duration::from_secs(1u64 << (attempt - 1));
println!(
"cargo:warning=Transient download failure for {url} (attempt {attempt}/{}): {} — retrying in {}s",
MAX_RETRIES + 1,
err.message,
backoff.as_secs(),
);
std::thread::sleep(backoff);
}
Err(err) => panic!("Failed to download {url}: {}", err.message),
}
}
}
struct DownloadError {
message: String,
transient: bool,
}
fn try_download(url: &str) -> Result<Vec<u8>, DownloadError> {
let connector = native_tls::TlsConnector::new().map_err(|e| DownloadError {
message: format!("native-tls init error: {e}"),
transient: false,
})?;
let agent = ureq::AgentBuilder::new()
.tls_connector(std::sync::Arc::new(connector))
.timeout_connect(Duration::from_secs(30))
.timeout_read(Duration::from_secs(120))
.build();
match agent.get(url).call() {
Ok(response) => {
let mut bytes = Vec::new();
response
.into_reader()
.read_to_end(&mut bytes)
.map_err(|e| DownloadError {
message: format!("read error: {e}"),
transient: true,
})?;
Ok(bytes)
}
Err(ureq::Error::Status(code, response)) if (500..600).contains(&code) => {
Err(DownloadError {
message: format!("HTTP {code} {}", response.status_text()),
transient: true,
})
}
Err(ureq::Error::Status(code, response)) => Err(DownloadError {
message: format!("HTTP {code} {}", response.status_text()),
transient: false,
}),
Err(ureq::Error::Transport(t)) => Err(DownloadError {
message: format!("transport error: {t}"),
transient: true,
}),
}
}
fn verify_binary_present_in_archive(archive: &[u8], binary_name: &str, package_name: &str) {
let found = archive_contains_tar_entry(archive, binary_name);
if !found {
panic!(
"Copilot CLI package `{package_name}` does not contain an entry named `{binary_name}`. \
The package layout may have changed; runtime extraction would fail. \
Update `verify_binary_present_in_archive` in build.rs and the matching `extract_binary` in src/embeddedcli.rs."
);
}
}
fn archive_contains_tar_entry(targz: &[u8], binary_name: &str) -> bool {
let gz = flate2::read::GzDecoder::new(targz);
let mut archive = tar::Archive::new(gz);
let Ok(entries) = archive.entries() else {
return false;
};
for entry in entries.flatten() {
let Ok(path) = entry.path() else {
continue;
};
let name = path.to_string_lossy();
if name == binary_name || name.ends_with(&format!("/{binary_name}")) {
return true;
}
}
false
}
fn verify_integrity(data: &[u8], integrity: &str) -> bool {
let Some(encoded) = integrity.strip_prefix("sha512-") else {
return false;
};
let Ok(expected) = base64::engine::general_purpose::STANDARD.decode(encoded) else {
return false;
};
let mut hasher = sha2::Sha512::new();
hasher.update(data);
hasher.finalize().as_slice() == expected
}