diff --git a/strix/interface/main.py b/strix/interface/main.py index cd36b778..4fa8082d 100644 --- a/strix/interface/main.py +++ b/strix/interface/main.py @@ -872,8 +872,10 @@ def main() -> None: apply_config_override(validate_config_file(args.config)) start_background_check() - if prompt_update_if_available(Console()) and is_binary_install() and sys.platform != "win32": - os.execv(sys.executable, sys.argv) # noqa: S606 # nosec B606 + if prompt_update_if_available(Console()): + if is_binary_install() and sys.platform != "win32": + os.execv(sys.executable, sys.argv) # noqa: S606 # nosec B606 + sys.exit(0) check_docker_installed() pull_docker_image() diff --git a/strix/interface/update_check.py b/strix/interface/update_check.py index b3260fe0..56ed4cd5 100644 --- a/strix/interface/update_check.py +++ b/strix/interface/update_check.py @@ -16,6 +16,7 @@ import os import platform import shutil import stat +import subprocess import sys import tarfile import tempfile @@ -27,7 +28,7 @@ from typing import cast import requests from rich.console import Console -from rich.prompt import Confirm +from rich.prompt import Prompt from strix.telemetry._common import get_version @@ -152,21 +153,25 @@ def _read_cache() -> dict[str, object]: return {} -def _write_cache(latest_version: str) -> None: +def _write_cache(**fields: object) -> None: try: + cache = _read_cache() + cache.update(fields) _CACHE_PATH.parent.mkdir(parents=True, exist_ok=True) - _CACHE_PATH.write_text( - json.dumps({"latest_version": latest_version, "checked_at": time.time()}), - encoding="utf-8", - ) + _CACHE_PATH.write_text(json.dumps(cache), encoding="utf-8") except Exception: # noqa: BLE001, S110 pass # nosec B110 +def skip_version(version: str) -> None: + """Remember not to prompt again for this version (newer releases still notify).""" + _write_cache(skipped_version=version) + + def _refresh_cache() -> None: latest = _fetch_latest_version() if latest: - _write_cache(latest) + _write_cache(latest_version=latest, checked_at=time.time()) def start_background_check() -> None: @@ -182,17 +187,20 @@ def start_background_check() -> None: _background_thread.start() -def get_available_update() -> str | None: +def get_available_update(*, respect_skip: bool = True) -> str | None: """Return the newer version from the cache, or None if up to date / unknown.""" if _is_disabled(): return None if _background_thread is not None: _background_thread.join(timeout=0.2) - latest = _read_cache().get("latest_version") + cache = _read_cache() + latest = cache.get("latest_version") current = get_version() - if isinstance(latest, str) and current != "unknown" and _is_newer(latest, current): - return latest - return None + if not isinstance(latest, str) or current == "unknown" or not _is_newer(latest, current): + return None + if respect_skip and cache.get("skipped_version") == latest: + return None + return latest def notify_update(console: Console) -> None: @@ -207,10 +215,29 @@ def notify_update(console: Console) -> None: console.print() -def prompt_update_if_available(console: Console) -> bool: - """Offer an interactive self-update before a scan starts. +def run_package_upgrade(console: Console, method: str) -> bool: + """Upgrade a package-manager install by running its upgrade command.""" + command = get_upgrade_command(method).split() + console.print(f"[dim]Running[/] [#60a5fa]{' '.join(command)}[/]") + try: + result = subprocess.run(command, check=False) # noqa: S603 + except OSError as e: + console.print(f"[bold red]Update failed:[/] {e}") + return False + if result.returncode != 0: + console.print( + f"[bold red]Update failed[/] [dim](exit code {result.returncode}).[/] " + f"Run it manually: [#60a5fa]{get_upgrade_command(method)}[/]" + ) + return False + console.print("[#22c55e]✓ strix updated — restart the scan to use the new version[/]") + return True - Returns True if the binary was updated (caller should re-exec). + +def prompt_update_if_available(console: Console) -> bool: + """Offer an interactive update before a scan starts. + + Returns True if strix was updated (caller should re-exec / exit). """ latest = get_available_update() if not latest or not sys.stdin.isatty() or not sys.stdout.isatty(): @@ -220,14 +247,20 @@ def prompt_update_if_available(console: Console) -> bool: f"[#eab308]A new version of strix is available:[/] " f"[dim]{get_version()}[/] [dim]→[/] [bold #22c55e]{latest}[/]" ) - if not is_binary_install(): - console.print(f"[dim]Upgrade with:[/] [#60a5fa]{get_upgrade_command()}[/]") - console.print() + console.print( + "[dim] y — update now n — not now (ask again next run) s — skip this version[/]" + ) + choice = Prompt.ask("Update strix?", choices=["y", "n", "s"], default="n") + console.print() + if choice == "s": + skip_version(latest) return False - if not Confirm.ask("Update now?", default=False): - console.print() + if choice != "y": return False - return self_update(console, version=latest) + method = get_install_method() + if method == "binary": + return self_update(console, version=latest) + return run_package_upgrade(console, method) def _release_target() -> str | None: @@ -351,6 +384,6 @@ def self_update(console: Console | None = None, version: str | None = None) -> b ) return False - _write_cache(latest) + _write_cache(latest_version=latest, checked_at=time.time()) console.print(f"[#22c55e]✓ Updated strix to {latest}[/]") return True diff --git a/tests/test_update_check.py b/tests/test_update_check.py index 6ef40ac8..b602f285 100644 --- a/tests/test_update_check.py +++ b/tests/test_update_check.py @@ -105,6 +105,33 @@ def test_background_check_runs_when_stale(monkeypatch: pytest.MonkeyPatch) -> No assert cache["latest_version"] == "1.2.3" +def test_skipped_version_suppresses_update(monkeypatch: pytest.MonkeyPatch) -> None: + update_check._CACHE_PATH.write_text( + json.dumps({"latest_version": "9.9.9", "checked_at": time.time()}) + ) + monkeypatch.setattr(update_check, "get_version", lambda: "1.0.0") + update_check.skip_version("9.9.9") + assert update_check.get_available_update() is None + assert update_check.get_available_update(respect_skip=False) == "9.9.9" + + +def test_newer_release_overrides_skipped_version(monkeypatch: pytest.MonkeyPatch) -> None: + update_check._CACHE_PATH.write_text( + json.dumps( + {"latest_version": "9.9.10", "checked_at": time.time(), "skipped_version": "9.9.9"} + ) + ) + monkeypatch.setattr(update_check, "get_version", lambda: "1.0.0") + assert update_check.get_available_update() == "9.9.10" + + +def test_write_cache_preserves_existing_fields() -> None: + update_check.skip_version("9.9.9") + update_check._write_cache(latest_version="1.2.3", checked_at=123.0) + cache = json.loads(update_check._CACHE_PATH.read_text()) + assert cache == {"latest_version": "1.2.3", "checked_at": 123.0, "skipped_version": "9.9.9"} + + def test_get_upgrade_command_all_methods() -> None: assert update_check.get_upgrade_command("binary") == "strix --update" assert update_check.get_upgrade_command("pipx") == "pipx upgrade strix-agent"