diff --git a/tests/pm/test_security_consumers.py b/tests/pm/test_security_consumers.py index 1fa8441764..bc9cf77639 100644 --- a/tests/pm/test_security_consumers.py +++ b/tests/pm/test_security_consumers.py @@ -314,6 +314,8 @@ def test_tirith_failed_cold_scans_make_one_attempt_then_explicit_can_retry(consu real_get(handler) monkeypatch.setattr(RangeHandler, "do_GET", record) tirith.check_command_security("echo hello") + for thread in tirith._install_threads.values(): + thread.join(10) first_attempt = list(requests) assert first_attempt tirith.check_command_security("echo hello") diff --git a/tests/tools/test_tirith_security.py b/tests/tools/test_tirith_security.py index c99fbdf991..56ea702016 100644 --- a/tests/tools/test_tirith_security.py +++ b/tests/tools/test_tirith_security.py @@ -348,8 +348,11 @@ class TestPmInstall: ensure.side_effect = lambda *_a, **_k: setattr( installed, "return_value", MagicMock(binary="/pm/tirith")) - assert _tirith_mod._resolve_tirith_path("tirith") == "/pm/tirith" + assert _tirith_mod._resolve_tirith_path("tirith") == "tirith" + for thread in _tirith_mod._install_threads.values(): + thread.join(5) ensure.assert_called_once_with("tirith") + assert _tirith_mod._resolve_tirith_path("tirith") == "/pm/tirith" def test_failed_install_is_not_retried(self, pm_tirith): """After a failed install, subsequent resolves fall back without retrying.""" @@ -357,6 +360,8 @@ class TestPmInstall: ensure.side_effect = RuntimeError("download failed") assert _tirith_mod._resolve_tirith_path("tirith") == "tirith" + for thread in _tirith_mod._install_threads.values(): + thread.join(5) assert _tirith_mod._resolve_tirith_path("tirith") == "tirith" assert ensure.call_count == 1 diff --git a/tools/tirith_security.py b/tools/tirith_security.py index 5dd3ce5083..e90a744b74 100644 --- a/tools/tirith_security.py +++ b/tools/tirith_security.py @@ -64,7 +64,7 @@ _circuit_open_at: float = 0.0 _breaker_lock = threading.Lock() # Warn-once: spawn/path warnings sit in the hot path and would otherwise repeat once per -# terminal command while tirith is unavailable (e.g. install thread still running). +# terminal command while tirith is unavailable. _warned_messages: set[str] = set() _warned_lock = threading.Lock() @@ -144,6 +144,11 @@ def _claim_install_attempt() -> bool: return True +def _install_in_flight() -> threading.Thread | None: + thread = _install_threads.get(hermes_home_key()) + return thread if thread is not None and thread.is_alive() else None + + def is_platform_supported() -> bool: """Whether PM has a managed Tirith build for this host.""" import pm @@ -174,15 +179,8 @@ def _resolve_tirith_path(configured_path: str) -> str: if configured_path == "tirith": import pm - if not pm.lazy_installs_allowed() or not _claim_install_attempt(): - return os.path.expanduser(configured_path) - try: - pm.ensure("tirith") - selected = pm.installed_package("tirith") - if selected and selected.binary: - return str(selected.binary) - except Exception as exc: - _warn_once("tirith_install", "tirith install unavailable: %s", exc) + if pm.lazy_installs_allowed(): + _start_background_install(log_failures=True) return os.path.expanduser(configured_path) @@ -196,6 +194,17 @@ def _background_install(*, log_failures: bool) -> None: log("tirith install failed: %s", exc) +def _start_background_install(*, log_failures: bool) -> None: + if _claim_install_attempt(): + context = copy_context() + thread = threading.Thread( + target=context.run, args=(_background_install,), + kwargs={"log_failures": log_failures}, daemon=True, + ) + _install_threads[hermes_home_key()] = thread + thread.start() + + def ensure_installed(*, log_failures: bool = True, explicit: bool = False): """Opt-in startup is non-blocking. Explicit setup waits and reports errors. @@ -218,14 +227,7 @@ def ensure_installed(*, log_failures: bool = True, explicit: bool = False): return found if not is_platform_supported() or not pm.lazy_installs_allowed(): return None - if _claim_install_attempt(): - context = copy_context() - thread = threading.Thread( - target=context.run, args=(_background_install,), - kwargs={"log_failures": log_failures}, daemon=True, - ) - _install_threads[hermes_home_key()] = thread - thread.start() + _start_background_install(log_failures=log_failures) return None @@ -241,8 +243,7 @@ def missing_is_expected() -> bool: configured = _load_security_config()["tirith_path"] if configured != "tirith": return False - thread = _install_threads.get(hermes_home_key()) - if thread is not None and thread.is_alive(): + if _install_in_flight(): return True return _local_tirith(configured) is not None or not pm.lazy_installs_allowed() @@ -310,6 +311,11 @@ def check_command_security(command: str) -> dict: if tirith_path is None: _warn_once("tirith_path_none", "tirith path resolved to None; scanning disabled") return _fail(fail_open, "tirith path unavailable", "tirith path unavailable (fail-closed)") + if tirith_path == "tirith" and (install := _install_in_flight()): + if fail_open: + return _verdict("allow", "tirith installing") + install.join() + tirith_path = _resolve_tirith_path(cfg["tirith_path"]) try: result = subprocess.run( [tirith_path, "check", "--json", "--non-interactive", "--shell", "posix", "--", command],