use std::ffi::OsString;
use std::io::{ErrorKind, Write};
use self_update::{backends::github::Update, cargo_crate_version, errors::Error as UpdError};
use semver::Version;
use tracing::error;
use tracing::warn;
use tokio::task;
use crate::{cli::global::GlobalArgs, reporter::styles::Styles};
#[derive(Clone, Debug, PartialEq, Eq)]
pub enum UpdateCheckStatus {
Disabled,
Failed,
Ok,
}
impl UpdateCheckStatus {
pub fn as_str(&self) -> &'static str {
match self {
UpdateCheckStatus::Disabled => "disabled",
UpdateCheckStatus::Failed => "failed",
UpdateCheckStatus::Ok => "ok",
}
}
}
#[derive(Clone, Debug, PartialEq, Eq)]
pub struct UpdateStatus {
pub message: Option<String>,
pub styled_message: Option<String>,
pub is_outdated: bool,
pub running_version: String,
pub latest_version: Option<String>,
pub check_status: UpdateCheckStatus,
pub was_self_updated: bool,
}
impl Default for UpdateStatus {
fn default() -> Self {
UpdateStatus {
message: None,
styled_message: None,
is_outdated: false,
running_version: cargo_crate_version!().to_string(),
latest_version: None,
check_status: UpdateCheckStatus::Disabled,
was_self_updated: false,
}
}
}
fn styled_heading(styles: &Styles, text: &str) -> String {
styles.style_finding_active_heading.apply_to(text).to_string()
}
pub fn check_for_update(global_args: &GlobalArgs, base_url: Option<&str>) -> UpdateStatus {
let running_version = cargo_crate_version!().to_string();
if global_args.no_update_check {
return UpdateStatus {
message: Some("Update check disabled (--no-update-check)".to_string()),
styled_message: None,
is_outdated: false,
running_version,
latest_version: None,
check_status: UpdateCheckStatus::Disabled,
was_self_updated: false,
};
}
let use_color = !global_args.quiet && global_args.use_color(std::io::stderr());
let styles = Styles::new(use_color);
let mut builder = Update::configure();
builder
.repo_owner("mongodb")
.repo_name("kingfisher")
.bin_name("kingfisher")
.show_download_progress(false)
.no_confirm(true) .current_version(cargo_crate_version!());
if let Some(url) = base_url {
builder.api_base_url(url);
}
#[cfg(all(target_os = "linux", target_arch = "aarch64"))]
builder.target("linux-arm64");
#[cfg(all(target_os = "linux", target_arch = "x86_64"))]
builder.target("linux-x64");
#[cfg(all(target_os = "macos", target_arch = "aarch64"))]
builder.target("darwin-arm64");
#[cfg(all(target_os = "macos", target_arch = "x86_64"))]
builder.target("darwin-x64");
#[cfg(all(target_os = "windows", target_arch = "x86_64"))]
builder.target("windows-x64");
#[cfg(all(target_os = "windows", target_arch = "aarch64"))]
builder.target("windows-arm64");
#[cfg(target_os = "windows")]
builder.asset_identifier("zip");
#[cfg(not(target_os = "windows"))]
builder.asset_identifier("tgz");
let Ok(updater) = builder.build() else {
let plain = "Failed to configure update checker".to_string();
let styled_message = styled_heading(&styles, &plain);
let _ = writeln!(std::io::stderr(), "{}", styled_message);
return UpdateStatus {
message: Some(plain),
styled_message: Some(styled_message),
is_outdated: false,
running_version,
latest_version: None,
check_status: UpdateCheckStatus::Failed,
was_self_updated: false,
};
};
let Ok(releases) = updater.get_latest_release() else {
let plain = "Failed to check for updates".to_string();
let styled_message = styled_heading(&styles, &plain);
let _ = writeln!(std::io::stderr(), "{}", styled_message);
return UpdateStatus {
message: Some(plain),
styled_message: Some(styled_message),
is_outdated: false,
running_version,
latest_version: None,
check_status: UpdateCheckStatus::Failed,
was_self_updated: false,
};
};
let Some(release) = releases.latest() else {
let plain = "Failed to check for updates".to_string();
let styled_message = styled_heading(&styles, &plain);
let _ = writeln!(std::io::stderr(), "{}", styled_message);
return UpdateStatus {
message: Some(plain),
styled_message: Some(styled_message),
is_outdated: false,
running_version,
latest_version: None,
check_status: UpdateCheckStatus::Failed,
was_self_updated: false,
};
};
if release.version() == running_version {
let plain = format!("Kingfisher {running_version} is up to date");
let _ = writeln!(std::io::stderr(), "{plain}");
return UpdateStatus {
message: Some(plain.clone()),
styled_message: Some(plain),
is_outdated: false,
running_version,
latest_version: Some(release.version().to_string()),
check_status: UpdateCheckStatus::Ok,
was_self_updated: false,
};
}
if let (Ok(curr), Ok(latest)) =
(Version::parse(&running_version), Version::parse(release.version()))
{
if curr > latest {
let plain =
format!("Running Kingfisher {curr} which is newer than latest released {latest}");
let styled_message = styled_heading(&styles, &plain);
let _ = writeln!(std::io::stderr(), "{}", styled_message);
return UpdateStatus {
message: Some(plain),
styled_message: Some(styled_message),
is_outdated: false,
running_version,
latest_version: Some(release.version().to_string()),
check_status: UpdateCheckStatus::Ok,
was_self_updated: false,
};
}
}
let plain = format!("New Kingfisher release {} available", release.version());
let styled_message = styled_heading(&styles, &plain);
let _ = writeln!(std::io::stderr(), "{}", styled_message);
let mut was_self_updated = false;
if global_args.self_update {
match updater.update() {
Ok(status) => {
if status.is_updated() {
let message = format!("Updated to version {}", status.version());
let _ = writeln!(std::io::stderr(), "{}", styled_heading(&styles, &message));
was_self_updated = true;
} else {
let _ = writeln!(
std::io::stderr(),
"{}",
styled_heading(
&styles,
&format!("Already at version {} — no update applied", status.version()),
)
);
}
}
Err(e) => match e {
UpdError::Io(ref io_err) => match io_err.kind() {
ErrorKind::PermissionDenied => {
let _ = writeln!(
std::io::stderr(),
"{}",
styled_heading(
&styles,
"Cannot replace the current binary - permission denied.\n\
If you installed via a package manager, run its upgrade command.\n\
Otherwise reinstall to a user-writable directory or re-run with sudo."
)
);
}
ErrorKind::NotFound => {
let _ = writeln!(
std::io::stderr(),
"{}",
styled_heading(
&styles,
"Cannot replace the current binary - file not found.\n\
If you installed via a package manager, run its upgrade command.\n\
Otherwise reinstall to a user-writable directory."
)
);
}
_ => error!("Failed to update: {e}"),
},
_ => error!("Failed to update: {e}"),
},
}
}
UpdateStatus {
message: Some(plain),
styled_message: Some(styled_message),
is_outdated: true,
running_version,
latest_version: Some(release.version().to_string()),
check_status: UpdateCheckStatus::Ok,
was_self_updated,
}
}
pub async fn check_for_update_async(
global_args: &GlobalArgs,
base_url: Option<&str>,
) -> UpdateStatus {
let args = global_args.clone();
let base = base_url.map(str::to_owned);
match task::spawn_blocking(move || check_for_update(&args, base.as_deref())).await {
Ok(status) => status,
Err(err) => {
warn!("Update check task cancelled: {err}");
UpdateStatus::default()
}
}
}
pub fn rewrite_argv_for_reexec(argv: impl IntoIterator<Item = OsString>) -> Vec<OsString> {
fn os_starts_with(tok: &OsString, prefix: &[u8]) -> bool {
#[cfg(unix)]
{
use std::os::unix::ffi::OsStrExt;
tok.as_os_str().as_bytes().starts_with(prefix)
}
#[cfg(windows)]
{
use std::os::windows::ffi::OsStrExt;
let prefix_wide: Vec<u16> = prefix.iter().map(|&b| b as u16).collect();
let tok_wide: Vec<u16> = tok.as_os_str().encode_wide().collect();
tok_wide.starts_with(&prefix_wide)
}
#[cfg(not(any(unix, windows)))]
{
tok.to_str().map(|s| s.as_bytes().starts_with(prefix)).unwrap_or(false)
}
}
let mut iter = argv.into_iter();
let mut out: Vec<OsString> = Vec::new();
let mut already_has_no_update_check = false;
let mut hit_double_dash = false;
let mut double_dash_idx: Option<usize> = None;
let had_argv0;
if let Some(argv0) = iter.next() {
out.push(argv0);
had_argv0 = true;
} else {
had_argv0 = false;
}
for tok in iter {
if hit_double_dash {
out.push(tok);
continue;
}
if tok == "--" {
hit_double_dash = true;
double_dash_idx = Some(out.len());
out.push(tok);
continue;
}
if tok == "--self-update" || tok == "--update" {
continue;
}
if os_starts_with(&tok, b"--self-update=") || os_starts_with(&tok, b"--update=") {
continue;
}
if tok == "--no-update-check" {
already_has_no_update_check = true;
}
out.push(tok);
}
if had_argv0 && !already_has_no_update_check {
let flag = OsString::from("--no-update-check");
match double_dash_idx {
Some(idx) => out.insert(idx, flag),
None => out.push(flag),
}
}
out
}
#[cfg(test)]
mod tests {
use super::*;
fn os(s: &str) -> OsString {
OsString::from(s)
}
fn argv(args: &[&str]) -> Vec<OsString> {
args.iter().map(|s| os(s)).collect()
}
#[test]
fn rewrite_argv_strips_self_update() {
let result = rewrite_argv_for_reexec(argv(&["kingfisher", "scan", ".", "--self-update"]));
assert_eq!(result, argv(&["kingfisher", "scan", ".", "--no-update-check"]));
}
#[test]
fn rewrite_argv_strips_update_alias() {
let result = rewrite_argv_for_reexec(argv(&["kingfisher", "scan", "foo", "--update"]));
assert_eq!(result, argv(&["kingfisher", "scan", "foo", "--no-update-check"]));
}
#[test]
fn rewrite_argv_strips_eq_form() {
let result =
rewrite_argv_for_reexec(argv(&["kingfisher", "--self-update=true", "scan", "foo"]));
assert_eq!(result, argv(&["kingfisher", "scan", "foo", "--no-update-check"]));
let result = rewrite_argv_for_reexec(argv(&["kingfisher", "--update=true", "scan", "foo"]));
assert_eq!(result, argv(&["kingfisher", "scan", "foo", "--no-update-check"]));
}
#[test]
fn rewrite_argv_appends_no_update_check_when_absent() {
let result = rewrite_argv_for_reexec(argv(&["kingfisher", "scan", "."]));
assert_eq!(result, argv(&["kingfisher", "scan", ".", "--no-update-check"]));
}
#[test]
fn rewrite_argv_idempotent_when_no_update_check_already_present() {
let result = rewrite_argv_for_reexec(argv(&[
"kingfisher",
"scan",
".",
"--no-update-check",
"--self-update",
]));
assert_eq!(result, argv(&["kingfisher", "scan", ".", "--no-update-check"]));
}
#[test]
fn rewrite_argv_preserves_argv0() {
let result = rewrite_argv_for_reexec(argv(&[
"/weird path/kingfisher-bin",
"scan",
".",
"--self-update",
]));
assert_eq!(result, argv(&["/weird path/kingfisher-bin", "scan", ".", "--no-update-check"]));
}
#[test]
fn rewrite_argv_preserves_tokens_after_double_dash() {
let result = rewrite_argv_for_reexec(argv(&[
"kingfisher",
"scan",
"--self-update",
"--",
"--self-update",
"--update",
]));
assert_eq!(
result,
argv(&["kingfisher", "scan", "--no-update-check", "--", "--self-update", "--update"])
);
}
#[test]
fn rewrite_argv_does_not_duplicate_no_update_check_when_already_present_before_double_dash() {
let result = rewrite_argv_for_reexec(argv(&[
"kingfisher",
"scan",
"--no-update-check",
"--self-update",
"--",
"--self-update",
]));
assert_eq!(
result,
argv(&["kingfisher", "scan", "--no-update-check", "--", "--self-update"])
);
}
#[test]
fn rewrite_argv_empty_input_returns_empty() {
let result: Vec<OsString> = rewrite_argv_for_reexec(Vec::<OsString>::new());
assert!(result.is_empty(), "empty input must produce empty output, got {:?}", result);
}
#[test]
fn rewrite_argv_does_not_strip_unrelated_update_prefixed_flags() {
let result = rewrite_argv_for_reexec(argv(&[
"kingfisher",
"rules",
"--update-rules",
"--self-updateable=ignored",
]));
assert_eq!(
result,
argv(&[
"kingfisher",
"rules",
"--update-rules",
"--self-updateable=ignored",
"--no-update-check"
])
);
}
#[cfg(unix)]
#[test]
fn rewrite_argv_handles_non_utf8_value_in_eq_form() {
use std::os::unix::ffi::OsStringExt;
let mut bad = b"--self-update=".to_vec();
bad.extend_from_slice(&[0xff, 0xfe]); let bad_os = OsString::from_vec(bad);
let input: Vec<OsString> = vec![os("kingfisher"), os("scan"), bad_os, os(".")];
let result = rewrite_argv_for_reexec(input);
assert_eq!(result, argv(&["kingfisher", "scan", ".", "--no-update-check"]));
}
}