From fae5de890d25d330d494db359f8895b7b7390237 Mon Sep 17 00:00:00 2001 From: Charlie Masters Date: Tue, 29 Sep 2026 18:14:35 +0100 Subject: [PATCH 01/35] wip(sdk): add agent placement and cancellable local execution --- src/hai_agents/__init__.py | 3 + src/hai_agents/client.py | 146 ++++++++++ src/hai_agents/inference.py | 37 +++ src/hai_agents/local/__init__.py | 25 ++ src/hai_agents/local/errors.py | 27 ++ src/hai_agents/local/install.py | 184 ++++++++++++ src/hai_agents/local/manifest.py | 69 +++++ src/hai_agents/local/process.py | 148 ++++++++++ src/hai_agents/local/runtime.py | 436 ++++++++++++++++++++++++++++ src/hai_agents/local/state.py | 90 ++++++ src/hai_agents_local/bridge.py | 17 +- src/hai_agents_local/manager.py | 22 +- src/hai_agents_local/sessions.py | 55 ++++ src/hai_agents_local/workstation.py | 6 + tests/test_local.py | 81 ++++++ tests/test_runtime_placement.py | 91 ++++++ tests/test_runtime_state.py | 43 +++ 17 files changed, 1473 insertions(+), 7 deletions(-) create mode 100644 src/hai_agents/inference.py create mode 100644 src/hai_agents/local/__init__.py create mode 100644 src/hai_agents/local/errors.py create mode 100644 src/hai_agents/local/install.py create mode 100644 src/hai_agents/local/manifest.py create mode 100644 src/hai_agents/local/process.py create mode 100644 src/hai_agents/local/runtime.py create mode 100644 src/hai_agents/local/state.py create mode 100644 tests/test_runtime_placement.py create mode 100644 tests/test_runtime_state.py diff --git a/src/hai_agents/__init__.py b/src/hai_agents/__init__.py index 86cfeec..9850f8a 100644 --- a/src/hai_agents/__init__.py +++ b/src/hai_agents/__init__.py @@ -735,3 +735,6 @@ def __dir__(): "wait_for_session", "webhooks", ] + +from .inference import Inference +__all__.append("Inference") diff --git a/src/hai_agents/client.py b/src/hai_agents/client.py index 82e5354..f736416 100644 --- a/src/hai_agents/client.py +++ b/src/hai_agents/client.py @@ -7,11 +7,13 @@ from __future__ import annotations +import asyncio import typing import typing_extensions from .base_client import AsyncBaseClient, BaseClient +from .inference import Inference from .polling import ( AnswerT, AsyncSessionHandle, @@ -29,6 +31,76 @@ class Client(BaseClient): + def __init__( + self, + *, + mode: typing.Literal["local", "remote"] = "remote", + inference: typing.Optional[Inference] = None, + auto_bridges: bool = True, + runtime: typing.Any = None, + local_options: typing.Optional[typing.Dict[str, typing.Any]] = None, + **kwargs: typing.Any, + ) -> None: + if mode not in {"local", "remote"}: + raise ValueError("mode must be local or remote") + self._auto_bridges = auto_bridges + self.mode = mode + self.local_runtime = None + self._owns_runtime = mode == "local" and runtime is None + self._owns_http = kwargs.get("httpx_client") is None + if mode == "remote": + if runtime is not None or local_options is not None: + raise ValueError("runtime and local_options require mode='local'") + if inference is not None and inference.base_url is not None: + raise ValueError("self-hosted inference currently requires a local agent") + else: + if "base_url" in kwargs or "api_key" in kwargs: + raise ValueError( + "local API credentials come from runtime; pass inference credentials via local_options" + ) + if runtime is not None and (local_options is not None or inference is not None): + raise ValueError("an attached runtime owns its inference and launch configuration") + if runtime is None: + from .local.runtime import LocalRuntime + + options = dict(local_options or {}) + options["required_recipe"] = "shared" + options["spawn_env"] = {"HAI_AGENT_RUNTIME_RECIPE": "shared", **options.get("spawn_env", {})} + if inference is not None: + options["spawn_env"] = inference.runtime_env(options.get("spawn_env")) + options["inherit_env"] = False + runtime = LocalRuntime.ensure_started(**options) + if inference is not None and not runtime.owned: + raise ValueError("inference selection cannot reconfigure an existing runtime; choose a free local port") + runtime.require_recipe("shared") + self.local_runtime = runtime + kwargs.update(base_url=runtime.base_url, api_key=runtime.api_key) + try: + super().__init__(**kwargs) + except BaseException: + if self._owns_runtime and self.local_runtime is not None and self.local_runtime.owned: + self.local_runtime.shutdown() + raise + + def close(self) -> None: + """Release this client's connections and any runtime it started; borrowed runtimes stay alive.""" + try: + if self._sessions is not None and hasattr(self._sessions, "close"): + self._sessions.close() + finally: + try: + if self._owns_runtime and self.local_runtime is not None and self.local_runtime.owned: + self.local_runtime.shutdown() + finally: + if self._owns_http: + self._client_wrapper.httpx_client.httpx_client.close() + + def __enter__(self) -> "Client": + return self + + def __exit__(self, *exc: typing.Any) -> None: + self.close() + def run_session( self, *, @@ -78,6 +150,8 @@ def session(self, id: str) -> SessionHandle: @property def sessions(self) -> SessionsClient: + if not self._auto_bridges: + return super().sessions if self._sessions is None: from hai_agents_local.sessions import LocalSessionsClient @@ -86,6 +160,76 @@ def sessions(self) -> SessionsClient: class AsyncClient(AsyncBaseClient): + def __init__( + self, + *, + mode: typing.Literal["local", "remote"] = "remote", + inference: typing.Optional[Inference] = None, + auto_bridges: bool = True, + runtime: typing.Any = None, + local_options: typing.Optional[typing.Dict[str, typing.Any]] = None, + **kwargs: typing.Any, + ) -> None: + if mode not in {"local", "remote"}: + raise ValueError("mode must be local or remote") + self._auto_bridges = auto_bridges + self.mode = mode + self.local_runtime = None + self._owns_runtime = mode == "local" and runtime is None + self._owns_http = kwargs.get("httpx_client") is None + if mode == "remote": + if runtime is not None or local_options is not None: + raise ValueError("runtime and local_options require mode='local'") + if inference is not None and inference.base_url is not None: + raise ValueError("self-hosted inference currently requires a local agent") + else: + if "base_url" in kwargs or "api_key" in kwargs: + raise ValueError( + "local API credentials come from runtime; pass inference credentials via local_options" + ) + if runtime is not None and (local_options is not None or inference is not None): + raise ValueError("an attached runtime owns its inference and launch configuration") + if runtime is None: + from .local.runtime import LocalRuntime + + options = dict(local_options or {}) + options["required_recipe"] = "shared" + options["spawn_env"] = {"HAI_AGENT_RUNTIME_RECIPE": "shared", **options.get("spawn_env", {})} + if inference is not None: + options["spawn_env"] = inference.runtime_env(options.get("spawn_env")) + options["inherit_env"] = False + runtime = LocalRuntime.ensure_started(**options) + if inference is not None and not runtime.owned: + raise ValueError("inference selection cannot reconfigure an existing runtime; choose a free local port") + runtime.require_recipe("shared") + self.local_runtime = runtime + kwargs.update(base_url=runtime.base_url, api_key=runtime.api_key) + try: + super().__init__(**kwargs) + except BaseException: + if self._owns_runtime and self.local_runtime is not None and self.local_runtime.owned: + self.local_runtime.shutdown() + raise + + async def aclose(self) -> None: + """Release this client's connections and any runtime it started; borrowed runtimes stay alive.""" + try: + if self._sessions is not None and hasattr(self._sessions, "aclose"): + await self._sessions.aclose() + finally: + try: + if self._owns_runtime and self.local_runtime is not None and self.local_runtime.owned: + await asyncio.to_thread(self.local_runtime.shutdown) + finally: + if self._owns_http: + await self._client_wrapper.httpx_client.httpx_client.aclose() + + async def __aenter__(self) -> "AsyncClient": + return self + + async def __aexit__(self, *exc: typing.Any) -> None: + await self.aclose() + async def run_session( self, *, @@ -135,6 +279,8 @@ def session(self, id: str) -> AsyncSessionHandle: @property def sessions(self) -> AsyncSessionsClient: + if not self._auto_bridges: + return super().sessions if self._sessions is None: from hai_agents_local.sessions import LocalAsyncSessionsClient diff --git a/src/hai_agents/inference.py b/src/hai_agents/inference.py new file mode 100644 index 0000000..c843c36 --- /dev/null +++ b/src/hai_agents/inference.py @@ -0,0 +1,37 @@ +"""Inference placement, independent of agent and environment placement.""" + +from __future__ import annotations + +import os +from dataclasses import dataclass +from typing import Dict, Optional +from urllib.parse import urlsplit + + +@dataclass(frozen=True) +class Inference: + base_url: Optional[str] = None + model: Optional[str] = None + + @classmethod + def cloud(cls) -> "Inference": + return cls() + + @classmethod + def self_hosted(cls, base_url: str, *, model: str) -> "Inference": + parsed = urlsplit(base_url) + if parsed.scheme not in {"http", "https"} or not parsed.hostname or parsed.username or not model: + raise ValueError("self-hosted inference requires an HTTP(S) endpoint and model") + return cls(base_url=base_url, model=model) + + def runtime_env(self, overrides: Optional[Dict[str, str]] = None) -> Dict[str, str]: + env = {**os.environ, **(overrides or {})} + if self.base_url is not None: + # Never forward a hosted inference credential to a user-selected endpoint. + env.pop("HAI_API_KEY", None) + env["HAI_AGENT_RUNTIME_BASE_URL"] = self.base_url + else: + env.pop("HAI_AGENT_RUNTIME_BASE_URL", None) + if self.model is not None: + env["HAI_AGENT_RUNTIME_MODEL"] = self.model + return env diff --git a/src/hai_agents/local/__init__.py b/src/hai_agents/local/__init__.py new file mode 100644 index 0000000..215e46d --- /dev/null +++ b/src/hai_agents/local/__init__.py @@ -0,0 +1,25 @@ +"""Local-mode runtime management: install/find/start a hai-agent-runtime binary. + +Never imported by the base ``hai_agents`` package; ``Client.local`` pulls it in +lazily so remote-only users pay nothing for it. +""" + +from .errors import ( + BinaryIncompatibleError, + BinaryNotFoundError, + DownloadVerificationError, + LocalRuntimeError, + RuntimeStartTimeoutError, + RuntimeUnhealthyError, +) +from .runtime import LocalRuntime + +__all__ = [ + "BinaryIncompatibleError", + "BinaryNotFoundError", + "DownloadVerificationError", + "LocalRuntime", + "LocalRuntimeError", + "RuntimeStartTimeoutError", + "RuntimeUnhealthyError", +] diff --git a/src/hai_agents/local/errors.py b/src/hai_agents/local/errors.py new file mode 100644 index 0000000..dfff655 --- /dev/null +++ b/src/hai_agents/local/errors.py @@ -0,0 +1,27 @@ +"""Error types for hai_agents.local.""" + +from __future__ import annotations + + +class LocalRuntimeError(RuntimeError): + """Base error for local hai-agent-runtime management.""" + + +class BinaryNotFoundError(LocalRuntimeError): + """No runtime binary: no override, nothing on PATH, no managed install, and download disabled.""" + + +class BinaryIncompatibleError(LocalRuntimeError): + """The pinned manifest has no verifiable artifact for this platform or requested version.""" + + +class RuntimeUnhealthyError(LocalRuntimeError): + """The runtime process exited, or /health is not answering with a 200.""" + + +class RuntimeStartTimeoutError(LocalRuntimeError): + """The spawned runtime did not become healthy before the timeout.""" + + +class DownloadVerificationError(LocalRuntimeError): + """A runtime download could not be sha256-verified (mismatch, or no digest to verify against).""" diff --git a/src/hai_agents/local/install.py b/src/hai_agents/local/install.py new file mode 100644 index 0000000..a662c8d --- /dev/null +++ b/src/hai_agents/local/install.py @@ -0,0 +1,184 @@ +"""Verified download and atomic install of the hai-agent-runtime binary. + +Port of holo_desktop.agent_client.runtime_install with the TTY prompt and rich +progress removed: the SDK is a library, so consent is the caller's +``download=True`` and progress is plain logging. +""" + +from __future__ import annotations + +import hashlib +import logging +import os +import pathlib +import shutil +import tempfile +import typing +import zipfile +from urllib.parse import urlsplit + +import httpx + +from .errors import BinaryIncompatibleError, DownloadVerificationError, LocalRuntimeError +from .manifest import ( + BINARY_NAME, + MANIFEST, + PINNED_RUNTIME_VERSION, + PLACEHOLDER_SHA256, + UNIMPLEMENTED_PLATFORMS, + RuntimeArtifact, + platform_key, +) +from .state import resolve_cache_dir + +logger = logging.getLogger(__name__) + +DOWNLOAD_URL_ENV = "HAI_AGENT_RUNTIME_DOWNLOAD_URL" +DOWNLOAD_SHA256_ENV = "HAI_AGENT_RUNTIME_DOWNLOAD_SHA256" +_DOWNLOAD_TIMEOUT = httpx.Timeout(30.0, read=600.0) +# Generous ceiling (the runtime is hundreds of MB); guards against a lying/absent Content-Length. +MAX_DOWNLOAD_BYTES = 1024 * 1024 * 1024 +_LOOPBACK_HOSTS = frozenset({"127.0.0.1", "localhost", "::1"}) + +_PathInput = typing.Union[str, "os.PathLike[str]"] + + +def _require_secure_url(url: str) -> None: + """Allow https anywhere, plain http only against loopback (test/ops overrides); reject the rest.""" + parsed = urlsplit(url) + if parsed.scheme == "https" or (parsed.scheme == "http" and parsed.hostname in _LOOPBACK_HOSTS): + return + raise LocalRuntimeError(f"refusing insecure hai-agent-runtime download URL (need https): {url}") + + +def bin_dir(version: str, *, cache_dir: typing.Optional[_PathInput] = None) -> pathlib.Path: + return resolve_cache_dir(cache_dir) / "bin" / version + + +def _find_binary(root: pathlib.Path) -> typing.Optional[pathlib.Path]: + direct = root / BINARY_NAME + if direct.is_file(): + return direct + # macOS app-bundle shape: /.app/Contents/MacOS/hai-agent-runtime + for candidate in sorted(root.glob("*.app/Contents/MacOS/hai-agent-runtime")): + if candidate.is_file(): + return candidate + return None + + +def installed_binary(version: str, *, cache_dir: typing.Optional[_PathInput] = None) -> typing.Optional[pathlib.Path]: + """The managed install's executable for `version`, or None if absent/incomplete.""" + return _find_binary(bin_dir(version, cache_dir=cache_dir)) + + +def pinned_artifact() -> RuntimeArtifact: + """Artifact to install: env override (tests/ops) or the pinned per-platform manifest entry.""" + override_url = os.environ.get(DOWNLOAD_URL_ENV, "").strip() + if override_url: + override_sha = os.environ.get(DOWNLOAD_SHA256_ENV, "").strip() + if not override_sha: + raise DownloadVerificationError( + f"{DOWNLOAD_URL_ENV} is set but {DOWNLOAD_SHA256_ENV} is not; refusing an unverified download" + ) + return RuntimeArtifact(url=override_url, sha256=override_sha) + key = platform_key() + if key in UNIMPLEMENTED_PLATFORMS: + raise BinaryIncompatibleError( + f"{UNIMPLEMENTED_PLATFORMS[key]}; put hai-agent-runtime on PATH, " + f"or set {DOWNLOAD_URL_ENV} + {DOWNLOAD_SHA256_ENV} to a trusted build" + ) + artifact = MANIFEST.get(key) + if artifact is None: + raise BinaryIncompatibleError(f"no hai-agent-runtime release artifact for platform {key}") + if artifact.sha256 == PLACEHOLDER_SHA256: + raise BinaryIncompatibleError( + f"hai-agent-runtime v{PINNED_RUNTIME_VERSION} has no published artifact for {key} yet; " + f"put hai-agent-runtime on PATH, or set {DOWNLOAD_URL_ENV} + {DOWNLOAD_SHA256_ENV} to a trusted build" + ) + return artifact + + +def install_runtime( + artifact: RuntimeArtifact, *, version: str, cache_dir: typing.Optional[_PathInput] = None +) -> pathlib.Path: + """Download, sha256-verify, and atomically install `artifact` as `version`; returns the executable path.""" + root = resolve_cache_dir(cache_dir) / "bin" + root.mkdir(parents=True, exist_ok=True) + version_dir = root / version + + # Stage on the same filesystem as the final location so os.replace stays atomic. + with tempfile.TemporaryDirectory(dir=root, prefix=".staging-") as staging_str: + staging = pathlib.Path(staging_str) + download_path = staging / "artifact" + actual_sha256 = _download_to(artifact.url, download_path) + if actual_sha256 != artifact.sha256.lower(): + raise DownloadVerificationError( + f"hai-agent-runtime download failed sha256 verification: expected {artifact.sha256}, " + f"got {actual_sha256} (url: {artifact.url})" + ) + + staged_version = staging / "version" + staged_version.mkdir() + if artifact.url.endswith(".zip"): + # Contents are sha256-verified above, so extraction is trusted. + with zipfile.ZipFile(download_path) as archive: + archive.extractall(staged_version) + else: + shutil.move(str(download_path), str(staged_version / BINARY_NAME)) + + binary = _find_binary(staged_version) + if binary is None: + raise LocalRuntimeError( + f"downloaded artifact contains no hai-agent-runtime executable (url: {artifact.url})" + ) + binary.chmod(0o755) # zipfile does not preserve the exec bit + + try: + os.replace(staged_version, version_dir) + except OSError: + # Target occupied: a concurrent installer won the race, or a half-finished dir is in the way. + existing = installed_binary(version, cache_dir=cache_dir) + if existing is not None: + logger.info("hai-agent-runtime %s already installed by a concurrent run", version) + return existing + shutil.rmtree(version_dir, ignore_errors=True) + os.replace(staged_version, version_dir) + + installed = installed_binary(version, cache_dir=cache_dir) + assert installed is not None, "atomic rename just published the staged install" + logger.info("installed hai-agent-runtime %s at %s", version, installed) + return installed + + +def _download_to(url: str, dest: pathlib.Path) -> str: + """Stream `url` into `dest`; returns the sha256 hex digest of the bytes written.""" + _require_secure_url(url) + digest = hashlib.sha256() + written = 0 + logger.info("downloading hai-agent-runtime from %s", url) + try: + with ( + httpx.Client(follow_redirects=True, timeout=_DOWNLOAD_TIMEOUT) as client, + client.stream("GET", url) as response, + ): + _require_secure_url(str(response.url)) # a redirect must not downgrade to plain http + if response.status_code != 200: + raise LocalRuntimeError(f"hai-agent-runtime download failed: HTTP {response.status_code} from {url}") + total = int(response.headers.get("Content-Length", "0")) or None + if total is not None and total > MAX_DOWNLOAD_BYTES: + raise LocalRuntimeError( + f"hai-agent-runtime download too large: {total} bytes exceeds {MAX_DOWNLOAD_BYTES}" + ) + with dest.open("wb") as fh: + for chunk in response.iter_bytes(): + written += len(chunk) + if written > MAX_DOWNLOAD_BYTES: + raise LocalRuntimeError( + f"hai-agent-runtime download exceeded {MAX_DOWNLOAD_BYTES} bytes; aborting" + ) + digest.update(chunk) + fh.write(chunk) + except httpx.HTTPError as exc: + raise LocalRuntimeError(f"hai-agent-runtime download failed: {exc} (url: {url})") from exc + logger.info("download complete (%d bytes)", written) + return digest.hexdigest() diff --git a/src/hai_agents/local/manifest.py b/src/hai_agents/local/manifest.py new file mode 100644 index 0000000..e9588f6 --- /dev/null +++ b/src/hai_agents/local/manifest.py @@ -0,0 +1,69 @@ +"""Pinned hai-agent-runtime version and per-platform artifact digests. + +This module is the SDK's single runtime pin: a runtime release bumps +PINNED_RUNTIME_VERSION and MANIFEST here (via the retargeted +release-hai-agent-runtime.yaml pin PR) and nothing else. Artifacts live under an +immutable version-scoped CDN prefix, so an edge can never serve stale bytes. +Cross-ref: eng_plans/14-06-2026-holodesktop-binary-versioning-autoupdate. +""" + +from __future__ import annotations + +import dataclasses +import platform +import sys +import typing + +# SHIP-GATE: repoint to the plan-005 release before merge +PINNED_RUNTIME_VERSION = "0.1.8" +RUNTIME_CDN_BASE = "https://assets.hcompanyprod.fr/hai-agent-runtime" +# Guard value: published manifest entries must never use it (every download would fail verification). +PLACEHOLDER_SHA256 = "0" * 64 +BINARY_NAME = "hai-agent-runtime.exe" if sys.platform == "win32" else "hai-agent-runtime" + + +@dataclasses.dataclass(frozen=True) +class RuntimeArtifact: + url: str + sha256: str + + +def _artifact(filename: str, sha256: str) -> RuntimeArtifact: + """A published release file resolved to its pinned, version-scoped CDN URL.""" + return RuntimeArtifact(url=f"{RUNTIME_CDN_BASE}/{PINNED_RUNTIME_VERSION}/{filename}", sha256=sha256) + + +MANIFEST: typing.Dict[str, RuntimeArtifact] = { + "darwin-arm64": _artifact( + "hai-agent-runtime-darwin-arm64.zip", + # SHIP-GATE: repoint to the plan-005 release before merge + "1aed0055898116732aee031dc4a1235782b2909ee51e0367e2d50bb3be6671c9", + ), + "windows-x86_64": _artifact( + "hai-agent-runtime-windows-x86_64.zip", + # SHIP-GATE: repoint to the plan-005 release before merge + "4e6b2bcd42af2bb6b22197fcde947327497f5c62fd60d48bc9037730d80dc691", + ), +} + +UNIMPLEMENTED_PLATFORMS: typing.Dict[str, str] = { + "darwin-x86_64": "hai-agent-runtime is not published for macOS Intel yet", + "linux-x86_64": "hai-agent-runtime is not published for Linux yet", +} + + +def platform_key() -> str: + """`-` manifest key for the current host, e.g. darwin-arm64.""" + if sys.platform == "darwin": + system = "darwin" + elif sys.platform.startswith("linux"): + system = "linux" + elif sys.platform == "win32": + system = "windows" + else: + raise RuntimeError(f"unsupported platform for hai-agent-runtime: {sys.platform}") + machine = platform.machine().lower() + arch = {"arm64": "arm64", "aarch64": "arm64", "x86_64": "x86_64", "amd64": "x86_64"}.get(machine) + if arch is None: + raise RuntimeError(f"unsupported architecture for hai-agent-runtime: {machine}") + return f"{system}-{arch}" diff --git a/src/hai_agents/local/process.py b/src/hai_agents/local/process.py new file mode 100644 index 0000000..50b45bb --- /dev/null +++ b/src/hai_agents/local/process.py @@ -0,0 +1,148 @@ +"""Spawn, health-check, and stop hai-agent-runtime processes (sync, loopback-only).""" + +from __future__ import annotations + +import logging +import os +import pathlib +import signal +import subprocess +import threading +import time +import typing + +import httpx + +from .errors import RuntimeStartTimeoutError, RuntimeUnhealthyError + +logger = logging.getLogger(__name__) + +LOOPBACK_HOST = "127.0.0.1" +SPAWN_TIMEOUT_S = 45.0 +HEALTH_POLL_INTERVAL_S = 0.25 +TERM_GRACE_S = 2.0 +LOG_TAIL_CHARS = 4000 + + +def probe_health(base_url: str) -> typing.Optional[typing.Dict[str, typing.Any]]: + """The /health JSON body on a 200 ({} for non-JSON bodies); None when unreachable/unhealthy.""" + try: + response = httpx.get(f"{base_url}/health", timeout=2.0) + except httpx.HTTPError: + return None + if response.status_code != 200: + return None + try: + payload = response.json() + except ValueError: + return {} + return payload if isinstance(payload, dict) else {} + + +def spawn(cmd: typing.List[str], *, env: typing.Dict[str, str], log_path: pathlib.Path) -> subprocess.Popen: + """Start the runtime in its own process group with stderr to `log_path`. + + stderr goes to a file, not a pipe: nobody drains a pipe after spawn, so the + buffer would fill and block. Own process group so we can reap grandchildren + (e.g. desktop helpers) the binary may spawn. + """ + log_path.parent.mkdir(parents=True, exist_ok=True) + logger.info("spawning hai-agent-runtime: %s (stderr -> %s)", " ".join(cmd), log_path) + with log_path.open("wb") as log_file: # child inherits the fd; the parent handle can close right away + return subprocess.Popen( + cmd, + stdout=subprocess.DEVNULL, + stderr=log_file, + env=env, + start_new_session=os.name == "posix", + creationflags=0 if os.name == "posix" else getattr(subprocess, "CREATE_NEW_PROCESS_GROUP", 0), + ) + + +def wait_healthy( + base_url: str, + proc: subprocess.Popen, + *, + timeout_s: float, + log_path: pathlib.Path, + cancel_event: typing.Optional[threading.Event] = None, +) -> typing.Dict[str, typing.Any]: + """Poll /health until 200; raises RuntimeUnhealthyError (child exited) or RuntimeStartTimeoutError.""" + deadline = time.monotonic() + timeout_s + while True: + if cancel_event is not None and cancel_event.is_set(): + raise RuntimeUnhealthyError("runtime startup cancelled") + payload = probe_health(base_url) + if payload is not None: + logger.info("hai-agent-runtime ready (pid %d)", proc.pid) + return payload + if proc.poll() is not None: + raise RuntimeUnhealthyError( + f"hai-agent-runtime exited with code {proc.returncode}: {log_tail(log_path)} (full log: {log_path})" + ) + if time.monotonic() >= deadline: + raise RuntimeStartTimeoutError( + f"hai-agent-runtime did not become healthy within {timeout_s:.0f}s (see {log_path})" + ) + time.sleep(HEALTH_POLL_INTERVAL_S) + + +def log_tail(path: pathlib.Path) -> str: + try: + text = path.read_text(encoding="utf-8", errors="replace").strip() + except OSError: + return "(stderr log unreadable)" + if not text: + return "(no stderr output)" + return text[-LOG_TAIL_CHARS:] + + +def _killpg_posix(pid: int, sig: int) -> bool: + """Send `sig` to `pid`'s process group; False if the process/group is already gone.""" + try: + os.killpg(os.getpgid(pid), sig) + except (OSError, ProcessLookupError): + return False + return True + + +def kill_process_group(pid: int) -> bool: + """Force-kill the runtime's process group by pid; False if it was already gone.""" + if os.name == "posix": + return _killpg_posix(pid, signal.SIGKILL) + try: + subprocess.run(["taskkill", "/F", "/T", "/PID", str(pid)], check=True, capture_output=True) + except (OSError, subprocess.CalledProcessError): + return False + return True + + +def _signal(proc: subprocess.Popen, *, force: bool) -> bool: + """Signal the runtime's whole process tree; False if it was already gone.""" + if os.name == "posix": + return _killpg_posix(proc.pid, signal.SIGKILL if force else signal.SIGTERM) + try: + # Windows has no portable graceful process-group signal. /F is required + # to ensure an executable launched through a .cmd shim cannot outlive it. + subprocess.run(["taskkill", "/F", "/T", "/PID", str(proc.pid)], check=True, capture_output=True) + except (OSError, subprocess.CalledProcessError): + return False + return True + + +def terminate(proc: subprocess.Popen) -> None: + """Graceful stop: SIGTERM the group, wait the grace period, then SIGKILL the group.""" + if proc.poll() is not None: + return + if not _signal(proc, force=False): + return + try: + proc.wait(timeout=TERM_GRACE_S) + return + except subprocess.TimeoutExpired: + pass + if _signal(proc, force=True): + try: + proc.wait(timeout=TERM_GRACE_S) + except subprocess.TimeoutExpired: + logger.warning("hai-agent-runtime (pid %d) did not exit after forced kill", proc.pid) diff --git a/src/hai_agents/local/runtime.py b/src/hai_agents/local/runtime.py new file mode 100644 index 0000000..664f4e9 --- /dev/null +++ b/src/hai_agents/local/runtime.py @@ -0,0 +1,436 @@ +"""SDK-managed local hai-agent-runtime: install/find/start/attach/stop.""" + +from __future__ import annotations + +import asyncio +import contextlib +import logging +import os +import pathlib +import secrets +import shutil +import subprocess +import threading +import time +import typing +from urllib.parse import urlsplit + +import httpx + +from .errors import ( + BinaryIncompatibleError, + BinaryNotFoundError, + LocalRuntimeError, + RuntimeUnhealthyError, +) +from .install import DOWNLOAD_SHA256_ENV, DOWNLOAD_URL_ENV, install_runtime, installed_binary, pinned_artifact +from .manifest import PINNED_RUNTIME_VERSION +from .process import ( + LOOPBACK_HOST, + SPAWN_TIMEOUT_S, + probe_health, + spawn, + terminate, + wait_healthy, +) +from .state import ( + DEFAULT_PORT, + pid_file_path, + read_pid, + read_state_file, + resolve_cache_dir, + runtime_log_path, + token_file_path, + unlink_if_content, + write_owner_only, +) + +logger = logging.getLogger(__name__) + +BINARY_PATH_ENV = "HAI_AGENT_LOCAL_BINARY_PATH" +BINARY_VERSION_ENV = "HAI_AGENT_LOCAL_BINARY_VERSION" +BASE_URL_ENV = "HAI_AGENT_LOCAL_BASE_URL" +PORT_ENV = "HAI_AGENT_RUNTIME_PORT" +AUTH_TOKEN_ENV = "HAI_AGENT_RUNTIME_API_TOKEN" + +_PathInput = typing.Union[str, "os.PathLike[str]"] + + +def _warn_on_version_skew(version: typing.Optional[str]) -> None: + """Warn (not fail) on a client/runtime version skew; PATH/override dev binaries stay usable.""" + if version is not None and version != PINNED_RUNTIME_VERSION: + logger.warning( + "hai-agent-runtime version skew: server reports %s, this SDK pins %s; " + "wire-contract drift may cause subtle failures", + version, + PINNED_RUNTIME_VERSION, + ) + + +def _port_of(base_url: str) -> int: + return urlsplit(base_url).port or DEFAULT_PORT + + +class LocalRuntime: + """A reachable local agent runtime: where it is, how to authenticate, and (if ours) the process.""" + + def __init__( + self, + *, + base_url: str, + api_key: str, + pid: typing.Optional[int], + version: typing.Optional[str], + log_path: typing.Optional[pathlib.Path], + owned: bool, + cache_dir: pathlib.Path, + port: int, + proc: typing.Optional[subprocess.Popen] = None, + token_file: typing.Optional[pathlib.Path] = None, + pid_file: typing.Optional[pathlib.Path] = None, + ) -> None: + self.base_url = base_url + self.api_key = api_key + self.pid = pid + self.version = version + self.log_path = log_path + self.owned = owned + self._cache_dir = cache_dir + self._port = port + self._proc = proc + # Set only on the spawner that published generated state; attachers never own the files. + self._token_file = token_file + self._pid_file = pid_file + + @classmethod + def ensure_started( + cls, + *, + required_recipe: typing.Optional[str] = None, + command: typing.Optional[typing.Sequence[str]] = None, + binary_path: typing.Optional[_PathInput] = None, + version: typing.Optional[str] = None, + cache_dir: typing.Optional[_PathInput] = None, + port: typing.Optional[int] = None, + spawn_env: typing.Optional[typing.Dict[str, str]] = None, + inherit_env: bool = True, + download: bool = True, + timeout_s: float = SPAWN_TIMEOUT_S, + _cancel_event: typing.Optional[threading.Event] = None, + ) -> "LocalRuntime": + """Return a reachable LocalRuntime, attaching to an existing one or spawning the binary.""" + resolved_cache = resolve_cache_dir(cache_dir) + base_override = os.environ.get(BASE_URL_ENV, "").strip() + if base_override: + attached = cls._attach(base_url=base_override.rstrip("/"), cache_dir=resolved_cache) + if attached is None: + raise RuntimeUnhealthyError( + f"{BASE_URL_ENV} is set to {base_override} but /health is not answering there" + ) + attached.require_recipe(required_recipe) + return attached + + resolved_port = port if port is not None else int(os.environ.get(PORT_ENV, "").strip() or DEFAULT_PORT) + with _startup_lock(resolved_cache, resolved_port, timeout_s): + base_url = f"http://{LOOPBACK_HOST}:{resolved_port}" + attached = cls._attach(base_url=base_url, cache_dir=resolved_cache) + if attached is not None: + attached.require_recipe(required_recipe) + return attached + + if command is not None and (not command or binary_path is not None): + raise ValueError("command must be nonempty and cannot be combined with binary_path") + cmd = ( + list(command) + if command is not None + else cls._resolve_command( + binary_path=binary_path, version=version, cache_dir=resolved_cache, download=download + ) + ) + if _cancel_event is not None and _cancel_event.is_set(): + raise RuntimeUnhealthyError("runtime startup cancelled") + explicit_token = ( + (spawn_env or {}).get(AUTH_TOKEN_ENV, os.environ.get(AUTH_TOKEN_ENV, "") if inherit_env else "").strip() + ) + token = explicit_token or secrets.token_urlsafe(32) + # Publish the token before the health wait so a client racing our probe can authenticate. + token_file = ( + None + if explicit_token + else write_owner_only(token_file_path(resolved_port, cache_dir=resolved_cache), token) + ) + log_path = runtime_log_path(resolved_port, cache_dir=resolved_cache) + proc = None + try: + proc = spawn( + cmd, + env=cls._child_env(port=resolved_port, token=token, spawn_env=spawn_env, inherit_env=inherit_env), + log_path=log_path, + ) + payload = wait_healthy( + base_url, proc, timeout_s=timeout_s, log_path=log_path, cancel_event=_cancel_event + ) + response = httpx.get( + f"{base_url}/api/v2/sessions", + headers={"Authorization": f"Bearer {token}"}, + params={"size": 1}, + timeout=2.0, + follow_redirects=False, + ) + if required_recipe is not None and payload.get("recipe") != required_recipe: + raise BinaryIncompatibleError( + f"runtime must support recipe {required_recipe!r}; use a compatible source command or binary" + ) + if response.status_code != 200 or proc.poll() is not None: + raise RuntimeUnhealthyError("spawned runtime failed authenticated readiness probe") + pid_file = write_owner_only(pid_file_path(resolved_port, cache_dir=resolved_cache), str(proc.pid)) + except BaseException: + # Covers KeyboardInterrupt mid-spawn: never leak the child or its token file. + if proc is not None: + terminate(proc) + if token_file is not None: + unlink_if_content(token_file, token) + raise + reported = payload.get("version") + reported_version = reported if isinstance(reported, str) else None + _warn_on_version_skew(reported_version) + return cls( + base_url=base_url, + api_key=token, + pid=proc.pid, + version=reported_version, + log_path=log_path, + owned=True, + cache_dir=resolved_cache, + port=resolved_port, + proc=proc, + token_file=token_file, + pid_file=pid_file, + ) + + @classmethod + async def ensure_started_async(cls, **options: typing.Any) -> "LocalRuntime": + """Start off the event loop; cancellation cleans up an owned child before returning.""" + cancelled = threading.Event() + startup = asyncio.create_task(asyncio.to_thread(cls.ensure_started, _cancel_event=cancelled, **options)) + try: + return await asyncio.shield(startup) + except asyncio.CancelledError: + cancelled.set() + # The worker may have finished between shield cancellation and setting the event. + with contextlib.suppress(Exception): + runtime = await startup + if runtime.owned: + await asyncio.to_thread(runtime.shutdown) + raise + + def require_recipe(self, recipe: typing.Optional[str]) -> None: + """Fail closed when an old binary or a differently configured daemon answers.""" + if recipe is None: + return + payload = probe_health(self.base_url) + if payload is None or payload.get("recipe") != recipe: + raise BinaryIncompatibleError( + f"runtime must support recipe {recipe!r}; use a compatible source command or binary" + ) + + @classmethod + def attach( + cls, *, port: typing.Optional[int] = None, cache_dir: typing.Optional[_PathInput] = None + ) -> typing.Optional["LocalRuntime"]: + """A LocalRuntime for an already-running local runtime, or None when nothing answers /health.""" + resolved_cache = resolve_cache_dir(cache_dir) + base_override = os.environ.get(BASE_URL_ENV, "").strip() + if base_override: + return cls._attach(base_url=base_override.rstrip("/"), cache_dir=resolved_cache) + resolved_port = port if port is not None else int(os.environ.get(PORT_ENV, "").strip() or DEFAULT_PORT) + return cls._attach(base_url=f"http://{LOOPBACK_HOST}:{resolved_port}", cache_dir=resolved_cache) + + @classmethod + def _attach(cls, *, base_url: str, cache_dir: pathlib.Path) -> typing.Optional["LocalRuntime"]: + parsed = urlsplit(base_url) + if parsed.scheme != "http" or parsed.hostname not in {"127.0.0.1", "localhost", "::1"} or parsed.username: + raise LocalRuntimeError("local runtime attachment requires a loopback HTTP URL") + port = _port_of(base_url) + payload = probe_health(base_url) + if payload is None: + return None + token = os.environ.get(AUTH_TOKEN_ENV, "").strip() or read_state_file( + token_file_path(port, cache_dir=cache_dir) + ) + if not token: + raise LocalRuntimeError( + f"an agent runtime is answering at {base_url} but no credentials were found: " + f"{AUTH_TOKEN_ENV} is not set and {token_file_path(port, cache_dir=cache_dir)} does not exist, " + "so this client cannot authenticate. Export the token or stop that runtime." + ) + response = httpx.get( + f"{base_url}/api/v2/sessions", + headers={"Authorization": f"Bearer {token}"}, + params={"size": 1}, + timeout=2.0, + follow_redirects=False, + ) + if response.status_code != 200: + raise LocalRuntimeError("runtime attachment failed authenticated session probe") + reported = payload.get("version") + reported_version = reported if isinstance(reported, str) else None + _warn_on_version_skew(reported_version) + log_path = runtime_log_path(port, cache_dir=cache_dir) + return cls( + base_url=base_url, + api_key=token, + pid=read_pid(port, cache_dir=cache_dir), + version=reported_version, + log_path=log_path if log_path.exists() else None, + owned=False, + cache_dir=cache_dir, + port=port, + ) + + @staticmethod + def _child_env( + *, + port: int, + token: str, + spawn_env: typing.Optional[typing.Dict[str, str]], + inherit_env: bool, + ) -> typing.Dict[str, str]: + """Child env: inherited-plus-overlay by default, caller-verbatim with inherit_env=False. + + Inheriting os.environ passes the model-gateway HAI_API_KEY / HAI_BASE_URL through to the + binary (without them local sessions cannot run inference) and forwards caller flags such as + HAI_AGENT_RUNTIME_MODEL/FAKE/FAST/RUNS_DIR. inherit_env=False takes spawn_env as the + complete base environment instead — for callers that must *remove* inherited keys, which an + overlay cannot express (HoloDesktop strips HAI_API_KEY for self-hosted base URLs). The + generated local bearer and the cloud HAI_API_KEY are different credentials: the token below + is the only local bearer, and the cloud key is never used to authenticate against the local + runtime. Port and token are set last in both modes so caller input never clobbers them. + """ + env = {**os.environ, **(spawn_env or {})} if inherit_env else dict(spawn_env or {}) + env[PORT_ENV] = str(port) + env[AUTH_TOKEN_ENV] = token + return env + + @staticmethod + def _resolve_command( + *, + binary_path: typing.Optional[_PathInput], + version: typing.Optional[str], + cache_dir: pathlib.Path, + download: bool, + ) -> typing.List[str]: + """Explicit path > HAI_AGENT_LOCAL_BINARY_PATH > PATH > managed install > verified download.""" + explicit = str(binary_path) if binary_path is not None else os.environ.get(BINARY_PATH_ENV, "").strip() + if explicit: + candidate = pathlib.Path(explicit).expanduser() + if not candidate.is_file(): + raise BinaryNotFoundError(f"binary_path / {BINARY_PATH_ENV} points at a missing file: {candidate}") + return [str(candidate)] + found = shutil.which("hai-agent-runtime") + if found: + logger.info("resolved hai-agent-runtime from PATH: %s", found) + return [found] + pinned = version or os.environ.get(BINARY_VERSION_ENV, "").strip() or PINNED_RUNTIME_VERSION + managed = installed_binary(pinned, cache_dir=cache_dir) + if managed is not None: + logger.info("resolved hai-agent-runtime from managed install v%s: %s", pinned, managed) + return [str(managed)] + if not download: + raise BinaryNotFoundError( + "hai-agent-runtime not found: not on PATH, no managed install under " + f"{cache_dir / 'bin'}, and download=False. Pass binary_path=, set {BINARY_PATH_ENV}, " + "or allow download=True." + ) + if pinned != PINNED_RUNTIME_VERSION and not os.environ.get(DOWNLOAD_URL_ENV, "").strip(): + raise BinaryIncompatibleError( + f"cannot download hai-agent-runtime {pinned}: this SDK pins sha256 digests for " + f"{PINNED_RUNTIME_VERSION} only. Install {pinned} yourself, or set " + f"{DOWNLOAD_URL_ENV} + {DOWNLOAD_SHA256_ENV} to a trusted build." + ) + installed = install_runtime(pinned_artifact(), version=pinned, cache_dir=cache_dir) + logger.info("resolved hai-agent-runtime from fresh download v%s: %s", pinned, installed) + return [str(installed)] + + def health(self) -> typing.Dict[str, typing.Any]: + """The /health JSON body; raises RuntimeUnhealthyError when the runtime is not answering.""" + payload = probe_health(self.base_url) + if payload is None: + raise RuntimeUnhealthyError(f"hai-agent-runtime at {self.base_url} is not answering /health") + return payload + + def shutdown(self) -> None: + """Gracefully stop the runtime this LocalRuntime spawned (SIGTERM group, grace, SIGKILL group).""" + if not self.owned or self._proc is None: + raise LocalRuntimeError("shutdown() only stops runtimes this LocalRuntime spawned") + terminate(self._proc) + self._cleanup_state_files() + + # Statuses that mean the runtime still holds live session state a shutdown would destroy. + # ("idle" sessions await user input but keep runtime state.) + ACTIVE_SESSION_STATUSES: typing.ClassVar[typing.Tuple[str, ...]] = ( + "pending", + "running", + "paused", + "idle", + "awaiting_tool_results", + ) + + def shutdown_if_idle(self) -> bool: + """Stop the owned runtime only when it hosts no active sessions; True when it was stopped.""" + from ..client import Client # runtime-time import: client.py imports this module lazily too + + client = Client(base_url=self.base_url, api_key=self.api_key) + page = client.sessions.list_sessions(status=list(self.ACTIVE_SESSION_STATUSES), size=1) + if page.items: + return False + self.shutdown() + return True + + def force_kill(self) -> None: + """Stop only the process this manager spawned; never trust a saved PID to claim ownership.""" + if not self.owned or self._proc is None: + raise LocalRuntimeError("cannot force-kill a borrowed runtime from a persisted PID") + terminate(self._proc) + self._cleanup_state_files() + + def _cleanup_state_files(self) -> None: + if self._token_file is not None: + unlink_if_content(self._token_file, self.api_key) + self._token_file = None + if self._pid_file is not None: + unlink_if_content(self._pid_file, str(self.pid)) + self._pid_file = None + + +@contextlib.contextmanager +def _startup_lock(cache_dir: pathlib.Path, port: int, timeout_s: float): + path = cache_dir / "state" / f"startup-{port}.lock" + path.parent.mkdir(parents=True, exist_ok=True, mode=0o700) + fd = os.open(path, os.O_RDWR | os.O_CREAT | getattr(os, "O_NOFOLLOW", 0), 0o600) + with os.fdopen(fd, "r+b") as handle: + deadline = time.monotonic() + timeout_s + while True: + try: + if os.name == "posix": + import fcntl + + fcntl.flock(handle, fcntl.LOCK_EX | fcntl.LOCK_NB) + else: + import msvcrt + + handle.seek(0) + msvcrt.locking(handle.fileno(), msvcrt.LK_NBLCK, 1) + break + except (BlockingIOError, OSError): + if time.monotonic() >= deadline: + raise LocalRuntimeError("timed out waiting for local runtime startup lock") + time.sleep(0.05) + try: + yield + finally: + if os.name == "posix": + fcntl.flock(handle, fcntl.LOCK_UN) + else: + handle.seek(0) + msvcrt.locking(handle.fileno(), msvcrt.LK_UNLCK, 1) diff --git a/src/hai_agents/local/state.py b/src/hai_agents/local/state.py new file mode 100644 index 0000000..1af0150 --- /dev/null +++ b/src/hai_agents/local/state.py @@ -0,0 +1,90 @@ +"""Owner-only on-disk discovery state for locally spawned hai-agent-runtime processes. + +A spawner persists the generated bearer token and the runtime pid under the SDK +cache dir so a second process can attach (token) or force-kill (pid) without any +IPC. Both files drive privileged actions, so they are 0600 from the first byte +and refuse pre-planted symlinks (port of holo_desktop launcher._write_owner_only). +""" + +from __future__ import annotations + +import contextlib +import os +import pathlib +import typing + +CACHE_DIR_ENV = "HAI_AGENT_LOCAL_CACHE_DIR" +DEFAULT_CACHE_DIR = pathlib.Path.home() / ".hai" / "agent-runtime" +# HoloDesktop's AGENT_API_DEFAULT_PORT: the shared well-known local runtime port. +DEFAULT_PORT = 18795 + +_PathInput = typing.Union[str, "os.PathLike[str]"] + + +def resolve_cache_dir(cache_dir: typing.Optional[_PathInput] = None) -> pathlib.Path: + """Explicit argument > HAI_AGENT_LOCAL_CACHE_DIR > ~/.hai/agent-runtime.""" + if cache_dir is not None: + return pathlib.Path(cache_dir).expanduser() + override = os.environ.get(CACHE_DIR_ENV, "").strip() + if override: + return pathlib.Path(override).expanduser() + return DEFAULT_CACHE_DIR + + +def state_dir(cache_dir: typing.Optional[_PathInput] = None) -> pathlib.Path: + return resolve_cache_dir(cache_dir) / "state" + + +def token_file_path(port: int, *, cache_dir: typing.Optional[_PathInput] = None) -> pathlib.Path: + """Where a spawner publishes its generated bearer token for other local clients.""" + return state_dir(cache_dir) / f"agent-token-{port}" + + +def pid_file_path(port: int, *, cache_dir: typing.Optional[_PathInput] = None) -> pathlib.Path: + """Where a spawner publishes the runtime pid so force_kill() works from another process.""" + return state_dir(cache_dir) / f"agent-pid-{port}" + + +def runtime_log_path(port: int, *, cache_dir: typing.Optional[_PathInput] = None) -> pathlib.Path: + """Where the runtime spawned on `port` writes its stderr.""" + return resolve_cache_dir(cache_dir) / "logs" / f"hai-agent-runtime-{port}.log" + + +def write_owner_only(path: pathlib.Path, content: str) -> pathlib.Path: + """Write `content` to `path` owner-only (0600), refusing a pre-existing symlink at the path.""" + path.parent.mkdir(parents=True, exist_ok=True) + with contextlib.suppress(OSError): + path.parent.chmod(0o700) # owner-only state dir; no-op on Windows + flags = os.O_WRONLY | os.O_CREAT | os.O_TRUNC | getattr(os, "O_NOFOLLOW", 0) + fd = os.open(path, flags, 0o600) + with os.fdopen(fd, "w", encoding="utf-8") as fh: + if os.name == "posix": + os.fchmod(fd, 0o600) # enforce owner-only even if the file pre-existed + fh.write(content) + return path + + +def read_state_file(path: pathlib.Path) -> typing.Optional[str]: + """The stripped file contents, or None when missing/unreadable/empty.""" + try: + return path.read_text(encoding="utf-8").strip() or None + except OSError: + return None + + +def read_pid(port: int, *, cache_dir: typing.Optional[_PathInput] = None) -> typing.Optional[int]: + """The persisted runtime pid for `port`, or None when absent or malformed.""" + raw = read_state_file(pid_file_path(port, cache_dir=cache_dir)) + if raw is None: + return None + try: + return int(raw) + except ValueError: + return None + + +def unlink_if_content(path: pathlib.Path, content: str) -> None: + """Remove our state file, but never one a concurrent spawner already replaced.""" + with contextlib.suppress(OSError): + if path.read_text(encoding="utf-8").strip() == content: + path.unlink(missing_ok=True) diff --git a/src/hai_agents_local/bridge.py b/src/hai_agents_local/bridge.py index 3f4ea4b..a764106 100644 --- a/src/hai_agents_local/bridge.py +++ b/src/hai_agents_local/bridge.py @@ -110,6 +110,9 @@ def request_stop(self) -> None: """Signal the poll loop to stop; safe to call from a signal handler.""" self._stop_event.set() + async def interrupt_driver(self) -> None: + """Stop run-owned work; drivers without owned processes need no special action.""" + async def run(self) -> None: """Serve commands until stopped; raises AuthError on a bad key.""" # An asyncio.Event binds to the loop it is first awaited on; a restarted bridge runs on a new loop. @@ -256,7 +259,19 @@ async def _process_commands(self, exchange: CommandExchange, commands: list[Comm self._results.move_to_end(cmd.command_uid) result, error = self._results[cmd.command_uid] else: - result, error = await asyncio.to_thread(self._dispatch, cmd.name, cmd.args) + dispatch = asyncio.create_task(asyncio.to_thread(self._dispatch, cmd.name, cmd.args)) + stopped = asyncio.create_task(self._stop_event.wait()) + try: + await asyncio.wait((dispatch, stopped), return_when=asyncio.FIRST_COMPLETED) + if self._stop_event.is_set(): + await self.interrupt_driver() + await dispatch + return + result, error = await dispatch + finally: + stopped.cancel() + with contextlib.suppress(asyncio.CancelledError): + await stopped self._results[cmd.command_uid] = (result, error) while len(self._results) > RESULT_CACHE_SIZE: self._results.popitem(last=False) diff --git a/src/hai_agents_local/manager.py b/src/hai_agents_local/manager.py index b10a388..7a2570c 100644 --- a/src/hai_agents_local/manager.py +++ b/src/hai_agents_local/manager.py @@ -97,16 +97,24 @@ def _displace_kind_locked(self, bridge: LocalBridge) -> list[_Runner]: def stop(self, session_ids: Sequence[str]) -> None: with self._lock: - stopping = [self._runners.pop(sid) for sid in session_ids if sid in self._runners] + stopping = [self._runners[sid] for sid in session_ids if sid in self._runners] + failures = [] for runner in stopping: - runner.stop() + try: + runner.stop() + except TimeoutError as error: + failures.append(str(error)) + else: + with self._lock: + if self._runners.get(runner.bridge.session_id) is runner: + del self._runners[runner.bridge.session_id] + if failures: + raise TimeoutError("; ".join(failures)) def stop_all(self) -> None: with self._lock: - stopping = list(self._runners.values()) - self._runners.clear() - for runner in stopping: - runner.stop() + session_ids = list(self._runners) + self.stop(session_ids) class _Runner: @@ -152,6 +160,8 @@ def stop(self) -> None: # A bridge's loss handler runs on its own runner thread; a thread cannot join itself. if threading.current_thread() is not self.thread: self.thread.join(timeout=STOP_JOIN_TIMEOUT_S) + if self.thread.is_alive(): + raise TimeoutError(f"Could not confirm stop of local {self.bridge.environment_kind} bridge") # Process-wide manager behind ensure_bridges/stop_bridges; cleaned up at interpreter exit. diff --git a/src/hai_agents_local/sessions.py b/src/hai_agents_local/sessions.py index d8dccf5..a02a76f 100644 --- a/src/hai_agents_local/sessions.py +++ b/src/hai_agents_local/sessions.py @@ -193,6 +193,30 @@ def _cancel_sessions_at_exit() -> None: class LocalSessionsClient(SessionsClient): + def close(self) -> None: + failures = [] + for session_id in list(getattr(self, "_owned_bridges", {})): + try: + self.cancel_session(session_id) + except Exception as error: + failures.append(error) + if failures: + raise RuntimeError("Could not confirm all client-owned sessions stopped") from failures[0] + + def cancel_session(self, session_id: str, **kwargs: typing.Any) -> typing.Any: + # Stop local execution even if the remote cancellation cannot be delivered. + owned = getattr(self, "_owned_bridges", {}).get(str(session_id), []) + try: + try: + if owned: + stop_bridges(owned) + self._owned_bridges.pop(str(session_id), None) + finally: + response = super().cancel_session(session_id, **kwargs) + finally: + _deregister_exit_cancel(str(session_id)) + return response + @functools.wraps(SessionsClient.create_session) def create_session(self, **kwargs: typing.Any) -> typing.Any: wrapper = self._raw_client._client_wrapper @@ -217,10 +241,37 @@ def create_session(self, **kwargs: typing.Any) -> typing.Any: if stop_watcher is not None and not stop_watcher.active: # A stop was filed while bridges or the session were starting; apply it now. _panic_stop() + if started: + if not hasattr(self, "_owned_bridges"): + self._owned_bridges = {} + self._owned_bridges[str(session.id)] = started return session class LocalAsyncSessionsClient(AsyncSessionsClient): + async def aclose(self) -> None: + failures = [] + for session_id in list(getattr(self, "_owned_bridges", {})): + try: + await self.cancel_session(session_id) + except Exception as error: + failures.append(error) + if failures: + raise RuntimeError("Could not confirm all client-owned sessions stopped") from failures[0] + + async def cancel_session(self, session_id: str, **kwargs: typing.Any) -> typing.Any: + owned = getattr(self, "_owned_bridges", {}).get(str(session_id), []) + try: + try: + if owned: + await asyncio.to_thread(stop_bridges, owned) + self._owned_bridges.pop(str(session_id), None) + finally: + response = await super().cancel_session(session_id, **kwargs) + finally: + _deregister_exit_cancel(str(session_id)) + return response + @functools.wraps(AsyncSessionsClient.create_session) async def create_session(self, **kwargs: typing.Any) -> typing.Any: wrapper = self._raw_client._client_wrapper @@ -246,4 +297,8 @@ async def create_session(self, **kwargs: typing.Any) -> typing.Any: if stop_watcher is not None and not stop_watcher.active: # A stop was filed while bridges or the session were starting; apply it now. await asyncio.to_thread(_panic_stop) + if started: + if not hasattr(self, "_owned_bridges"): + self._owned_bridges = {} + self._owned_bridges[str(session.id)] = started return session diff --git a/src/hai_agents_local/workstation.py b/src/hai_agents_local/workstation.py index 9fe332e..f2724be 100644 --- a/src/hai_agents_local/workstation.py +++ b/src/hai_agents_local/workstation.py @@ -2,6 +2,7 @@ from __future__ import annotations +import asyncio import logging import os import shutil @@ -62,6 +63,11 @@ def create_driver(self) -> ManagedCodeSandboxInterface: }, ) + async def interrupt_driver(self) -> None: + # This bridge creates one sandbox per run; close owns only that run's children. + if self._driver is not None: + await asyncio.to_thread(self._driver.close) + def driver_interface(self) -> type: from hai_drivers.code_sandbox.interface import ManagedCodeSandboxInterface diff --git a/tests/test_local.py b/tests/test_local.py index d485639..c6ef7d5 100644 --- a/tests/test_local.py +++ b/tests/test_local.py @@ -869,3 +869,84 @@ def test_missing_login_fails_with_fix_and_skips_the_platform_check(self, monkeyp assert not by_name["login"].ok and by_name["login"].fix is not None assert "platform" not in by_name assert {"browser", "desktop"} <= set(by_name) + + +async def test_stopping_workstation_interrupts_running_command_and_skips_queue(tmp_path): + import asyncio + + from hai_drivers.code_sandbox.local.driver import LocalCodeSandbox + + bridge = WorkstationBridge(api_key="test", workspace=str(tmp_path)) + bridge._driver = LocalCodeSandbox(str(tmp_path)) + + class Exchange: + async def post_result(self, *args, **kwargs): + pytest.fail("A stopped execution must not report a successful tool result") + + commands = [ + Command(id="one", command_uid="one", name="execute", args={"command": "echo $$ > pid; sleep 60"}), + Command(id="two", command_uid="two", name="execute", args={"command": "touch queued"}), + ] + dispatch = asyncio.create_task(bridge._process_commands(Exchange(), commands)) + try: + async with asyncio.timeout(5): + while not (tmp_path / "pid").exists(): + await asyncio.sleep(0.01) + bridge.request_stop() + await asyncio.wait_for(dispatch, 5) + assert not (tmp_path / "queued").exists() + finally: + await asyncio.to_thread(bridge._driver.close) + await dispatch + + +@pytest.mark.asyncio +@pytest.mark.parametrize("asynchronous", [False, True]) +async def test_cancel_reaches_api_when_local_stop_cannot_be_confirmed(monkeypatch, asynchronous): + from hai_agents import AsyncClient + from hai_agents.sessions.client import AsyncSessionsClient + + cancelled = [] + + def refuse_stop(ids): + raise TimeoutError("local command still running") + + def cancel(self, session_id, **kwargs): + cancelled.append(session_id) + + async def async_cancel(self, session_id, **kwargs): + cancel(self, session_id, **kwargs) + + monkeypatch.setattr("hai_agents_local.sessions.stop_bridges", refuse_stop) + monkeypatch.setattr(SessionsClient, "cancel_session", cancel) + monkeypatch.setattr(AsyncSessionsClient, "cancel_session", async_cancel) + client = (AsyncClient if asynchronous else Client)(api_key=API_KEY) + sessions = client.sessions + sessions._owned_bridges = {"run": ["device"]} + with pytest.raises(TimeoutError, match="still running"): + if asynchronous: + await sessions.cancel_session("run") + else: + sessions.cancel_session("run") + assert cancelled == ["run"] + assert sessions._owned_bridges == {"run": ["device"]}, "An unconfirmed stop must remain retryable" + + +def test_manager_reports_unconfirmed_stop_and_allows_retry(monkeypatch): + import hai_agents_local.manager as manager_module + + class StubbornBridge(ServingBridge): + def request_stop(self): + pass + + monkeypatch.setattr(manager_module, "STOP_JOIN_TIMEOUT_S", 0.01) + manager = BridgeManager() + bridge = StubbornBridge(api_key=API_KEY) + manager.ensure([bridge]) + try: + with pytest.raises(TimeoutError, match="confirm stop"): + manager.stop([bridge.session_id]) + finally: + bridge.request_stop = lambda: ServingBridge.request_stop(bridge) + manager.stop([bridge.session_id]) + assert not manager._runners diff --git a/tests/test_runtime_placement.py b/tests/test_runtime_placement.py new file mode 100644 index 0000000..8f9e46b --- /dev/null +++ b/tests/test_runtime_placement.py @@ -0,0 +1,91 @@ +"""Placement boundaries: authentication, compatibility and executor ownership.""" + +import httpx +import pytest + +from hai_agents import AsyncClient, Client, Inference +from hai_agents.local.errors import BinaryIncompatibleError, LocalRuntimeError +from hai_agents.local.runtime import LocalRuntime +from hai_agents.local.state import token_file_path, write_owner_only +from hai_agents.sessions.client import AsyncSessionsClient, SessionsClient + + +def test_both_clients_respect_product_owned_execution(): + assert type(Client(api_key="test", auto_bridges=False).sessions) is SessionsClient + assert type(AsyncClient(api_key="test", auto_bridges=False).sessions) is AsyncSessionsClient + + +def test_self_hosted_inference_does_not_receive_hosted_key(monkeypatch): + monkeypatch.setenv("HAI_API_KEY", "hosted-secret") + env = Inference.self_hosted("http://127.0.0.1:8000/v1", model="my-model").runtime_env() + assert "HAI_API_KEY" not in env + assert env["HAI_AGENT_RUNTIME_MODEL"] == "my-model" + monkeypatch.setenv("HAI_AGENT_RUNTIME_BASE_URL", "http://localhost:8000/v1") + assert "HAI_AGENT_RUNTIME_BASE_URL" not in Inference.cloud().runtime_env() + + +def test_attachment_authenticates_and_checks_recipe_before_use(tmp_path, monkeypatch): + from hai_agents.local import runtime as module + + write_owner_only(token_file_path(18795, cache_dir=tmp_path), "local-token") + monkeypatch.delenv("HAI_AGENT_RUNTIME_API_TOKEN", raising=False) + monkeypatch.setattr(module, "probe_health", lambda url: {"version": "old", "recipe": "desktop"}) + calls = [] + + def get(url, **kwargs): + calls.append(kwargs["headers"]) + return httpx.Response(200) + + monkeypatch.setattr(module.httpx, "get", get) + attached = LocalRuntime.attach(cache_dir=tmp_path) + assert calls == [{"Authorization": "Bearer local-token"}] + with pytest.raises(BinaryIncompatibleError): + attached.require_recipe("shared") + with pytest.raises(LocalRuntimeError): + attached.force_kill() + assert token_file_path(18795, cache_dir=tmp_path).exists() + monkeypatch.setattr(module.httpx, "get", lambda *a, **k: httpx.Response(401)) + with pytest.raises(LocalRuntimeError, match="authenticated"): + LocalRuntime.attach(cache_dir=tmp_path) + + +@pytest.mark.parametrize("url", ["https://remote.example", "http://user:pass@localhost:80"]) +def test_local_attach_rejects_remote_or_credential_urls(tmp_path, monkeypatch, url): + monkeypatch.setenv("HAI_AGENT_LOCAL_BASE_URL", url) + with pytest.raises(LocalRuntimeError, match="loopback"): + LocalRuntime.attach(cache_dir=tmp_path) + + +@pytest.mark.asyncio +@pytest.mark.parametrize("asynchronous", [False, True]) +@pytest.mark.parametrize("borrowed", [False, True]) +async def test_client_close_releases_only_owned_runtime_and_http(monkeypatch, asynchronous, borrowed): + class Runtime: + base_url = "http://127.0.0.1:18795" + api_key = "local-token" + owned = True + stopped = False + + def require_recipe(self, recipe): + assert recipe == "shared" + + def shutdown(self): + self.stopped = True + + runtime = Runtime() + monkeypatch.setattr(LocalRuntime, "ensure_started", lambda **options: runtime) + http = httpx.AsyncClient() if asynchronous else httpx.Client() + client_type = AsyncClient if asynchronous else Client + options = {"runtime": runtime, "httpx_client": http} if borrowed else {} + client = client_type(mode="local", **options) + if asynchronous: + await client.aclose() + else: + client.close() + assert runtime.stopped is (not borrowed) + if borrowed: + assert not http.is_closed + if asynchronous: + await http.aclose() + else: + http.close() diff --git a/tests/test_runtime_state.py b/tests/test_runtime_state.py new file mode 100644 index 0000000..9adc2cd --- /dev/null +++ b/tests/test_runtime_state.py @@ -0,0 +1,43 @@ +"""Behavioural tests for the runtime pid file's on-disk hardening. + +The pid file's contents drive a privileged operation: ``holo stop --force`` reads it and +``os.killpg(..., SIGKILL)`` the pid it finds. It therefore earns the same protections as the bearer +token — owner-only permissions and a refusal to follow a symlink planted at its path — so it can never +be steered into killing an arbitrary process group. +""" + +from __future__ import annotations + +import os +import stat +import sys +from pathlib import Path + +import pytest + +from hai_agents.local.state import pid_file_path, write_owner_only + + +@pytest.mark.skipif(sys.platform == "win32", reason="POSIX file-mode semantics") +def test_pid_file_written_owner_only(tmp_path: Path, monkeypatch: pytest.MonkeyPatch) -> None: + path = write_owner_only(pid_file_path(4242, cache_dir=tmp_path), "31337") + + assert path is not None + assert path.read_text(encoding="utf-8") == "31337" + assert stat.S_IMODE(path.stat().st_mode) == 0o600 + + +@pytest.mark.skipif(sys.platform == "win32", reason="O_NOFOLLOW is POSIX-only") +def test_pid_file_write_refuses_symlink_at_path(tmp_path: Path, monkeypatch: pytest.MonkeyPatch) -> None: + # An attacker pre-plants a symlink at the pid path pointing at a file they want clobbered (and, + # later, read back by `holo stop --force`). The hardened write must refuse to follow it. + victim = tmp_path / "victim" + victim.write_text("untouched", encoding="utf-8") + pid_path = pid_file_path(4242, cache_dir=tmp_path) + pid_path.parent.mkdir(parents=True, exist_ok=True) + os.symlink(victim, pid_path) + + with pytest.raises(OSError): + write_owner_only(pid_path, "31337") + assert victim.read_text(encoding="utf-8") == "untouched", "symlink target must not be clobbered" + assert pid_path.is_symlink(), "the planted symlink itself is left as-is, never written through" From f94841b5349e5b335437a026fb489e25917016af Mon Sep 17 00:00:00 2001 From: Charlie Masters Date: Tue, 29 Sep 2026 22:53:45 +0100 Subject: [PATCH 02/35] docs: explain candidate local agent setup and lifecycle limits --- README.md | 51 +++++++++++++++++++++++++++++++++++++++++++++++++++ 1 file changed, 51 insertions(+) diff --git a/README.md b/README.md index c3e7a7b..2bfe302 100644 --- a/README.md +++ b/README.md @@ -67,6 +67,57 @@ print(result.answer) `result` is a `SessionRunResult`: `id`, `status`, `answer`, the accumulated `events`, and `final_changes`. +## Candidate local-agent support + +The placement branch adds `Client(mode="local")` and the same option on +`AsyncClient`. Agent placement and environment placement are separate: the client +selects where the agent runs; each agent environment selects `host="user_device"` +or `host="cloud"`. `Client()` continues to use the hosted Agents API. + +For development, install the candidate SDK and use a prepared HAI source checkout: + +```python +from hai_agents import Client + +with Client( + mode="local", + local_options={ + "command": ["/path/to/hai/.venv/bin/python", "-m", "hai_agent_runtime"], + "download": False, + }, +) as client: + session = client.sessions.create_session( + agent={ + "name": "local-example", + "description": "Local workstation example", + "instructions": "Answer the user's task using the workstation tools.", + "environments": [ + {"id": "workstation", "kind": "workstation", "host": "user_device"} + ], + }, + messages=[{"type": "user_message", "message": "Print hello using the shell."}], + max_steps=8, + max_time_s=120, + ) + # Poll or steer the session here, before leaving the client context. +``` + +The source runtime must support the `shared` recipe. The current pinned binary +0.1.8 does not; candidate source or an explicitly supplied compatible +`local_options["binary_path"]` is required until a compatible release is pinned. +The runtime inherits `HAI_API_KEY` for hosted inference. Local agent execution +does not by itself imply local inference: `Inference.self_hosted(url, model=...)` +selects a model endpoint for a newly started local runtime. Hosted agents do not +currently accept that override. + +Closing the client shuts down a runtime it started; an attached runtime remains +owned by its caller. `cancel()` ends the agent session. For a cloud workstation, +an explicit `session_id` attaches to a caller-owned environment, which the caller +must eventually release. Automatically provisioned cloud environments are torn +down with their agent adapter. Keeping those environments across cancellation is +still an open lifecycle requirement. Local resource-upload endpoints are not yet +implemented; driver-level file transfer is a separate capability. + ## How a session works A session is one run of an agent against a task. It moves through a small set of states: `pending`, `running`, and then a settled state such as `completed`, `idle`, `failed`, `timed_out`, or `interrupted`. From 48b35eeaa12a071c98468d6980073faeabadfe4d Mon Sep 17 00:00:00 2001 From: Charlie Masters Date: Tue, 29 Sep 2026 23:05:17 +0100 Subject: [PATCH 03/35] docs: describe local shared file support and lifetime --- README.md | 7 +++++-- 1 file changed, 5 insertions(+), 2 deletions(-) diff --git a/README.md b/README.md index 2bfe302..fd260d7 100644 --- a/README.md +++ b/README.md @@ -115,8 +115,11 @@ owned by its caller. `cancel()` ends the agent session. For a cloud workstation, an explicit `session_id` attaches to a caller-owned environment, which the caller must eventually release. Automatically provisioned cloud environments are torn down with their agent adapter. Keeping those environments across cancellation is -still an open lifecycle requirement. Local resource-upload endpoints are not yet -implemented; driver-level file transfer is a separate capability. +still an open lifecycle requirement. The candidate source runtime accepts base64 +message attachments and exposes files shared by the agent through +`sessions.get_session_resource(id, "local", key)`. Download them before closing +the runtime: local shared resources expire with the retained session. The limits +are 50 MiB per file, 64 MiB and 128 shared files per session. ## How a session works From e2a3739b008072d1fc516f3d704393fcfc4906f7 Mon Sep 17 00:00:00 2001 From: Charlie Masters Date: Tue, 29 Sep 2026 23:29:46 +0100 Subject: [PATCH 04/35] docs: describe candidate cloud workstation retention --- README.md | 12 +++++++++--- 1 file changed, 9 insertions(+), 3 deletions(-) diff --git a/README.md b/README.md index fd260d7..a62948a 100644 --- a/README.md +++ b/README.md @@ -113,9 +113,15 @@ currently accept that override. Closing the client shuts down a runtime it started; an attached runtime remains owned by its caller. `cancel()` ends the agent session. For a cloud workstation, an explicit `session_id` attaches to a caller-owned environment, which the caller -must eventually release. Automatically provisioned cloud environments are torn -down with their agent adapter. Keeping those environments across cancellation is -still an open lifecycle requirement. The candidate source runtime accepts base64 +must eventually release. In the candidate shared recipe, automatically provisioned +cloud workstations survive agent cancellation and expire through the environment +manager after 30 minutes without commands (the runner's fixed deadline still +applies). Reattach using the `RunnerSessionEvent` ID in a new agent session; +cancellation does not revive the old agent session. The existing environment +manager API can release the workstation earlier. A manager that cannot confirm +the requested expiry is rejected and the newly created runner is deleted. This +is temporary compute retention, not a durable-storage or pause/resume guarantee. +The candidate source runtime accepts base64 message attachments and exposes files shared by the agent through `sessions.get_session_resource(id, "local", key)`. Download them before closing the runtime: local shared resources expire with the retained session. The limits From ac92c76e95b105ab887b789e8050bc2223c37359 Mon Sep 17 00:00:00 2001 From: Charlie Masters Date: Tue, 29 Sep 2026 23:53:42 +0100 Subject: [PATCH 05/35] Interrupt desktop commands when the SDK bridge stops --- src/hai_agents_local/desktop.py | 5 +++ tests/test_desktop_stop.py | 76 +++++++++++++++++++++++++++++++++ 2 files changed, 81 insertions(+) create mode 100644 tests/test_desktop_stop.py diff --git a/src/hai_agents_local/desktop.py b/src/hai_agents_local/desktop.py index e0b1ac0..af65168 100644 --- a/src/hai_agents_local/desktop.py +++ b/src/hai_agents_local/desktop.py @@ -1,5 +1,6 @@ from __future__ import annotations +import asyncio import sys from typing import TYPE_CHECKING, Literal @@ -98,6 +99,10 @@ def create_driver(self) -> DesktopDriverInterface: quality=self.quality, ) + async def interrupt_driver(self) -> None: + if self._driver is not None: + await asyncio.to_thread(self._driver.close) + def driver_interface(self) -> type: # Runtime import: hai-drivers is absent unless installed with hai-agents[desktop]. from hai_drivers.desktop.interface import DesktopDriverInterface diff --git a/tests/test_desktop_stop.py b/tests/test_desktop_stop.py new file mode 100644 index 0000000..679557f --- /dev/null +++ b/tests/test_desktop_stop.py @@ -0,0 +1,76 @@ +"""Stop crosses the SDK bridge and scaled driver and kills only the active command tree.""" + +import asyncio +import json +import os +import subprocess +import sys +from types import SimpleNamespace + +import pytest +from hai_drivers.desktop.scaled import ScaledDesktopDriver +from hai_drivers.desktop.utils import DesktopCommandRunner + +from hai_agents_local.desktop import PyautoguiDesktopBridge +from hai_agents_local.transport import Command + + +@pytest.mark.asyncio +@pytest.mark.skipif(os.name != "posix", reason="Process-group liveness requires POSIX; Windows needs native QA") +async def test_desktop_stop_kills_owned_command_tree_and_rejects_queued_work(tmp_path): + runner = DesktopCommandRunner() + + def run_command(command, timeout=60, env=None, cwd=None, detach=False, ignore_errors=False): + return runner.run(command, timeout=timeout, env=env, cwd=cwd, detach=detach) + + # Only the UI surface is absent; command execution, scale forwarding and SDK Stop are real. + desktop = SimpleNamespace(run_command=run_command, close=runner.close) + bridge = PyautoguiDesktopBridge(api_key="test") + bridge._driver = ScaledDesktopDriver(desktop, max_width=1920) + pid_file = tmp_path / "pids.json" + queued = tmp_path / "queued" + program = ( + "import subprocess,sys,os,json,time;from pathlib import Path;" + "child=subprocess.Popen([sys.executable,'-c','import time;time.sleep(60)']);" + f"Path({str(pid_file)!r}).write_text(json.dumps([os.getpid(),child.pid]));time.sleep(60)" + ) + commands = [ + Command(id="one", command_uid="one", name="run_command", args={"command": [sys.executable, "-c", program]}), + Command( + id="two", + command_uid="two", + name="run_command", + args={"command": [sys.executable, "-c", f"from pathlib import Path;Path({str(queued)!r}).touch()"]}, + ), + ] + + class Exchange: + async def post_result(self, *args, **kwargs): + pytest.fail("Stopped work must not report a successful result") + + observer = subprocess.Popen([sys.executable, "-c", "import time;time.sleep(60)"]) + dispatch = asyncio.create_task(bridge._process_commands(Exchange(), commands)) + try: + async with asyncio.timeout(5): + while not pid_file.exists(): + await asyncio.sleep(0.01) + owned = json.loads(pid_file.read_text()) + bridge.request_stop() + await asyncio.wait_for(dispatch, 5) + assert observer.poll() is None + assert not queued.exists() + for pid in owned: + async with asyncio.timeout(5): + while True: + try: + os.kill(pid, 0) + except ProcessLookupError: + break + await asyncio.sleep(0.02) + with pytest.raises(RuntimeError, match="closed"): + runner.run([sys.executable, "-c", f"from pathlib import Path;Path({str(queued)!r}).touch()"]) + finally: + await asyncio.to_thread(runner.close) + observer.terminate() + observer.wait(timeout=5) + await dispatch From 0b58d7e0241dc5abbeac6cdde5ba303673a6b204 Mon Sep 17 00:00:00 2001 From: Charlie Masters Date: Wed, 30 Sep 2026 00:27:18 +0100 Subject: [PATCH 06/35] Keep async desktop permission prompts on the caller thread --- src/hai_agents_local/desktop.py | 2 ++ src/hai_agents_local/sessions.py | 3 +++ tests/test_local.py | 41 +++++++++++++++++++++++++++++++- 3 files changed, 45 insertions(+), 1 deletion(-) diff --git a/src/hai_agents_local/desktop.py b/src/hai_agents_local/desktop.py index af65168..0339b98 100644 --- a/src/hai_agents_local/desktop.py +++ b/src/hai_agents_local/desktop.py @@ -2,6 +2,7 @@ import asyncio import sys +import threading from typing import TYPE_CHECKING, Literal from .bridge import LocalBridge, TokenSource @@ -29,6 +30,7 @@ def ensure_macos_input_permissions(prompt: bool = True) -> None: from ApplicationServices import AXIsProcessTrustedWithOptions, kAXTrustedCheckOptionPrompt from Quartz import CGPreflightScreenCaptureAccess, CGRequestScreenCaptureAccess + prompt = prompt and threading.current_thread() is threading.main_thread() missing = [] if not AXIsProcessTrustedWithOptions({kAXTrustedCheckOptionPrompt: prompt}): missing.append("Accessibility (moves the mouse and types)") diff --git a/src/hai_agents_local/sessions.py b/src/hai_agents_local/sessions.py index a02a76f..ae5a6ef 100644 --- a/src/hai_agents_local/sessions.py +++ b/src/hai_agents_local/sessions.py @@ -278,6 +278,9 @@ async def create_session(self, **kwargs: typing.Any) -> typing.Any: bridges = _localize(wrapper, kwargs) if bridges: _apply_runaway_budgets(kwargs) + # Native permission prompts must run before bridge startup moves to a worker. + for bridge in bridges: + bridge.preflight() stop_watcher = _ensure_stop_watcher() if bridges else None watcher = _LossWatcher(bridges) started = await asyncio.to_thread(ensure_bridges, bridges) diff --git a/tests/test_local.py b/tests/test_local.py index c6ef7d5..47b9cc3 100644 --- a/tests/test_local.py +++ b/tests/test_local.py @@ -475,12 +475,13 @@ def _fake_frameworks(self, monkeypatch, *, ax: bool, screen: bool): import sys as _sys import types - calls = {"ax_prompts": [], "screen_requests": 0} + calls = {"ax_prompts": [], "ax_main_thread": [], "screen_requests": 0} apps = types.ModuleType("ApplicationServices") apps.kAXTrustedCheckOptionPrompt = "AXTrustedCheckOptionPrompt" def ax_check(options): calls["ax_prompts"].append(options["AXTrustedCheckOptionPrompt"]) + calls["ax_main_thread"].append(threading.current_thread() is threading.main_thread()) return ax def screen_request(): @@ -521,6 +522,44 @@ def test_prompt_false_never_triggers_dialogs(self, monkeypatch): assert calls["ax_prompts"] == [False] assert calls["screen_requests"] == 0 + @pytest.mark.asyncio + @pytest.mark.parametrize("granted", [False, True]) + async def test_async_session_prompts_before_worker_startup(self, monkeypatch, granted): + from hai_agents import AsyncClient + from hai_agents.sessions.client import AsyncSessionsClient + + monkeypatch.setenv(AUTO_BRIDGE_ENV_VAR, "1") + monkeypatch.setattr(sys, "platform", "darwin") + calls = self._fake_frameworks(monkeypatch, ax=granted, screen=granted) + requested = [] + + def start(bridges): + assert threading.current_thread() is not threading.main_thread() + for bridge in bridges: + bridge.preflight() # The manager re-checks before starting its driver thread. + return [] + + async def create(self, **kwargs): + requested.append(kwargs) + return types.SimpleNamespace(id=None) + + monkeypatch.setattr("hai_agents_local.sessions.ensure_bridges", start) + monkeypatch.setattr("hai_agents_local.sessions._ensure_stop_watcher", lambda: None) + monkeypatch.setattr(AsyncSessionsClient, "create_session", create) + async with AsyncClient(api_key=API_KEY) as client: + request = client.sessions.create_session( + agent={"name": "qa", "environments": [{"id": "desktop", "kind": "desktop", "host": "user_device"}]}, + messages="test", + ) + if granted: + await request + else: + with pytest.raises(PermissionError): + await request + assert calls["ax_prompts"] == ([True, False] if granted else [True]) + assert calls["ax_main_thread"] == ([True, False] if granted else [True]) + assert bool(requested) is granted + class TestBridgeProtocol: def test_bytes_results_are_base64(self): From 0f6250d27fe08ad6feb3cfe37e8ac995b8c187d9 Mon Sep 17 00:00:00 2001 From: Charlie Masters Date: Tue, 29 Sep 2026 23:42:25 +0100 Subject: [PATCH 07/35] build: prepare SDK-owned runtime pin updates --- README.md | 9 +++ scripts/bump_runtime.py | 102 +++++++++++++++++++++++++++++++ src/hai_agents/local/manifest.py | 4 +- tests/test_bump_runtime.py | 38 ++++++++++++ 4 files changed, 151 insertions(+), 2 deletions(-) create mode 100644 scripts/bump_runtime.py create mode 100644 tests/test_bump_runtime.py diff --git a/README.md b/README.md index a62948a..271ef30 100644 --- a/README.md +++ b/README.md @@ -127,6 +127,15 @@ message attachments and exposes files shared by the agent through the runtime: local shared resources expire with the retained session. The limits are 50 MiB per file, 64 MiB and 128 shared files per session. +## Runtime release maintenance + +The SDK runtime manifest is maintained independently of schema generation. After +publishing and verifying a compatible runtime, run `scripts/bump_runtime.py` with +`--version` and one `--sha PLATFORM=SHA256` for every platform already in the +manifest. Partial updates are rejected so a new URL cannot retain an old digest. +The generator preserves this SDK-owned file. The release workflow still opens +legacy CLI pin PRs; retarget it only after the SDK/CLI migration has shipped. + ## How a session works A session is one run of an agent against a task. It moves through a small set of states: `pending`, `running`, and then a settled state such as `completed`, `idle`, `failed`, `timed_out`, or `interrupted`. diff --git a/scripts/bump_runtime.py b/scripts/bump_runtime.py new file mode 100644 index 0000000..47a4732 --- /dev/null +++ b/scripts/bump_runtime.py @@ -0,0 +1,102 @@ +"""Rewrite the pinned hai-agent-runtime version + per-platform sha256 in the SDK runtime manifest.""" + +from __future__ import annotations + +import argparse +import re +import sys +from dataclasses import dataclass +from pathlib import Path + +# Stdlib only on purpose: this runs as `python scripts/bump_runtime.py` in a +# checkout with no dependencies installed, so a third-party import would break +# the bump step. +RUNTIME_INSTALL = Path(__file__).parents[1] / "src" / "hai_agents" / "local" / "manifest.py" +_SHA256_RE = re.compile(r"^[0-9a-f]{64}$") +_MANIFEST_FILENAME_RE = re.compile(r'"hai-agent-runtime-([^".]+)\.zip"') + + +@dataclass(frozen=True) +class RuntimeBump: + version: str + shas: dict[str, str] # platform key (e.g. darwin-arm64) -> sha256 hex + + def __post_init__(self) -> None: + if not re.fullmatch(r"[0-9]+\.[0-9]+\.[0-9]+(?:[-+][0-9A-Za-z.-]+)?", self.version): + raise ValueError("runtime version must be a release version, e.g. 0.1.13") + for platform, sha in self.shas.items(): + if not _SHA256_RE.fullmatch(sha) or sha == "0" * 64: + raise ValueError(f"{platform}: {sha!r} is not a lowercase 64-char sha256") + + +def _filename_for(platform: str) -> str: + return f"hai-agent-runtime-{platform}.zip" + + +def _manifest_platforms(source: str) -> set[str]: + """Platform keys that have a published artifact literal in `source`.""" + return set(_MANIFEST_FILENAME_RE.findall(source)) + + +def apply_bump(source: str, bump: RuntimeBump) -> str: + """Return `source` with PINNED_RUNTIME_VERSION and the manifest digests replaced; raises if any anchor is missing.""" + published = _manifest_platforms(source) + extra = bump.shas.keys() - published + if extra: + raise ValueError(f"no manifest entry for platform(s): {sorted(extra)}") + # The version is a single literal feeding every derived URL, so any published + # platform left without a fresh sha would keep a stale digest at the new + # version's URL and fail verification on download. Refuse the partial bump. + missing = published - bump.shas.keys() + if missing: + raise ValueError(f"missing sha for published platform(s): {sorted(missing)}") + + updated, count = re.subn( + r'PINNED_RUNTIME_VERSION = "[^"]*"', + f'PINNED_RUNTIME_VERSION = "{bump.version}"', + source, + ) + if count != 1: + raise ValueError(f"expected exactly one PINNED_RUNTIME_VERSION assignment, found {count}") + + for platform, sha in bump.shas.items(): + filename = _filename_for(platform) + pattern = re.compile(rf'("{re.escape(filename)}",\s*(?:#[^\n]*\n\s*)*")[0-9a-fA-F]{{64}}(")') + updated, count = pattern.subn(rf"\g<1>{sha}\g<2>", updated) + if count != 1: + raise ValueError(f"expected exactly one sha256 literal for {filename}, found {count}") + return updated + + +def _parse_args(argv: list[str]) -> RuntimeBump: + parser = argparse.ArgumentParser(description=__doc__) + parser.add_argument("--version", required=True) + parser.add_argument( + "--sha", + action="append", + required=True, + metavar="PLATFORM=SHA256", + help="per-platform digest, e.g. darwin-arm64= (repeatable)", + ) + args = parser.parse_args(argv) + shas: dict[str, str] = {} + for entry in args.sha: + platform, _, sha = entry.partition("=") + if not platform or not sha: + parser.error(f"--sha must be PLATFORM=SHA256, got {entry!r}") + if platform in shas: + parser.error(f"duplicate --sha for {platform}") + shas[platform] = sha.lower() + return RuntimeBump(version=args.version, shas=shas) + + +def main(argv: list[str]) -> int: + bump = _parse_args(argv) + source = RUNTIME_INSTALL.read_text() + RUNTIME_INSTALL.write_text(apply_bump(source, bump)) + print(f"bumped runtime to {bump.version} ({', '.join(sorted(bump.shas))})") + return 0 + + +if __name__ == "__main__": + raise SystemExit(main(sys.argv[1:])) diff --git a/src/hai_agents/local/manifest.py b/src/hai_agents/local/manifest.py index e9588f6..a5e83a5 100644 --- a/src/hai_agents/local/manifest.py +++ b/src/hai_agents/local/manifest.py @@ -1,8 +1,8 @@ """Pinned hai-agent-runtime version and per-platform artifact digests. This module is the SDK's single runtime pin: a runtime release bumps -PINNED_RUNTIME_VERSION and MANIFEST here (via the retargeted -release-hai-agent-runtime.yaml pin PR) and nothing else. Artifacts live under an +PINNED_RUNTIME_VERSION and MANIFEST here with scripts/bump_runtime.py. +The release workflow targets the legacy CLI until the consumer migration ships. Artifacts live under an immutable version-scoped CDN prefix, so an edge can never serve stale bytes. Cross-ref: eng_plans/14-06-2026-holodesktop-binary-versioning-autoupdate. """ diff --git a/tests/test_bump_runtime.py b/tests/test_bump_runtime.py new file mode 100644 index 0000000..263311f --- /dev/null +++ b/tests/test_bump_runtime.py @@ -0,0 +1,38 @@ +"""The release updater must never move URLs while retaining a platform's old digest.""" + +import runpy +import shutil +import subprocess +import sys +from pathlib import Path + +import pytest + +ROOT = Path(__file__).resolve().parents[1] + + +@pytest.mark.parametrize("complete", [True, False]) +def test_release_pin_update_is_complete_or_leaves_manifest_unchanged(tmp_path, complete): + script = tmp_path / "scripts" / "bump_runtime.py" + manifest = tmp_path / "src" / "hai_agents" / "local" / "manifest.py" + script.parent.mkdir(parents=True) + manifest.parent.mkdir(parents=True) + shutil.copyfile(ROOT / "scripts" / "bump_runtime.py", script) + shutil.copyfile(ROOT / "src" / "hai_agents" / "local" / "manifest.py", manifest) + before = manifest.read_bytes() + original = runpy.run_path(str(manifest))["MANIFEST"] + shas = {platform: f"{index:064x}" for index, platform in enumerate(original, start=1)} + args = [sys.executable, str(script), "--version", "9.8.7"] + for platform, sha in list(shas.items())[: len(shas) if complete else -1]: + args.extend(["--sha", f"{platform}={sha}"]) + result = subprocess.run(args, capture_output=True, text=True) + if not complete: + assert result.returncode != 0 + assert manifest.read_bytes() == before + return + assert result.returncode == 0, result.stderr + updated = runpy.run_path(str(manifest))["MANIFEST"] + assert set(updated) == set(original) + for platform, artifact in updated.items(): + assert artifact.sha256 == shas[platform] + assert artifact.url.endswith(f"/9.8.7/hai-agent-runtime-{platform}.zip") From 5fcc24b4eaa4bbd9759c24ec767e23539103d497 Mon Sep 17 00:00:00 2001 From: Charlie Masters Date: Wed, 30 Sep 2026 16:27:21 +0100 Subject: [PATCH 08/35] style(sdk): format shared client exports --- src/hai_agents/__init__.py | 1 + 1 file changed, 1 insertion(+) diff --git a/src/hai_agents/__init__.py b/src/hai_agents/__init__.py index 9850f8a..e454f8d 100644 --- a/src/hai_agents/__init__.py +++ b/src/hai_agents/__init__.py @@ -737,4 +737,5 @@ def __dir__(): ] from .inference import Inference + __all__.append("Inference") From 486bf463885ccc6067c3898fdf4e233605d0a16f Mon Sep 17 00:00:00 2001 From: Charlie Masters Date: Wed, 30 Sep 2026 16:37:18 +0100 Subject: [PATCH 09/35] fix(sdk): release download lock scope and close idle-probe clients --- src/hai_agents/local/runtime.py | 35 ++++++++++++-------- tests/test_runtime_placement.py | 58 +++++++++++++++++++++++++++++++++ 2 files changed, 79 insertions(+), 14 deletions(-) diff --git a/src/hai_agents/local/runtime.py b/src/hai_agents/local/runtime.py index 664f4e9..072e4a2 100644 --- a/src/hai_agents/local/runtime.py +++ b/src/hai_agents/local/runtime.py @@ -131,22 +131,29 @@ def ensure_started( return attached resolved_port = port if port is not None else int(os.environ.get(PORT_ENV, "").strip() or DEFAULT_PORT) + base_url = f"http://{LOOPBACK_HOST}:{resolved_port}" + attached = cls._attach(base_url=base_url, cache_dir=resolved_cache) + if attached is not None: + attached.require_recipe(required_recipe) + return attached + + if command is not None and (not command or binary_path is not None): + raise ValueError("command must be nonempty and cannot be combined with binary_path") + # Downloads are atomically installed and can outlast the process health budget. + # Keep them outside the port lock; recheck attachment before spawning. + cmd = ( + list(command) + if command is not None + else cls._resolve_command( + binary_path=binary_path, version=version, cache_dir=resolved_cache, download=download + ) + ) with _startup_lock(resolved_cache, resolved_port, timeout_s): - base_url = f"http://{LOOPBACK_HOST}:{resolved_port}" attached = cls._attach(base_url=base_url, cache_dir=resolved_cache) if attached is not None: attached.require_recipe(required_recipe) return attached - if command is not None and (not command or binary_path is not None): - raise ValueError("command must be nonempty and cannot be combined with binary_path") - cmd = ( - list(command) - if command is not None - else cls._resolve_command( - binary_path=binary_path, version=version, cache_dir=resolved_cache, download=download - ) - ) if _cancel_event is not None and _cancel_event.is_set(): raise RuntimeUnhealthyError("runtime startup cancelled") explicit_token = ( @@ -380,10 +387,10 @@ def shutdown_if_idle(self) -> bool: """Stop the owned runtime only when it hosts no active sessions; True when it was stopped.""" from ..client import Client # runtime-time import: client.py imports this module lazily too - client = Client(base_url=self.base_url, api_key=self.api_key) - page = client.sessions.list_sessions(status=list(self.ACTIVE_SESSION_STATUSES), size=1) - if page.items: - return False + with Client(base_url=self.base_url, api_key=self.api_key, auto_bridges=False) as client: + page = client.sessions.list_sessions(status=list(self.ACTIVE_SESSION_STATUSES), size=1) + if page.items: + return False self.shutdown() return True diff --git a/tests/test_runtime_placement.py b/tests/test_runtime_placement.py index 8f9e46b..dfd9c61 100644 --- a/tests/test_runtime_placement.py +++ b/tests/test_runtime_placement.py @@ -89,3 +89,61 @@ def shutdown(self): await http.aclose() else: http.close() + + +def test_binary_resolution_does_not_hold_the_port_startup_lock(tmp_path, monkeypatch): + from hai_agents.local import runtime as module + from hai_agents.local.errors import BinaryNotFoundError + + monkeypatch.delenv("HAI_AGENT_LOCAL_BASE_URL", raising=False) + monkeypatch.setattr(LocalRuntime, "_attach", lambda **kwargs: None) + + def resolve(**kwargs): + # A competing client can acquire this actual lock while resolution/download is pending. + with module._startup_lock(tmp_path, 18795, 0.1): + raise BinaryNotFoundError("candidate not installed") + + monkeypatch.setattr(LocalRuntime, "_resolve_command", resolve) + with pytest.raises(BinaryNotFoundError, match="candidate not installed"): + LocalRuntime.ensure_started(cache_dir=tmp_path, port=18795) + + +@pytest.mark.parametrize("status", [200, 400]) +def test_idle_probe_uses_runtime_http_and_closes_its_pool(tmp_path, monkeypatch, status): + from hai_agents.core.api_error import ApiError + + clients, requests, stopped = [], [], [] + original_client = httpx.Client + + def respond(request): + requests.append(request) + assert request.url.path == "/api/v2/sessions" + assert request.headers["Authorization"] == "Bearer local-token" + return httpx.Response(status, json={"items": [], "total": 0, "page": 1, "size": 1}) + + def http_client(**kwargs): + client = original_client(transport=httpx.MockTransport(respond), **kwargs) + clients.append(client) + return client + + monkeypatch.setattr(httpx, "Client", http_client) + runtime = LocalRuntime( + base_url="http://127.0.0.1:18795", + api_key="local-token", + pid=123, + version=None, + log_path=None, + owned=True, + cache_dir=tmp_path, + port=18795, + ) + monkeypatch.setattr(runtime, "shutdown", lambda: stopped.append(True)) + if status == 200: + assert runtime.shutdown_if_idle() + assert stopped == [True] + else: + with pytest.raises(ApiError): + runtime.shutdown_if_idle() + assert stopped == [] + assert requests + assert all(client.is_closed for client in clients) From 6f332031f469606c3a4d2eed0821fc9559cfb39b Mon Sep 17 00:00:00 2001 From: Charlie Masters Date: Wed, 30 Sep 2026 16:41:36 +0100 Subject: [PATCH 10/35] fix(sdk): preserve cancellation retries and safe asynchronous lifecycle --- README.md | 5 ++- src/hai_agents/client.py | 46 +++++++++++++++++----- src/hai_agents/local/runtime.py | 14 ++++--- src/hai_agents/local/state.py | 2 +- src/hai_agents_local/manager.py | 10 ++++- src/hai_agents_local/sessions.py | 25 ++++++------ tests/test_local.py | 47 +++++++++++++++++++++- tests/test_runtime_placement.py | 67 +++++++++++++++++++++++++++++++- 8 files changed, 180 insertions(+), 36 deletions(-) diff --git a/README.md b/README.md index a62948a..f88b213 100644 --- a/README.md +++ b/README.md @@ -69,8 +69,9 @@ print(result.answer) ## Candidate local-agent support -The placement branch adds `Client(mode="local")` and the same option on -`AsyncClient`. Agent placement and environment placement are separate: the client +The placement branch adds `Client(mode="local")` and `await AsyncClient.local()` +for nonblocking asynchronous startup. `AsyncClient(mode="local", runtime=...)` accepts +an already prepared runtime. Agent placement and environment placement are separate: the client selects where the agent runs; each agent environment selects `host="user_device"` or `host="cloud"`. `Client()` continues to use the hosted Agents API. diff --git a/src/hai_agents/client.py b/src/hai_agents/client.py index f736416..e8411d2 100644 --- a/src/hai_agents/client.py +++ b/src/hai_agents/client.py @@ -8,6 +8,7 @@ from __future__ import annotations import asyncio +import contextlib import typing import typing_extensions @@ -190,15 +191,7 @@ def __init__( if runtime is not None and (local_options is not None or inference is not None): raise ValueError("an attached runtime owns its inference and launch configuration") if runtime is None: - from .local.runtime import LocalRuntime - - options = dict(local_options or {}) - options["required_recipe"] = "shared" - options["spawn_env"] = {"HAI_AGENT_RUNTIME_RECIPE": "shared", **options.get("spawn_env", {})} - if inference is not None: - options["spawn_env"] = inference.runtime_env(options.get("spawn_env")) - options["inherit_env"] = False - runtime = LocalRuntime.ensure_started(**options) + raise ValueError("Use await AsyncClient.local() to start a runtime without blocking the event loop") if inference is not None and not runtime.owned: raise ValueError("inference selection cannot reconfigure an existing runtime; choose a free local port") runtime.require_recipe("shared") @@ -211,6 +204,41 @@ def __init__( self.local_runtime.shutdown() raise + @classmethod + async def local( + cls, + *, + inference: typing.Optional[Inference] = None, + local_options: typing.Optional[typing.Dict[str, typing.Any]] = None, + **kwargs: typing.Any, + ) -> "AsyncClient": + """Start or attach off the event loop; the returned client owns any runtime it starts.""" + from .local.runtime import LocalRuntime + + options = dict(local_options or {}) + options["required_recipe"] = "shared" + options["spawn_env"] = {"HAI_AGENT_RUNTIME_RECIPE": "shared", **options.get("spawn_env", {})} + if inference is not None: + options["spawn_env"] = inference.runtime_env(options.get("spawn_env")) + options["inherit_env"] = False + runtime = await LocalRuntime.ensure_started_async(**options) + construction = None + try: + if inference is not None and not runtime.owned: + raise ValueError("inference selection cannot reconfigure an existing runtime; choose a free local port") + construction = asyncio.create_task(asyncio.to_thread(cls, mode="local", runtime=runtime, **kwargs)) + client = await asyncio.shield(construction) + client._owns_runtime = True + return client + except BaseException: + if construction is not None: + with contextlib.suppress(Exception): + client = await construction + await client.aclose() + if runtime.owned: + await asyncio.to_thread(runtime.shutdown) + raise + async def aclose(self) -> None: """Release this client's connections and any runtime it started; borrowed runtimes stay alive.""" try: diff --git a/src/hai_agents/local/runtime.py b/src/hai_agents/local/runtime.py index 072e4a2..b906bd3 100644 --- a/src/hai_agents/local/runtime.py +++ b/src/hai_agents/local/runtime.py @@ -402,12 +402,14 @@ def force_kill(self) -> None: self._cleanup_state_files() def _cleanup_state_files(self) -> None: - if self._token_file is not None: - unlink_if_content(self._token_file, self.api_key) - self._token_file = None - if self._pid_file is not None: - unlink_if_content(self._pid_file, str(self.pid)) - self._pid_file = None + # Serialize compare-and-unlink with publication of a replacement runtime's state. + with _startup_lock(self._cache_dir, self._port, SPAWN_TIMEOUT_S): + if self._token_file is not None: + unlink_if_content(self._token_file, self.api_key) + self._token_file = None + if self._pid_file is not None: + unlink_if_content(self._pid_file, str(self.pid)) + self._pid_file = None @contextlib.contextmanager diff --git a/src/hai_agents/local/state.py b/src/hai_agents/local/state.py index 1af0150..f3eae96 100644 --- a/src/hai_agents/local/state.py +++ b/src/hai_agents/local/state.py @@ -84,7 +84,7 @@ def read_pid(port: int, *, cache_dir: typing.Optional[_PathInput] = None) -> typ def unlink_if_content(path: pathlib.Path, content: str) -> None: - """Remove our state file, but never one a concurrent spawner already replaced.""" + """Remove matching state; callers must hold the startup lock against concurrent replacement.""" with contextlib.suppress(OSError): if path.read_text(encoding="utf-8").strip() == content: path.unlink(missing_ok=True) diff --git a/src/hai_agents_local/manager.py b/src/hai_agents_local/manager.py index 7a2570c..5682c42 100644 --- a/src/hai_agents_local/manager.py +++ b/src/hai_agents_local/manager.py @@ -35,7 +35,10 @@ def ensure(self, bridges: Sequence[LocalBridge]) -> list[str]: if self._ensure_one(bridge): started.append(bridge.session_id) except BaseException: - self.stop(started) + try: + self.stop(started) + except Exception: + logger.exception("Failed to clean up bridges after startup failure") raise return started @@ -71,7 +74,10 @@ def _ensure_one(self, bridge: LocalBridge) -> bool: with self._lock: if self._runners.get(bridge.session_id) is runner: del self._runners[bridge.session_id] - runner.stop() + try: + runner.stop() + except Exception: + logger.exception("Failed to clean up bridge after startup failure") raise return started diff --git a/src/hai_agents_local/sessions.py b/src/hai_agents_local/sessions.py index ae5a6ef..455eacb 100644 --- a/src/hai_agents_local/sessions.py +++ b/src/hai_agents_local/sessions.py @@ -207,14 +207,13 @@ def cancel_session(self, session_id: str, **kwargs: typing.Any) -> typing.Any: # Stop local execution even if the remote cancellation cannot be delivered. owned = getattr(self, "_owned_bridges", {}).get(str(session_id), []) try: - try: - if owned: - stop_bridges(owned) - self._owned_bridges.pop(str(session_id), None) - finally: - response = super().cancel_session(session_id, **kwargs) + if owned: + stop_bridges(owned) + self._owned_bridges.pop(str(session_id), None) finally: - _deregister_exit_cancel(str(session_id)) + response = super().cancel_session(session_id, **kwargs) + # Keep the exit retry registered until cancellation is confirmed. + _deregister_exit_cancel(str(session_id)) return response @functools.wraps(SessionsClient.create_session) @@ -262,14 +261,12 @@ async def aclose(self) -> None: async def cancel_session(self, session_id: str, **kwargs: typing.Any) -> typing.Any: owned = getattr(self, "_owned_bridges", {}).get(str(session_id), []) try: - try: - if owned: - await asyncio.to_thread(stop_bridges, owned) - self._owned_bridges.pop(str(session_id), None) - finally: - response = await super().cancel_session(session_id, **kwargs) + if owned: + await asyncio.to_thread(stop_bridges, owned) + self._owned_bridges.pop(str(session_id), None) finally: - _deregister_exit_cancel(str(session_id)) + response = await super().cancel_session(session_id, **kwargs) + _deregister_exit_cancel(str(session_id)) return response @functools.wraps(AsyncSessionsClient.create_session) diff --git a/tests/test_local.py b/tests/test_local.py index 47b9cc3..c734110 100644 --- a/tests/test_local.py +++ b/tests/test_local.py @@ -764,13 +764,24 @@ def manager(self): yield manager manager.stop_all() - def test_startup_failure_surfaces_to_caller_without_firing_on_crash(self, manager): + @pytest.mark.parametrize("cleanup_fails", [False, True]) + def test_startup_failure_surfaces_to_caller_without_firing_on_crash(self, manager, monkeypatch, cleanup_fails): crashed = threading.Event() class FailingBridge(ServingBridge): async def run(self): raise AuthError("bad key") + if cleanup_fails: + from hai_agents_local.manager import _Runner + + original_stop = _Runner.stop + + def failed_cleanup(runner): + original_stop(runner) + raise TimeoutError("cleanup also failed") + + monkeypatch.setattr(_Runner, "stop", failed_cleanup) bridge = FailingBridge(api_key="k") bridge.on_crash = crashed.set with pytest.raises(RuntimeError) as exc_info: @@ -989,3 +1000,37 @@ def request_stop(self): bridge.request_stop = lambda: ServingBridge.request_stop(bridge) manager.stop([bridge.session_id]) assert not manager._runners + + +@pytest.mark.asyncio +@pytest.mark.parametrize("asynchronous", [False, True]) +async def test_failed_api_cancel_keeps_interpreter_exit_retry(monkeypatch, asynchronous): + from hai_agents import AsyncClient + from hai_agents.sessions.client import AsyncSessionsClient + from hai_agents_local import sessions as module + + retried = [] + monkeypatch.setattr(module, "_exit_cancels", {"run": lambda: retried.append("run")}) + + def cancel(*args, **kwargs): + raise RuntimeError("remote cancel unavailable") + + async def async_cancel(*args, **kwargs): + cancel() + + monkeypatch.setattr(SessionsClient, "cancel_session", cancel) + monkeypatch.setattr(AsyncSessionsClient, "cancel_session", async_cancel) + client = (AsyncClient if asynchronous else Client)(api_key=API_KEY) + try: + with pytest.raises(RuntimeError, match="remote cancel unavailable"): + if asynchronous: + await client.sessions.cancel_session("run") + else: + client.sessions.cancel_session("run") + module._cancel_sessions_at_exit() + assert retried == ["run"] + finally: + if asynchronous: + await client.aclose() + else: + client.close() diff --git a/tests/test_runtime_placement.py b/tests/test_runtime_placement.py index dfd9c61..5de6f67 100644 --- a/tests/test_runtime_placement.py +++ b/tests/test_runtime_placement.py @@ -77,7 +77,7 @@ def shutdown(self): http = httpx.AsyncClient() if asynchronous else httpx.Client() client_type = AsyncClient if asynchronous else Client options = {"runtime": runtime, "httpx_client": http} if borrowed else {} - client = client_type(mode="local", **options) + client = await AsyncClient.local() if asynchronous and not borrowed else client_type(mode="local", **options) if asynchronous: await client.aclose() else: @@ -147,3 +147,68 @@ def http_client(**kwargs): assert stopped == [] assert requests assert all(client.is_closed for client in clients) + + +@pytest.mark.asyncio +async def test_async_local_startup_keeps_loop_responsive_and_cancellation_cleans_child(monkeypatch): + import asyncio + import threading + from types import SimpleNamespace + + entered, release = threading.Event(), threading.Event() + stopped = [] + runtime = SimpleNamespace(owned=True, shutdown=lambda: stopped.append(True)) + + def start(**kwargs): + entered.set() + assert release.wait(5) + return runtime + + monkeypatch.setattr(LocalRuntime, "ensure_started", start) + startup = asyncio.create_task(AsyncClient.local()) + try: + assert await asyncio.to_thread(entered.wait, 1), "startup blocked the event loop" + startup.cancel() + release.set() + with pytest.raises(asyncio.CancelledError): + await startup + assert stopped == [True] + finally: + release.set() + if not startup.done(): + await startup + + +def test_state_cleanup_waits_for_startup_and_preserves_replacement(tmp_path): + import threading + from concurrent.futures import ThreadPoolExecutor + + from hai_agents.local import runtime as module + + token_file = write_owner_only(token_file_path(18795, cache_dir=tmp_path), "old-token") + runtime = LocalRuntime( + base_url="http://127.0.0.1:18795", + api_key="old-token", + pid=123, + version=None, + log_path=None, + owned=True, + cache_dir=tmp_path, + port=18795, + token_file=token_file, + ) + entered = threading.Event() + + def cleanup(): + entered.set() + runtime._cleanup_state_files() + + with ThreadPoolExecutor() as pool: + with module._startup_lock(tmp_path, 18795, 1): + cleaning = pool.submit(cleanup) + assert entered.wait(1) + with pytest.raises(TimeoutError): + cleaning.result(timeout=0.05) + write_owner_only(token_file, "replacement-token") + cleaning.result(timeout=2) + assert token_file.read_text() == "replacement-token" From f75842fc28fc370250130af599083036b779e9a3 Mon Sep 17 00:00:00 2001 From: Charlie Masters Date: Wed, 30 Sep 2026 16:59:08 +0100 Subject: [PATCH 11/35] Clean up owned runtime when client validation fails --- src/hai_agents/client.py | 10 +++++---- src/hai_agents/local/runtime.py | 17 ++++++++------- tests/test_runtime_placement.py | 37 +++++++++++++++++++++++++++++++-- 3 files changed, 51 insertions(+), 13 deletions(-) diff --git a/src/hai_agents/client.py b/src/hai_agents/client.py index e8411d2..2fb6b09 100644 --- a/src/hai_agents/client.py +++ b/src/hai_agents/client.py @@ -73,10 +73,11 @@ def __init__( runtime = LocalRuntime.ensure_started(**options) if inference is not None and not runtime.owned: raise ValueError("inference selection cannot reconfigure an existing runtime; choose a free local port") - runtime.require_recipe("shared") self.local_runtime = runtime - kwargs.update(base_url=runtime.base_url, api_key=runtime.api_key) try: + if self.local_runtime is not None: + self.local_runtime.require_recipe("shared") + kwargs.update(base_url=self.local_runtime.base_url, api_key=self.local_runtime.api_key) super().__init__(**kwargs) except BaseException: if self._owns_runtime and self.local_runtime is not None and self.local_runtime.owned: @@ -194,10 +195,11 @@ def __init__( raise ValueError("Use await AsyncClient.local() to start a runtime without blocking the event loop") if inference is not None and not runtime.owned: raise ValueError("inference selection cannot reconfigure an existing runtime; choose a free local port") - runtime.require_recipe("shared") self.local_runtime = runtime - kwargs.update(base_url=runtime.base_url, api_key=runtime.api_key) try: + if self.local_runtime is not None: + self.local_runtime.require_recipe("shared") + kwargs.update(base_url=self.local_runtime.base_url, api_key=self.local_runtime.api_key) super().__init__(**kwargs) except BaseException: if self._owns_runtime and self.local_runtime is not None and self.local_runtime.owned: diff --git a/src/hai_agents/local/runtime.py b/src/hai_agents/local/runtime.py index b906bd3..f5ea173 100644 --- a/src/hai_agents/local/runtime.py +++ b/src/hai_agents/local/runtime.py @@ -271,13 +271,16 @@ def _attach(cls, *, base_url: str, cache_dir: pathlib.Path) -> typing.Optional[" f"{AUTH_TOKEN_ENV} is not set and {token_file_path(port, cache_dir=cache_dir)} does not exist, " "so this client cannot authenticate. Export the token or stop that runtime." ) - response = httpx.get( - f"{base_url}/api/v2/sessions", - headers={"Authorization": f"Bearer {token}"}, - params={"size": 1}, - timeout=2.0, - follow_redirects=False, - ) + try: + response = httpx.get( + f"{base_url}/api/v2/sessions", + headers={"Authorization": f"Bearer {token}"}, + params={"size": 1}, + timeout=2.0, + follow_redirects=False, + ) + except httpx.HTTPError as exc: + raise LocalRuntimeError("runtime attachment failed authenticated session probe") from exc if response.status_code != 200: raise LocalRuntimeError("runtime attachment failed authenticated session probe") reported = payload.get("version") diff --git a/tests/test_runtime_placement.py b/tests/test_runtime_placement.py index 5de6f67..5ba76ed 100644 --- a/tests/test_runtime_placement.py +++ b/tests/test_runtime_placement.py @@ -24,7 +24,8 @@ def test_self_hosted_inference_does_not_receive_hosted_key(monkeypatch): assert "HAI_AGENT_RUNTIME_BASE_URL" not in Inference.cloud().runtime_env() -def test_attachment_authenticates_and_checks_recipe_before_use(tmp_path, monkeypatch): +@pytest.mark.parametrize("probe_error", [None, httpx.ConnectError("offline"), httpx.ReadTimeout("timeout")]) +def test_attachment_authenticates_and_checks_recipe_before_use(tmp_path, monkeypatch, probe_error): from hai_agents.local import runtime as module write_owner_only(token_file_path(18795, cache_dir=tmp_path), "local-token") @@ -44,7 +45,13 @@ def get(url, **kwargs): with pytest.raises(LocalRuntimeError): attached.force_kill() assert token_file_path(18795, cache_dir=tmp_path).exists() - monkeypatch.setattr(module.httpx, "get", lambda *a, **k: httpx.Response(401)) + + def failed_probe(*args, **kwargs): + if probe_error is not None: + raise probe_error + return httpx.Response(401) + + monkeypatch.setattr(module.httpx, "get", failed_probe) with pytest.raises(LocalRuntimeError, match="authenticated"): LocalRuntime.attach(cache_dir=tmp_path) @@ -212,3 +219,29 @@ def cleanup(): write_owner_only(token_file, "replacement-token") cleaning.result(timeout=2) assert token_file.read_text() == "replacement-token" + + +@pytest.mark.asyncio +@pytest.mark.parametrize("asynchronous", [False, True]) +@pytest.mark.parametrize("borrowed", [False, True]) +async def test_failed_client_recipe_check_releases_only_started_runtime(monkeypatch, asynchronous, borrowed): + stopped = [] + + class Runtime: + owned = True + + def require_recipe(self, recipe): + raise BinaryIncompatibleError("recipe changed") + + def shutdown(self): + stopped.append(True) + + runtime = Runtime() + monkeypatch.setattr(LocalRuntime, "ensure_started", lambda **options: runtime) + with pytest.raises(BinaryIncompatibleError, match="recipe changed"): + if asynchronous and not borrowed: + await AsyncClient.local() + else: + client_type = AsyncClient if asynchronous else Client + client_type(mode="local", **({"runtime": runtime} if borrowed else {})) + assert stopped == ([] if borrowed else [True]) From 82bf0e77f3d7acf24eea88f5e249be306fae5fe0 Mon Sep 17 00:00:00 2001 From: Charlie Masters Date: Wed, 30 Sep 2026 17:10:26 +0100 Subject: [PATCH 12/35] Notify displaced sessions even when bridge stop times out --- src/hai_agents_local/manager.py | 8 ++++++-- tests/test_local.py | 11 ++++++++++- 2 files changed, 16 insertions(+), 3 deletions(-) diff --git a/src/hai_agents_local/manager.py b/src/hai_agents_local/manager.py index 5682c42..044aeef 100644 --- a/src/hai_agents_local/manager.py +++ b/src/hai_agents_local/manager.py @@ -56,8 +56,12 @@ def _ensure_one(self, bridge: LocalBridge) -> bool: runner = _Runner(bridge) self._runners[bridge.session_id] = runner for other in displaced: - other.stop() - other.notify_lost() + try: + other.stop() + except TimeoutError: + logger.exception("Displaced bridge did not stop in time") + finally: + other.notify_lost() try: if not runner.bridge.ready.wait(READY_TIMEOUT_S): hint = f" ({bridge.startup_hint})" if bridge.startup_hint is not None else "" diff --git a/tests/test_local.py b/tests/test_local.py index c734110..b28702e 100644 --- a/tests/test_local.py +++ b/tests/test_local.py @@ -811,7 +811,8 @@ async def run(self): manager.ensure([NeverReadyBridge(api_key="k")]) assert manager._runners == {} - def test_newer_session_takes_over_the_kind_and_notifies_the_displaced(self, manager): + @pytest.mark.parametrize("stop_times_out", [False, True]) + def test_newer_session_takes_over_the_kind_and_notifies_the_displaced(self, manager, monkeypatch, stop_times_out): first = ServingBridge(api_key="k") second = ServingBridge(api_key="k") browser = BrowserServingBridge(api_key="k") @@ -820,6 +821,14 @@ def test_newer_session_takes_over_the_kind_and_notifies_the_displaced(self, mana second.on_crash = second_lost.set manager.ensure([first, browser]) first_runner = manager._runners[first.session_id] + if stop_times_out: + original_stop = first_runner.stop + + def timed_out_stop(): + original_stop() + raise TimeoutError("displaced bridge timeout") + + monkeypatch.setattr(first_runner, "stop", timed_out_stop) manager.ensure([second]) assert first.session_id not in manager._runners assert not first_runner.thread.is_alive() From d21f77fd5fe0e8358166fa607c11c7633132c734 Mon Sep 17 00:00:00 2001 From: abonneth Date: Thu, 1 Oct 2026 18:10:11 +0200 Subject: [PATCH 13/35] refactor(sdk): move local runtime into hai_agents_local and add Client.local() Hand-written runtime management lives outside the generated package, so a codegen sync cannot wipe it. Local clients are built with Client.local() and AsyncClient.local(), keeping the generated constructors and their typing. --- README.md | 16 +- src/hai_agents/__init__.py | 4 - src/hai_agents/client.py | 194 ++++++------------ src/hai_agents_local/config.py | 2 +- .../runtime}/__init__.py | 10 +- src/hai_agents_local/runtime/acquire.py | 67 ++++++ .../runtime}/errors.py | 2 +- .../runtime}/inference.py | 2 +- .../runtime}/install.py | 5 +- .../runtime}/manifest.py | 10 +- .../runtime}/process.py | 0 .../runtime}/runtime.py | 18 +- .../runtime}/state.py | 4 +- src/hai_agents_local/sessions.py | 35 +++- tests/test_runtime_placement.py | 120 ++++++----- tests/test_runtime_state.py | 2 +- 16 files changed, 261 insertions(+), 230 deletions(-) rename src/{hai_agents/local => hai_agents_local/runtime}/__init__.py (56%) create mode 100644 src/hai_agents_local/runtime/acquire.py rename src/{hai_agents/local => hai_agents_local/runtime}/errors.py (94%) rename src/{hai_agents => hai_agents_local/runtime}/inference.py (93%) rename src/{hai_agents/local => hai_agents_local/runtime}/install.py (97%) rename src/{hai_agents/local => hai_agents_local/runtime}/manifest.py (84%) rename src/{hai_agents/local => hai_agents_local/runtime}/process.py (100%) rename src/{hai_agents/local => hai_agents_local/runtime}/runtime.py (95%) rename src/{hai_agents/local => hai_agents_local/runtime}/state.py (95%) diff --git a/README.md b/README.md index f88b213..2e573c1 100644 --- a/README.md +++ b/README.md @@ -69,9 +69,9 @@ print(result.answer) ## Candidate local-agent support -The placement branch adds `Client(mode="local")` and `await AsyncClient.local()` -for nonblocking asynchronous startup. `AsyncClient(mode="local", runtime=...)` accepts -an already prepared runtime. Agent placement and environment placement are separate: the client +The placement branch adds `Client.local()` and `await AsyncClient.local()`, which +start a local agent runtime or, with `runtime=...`, use an already prepared one. +Agent placement and environment placement are separate: the client selects where the agent runs; each agent environment selects `host="user_device"` or `host="cloud"`. `Client()` continues to use the hosted Agents API. @@ -80,8 +80,7 @@ For development, install the candidate SDK and use a prepared HAI source checkou ```python from hai_agents import Client -with Client( - mode="local", +with Client.local( local_options={ "command": ["/path/to/hai/.venv/bin/python", "-m", "hai_agent_runtime"], "download": False, @@ -107,9 +106,10 @@ The source runtime must support the `shared` recipe. The current pinned binary 0.1.8 does not; candidate source or an explicitly supplied compatible `local_options["binary_path"]` is required until a compatible release is pinned. The runtime inherits `HAI_API_KEY` for hosted inference. Local agent execution -does not by itself imply local inference: `Inference.self_hosted(url, model=...)` -selects a model endpoint for a newly started local runtime. Hosted agents do not -currently accept that override. +does not by itself imply local inference: +`Client.local(inference=Inference.self_hosted(url, model=...))`, with `Inference` +from `hai_agents_local.runtime`, selects a model endpoint for a newly started +local runtime. Hosted agents do not currently accept that override. Closing the client shuts down a runtime it started; an attached runtime remains owned by its caller. `cancel()` ends the agent session. For a cloud workstation, diff --git a/src/hai_agents/__init__.py b/src/hai_agents/__init__.py index e454f8d..86cfeec 100644 --- a/src/hai_agents/__init__.py +++ b/src/hai_agents/__init__.py @@ -735,7 +735,3 @@ def __dir__(): "wait_for_session", "webhooks", ] - -from .inference import Inference - -__all__.append("Inference") diff --git a/src/hai_agents/client.py b/src/hai_agents/client.py index 2fb6b09..ae7ae55 100644 --- a/src/hai_agents/client.py +++ b/src/hai_agents/client.py @@ -8,13 +8,11 @@ from __future__ import annotations import asyncio -import contextlib import typing import typing_extensions from .base_client import AsyncBaseClient, BaseClient -from .inference import Inference from .polling import ( AnswerT, AsyncSessionHandle, @@ -30,74 +28,52 @@ from .sessions.client import AsyncSessionsClient, SessionsClient from .tools import ToolInput, as_tools +if typing.TYPE_CHECKING: + from hai_agents_local.runtime import Inference, LocalRuntime + class Client(BaseClient): - def __init__( - self, + local_runtime: typing.Optional[LocalRuntime] = None + _owns_runtime = False + _auto_bridges = True + + @classmethod + def local( + cls, *, - mode: typing.Literal["local", "remote"] = "remote", + runtime: typing.Optional[LocalRuntime] = None, inference: typing.Optional[Inference] = None, - auto_bridges: bool = True, - runtime: typing.Any = None, local_options: typing.Optional[typing.Dict[str, typing.Any]] = None, - **kwargs: typing.Any, - ) -> None: - if mode not in {"local", "remote"}: - raise ValueError("mode must be local or remote") - self._auto_bridges = auto_bridges - self.mode = mode - self.local_runtime = None - self._owns_runtime = mode == "local" and runtime is None - self._owns_http = kwargs.get("httpx_client") is None - if mode == "remote": - if runtime is not None or local_options is not None: - raise ValueError("runtime and local_options require mode='local'") - if inference is not None and inference.base_url is not None: - raise ValueError("self-hosted inference currently requires a local agent") - else: - if "base_url" in kwargs or "api_key" in kwargs: - raise ValueError( - "local API credentials come from runtime; pass inference credentials via local_options" - ) - if runtime is not None and (local_options is not None or inference is not None): - raise ValueError("an attached runtime owns its inference and launch configuration") - if runtime is None: - from .local.runtime import LocalRuntime + auto_bridges: bool = True, + timeout: typing.Optional[float] = None, + ) -> Client: + """A client on a local agent runtime: ``runtime`` if given, else one this client starts and owns.""" + from hai_agents_local.runtime import acquire_runtime - options = dict(local_options or {}) - options["required_recipe"] = "shared" - options["spawn_env"] = {"HAI_AGENT_RUNTIME_RECIPE": "shared", **options.get("spawn_env", {})} - if inference is not None: - options["spawn_env"] = inference.runtime_env(options.get("spawn_env")) - options["inherit_env"] = False - runtime = LocalRuntime.ensure_started(**options) - if inference is not None and not runtime.owned: - raise ValueError("inference selection cannot reconfigure an existing runtime; choose a free local port") - self.local_runtime = runtime + runtime, owned = acquire_runtime(runtime, inference=inference, local_options=local_options) try: - if self.local_runtime is not None: - self.local_runtime.require_recipe("shared") - kwargs.update(base_url=self.local_runtime.base_url, api_key=self.local_runtime.api_key) - super().__init__(**kwargs) + client = cls(base_url=runtime.base_url, api_key=runtime.api_key, httpx_client=runtime.http_client(timeout)) except BaseException: - if self._owns_runtime and self.local_runtime is not None and self.local_runtime.owned: - self.local_runtime.shutdown() + if owned: + runtime.shutdown() raise + client.local_runtime, client._owns_runtime, client._auto_bridges = runtime, owned, auto_bridges + return client def close(self) -> None: - """Release this client's connections and any runtime it started; borrowed runtimes stay alive.""" + """Stop sessions this client bridged; a local client also releases its runtime and connections.""" try: - if self._sessions is not None and hasattr(self._sessions, "close"): + if self._sessions is not None: self._sessions.close() finally: - try: - if self._owns_runtime and self.local_runtime is not None and self.local_runtime.owned: - self.local_runtime.shutdown() - finally: - if self._owns_http: + if self.local_runtime is not None: + try: + if self._owns_runtime: + self.local_runtime.shutdown() + finally: self._client_wrapper.httpx_client.httpx_client.close() - def __enter__(self) -> "Client": + def __enter__(self) -> Client: return self def __exit__(self, *exc: typing.Any) -> None: @@ -152,109 +128,59 @@ def session(self, id: str) -> SessionHandle: @property def sessions(self) -> SessionsClient: - if not self._auto_bridges: - return super().sessions if self._sessions is None: from hai_agents_local.sessions import LocalSessionsClient - self._sessions = LocalSessionsClient(client_wrapper=self._client_wrapper) + self._sessions = LocalSessionsClient( + client_wrapper=self._client_wrapper, runtime=self.local_runtime, auto_bridges=self._auto_bridges + ) return self._sessions class AsyncClient(AsyncBaseClient): - def __init__( - self, - *, - mode: typing.Literal["local", "remote"] = "remote", - inference: typing.Optional[Inference] = None, - auto_bridges: bool = True, - runtime: typing.Any = None, - local_options: typing.Optional[typing.Dict[str, typing.Any]] = None, - **kwargs: typing.Any, - ) -> None: - if mode not in {"local", "remote"}: - raise ValueError("mode must be local or remote") - self._auto_bridges = auto_bridges - self.mode = mode - self.local_runtime = None - self._owns_runtime = mode == "local" and runtime is None - self._owns_http = kwargs.get("httpx_client") is None - if mode == "remote": - if runtime is not None or local_options is not None: - raise ValueError("runtime and local_options require mode='local'") - if inference is not None and inference.base_url is not None: - raise ValueError("self-hosted inference currently requires a local agent") - else: - if "base_url" in kwargs or "api_key" in kwargs: - raise ValueError( - "local API credentials come from runtime; pass inference credentials via local_options" - ) - if runtime is not None and (local_options is not None or inference is not None): - raise ValueError("an attached runtime owns its inference and launch configuration") - if runtime is None: - raise ValueError("Use await AsyncClient.local() to start a runtime without blocking the event loop") - if inference is not None and not runtime.owned: - raise ValueError("inference selection cannot reconfigure an existing runtime; choose a free local port") - self.local_runtime = runtime - try: - if self.local_runtime is not None: - self.local_runtime.require_recipe("shared") - kwargs.update(base_url=self.local_runtime.base_url, api_key=self.local_runtime.api_key) - super().__init__(**kwargs) - except BaseException: - if self._owns_runtime and self.local_runtime is not None and self.local_runtime.owned: - self.local_runtime.shutdown() - raise + local_runtime: typing.Optional[LocalRuntime] = None + _owns_runtime = False + _auto_bridges = True @classmethod async def local( cls, *, + runtime: typing.Optional[LocalRuntime] = None, inference: typing.Optional[Inference] = None, local_options: typing.Optional[typing.Dict[str, typing.Any]] = None, - **kwargs: typing.Any, - ) -> "AsyncClient": - """Start or attach off the event loop; the returned client owns any runtime it starts.""" - from .local.runtime import LocalRuntime + auto_bridges: bool = True, + timeout: typing.Optional[float] = None, + ) -> AsyncClient: + """A client on a local agent runtime: ``runtime`` if given, else one this client starts and owns.""" + from hai_agents_local.runtime import acquire_runtime_async - options = dict(local_options or {}) - options["required_recipe"] = "shared" - options["spawn_env"] = {"HAI_AGENT_RUNTIME_RECIPE": "shared", **options.get("spawn_env", {})} - if inference is not None: - options["spawn_env"] = inference.runtime_env(options.get("spawn_env")) - options["inherit_env"] = False - runtime = await LocalRuntime.ensure_started_async(**options) - construction = None + runtime, owned = await acquire_runtime_async(runtime, inference=inference, local_options=local_options) try: - if inference is not None and not runtime.owned: - raise ValueError("inference selection cannot reconfigure an existing runtime; choose a free local port") - construction = asyncio.create_task(asyncio.to_thread(cls, mode="local", runtime=runtime, **kwargs)) - client = await asyncio.shield(construction) - client._owns_runtime = True - return client + client = cls( + base_url=runtime.base_url, api_key=runtime.api_key, httpx_client=runtime.async_http_client(timeout) + ) except BaseException: - if construction is not None: - with contextlib.suppress(Exception): - client = await construction - await client.aclose() - if runtime.owned: + if owned: await asyncio.to_thread(runtime.shutdown) raise + client.local_runtime, client._owns_runtime, client._auto_bridges = runtime, owned, auto_bridges + return client async def aclose(self) -> None: - """Release this client's connections and any runtime it started; borrowed runtimes stay alive.""" + """Stop sessions this client bridged; a local client also releases its runtime and connections.""" try: - if self._sessions is not None and hasattr(self._sessions, "aclose"): + if self._sessions is not None: await self._sessions.aclose() finally: - try: - if self._owns_runtime and self.local_runtime is not None and self.local_runtime.owned: - await asyncio.to_thread(self.local_runtime.shutdown) - finally: - if self._owns_http: + if self.local_runtime is not None: + try: + if self._owns_runtime: + await asyncio.to_thread(self.local_runtime.shutdown) + finally: await self._client_wrapper.httpx_client.httpx_client.aclose() - async def __aenter__(self) -> "AsyncClient": + async def __aenter__(self) -> AsyncClient: return self async def __aexit__(self, *exc: typing.Any) -> None: @@ -309,10 +235,10 @@ def session(self, id: str) -> AsyncSessionHandle: @property def sessions(self) -> AsyncSessionsClient: - if not self._auto_bridges: - return super().sessions if self._sessions is None: from hai_agents_local.sessions import LocalAsyncSessionsClient - self._sessions = LocalAsyncSessionsClient(client_wrapper=self._client_wrapper) + self._sessions = LocalAsyncSessionsClient( + client_wrapper=self._client_wrapper, runtime=self.local_runtime, auto_bridges=self._auto_bridges + ) return self._sessions diff --git a/src/hai_agents_local/config.py b/src/hai_agents_local/config.py index e057ba4..054cbc4 100644 --- a/src/hai_agents_local/config.py +++ b/src/hai_agents_local/config.py @@ -1,4 +1,4 @@ -"""Environment variables read by hai_agents.local.""" +"""Environment variables read by hai_agents_local.""" from __future__ import annotations diff --git a/src/hai_agents/local/__init__.py b/src/hai_agents_local/runtime/__init__.py similarity index 56% rename from src/hai_agents/local/__init__.py rename to src/hai_agents_local/runtime/__init__.py index 215e46d..10207f6 100644 --- a/src/hai_agents/local/__init__.py +++ b/src/hai_agents_local/runtime/__init__.py @@ -1,9 +1,9 @@ -"""Local-mode runtime management: install/find/start a hai-agent-runtime binary. +"""Local agent runtime management: install, find, start, attach to and verify a hai-agent-runtime binary. -Never imported by the base ``hai_agents`` package; ``Client.local`` pulls it in -lazily so remote-only users pay nothing for it. +Imported lazily by ``Client.local`` so remote-only users pay nothing for it. """ +from .acquire import acquire_runtime, acquire_runtime_async from .errors import ( BinaryIncompatibleError, BinaryNotFoundError, @@ -12,14 +12,18 @@ RuntimeStartTimeoutError, RuntimeUnhealthyError, ) +from .inference import Inference from .runtime import LocalRuntime __all__ = [ "BinaryIncompatibleError", "BinaryNotFoundError", "DownloadVerificationError", + "Inference", "LocalRuntime", "LocalRuntimeError", "RuntimeStartTimeoutError", "RuntimeUnhealthyError", + "acquire_runtime", + "acquire_runtime_async", ] diff --git a/src/hai_agents_local/runtime/acquire.py b/src/hai_agents_local/runtime/acquire.py new file mode 100644 index 0000000..376f054 --- /dev/null +++ b/src/hai_agents_local/runtime/acquire.py @@ -0,0 +1,67 @@ +"""The runtime behind ``Client.local``: a caller-provided one, or a shared-recipe runtime started for the client.""" + +from __future__ import annotations + +import asyncio +import typing + +from .inference import Inference +from .runtime import LocalRuntime + +SHARED_RECIPE = "shared" +RECIPE_ENV = "HAI_AGENT_RUNTIME_RECIPE" + + +def acquire_runtime( + runtime: typing.Optional[LocalRuntime], + *, + inference: typing.Optional[Inference], + local_options: typing.Optional[typing.Dict[str, typing.Any]], +) -> typing.Tuple[LocalRuntime, bool]: + """The runtime a local client talks to, and whether that client owns (and so stops) it.""" + if runtime is not None: + _require_unconfigured(inference, local_options) + runtime.require_recipe(SHARED_RECIPE) + return runtime, False + started = LocalRuntime.ensure_started(**_launch_options(inference, local_options)) + return started, _claim(started, inference) + + +async def acquire_runtime_async( + runtime: typing.Optional[LocalRuntime], + *, + inference: typing.Optional[Inference], + local_options: typing.Optional[typing.Dict[str, typing.Any]], +) -> typing.Tuple[LocalRuntime, bool]: + """``acquire_runtime`` off the event loop; cancellation stops a runtime started for the client.""" + if runtime is not None: + _require_unconfigured(inference, local_options) + await asyncio.to_thread(runtime.require_recipe, SHARED_RECIPE) + return runtime, False + started = await LocalRuntime.ensure_started_async(**_launch_options(inference, local_options)) + return started, _claim(started, inference) + + +def _require_unconfigured( + inference: typing.Optional[Inference], local_options: typing.Optional[typing.Dict[str, typing.Any]] +) -> None: + if inference is not None or local_options is not None: + raise ValueError("an attached runtime owns its inference and launch configuration") + + +def _launch_options( + inference: typing.Optional[Inference], local_options: typing.Optional[typing.Dict[str, typing.Any]] +) -> typing.Dict[str, typing.Any]: + options = dict(local_options or {}) + options["required_recipe"] = SHARED_RECIPE + options["spawn_env"] = {RECIPE_ENV: SHARED_RECIPE, **options.get("spawn_env", {})} + if inference is not None: + options["spawn_env"] = inference.runtime_env(options["spawn_env"]) + options["inherit_env"] = False + return options + + +def _claim(runtime: LocalRuntime, inference: typing.Optional[Inference]) -> bool: + if inference is not None and not runtime.owned: + raise ValueError("inference selection cannot reconfigure an existing runtime; choose a free local port") + return runtime.owned diff --git a/src/hai_agents/local/errors.py b/src/hai_agents_local/runtime/errors.py similarity index 94% rename from src/hai_agents/local/errors.py rename to src/hai_agents_local/runtime/errors.py index dfff655..970aa39 100644 --- a/src/hai_agents/local/errors.py +++ b/src/hai_agents_local/runtime/errors.py @@ -1,4 +1,4 @@ -"""Error types for hai_agents.local.""" +"""Error types for hai_agents_local.runtime.""" from __future__ import annotations diff --git a/src/hai_agents/inference.py b/src/hai_agents_local/runtime/inference.py similarity index 93% rename from src/hai_agents/inference.py rename to src/hai_agents_local/runtime/inference.py index c843c36..c1eac1a 100644 --- a/src/hai_agents/inference.py +++ b/src/hai_agents_local/runtime/inference.py @@ -1,4 +1,4 @@ -"""Inference placement, independent of agent and environment placement.""" +"""Inference placement for a local agent runtime, independent of environment placement.""" from __future__ import annotations diff --git a/src/hai_agents/local/install.py b/src/hai_agents_local/runtime/install.py similarity index 97% rename from src/hai_agents/local/install.py rename to src/hai_agents_local/runtime/install.py index a662c8d..d928013 100644 --- a/src/hai_agents/local/install.py +++ b/src/hai_agents_local/runtime/install.py @@ -1,8 +1,7 @@ """Verified download and atomic install of the hai-agent-runtime binary. -Port of holo_desktop.agent_client.runtime_install with the TTY prompt and rich -progress removed: the SDK is a library, so consent is the caller's -``download=True`` and progress is plain logging. +The SDK is a library, so consent is the caller's ``download=True`` and progress +is plain logging. """ from __future__ import annotations diff --git a/src/hai_agents/local/manifest.py b/src/hai_agents_local/runtime/manifest.py similarity index 84% rename from src/hai_agents/local/manifest.py rename to src/hai_agents_local/runtime/manifest.py index e9588f6..e1f782c 100644 --- a/src/hai_agents/local/manifest.py +++ b/src/hai_agents_local/runtime/manifest.py @@ -1,10 +1,8 @@ """Pinned hai-agent-runtime version and per-platform artifact digests. This module is the SDK's single runtime pin: a runtime release bumps -PINNED_RUNTIME_VERSION and MANIFEST here (via the retargeted -release-hai-agent-runtime.yaml pin PR) and nothing else. Artifacts live under an +PINNED_RUNTIME_VERSION and MANIFEST here and nothing else. Artifacts live under an immutable version-scoped CDN prefix, so an edge can never serve stale bytes. -Cross-ref: eng_plans/14-06-2026-holodesktop-binary-versioning-autoupdate. """ from __future__ import annotations @@ -14,7 +12,7 @@ import sys import typing -# SHIP-GATE: repoint to the plan-005 release before merge +# SHIP-GATE: pin a release that serves the shared recipe before merge PINNED_RUNTIME_VERSION = "0.1.8" RUNTIME_CDN_BASE = "https://assets.hcompanyprod.fr/hai-agent-runtime" # Guard value: published manifest entries must never use it (every download would fail verification). @@ -36,12 +34,12 @@ def _artifact(filename: str, sha256: str) -> RuntimeArtifact: MANIFEST: typing.Dict[str, RuntimeArtifact] = { "darwin-arm64": _artifact( "hai-agent-runtime-darwin-arm64.zip", - # SHIP-GATE: repoint to the plan-005 release before merge + # SHIP-GATE: pin a release that serves the shared recipe before merge "1aed0055898116732aee031dc4a1235782b2909ee51e0367e2d50bb3be6671c9", ), "windows-x86_64": _artifact( "hai-agent-runtime-windows-x86_64.zip", - # SHIP-GATE: repoint to the plan-005 release before merge + # SHIP-GATE: pin a release that serves the shared recipe before merge "4e6b2bcd42af2bb6b22197fcde947327497f5c62fd60d48bc9037730d80dc691", ), } diff --git a/src/hai_agents/local/process.py b/src/hai_agents_local/runtime/process.py similarity index 100% rename from src/hai_agents/local/process.py rename to src/hai_agents_local/runtime/process.py diff --git a/src/hai_agents/local/runtime.py b/src/hai_agents_local/runtime/runtime.py similarity index 95% rename from src/hai_agents/local/runtime.py rename to src/hai_agents_local/runtime/runtime.py index f5ea173..d139e24 100644 --- a/src/hai_agents/local/runtime.py +++ b/src/hai_agents_local/runtime/runtime.py @@ -17,6 +17,8 @@ import httpx +from hai_agents.base_client import BaseClient + from .errors import ( BinaryIncompatibleError, BinaryNotFoundError, @@ -52,6 +54,7 @@ BASE_URL_ENV = "HAI_AGENT_LOCAL_BASE_URL" PORT_ENV = "HAI_AGENT_RUNTIME_PORT" AUTH_TOKEN_ENV = "HAI_AGENT_RUNTIME_API_TOKEN" +CLIENT_TIMEOUT_S = 60.0 _PathInput = typing.Union[str, "os.PathLike[str]"] @@ -312,7 +315,7 @@ def _child_env( binary (without them local sessions cannot run inference) and forwards caller flags such as HAI_AGENT_RUNTIME_MODEL/FAKE/FAST/RUNS_DIR. inherit_env=False takes spawn_env as the complete base environment instead — for callers that must *remove* inherited keys, which an - overlay cannot express (HoloDesktop strips HAI_API_KEY for self-hosted base URLs). The + overlay cannot express (e.g. stripping HAI_API_KEY for self-hosted base URLs). The generated local bearer and the cloud HAI_API_KEY are different credentials: the token below is the only local bearer, and the cloud key is never used to authenticate against the local runtime. Port and token are set last in both modes so caller input never clobbers them. @@ -362,6 +365,14 @@ def _resolve_command( logger.info("resolved hai-agent-runtime from fresh download v%s: %s", pinned, installed) return [str(installed)] + def http_client(self, timeout: typing.Optional[float] = None) -> httpx.Client: + """An HTTP client for this runtime's API.""" + return httpx.Client(timeout=CLIENT_TIMEOUT_S if timeout is None else timeout, follow_redirects=True) + + def async_http_client(self, timeout: typing.Optional[float] = None) -> httpx.AsyncClient: + """An asynchronous HTTP client for this runtime's API.""" + return httpx.AsyncClient(timeout=CLIENT_TIMEOUT_S if timeout is None else timeout, follow_redirects=True) + def health(self) -> typing.Dict[str, typing.Any]: """The /health JSON body; raises RuntimeUnhealthyError when the runtime is not answering.""" payload = probe_health(self.base_url) @@ -388,9 +399,8 @@ def shutdown(self) -> None: def shutdown_if_idle(self) -> bool: """Stop the owned runtime only when it hosts no active sessions; True when it was stopped.""" - from ..client import Client # runtime-time import: client.py imports this module lazily too - - with Client(base_url=self.base_url, api_key=self.api_key, auto_bridges=False) as client: + with self.http_client() as http: + client = BaseClient(base_url=self.base_url, api_key=self.api_key, httpx_client=http) page = client.sessions.list_sessions(status=list(self.ACTIVE_SESSION_STATUSES), size=1) if page.items: return False diff --git a/src/hai_agents/local/state.py b/src/hai_agents_local/runtime/state.py similarity index 95% rename from src/hai_agents/local/state.py rename to src/hai_agents_local/runtime/state.py index f3eae96..ddf80b1 100644 --- a/src/hai_agents/local/state.py +++ b/src/hai_agents_local/runtime/state.py @@ -3,7 +3,7 @@ A spawner persists the generated bearer token and the runtime pid under the SDK cache dir so a second process can attach (token) or force-kill (pid) without any IPC. Both files drive privileged actions, so they are 0600 from the first byte -and refuse pre-planted symlinks (port of holo_desktop launcher._write_owner_only). +and refuse pre-planted symlinks. """ from __future__ import annotations @@ -15,7 +15,7 @@ CACHE_DIR_ENV = "HAI_AGENT_LOCAL_CACHE_DIR" DEFAULT_CACHE_DIR = pathlib.Path.home() / ".hai" / "agent-runtime" -# HoloDesktop's AGENT_API_DEFAULT_PORT: the shared well-known local runtime port. +# The shared well-known local runtime port. DEFAULT_PORT = 18795 _PathInput = typing.Union[str, "os.PathLike[str]"] diff --git a/src/hai_agents_local/sessions.py b/src/hai_agents_local/sessions.py index 455eacb..35786de 100644 --- a/src/hai_agents_local/sessions.py +++ b/src/hai_agents_local/sessions.py @@ -20,6 +20,9 @@ from .manager import ensure_bridges, stop_bridges from .routing import localize_agent +if typing.TYPE_CHECKING: + from .runtime import LocalRuntime + logger = logging.getLogger(__name__) # Runaway guards for sessions whose caller passes no budget: a local session left running (its @@ -193,9 +196,17 @@ def _cancel_sessions_at_exit() -> None: class LocalSessionsClient(SessionsClient): + def __init__( + self, *, client_wrapper: typing.Any, runtime: typing.Optional[LocalRuntime] = None, auto_bridges: bool = True + ) -> None: + super().__init__(client_wrapper=client_wrapper) + self._runtime = runtime + self._auto_bridges = auto_bridges + self._owned_bridges: typing.Dict[str, typing.List[str]] = {} + def close(self) -> None: failures = [] - for session_id in list(getattr(self, "_owned_bridges", {})): + for session_id in list(self._owned_bridges): try: self.cancel_session(session_id) except Exception as error: @@ -205,7 +216,7 @@ def close(self) -> None: def cancel_session(self, session_id: str, **kwargs: typing.Any) -> typing.Any: # Stop local execution even if the remote cancellation cannot be delivered. - owned = getattr(self, "_owned_bridges", {}).get(str(session_id), []) + owned = self._owned_bridges.get(str(session_id), []) try: if owned: stop_bridges(owned) @@ -219,7 +230,7 @@ def cancel_session(self, session_id: str, **kwargs: typing.Any) -> typing.Any: @functools.wraps(SessionsClient.create_session) def create_session(self, **kwargs: typing.Any) -> typing.Any: wrapper = self._raw_client._client_wrapper - bridges = _localize(wrapper, kwargs) + bridges = _localize(wrapper, kwargs) if self._auto_bridges else [] if bridges: _apply_runaway_budgets(kwargs) stop_watcher = _ensure_stop_watcher() if bridges else None @@ -241,16 +252,22 @@ def create_session(self, **kwargs: typing.Any) -> typing.Any: # A stop was filed while bridges or the session were starting; apply it now. _panic_stop() if started: - if not hasattr(self, "_owned_bridges"): - self._owned_bridges = {} self._owned_bridges[str(session.id)] = started return session class LocalAsyncSessionsClient(AsyncSessionsClient): + def __init__( + self, *, client_wrapper: typing.Any, runtime: typing.Optional[LocalRuntime] = None, auto_bridges: bool = True + ) -> None: + super().__init__(client_wrapper=client_wrapper) + self._runtime = runtime + self._auto_bridges = auto_bridges + self._owned_bridges: typing.Dict[str, typing.List[str]] = {} + async def aclose(self) -> None: failures = [] - for session_id in list(getattr(self, "_owned_bridges", {})): + for session_id in list(self._owned_bridges): try: await self.cancel_session(session_id) except Exception as error: @@ -259,7 +276,7 @@ async def aclose(self) -> None: raise RuntimeError("Could not confirm all client-owned sessions stopped") from failures[0] async def cancel_session(self, session_id: str, **kwargs: typing.Any) -> typing.Any: - owned = getattr(self, "_owned_bridges", {}).get(str(session_id), []) + owned = self._owned_bridges.get(str(session_id), []) try: if owned: await asyncio.to_thread(stop_bridges, owned) @@ -272,7 +289,7 @@ async def cancel_session(self, session_id: str, **kwargs: typing.Any) -> typing. @functools.wraps(AsyncSessionsClient.create_session) async def create_session(self, **kwargs: typing.Any) -> typing.Any: wrapper = self._raw_client._client_wrapper - bridges = _localize(wrapper, kwargs) + bridges = _localize(wrapper, kwargs) if self._auto_bridges else [] if bridges: _apply_runaway_budgets(kwargs) # Native permission prompts must run before bridge startup moves to a worker. @@ -298,7 +315,5 @@ async def create_session(self, **kwargs: typing.Any) -> typing.Any: # A stop was filed while bridges or the session were starting; apply it now. await asyncio.to_thread(_panic_stop) if started: - if not hasattr(self, "_owned_bridges"): - self._owned_bridges = {} self._owned_bridges[str(session.id)] = started return session diff --git a/tests/test_runtime_placement.py b/tests/test_runtime_placement.py index 5ba76ed..c16ae88 100644 --- a/tests/test_runtime_placement.py +++ b/tests/test_runtime_placement.py @@ -3,16 +3,60 @@ import httpx import pytest -from hai_agents import AsyncClient, Client, Inference -from hai_agents.local.errors import BinaryIncompatibleError, LocalRuntimeError -from hai_agents.local.runtime import LocalRuntime -from hai_agents.local.state import token_file_path, write_owner_only -from hai_agents.sessions.client import AsyncSessionsClient, SessionsClient +from hai_agents import AsyncClient, Client +from hai_agents_local.runtime import BinaryIncompatibleError, Inference, LocalRuntime, LocalRuntimeError +from hai_agents_local.runtime.state import token_file_path, write_owner_only -def test_both_clients_respect_product_owned_execution(): - assert type(Client(api_key="test", auto_bridges=False).sessions) is SessionsClient - assert type(AsyncClient(api_key="test", auto_bridges=False).sessions) is AsyncSessionsClient +class FakeRuntime: + base_url = "http://127.0.0.1:18795" + api_key = "local-token" + owned = True + + def __init__(self): + self.stopped = False + + def require_recipe(self, recipe): + assert recipe == "shared" + + def http_client(self, timeout=None): + return httpx.Client() + + def async_http_client(self, timeout=None): + return httpx.AsyncClient() + + def shutdown(self): + self.stopped = True + + +@pytest.mark.asyncio +@pytest.mark.parametrize("asynchronous", [False, True]) +async def test_local_client_without_auto_bridges_leaves_execution_to_the_product(monkeypatch, asynchronous): + from hai_agents.sessions.client import AsyncSessionsClient, SessionsClient + + monkeypatch.setenv("HAI_AUTO_BRIDGE", "1") + served = [] + monkeypatch.setattr("hai_agents_local.sessions.ensure_bridges", lambda bridges: served.extend(bridges) or []) + requested = [] + + def create(self, **kwargs): + requested.append(kwargs) + return type("Session", (), {"id": "run"})() + + async def async_create(self, **kwargs): + return create(self, **kwargs) + + monkeypatch.setattr(SessionsClient, "create_session", create) + monkeypatch.setattr(AsyncSessionsClient, "create_session", async_create) + agent = {"name": "qa", "environments": [{"id": "desktop", "kind": "desktop", "host": "user_device"}]} + if asynchronous: + async with await AsyncClient.local(runtime=FakeRuntime(), auto_bridges=False) as client: + await client.sessions.create_session(agent=agent, messages="test") + else: + with Client.local(runtime=FakeRuntime(), auto_bridges=False) as client: + client.sessions.create_session(agent=agent, messages="test") + assert served == [] + assert requested[0]["agent"] == agent def test_self_hosted_inference_does_not_receive_hosted_key(monkeypatch): @@ -26,7 +70,7 @@ def test_self_hosted_inference_does_not_receive_hosted_key(monkeypatch): @pytest.mark.parametrize("probe_error", [None, httpx.ConnectError("offline"), httpx.ReadTimeout("timeout")]) def test_attachment_authenticates_and_checks_recipe_before_use(tmp_path, monkeypatch, probe_error): - from hai_agents.local import runtime as module + from hai_agents_local.runtime import runtime as module write_owner_only(token_file_path(18795, cache_dir=tmp_path), "local-token") monkeypatch.delenv("HAI_AGENT_RUNTIME_API_TOKEN", raising=False) @@ -66,41 +110,23 @@ def test_local_attach_rejects_remote_or_credential_urls(tmp_path, monkeypatch, u @pytest.mark.asyncio @pytest.mark.parametrize("asynchronous", [False, True]) @pytest.mark.parametrize("borrowed", [False, True]) -async def test_client_close_releases_only_owned_runtime_and_http(monkeypatch, asynchronous, borrowed): - class Runtime: - base_url = "http://127.0.0.1:18795" - api_key = "local-token" - owned = True - stopped = False - - def require_recipe(self, recipe): - assert recipe == "shared" - - def shutdown(self): - self.stopped = True - - runtime = Runtime() +async def test_client_close_releases_only_owned_runtime(monkeypatch, asynchronous, borrowed): + runtime = FakeRuntime() monkeypatch.setattr(LocalRuntime, "ensure_started", lambda **options: runtime) - http = httpx.AsyncClient() if asynchronous else httpx.Client() - client_type = AsyncClient if asynchronous else Client - options = {"runtime": runtime, "httpx_client": http} if borrowed else {} - client = await AsyncClient.local() if asynchronous and not borrowed else client_type(mode="local", **options) + options = {"runtime": runtime} if borrowed else {} if asynchronous: + client = await AsyncClient.local(**options) await client.aclose() else: + client = Client.local(**options) client.close() assert runtime.stopped is (not borrowed) - if borrowed: - assert not http.is_closed - if asynchronous: - await http.aclose() - else: - http.close() + assert client._client_wrapper.httpx_client.httpx_client.is_closed def test_binary_resolution_does_not_hold_the_port_startup_lock(tmp_path, monkeypatch): - from hai_agents.local import runtime as module - from hai_agents.local.errors import BinaryNotFoundError + from hai_agents_local.runtime import BinaryNotFoundError + from hai_agents_local.runtime import runtime as module monkeypatch.delenv("HAI_AGENT_LOCAL_BASE_URL", raising=False) monkeypatch.setattr(LocalRuntime, "_attach", lambda **kwargs: None) @@ -190,7 +216,7 @@ def test_state_cleanup_waits_for_startup_and_preserves_replacement(tmp_path): import threading from concurrent.futures import ThreadPoolExecutor - from hai_agents.local import runtime as module + from hai_agents_local.runtime import runtime as module token_file = write_owner_only(token_file_path(18795, cache_dir=tmp_path), "old-token") runtime = LocalRuntime( @@ -223,25 +249,15 @@ def cleanup(): @pytest.mark.asyncio @pytest.mark.parametrize("asynchronous", [False, True]) -@pytest.mark.parametrize("borrowed", [False, True]) -async def test_failed_client_recipe_check_releases_only_started_runtime(monkeypatch, asynchronous, borrowed): - stopped = [] - - class Runtime: - owned = True - +async def test_attached_runtime_with_another_recipe_is_rejected_and_left_running(asynchronous): + class Runtime(FakeRuntime): def require_recipe(self, recipe): raise BinaryIncompatibleError("recipe changed") - def shutdown(self): - stopped.append(True) - runtime = Runtime() - monkeypatch.setattr(LocalRuntime, "ensure_started", lambda **options: runtime) with pytest.raises(BinaryIncompatibleError, match="recipe changed"): - if asynchronous and not borrowed: - await AsyncClient.local() + if asynchronous: + await AsyncClient.local(runtime=runtime) else: - client_type = AsyncClient if asynchronous else Client - client_type(mode="local", **({"runtime": runtime} if borrowed else {})) - assert stopped == ([] if borrowed else [True]) + Client.local(runtime=runtime) + assert not runtime.stopped diff --git a/tests/test_runtime_state.py b/tests/test_runtime_state.py index 9adc2cd..523c504 100644 --- a/tests/test_runtime_state.py +++ b/tests/test_runtime_state.py @@ -15,7 +15,7 @@ import pytest -from hai_agents.local.state import pid_file_path, write_owner_only +from hai_agents_local.runtime.state import pid_file_path, write_owner_only @pytest.mark.skipif(sys.platform == "win32", reason="POSIX file-mode semantics") From 0f5f9a4efce9cb5e4500fbcf62aea5310f179046 Mon Sep 17 00:00:00 2001 From: abonneth Date: Thu, 1 Oct 2026 18:10:53 +0200 Subject: [PATCH 14/35] fix(sdk): keep loopback runtime traffic off environment proxies HTTP(S)_PROXY no longer receives the local runtime bearer token. --- src/hai_agents_local/browser.py | 2 +- src/hai_agents_local/runtime/process.py | 2 +- src/hai_agents_local/runtime/runtime.py | 10 ++++++++-- 3 files changed, 10 insertions(+), 4 deletions(-) diff --git a/src/hai_agents_local/browser.py b/src/hai_agents_local/browser.py index 085ec9a..e522588 100644 --- a/src/hai_agents_local/browser.py +++ b/src/hai_agents_local/browser.py @@ -117,6 +117,6 @@ def _ensure_local_chrome(self) -> None: def _debugger_listening(port: int) -> bool: try: - return httpx.get(f"http://127.0.0.1:{port}/json/version", timeout=2.0).status_code == 200 + return httpx.get(f"http://127.0.0.1:{port}/json/version", timeout=2.0, trust_env=False).status_code == 200 except httpx.HTTPError: return False diff --git a/src/hai_agents_local/runtime/process.py b/src/hai_agents_local/runtime/process.py index 50b45bb..c2953ea 100644 --- a/src/hai_agents_local/runtime/process.py +++ b/src/hai_agents_local/runtime/process.py @@ -27,7 +27,7 @@ def probe_health(base_url: str) -> typing.Optional[typing.Dict[str, typing.Any]]: """The /health JSON body on a 200 ({} for non-JSON bodies); None when unreachable/unhealthy.""" try: - response = httpx.get(f"{base_url}/health", timeout=2.0) + response = httpx.get(f"{base_url}/health", timeout=2.0, trust_env=False) except httpx.HTTPError: return None if response.status_code != 200: diff --git a/src/hai_agents_local/runtime/runtime.py b/src/hai_agents_local/runtime/runtime.py index d139e24..df5a831 100644 --- a/src/hai_agents_local/runtime/runtime.py +++ b/src/hai_agents_local/runtime/runtime.py @@ -186,6 +186,7 @@ def ensure_started( params={"size": 1}, timeout=2.0, follow_redirects=False, + trust_env=False, ) if required_recipe is not None and payload.get("recipe") != required_recipe: raise BinaryIncompatibleError( @@ -281,6 +282,7 @@ def _attach(cls, *, base_url: str, cache_dir: pathlib.Path) -> typing.Optional[" params={"size": 1}, timeout=2.0, follow_redirects=False, + trust_env=False, ) except httpx.HTTPError as exc: raise LocalRuntimeError("runtime attachment failed authenticated session probe") from exc @@ -367,11 +369,15 @@ def _resolve_command( def http_client(self, timeout: typing.Optional[float] = None) -> httpx.Client: """An HTTP client for this runtime's API.""" - return httpx.Client(timeout=CLIENT_TIMEOUT_S if timeout is None else timeout, follow_redirects=True) + return httpx.Client( + timeout=CLIENT_TIMEOUT_S if timeout is None else timeout, follow_redirects=True, trust_env=False + ) def async_http_client(self, timeout: typing.Optional[float] = None) -> httpx.AsyncClient: """An asynchronous HTTP client for this runtime's API.""" - return httpx.AsyncClient(timeout=CLIENT_TIMEOUT_S if timeout is None else timeout, follow_redirects=True) + return httpx.AsyncClient( + timeout=CLIENT_TIMEOUT_S if timeout is None else timeout, follow_redirects=True, trust_env=False + ) def health(self) -> typing.Dict[str, typing.Any]: """The /health JSON body; raises RuntimeUnhealthyError when the runtime is not answering.""" From 547d2a3357ac9e1aa6c79843fb2ac0116d350b9b Mon Sep 17 00:00:00 2001 From: abonneth Date: Thu, 1 Oct 2026 18:17:15 +0200 Subject: [PATCH 15/35] fix(sdk): require the local runtime to prove its identity Every request to a local runtime carries a fresh X-Hai-Runtime-Challenge, and every response must carry X-Hai-Runtime-Proof, the HMAC-SHA256 of the challenge keyed by the runtime token. Unproven responses raise LocalRuntimeError, so a server squatting the runtime port never receives the bearer token and its commands never reach a device bridge. /health is proven before any bearer-authenticated request, both when spawning and when attaching. --- src/hai_agents_local/bridge.py | 16 +- src/hai_agents_local/browser.py | 5 +- src/hai_agents_local/desktop.py | 5 +- src/hai_agents_local/routing.py | 20 ++- src/hai_agents_local/runtime/identity.py | 60 ++++++++ src/hai_agents_local/runtime/process.py | 22 ++- src/hai_agents_local/runtime/runtime.py | 62 ++++---- src/hai_agents_local/sessions.py | 57 +++++--- src/hai_agents_local/workstation.py | 5 +- tests/test_runtime_placement.py | 178 +++++++++++++++-------- 10 files changed, 305 insertions(+), 125 deletions(-) create mode 100644 src/hai_agents_local/runtime/identity.py diff --git a/src/hai_agents_local/bridge.py b/src/hai_agents_local/bridge.py index a764106..96f37fa 100644 --- a/src/hai_agents_local/bridge.py +++ b/src/hai_agents_local/bridge.py @@ -18,6 +18,7 @@ from .config import default_base_url from .errors import RateLimitedError, SessionNotFoundError +from .runtime import identity from .transport import Command, CommandExchange, Json, deserialize_args, serialize_result logger = logging.getLogger(__name__) @@ -69,9 +70,12 @@ def __init__( api_key: TokenSource, base_url: str | None = None, session_id: str | None = None, + verify_runtime: bool = False, ) -> None: if not api_key: raise ValueError("api_key is required") + if verify_runtime and not isinstance(api_key, str): + raise ValueError("verify_runtime needs api_key to be the local runtime's token string") if session_id is not None: try: uuid.UUID(session_id) @@ -81,6 +85,7 @@ def __init__( self.api_key = api_key self.base_url = base_url or default_base_url() self.session_id = session_id or str(uuid.uuid4()) + self.verify_runtime = verify_runtime self.ready = threading.Event() self.on_crash: Callable[[], None] | None = None self._driver: DriverT | None = None @@ -117,9 +122,16 @@ async def run(self) -> None: """Serve commands until stopped; raises AuthError on a bad key.""" # An asyncio.Event binds to the loop it is first awaited on; a restarted bridge runs on a new loop. self._stop_event = asyncio.Event() + options: dict[str, Any] = { + "headers": {"Accept": "application/json"}, + "auth": _BearerAuth(self.api_key), + "follow_redirects": True, + } try: - async with httpx.AsyncClient( - headers={"Accept": "application/json"}, auth=_BearerAuth(self.api_key), follow_redirects=True + async with ( + identity.async_http_client(self.api_key, **options) + if self.verify_runtime + else httpx.AsyncClient(**options) ) as client: exchange = CommandExchange(client, self.base_url) if not await self._open_channel(exchange): diff --git a/src/hai_agents_local/browser.py b/src/hai_agents_local/browser.py index e522588..ca43a1a 100644 --- a/src/hai_agents_local/browser.py +++ b/src/hai_agents_local/browser.py @@ -54,8 +54,11 @@ def __init__( api_key: TokenSource, base_url: str | None = None, session_id: str | None = None, + verify_runtime: bool = False, ) -> None: - super().__init__(environment_id, api_key=api_key, base_url=base_url, session_id=session_id) + super().__init__( + environment_id, api_key=api_key, base_url=base_url, session_id=session_id, verify_runtime=verify_runtime + ) self.debugging_port = debugging_port def create_driver(self) -> SeleniumWebDriver: diff --git a/src/hai_agents_local/desktop.py b/src/hai_agents_local/desktop.py index 0339b98..fc4b1a8 100644 --- a/src/hai_agents_local/desktop.py +++ b/src/hai_agents_local/desktop.py @@ -62,12 +62,15 @@ def __init__( api_key: TokenSource, base_url: str | None = None, session_id: str | None = None, + verify_runtime: bool = False, max_width: int | None = DEFAULT_MAX_WIDTH, max_height: int | None = None, image_format: ImageFormat | None = DEFAULT_IMAGE_FORMAT, quality: int = DEFAULT_QUALITY, ) -> None: - super().__init__(environment_id, api_key=api_key, base_url=base_url, session_id=session_id) + super().__init__( + environment_id, api_key=api_key, base_url=base_url, session_id=session_id, verify_runtime=verify_runtime + ) self.max_width = max_width self.max_height = max_height self.image_format = image_format diff --git a/src/hai_agents_local/routing.py b/src/hai_agents_local/routing.py index 4c6f8f3..9056f1e 100644 --- a/src/hai_agents_local/routing.py +++ b/src/hai_agents_local/routing.py @@ -25,36 +25,42 @@ def localize_agent( - agent: AgentLike, *, api_key: TokenSource, base_url: str | None = None + agent: AgentLike, *, api_key: TokenSource, base_url: str | None = None, verify_runtime: bool = False ) -> tuple[AgentLike, list[LocalBridge]]: """Copy of the agent where every unclaimed user_device environment is stamped with the session id of a freshly built bridge, plus those bridges. Environments that already carry a session_id are assumed to be served elsewhere and left alone, as are string agent references.""" bridges: list[LocalBridge] = [] - return _localize_agent(agent, bridges, api_key, base_url), bridges + return _localize_agent(agent, bridges, api_key, base_url, verify_runtime), bridges def _localize_agent( - agent: AgentLike, bridges: list[LocalBridge], api_key: TokenSource, base_url: str | None + agent: AgentLike, bridges: list[LocalBridge], api_key: TokenSource, base_url: str | None, verify_runtime: bool ) -> AgentLike: if isinstance(agent, str): return agent changes: dict[str, Any] = {} environments = _read(agent, "environments") if isinstance(environments, (list, tuple)): - localized_envs = [_localize_environment(env, bridges, api_key, base_url) for env in environments] + localized_envs = [ + _localize_environment(env, bridges, api_key, base_url, verify_runtime) for env in environments + ] if _any_replaced(localized_envs, environments): changes["environments"] = localized_envs subagents = _read(agent, "subagents") if isinstance(subagents, (list, tuple)): - localized_subs = [_localize_agent(sub, bridges, api_key, base_url) for sub in subagents] + localized_subs = [_localize_agent(sub, bridges, api_key, base_url, verify_runtime) for sub in subagents] if _any_replaced(localized_subs, subagents): changes["subagents"] = localized_subs return _replace(agent, **changes) if changes else agent def _localize_environment( - env: EnvironmentLike, bridges: list[LocalBridge], api_key: TokenSource, base_url: str | None + env: EnvironmentLike, + bridges: list[LocalBridge], + api_key: TokenSource, + base_url: str | None, + verify_runtime: bool, ) -> EnvironmentLike: kind = _local_kind(env) if kind is None or _read(env, "session_id"): @@ -65,7 +71,7 @@ def _localize_environment( f"serve one local {kind}; give the extra environments an explicit session_id and serve each " "from its own machine with `hai local browser|desktop|workstation --session-id `" ) - bridge = BRIDGE_TYPES[kind](_read(env, "id"), api_key=api_key, base_url=base_url) + bridge = BRIDGE_TYPES[kind](_read(env, "id"), api_key=api_key, base_url=base_url, verify_runtime=verify_runtime) bridges.append(bridge) return _replace(env, kind=kind, session_id=bridge.session_id) diff --git a/src/hai_agents_local/runtime/identity.py b/src/hai_agents_local/runtime/identity.py new file mode 100644 index 0000000..5aa6b0c --- /dev/null +++ b/src/hai_agents_local/runtime/identity.py @@ -0,0 +1,60 @@ +"""Runtime identity: each request carries a fresh challenge the runtime answers with an HMAC keyed by its token.""" + +from __future__ import annotations + +import hashlib +import hmac +import secrets +import typing + +import httpx + +from .errors import LocalRuntimeError + +CHALLENGE_HEADER = "X-Hai-Runtime-Challenge" +PROOF_HEADER = "X-Hai-Runtime-Proof" + + +def challenge() -> str: + return secrets.token_urlsafe(24) + + +def verify(token: str, request: httpx.Request, response: httpx.Response) -> None: + """Raise unless ``response`` proves its server holds ``token`` for the challenge ``request`` sent.""" + sent = request.headers.get(CHALLENGE_HEADER, "") + expected = hmac.new(token.encode(), sent.encode(), hashlib.sha256).hexdigest() + received = response.headers.get(PROOF_HEADER, "") + if not sent or not hmac.compare_digest(received.encode(), expected.encode()): + raise LocalRuntimeError( + f"{request.url.scheme}://{request.url.netloc.decode()} is not the runtime this client started or attached to" + ) + + +def event_hooks(token: str) -> typing.Dict[str, typing.List[typing.Callable[..., None]]]: + def send_challenge(request: httpx.Request) -> None: + request.headers[CHALLENGE_HEADER] = challenge() + + def check_proof(response: httpx.Response) -> None: + verify(token, response.request, response) + + return {"request": [send_challenge], "response": [check_proof]} + + +def async_event_hooks(token: str) -> typing.Dict[str, typing.List[typing.Callable[..., typing.Awaitable[None]]]]: + async def send_challenge(request: httpx.Request) -> None: + request.headers[CHALLENGE_HEADER] = challenge() + + async def check_proof(response: httpx.Response) -> None: + verify(token, response.request, response) + + return {"request": [send_challenge], "response": [check_proof]} + + +def http_client(token: str, **kwargs: typing.Any) -> httpx.Client: + """A proxy-free client that rejects any response not proven by the runtime holding ``token``.""" + return httpx.Client(trust_env=False, event_hooks=event_hooks(token), **kwargs) + + +def async_http_client(token: str, **kwargs: typing.Any) -> httpx.AsyncClient: + """``http_client`` for asyncio callers.""" + return httpx.AsyncClient(trust_env=False, event_hooks=async_event_hooks(token), **kwargs) diff --git a/src/hai_agents_local/runtime/process.py b/src/hai_agents_local/runtime/process.py index c2953ea..650ff41 100644 --- a/src/hai_agents_local/runtime/process.py +++ b/src/hai_agents_local/runtime/process.py @@ -13,6 +13,7 @@ import httpx +from . import identity from .errors import RuntimeStartTimeoutError, RuntimeUnhealthyError logger = logging.getLogger(__name__) @@ -24,10 +25,20 @@ LOG_TAIL_CHARS = 4000 -def probe_health(base_url: str) -> typing.Optional[typing.Dict[str, typing.Any]]: - """The /health JSON body on a 200 ({} for non-JSON bodies); None when unreachable/unhealthy.""" +def responds(base_url: str) -> bool: + """Whether any HTTP server answers at `base_url`, proven or not.""" try: - response = httpx.get(f"{base_url}/health", timeout=2.0, trust_env=False) + httpx.get(f"{base_url}/health", timeout=2.0, trust_env=False) + except httpx.HTTPError: + return False + return True + + +def probe_health(base_url: str, token: str) -> typing.Optional[typing.Dict[str, typing.Any]]: + """The proven /health JSON body on a 200 ({} for non-JSON bodies); None when unreachable/unhealthy.""" + try: + with identity.http_client(token, timeout=2.0) as client: + response = client.get(f"{base_url}/health") except httpx.HTTPError: return None if response.status_code != 200: @@ -63,16 +74,17 @@ def wait_healthy( base_url: str, proc: subprocess.Popen, *, + token: str, timeout_s: float, log_path: pathlib.Path, cancel_event: typing.Optional[threading.Event] = None, ) -> typing.Dict[str, typing.Any]: - """Poll /health until 200; raises RuntimeUnhealthyError (child exited) or RuntimeStartTimeoutError.""" + """Poll /health until a proven 200; raises LocalRuntimeError if another server answers, or the child fails.""" deadline = time.monotonic() + timeout_s while True: if cancel_event is not None and cancel_event.is_set(): raise RuntimeUnhealthyError("runtime startup cancelled") - payload = probe_health(base_url) + payload = probe_health(base_url, token) if payload is not None: logger.info("hai-agent-runtime ready (pid %d)", proc.pid) return payload diff --git a/src/hai_agents_local/runtime/runtime.py b/src/hai_agents_local/runtime/runtime.py index df5a831..515fdeb 100644 --- a/src/hai_agents_local/runtime/runtime.py +++ b/src/hai_agents_local/runtime/runtime.py @@ -19,6 +19,7 @@ from hai_agents.base_client import BaseClient +from . import identity from .errors import ( BinaryIncompatibleError, BinaryNotFoundError, @@ -31,6 +32,7 @@ LOOPBACK_HOST, SPAWN_TIMEOUT_S, probe_health, + responds, spawn, terminate, wait_healthy, @@ -74,6 +76,18 @@ def _port_of(base_url: str) -> int: return urlsplit(base_url).port or DEFAULT_PORT +def _authenticated_probe(base_url: str, token: str) -> int: + """Status of a proven, bearer-authenticated session listing; call only after /health proved the server.""" + with identity.http_client(token, timeout=2.0) as client: + response = client.get( + f"{base_url}/api/v2/sessions", + headers={"Authorization": f"Bearer {token}"}, + params={"size": 1}, + follow_redirects=False, + ) + return response.status_code + + class LocalRuntime: """A reachable local agent runtime: where it is, how to authenticate, and (if ours) the process.""" @@ -178,21 +192,14 @@ def ensure_started( log_path=log_path, ) payload = wait_healthy( - base_url, proc, timeout_s=timeout_s, log_path=log_path, cancel_event=_cancel_event - ) - response = httpx.get( - f"{base_url}/api/v2/sessions", - headers={"Authorization": f"Bearer {token}"}, - params={"size": 1}, - timeout=2.0, - follow_redirects=False, - trust_env=False, + base_url, proc, token=token, timeout_s=timeout_s, log_path=log_path, cancel_event=_cancel_event ) + status = _authenticated_probe(base_url, token) if required_recipe is not None and payload.get("recipe") != required_recipe: raise BinaryIncompatibleError( f"runtime must support recipe {required_recipe!r}; use a compatible source command or binary" ) - if response.status_code != 200 or proc.poll() is not None: + if status != 200 or proc.poll() is not None: raise RuntimeUnhealthyError("spawned runtime failed authenticated readiness probe") pid_file = write_owner_only(pid_file_path(resolved_port, cache_dir=resolved_cache), str(proc.pid)) except BaseException: @@ -239,7 +246,7 @@ def require_recipe(self, recipe: typing.Optional[str]) -> None: """Fail closed when an old binary or a differently configured daemon answers.""" if recipe is None: return - payload = probe_health(self.base_url) + payload = probe_health(self.base_url, self.api_key) if payload is None or payload.get("recipe") != recipe: raise BinaryIncompatibleError( f"runtime must support recipe {recipe!r}; use a compatible source command or binary" @@ -263,30 +270,25 @@ def _attach(cls, *, base_url: str, cache_dir: pathlib.Path) -> typing.Optional[" if parsed.scheme != "http" or parsed.hostname not in {"127.0.0.1", "localhost", "::1"} or parsed.username: raise LocalRuntimeError("local runtime attachment requires a loopback HTTP URL") port = _port_of(base_url) - payload = probe_health(base_url) - if payload is None: - return None token = os.environ.get(AUTH_TOKEN_ENV, "").strip() or read_state_file( token_file_path(port, cache_dir=cache_dir) ) if not token: + if not responds(base_url): + return None raise LocalRuntimeError( f"an agent runtime is answering at {base_url} but no credentials were found: " f"{AUTH_TOKEN_ENV} is not set and {token_file_path(port, cache_dir=cache_dir)} does not exist, " "so this client cannot authenticate. Export the token or stop that runtime." ) + payload = probe_health(base_url, token) + if payload is None: + return None try: - response = httpx.get( - f"{base_url}/api/v2/sessions", - headers={"Authorization": f"Bearer {token}"}, - params={"size": 1}, - timeout=2.0, - follow_redirects=False, - trust_env=False, - ) + status = _authenticated_probe(base_url, token) except httpx.HTTPError as exc: raise LocalRuntimeError("runtime attachment failed authenticated session probe") from exc - if response.status_code != 200: + if status != 200: raise LocalRuntimeError("runtime attachment failed authenticated session probe") reported = payload.get("version") reported_version = reported if isinstance(reported, str) else None @@ -368,20 +370,20 @@ def _resolve_command( return [str(installed)] def http_client(self, timeout: typing.Optional[float] = None) -> httpx.Client: - """An HTTP client for this runtime's API.""" - return httpx.Client( - timeout=CLIENT_TIMEOUT_S if timeout is None else timeout, follow_redirects=True, trust_env=False + """An HTTP client for this runtime's API that rejects responses this runtime did not prove.""" + return identity.http_client( + self.api_key, timeout=CLIENT_TIMEOUT_S if timeout is None else timeout, follow_redirects=True ) def async_http_client(self, timeout: typing.Optional[float] = None) -> httpx.AsyncClient: - """An asynchronous HTTP client for this runtime's API.""" - return httpx.AsyncClient( - timeout=CLIENT_TIMEOUT_S if timeout is None else timeout, follow_redirects=True, trust_env=False + """``http_client`` for asyncio callers.""" + return identity.async_http_client( + self.api_key, timeout=CLIENT_TIMEOUT_S if timeout is None else timeout, follow_redirects=True ) def health(self) -> typing.Dict[str, typing.Any]: """The /health JSON body; raises RuntimeUnhealthyError when the runtime is not answering.""" - payload = probe_health(self.base_url) + payload = probe_health(self.base_url, self.api_key) if payload is None: raise RuntimeUnhealthyError(f"hai-agent-runtime at {self.base_url} is not answering /health") return payload diff --git a/src/hai_agents_local/sessions.py b/src/hai_agents_local/sessions.py index 35786de..43befbb 100644 --- a/src/hai_agents_local/sessions.py +++ b/src/hai_agents_local/sessions.py @@ -12,6 +12,9 @@ import threading import typing +import httpx + +from hai_agents.base_client import BaseClient from hai_agents.sessions.client import AsyncSessionsClient, SessionsClient from .bridge import LocalBridge, TokenSource @@ -30,6 +33,7 @@ # hits these. Explicit values, including None for unbounded, are respected. DEFAULT_LOCAL_MAX_STEPS = 150 DEFAULT_LOCAL_MAX_TIME_S = 1800.0 +REMOTE_CANCEL_TIMEOUT_S = 60.0 def _apply_runaway_budgets(kwargs: typing.Dict[str, typing.Any]) -> None: @@ -68,7 +72,9 @@ def _warn_if_overrides_target_user_device(kwargs: typing.Dict[str, typing.Any]) ) -def _localize(client_wrapper: typing.Any, kwargs: typing.Dict[str, typing.Any]) -> typing.List[LocalBridge]: +def _localize( + client_wrapper: typing.Any, runtime: typing.Optional[LocalRuntime], kwargs: typing.Dict[str, typing.Any] +) -> typing.List[LocalBridge]: """Spawn bridges for unclaimed user_device environments in an inline agent and stamp their session ids. String agent references are left alone: registered agents must carry an explicit session_id on their @@ -78,9 +84,12 @@ def _localize(client_wrapper: typing.Any, kwargs: typing.Dict[str, typing.Any]) if agent is None or isinstance(agent, str) or not auto_bridges_enabled(): return [] _warn_if_overrides_target_user_device(kwargs) - localized, bridges = localize_agent( - agent, api_key=_token_source(client_wrapper), base_url=client_wrapper.get_base_url() + credentials: typing.Dict[str, typing.Any] = ( + {"api_key": runtime.api_key, "base_url": runtime.base_url, "verify_runtime": True} + if runtime is not None + else {"api_key": _token_source(client_wrapper), "base_url": client_wrapper.get_base_url()} ) + localized, bridges = localize_agent(agent, **credentials) kwargs["agent"] = localized return bridges @@ -114,8 +123,11 @@ def attach(self, action: typing.Callable[[], None]) -> bool: return False +RemoteCancel = typing.Callable[[str], None] + + def _cancel_action( - client_wrapper: typing.Any, bridges: typing.Sequence[LocalBridge], session: typing.Any + cancel_remote: RemoteCancel, bridges: typing.Sequence[LocalBridge], session: typing.Any ) -> typing.Optional[typing.Callable[[], None]]: """A bridge that dies mid-session leaves the agent without local control; cancel the session then.""" session_id = getattr(session, "id", None) @@ -126,7 +138,7 @@ def cancel() -> None: logger.error("local bridge for session %s crashed; cancelling the session", session_id) _deregister_exit_cancel(session_id) try: - _cancel_remote_session(client_wrapper, session_id) + cancel_remote(session_id) except Exception: logger.exception("failed to cancel session %s after its local bridge crashed", session_id) finally: @@ -135,12 +147,15 @@ def cancel() -> None: return cancel -def _cancel_remote_session(client_wrapper: typing.Any, session_id: str) -> None: - from hai_agents.client import Client - - api_key = _resolve_token(_token_source(client_wrapper)) - client = Client(api_key=api_key, base_url=client_wrapper.get_base_url()) - client.sessions.cancel_session(session_id) +def _cancel_remote_session(client_wrapper: typing.Any, runtime: typing.Optional[LocalRuntime], session_id: str) -> None: + """Cancel over a fresh connection, so it works from any thread and at interpreter exit.""" + if runtime is not None: + http, base_url, api_key = runtime.http_client(), runtime.base_url, runtime.api_key + else: + http = httpx.Client(timeout=REMOTE_CANCEL_TIMEOUT_S, follow_redirects=True) + base_url, api_key = client_wrapper.get_base_url(), _resolve_token(_token_source(client_wrapper)) + with http: + BaseClient(base_url=base_url, api_key=api_key, httpx_client=http).sessions.cancel_session(session_id) # Sessions that depend on this process's bridges, cancelled at interpreter exit: the bridges die @@ -159,11 +174,11 @@ def _ensure_stop_watcher() -> StopWatcher: return _stop_watcher -def _register_exit_cancel(client_wrapper: typing.Any, session_id: str) -> None: +def _register_exit_cancel(cancel_remote: RemoteCancel, session_id: str) -> None: def cancel_quietly() -> None: # Best effort: the session may have finished long ago; the platform rejects the cancel then. try: - _cancel_remote_session(client_wrapper, session_id) + cancel_remote(session_id) logger.info("cancelled session %s at exit: its local bridge lives in this process", session_id) except Exception as exc: logger.debug("exit-time cancel of session %s skipped: %s", session_id, exc) @@ -202,6 +217,7 @@ def __init__( super().__init__(client_wrapper=client_wrapper) self._runtime = runtime self._auto_bridges = auto_bridges + self._cancel_remote: RemoteCancel = functools.partial(_cancel_remote_session, client_wrapper, runtime) self._owned_bridges: typing.Dict[str, typing.List[str]] = {} def close(self) -> None: @@ -229,8 +245,7 @@ def cancel_session(self, session_id: str, **kwargs: typing.Any) -> typing.Any: @functools.wraps(SessionsClient.create_session) def create_session(self, **kwargs: typing.Any) -> typing.Any: - wrapper = self._raw_client._client_wrapper - bridges = _localize(wrapper, kwargs) if self._auto_bridges else [] + bridges = _localize(self._raw_client._client_wrapper, self._runtime, kwargs) if self._auto_bridges else [] if bridges: _apply_runaway_budgets(kwargs) stop_watcher = _ensure_stop_watcher() if bridges else None @@ -242,12 +257,12 @@ def create_session(self, **kwargs: typing.Any) -> typing.Any: stop_bridges(started) raise if bridges: - cancel = _cancel_action(wrapper, bridges, session) + cancel = _cancel_action(self._cancel_remote, bridges, session) if cancel is not None: if watcher.attach(cancel): cancel() else: - _register_exit_cancel(wrapper, session.id) + _register_exit_cancel(self._cancel_remote, session.id) if stop_watcher is not None and not stop_watcher.active: # A stop was filed while bridges or the session were starting; apply it now. _panic_stop() @@ -263,6 +278,7 @@ def __init__( super().__init__(client_wrapper=client_wrapper) self._runtime = runtime self._auto_bridges = auto_bridges + self._cancel_remote: RemoteCancel = functools.partial(_cancel_remote_session, client_wrapper, runtime) self._owned_bridges: typing.Dict[str, typing.List[str]] = {} async def aclose(self) -> None: @@ -288,8 +304,7 @@ async def cancel_session(self, session_id: str, **kwargs: typing.Any) -> typing. @functools.wraps(AsyncSessionsClient.create_session) async def create_session(self, **kwargs: typing.Any) -> typing.Any: - wrapper = self._raw_client._client_wrapper - bridges = _localize(wrapper, kwargs) if self._auto_bridges else [] + bridges = _localize(self._raw_client._client_wrapper, self._runtime, kwargs) if self._auto_bridges else [] if bridges: _apply_runaway_budgets(kwargs) # Native permission prompts must run before bridge startup moves to a worker. @@ -304,13 +319,13 @@ async def create_session(self, **kwargs: typing.Any) -> typing.Any: await asyncio.to_thread(stop_bridges, started) raise if bridges: - cancel = _cancel_action(wrapper, bridges, session) + cancel = _cancel_action(self._cancel_remote, bridges, session) if cancel is not None: if watcher.attach(cancel): # cancel_session blocks on HTTP; keep it off the event loop thread. await asyncio.to_thread(cancel) else: - _register_exit_cancel(wrapper, session.id) + _register_exit_cancel(self._cancel_remote, session.id) if stop_watcher is not None and not stop_watcher.active: # A stop was filed while bridges or the session were starting; apply it now. await asyncio.to_thread(_panic_stop) diff --git a/src/hai_agents_local/workstation.py b/src/hai_agents_local/workstation.py index f2724be..ed079a9 100644 --- a/src/hai_agents_local/workstation.py +++ b/src/hai_agents_local/workstation.py @@ -36,8 +36,11 @@ def __init__( api_key: TokenSource, base_url: str | None = None, session_id: str | None = None, + verify_runtime: bool = False, ) -> None: - super().__init__(environment_id, api_key=api_key, base_url=base_url, session_id=session_id) + super().__init__( + environment_id, api_key=api_key, base_url=base_url, session_id=session_id, verify_runtime=verify_runtime + ) self.workspace = Path(workspace).expanduser() if workspace else Path.home() / "hai" / self.session_id def preflight(self) -> None: diff --git a/tests/test_runtime_placement.py b/tests/test_runtime_placement.py index c16ae88..98b3d12 100644 --- a/tests/test_runtime_placement.py +++ b/tests/test_runtime_placement.py @@ -1,5 +1,12 @@ """Placement boundaries: authentication, compatibility and executor ownership.""" +import hashlib +import hmac +import json +import threading +from http.server import BaseHTTPRequestHandler, ThreadingHTTPServer +from urllib.parse import urlsplit + import httpx import pytest @@ -8,6 +15,77 @@ from hai_agents_local.runtime.state import token_file_path, write_owner_only +class RuntimeServer(ThreadingHTTPServer): + """A loopback stand-in for the runtime: answers every response with an HMAC proof keyed by `proof_token`.""" + + def __init__(self, token): + super().__init__(("127.0.0.1", 0), _RuntimeHandler) + self.token = token + self.proof_token = token + self.requests = [] + self.active = [] + self.cancelled = [] + + @property + def port(self): + return self.server_address[1] + + +class _RuntimeHandler(BaseHTTPRequestHandler): + def log_message(self, *args): + pass + + def do_GET(self): + self._route() + + def do_DELETE(self): + self._route() + + def _route(self): + server = self.server + server.requests.append({name.lower() for name in self.headers}) + path = urlsplit(self.path).path + if path == "/health": + return self._reply(200, {"recipe": "shared", "version": "test"}) + if self.headers.get("Authorization") != f"Bearer {server.token}": + return self._reply(401, {"error": "unauthorized"}) + if self.command == "GET" and path == "/api/v2/sessions": + items = [{"id": sid, "status": "running", "created_at": "2026-01-01T00:00:00Z"} for sid in server.active] + return self._reply(200, {"items": items, "total": len(items), "page": 1}) + if self.command == "DELETE" and path.startswith("/api/v2/sessions/"): + session_id = path.rsplit("/", 1)[1] + server.cancelled.append(session_id) + if session_id not in server.active: + return self._reply(404, {"detail": "Session not found"}) + server.active.remove(session_id) + return self._reply(204, None) + self._reply(404, {"detail": "Not Found"}) + + def _reply(self, status, body): + payload = b"" if body is None else json.dumps(body).encode() + self.send_response(status) + challenge = self.headers.get("X-Hai-Runtime-Challenge") + if challenge and self.server.proof_token: + proof = hmac.new(self.server.proof_token.encode(), challenge.encode(), hashlib.sha256).hexdigest() + self.send_header("X-Hai-Runtime-Proof", proof) + self.send_header("Content-Type", "application/json") + self.send_header("Content-Length", str(len(payload))) + self.end_headers() + self.wfile.write(payload) + + +@pytest.fixture +def runtime_server(monkeypatch): + monkeypatch.delenv("HAI_AGENT_RUNTIME_API_TOKEN", raising=False) + monkeypatch.delenv("HAI_AGENT_LOCAL_BASE_URL", raising=False) + server = RuntimeServer("local-token") + thread = threading.Thread(target=server.serve_forever, daemon=True) + thread.start() + yield server + server.shutdown() + server.server_close() + + class FakeRuntime: base_url = "http://127.0.0.1:18795" api_key = "local-token" @@ -68,36 +146,39 @@ def test_self_hosted_inference_does_not_receive_hosted_key(monkeypatch): assert "HAI_AGENT_RUNTIME_BASE_URL" not in Inference.cloud().runtime_env() -@pytest.mark.parametrize("probe_error", [None, httpx.ConnectError("offline"), httpx.ReadTimeout("timeout")]) -def test_attachment_authenticates_and_checks_recipe_before_use(tmp_path, monkeypatch, probe_error): - from hai_agents_local.runtime import runtime as module - - write_owner_only(token_file_path(18795, cache_dir=tmp_path), "local-token") - monkeypatch.delenv("HAI_AGENT_RUNTIME_API_TOKEN", raising=False) - monkeypatch.setattr(module, "probe_health", lambda url: {"version": "old", "recipe": "desktop"}) - calls = [] - - def get(url, **kwargs): - calls.append(kwargs["headers"]) - return httpx.Response(200) - - monkeypatch.setattr(module.httpx, "get", get) - attached = LocalRuntime.attach(cache_dir=tmp_path) - assert calls == [{"Authorization": "Bearer local-token"}] +@pytest.mark.parametrize("proof_token", ["local-token", "squatter-token", None]) +def test_only_the_runtime_holding_the_token_ever_receives_it(tmp_path, runtime_server, proof_token): + write_owner_only(token_file_path(runtime_server.port, cache_dir=tmp_path), "local-token") + runtime_server.proof_token = proof_token + if proof_token != "local-token": + with pytest.raises(LocalRuntimeError, match="not the runtime"): + LocalRuntime.attach(port=runtime_server.port, cache_dir=tmp_path) + assert runtime_server.requests and not any("authorization" in seen for seen in runtime_server.requests) + return + attached = LocalRuntime.attach(port=runtime_server.port, cache_dir=tmp_path) with pytest.raises(BinaryIncompatibleError): - attached.require_recipe("shared") + attached.require_recipe("desktop") with pytest.raises(LocalRuntimeError): attached.force_kill() - assert token_file_path(18795, cache_dir=tmp_path).exists() + with Client.local(runtime=attached) as client: + assert client.sessions.list_sessions().items == [] + runtime_server.proof_token = "squatter-token" + with pytest.raises(LocalRuntimeError, match="not the runtime"): + client.sessions.list_sessions() - def failed_probe(*args, **kwargs): - if probe_error is not None: - raise probe_error - return httpx.Response(401) - monkeypatch.setattr(module.httpx, "get", failed_probe) - with pytest.raises(LocalRuntimeError, match="authenticated"): - LocalRuntime.attach(cache_dir=tmp_path) +@pytest.mark.asyncio +async def test_bridge_never_serves_an_unproven_runtime(runtime_server): + from hai_agents_local.routing import localize_agent + + runtime_server.proof_token = "squatter-token" + agent = {"environments": [{"id": "workstation", "kind": "workstation", "host": "user_device"}]} + _, [bridge] = localize_agent( + agent, api_key="local-token", base_url=f"http://127.0.0.1:{runtime_server.port}", verify_runtime=True + ) + bridge.create_driver = lambda: pytest.fail("a driver started for an unproven runtime") + with pytest.raises(LocalRuntimeError, match="not the runtime"): + await bridge.run() @pytest.mark.parametrize("url", ["https://remote.example", "http://user:pass@localhost:80"]) @@ -141,45 +222,28 @@ def resolve(**kwargs): LocalRuntime.ensure_started(cache_dir=tmp_path, port=18795) -@pytest.mark.parametrize("status", [200, 400]) -def test_idle_probe_uses_runtime_http_and_closes_its_pool(tmp_path, monkeypatch, status): - from hai_agents.core.api_error import ApiError - - clients, requests, stopped = [], [], [] - original_client = httpx.Client - - def respond(request): - requests.append(request) - assert request.url.path == "/api/v2/sessions" - assert request.headers["Authorization"] == "Bearer local-token" - return httpx.Response(status, json={"items": [], "total": 0, "page": 1, "size": 1}) - - def http_client(**kwargs): - client = original_client(transport=httpx.MockTransport(respond), **kwargs) - clients.append(client) - return client - - monkeypatch.setattr(httpx, "Client", http_client) +def _owned_runtime(server, cache_dir, monkeypatch, stopped): runtime = LocalRuntime( - base_url="http://127.0.0.1:18795", - api_key="local-token", + base_url=f"http://127.0.0.1:{server.port}", + api_key=server.token, pid=123, version=None, log_path=None, owned=True, - cache_dir=tmp_path, - port=18795, + cache_dir=cache_dir, + port=server.port, ) monkeypatch.setattr(runtime, "shutdown", lambda: stopped.append(True)) - if status == 200: - assert runtime.shutdown_if_idle() - assert stopped == [True] - else: - with pytest.raises(ApiError): - runtime.shutdown_if_idle() - assert stopped == [] - assert requests - assert all(client.is_closed for client in clients) + return runtime + + +@pytest.mark.parametrize("active", [[], ["other-client-run"]]) +def test_idle_probe_stops_only_an_unused_runtime(tmp_path, monkeypatch, runtime_server, active): + stopped = [] + runtime_server.active = list(active) + runtime = _owned_runtime(runtime_server, tmp_path, monkeypatch, stopped) + assert runtime.shutdown_if_idle() is (not active) + assert stopped == ([] if active else [True]) @pytest.mark.asyncio From 1446f8e32c248d936d124ad39737d6733ba355d0 Mon Sep 17 00:00:00 2001 From: abonneth Date: Thu, 1 Oct 2026 18:21:05 +0200 Subject: [PATCH 16/35] fix(sdk): publish the runtime token only after the child proves the port A spawner whose health check missed a live runtime used to replace that runtime's token file before discovering the port was taken, then delete it on failure, locking every client out of the live runtime. State files are now written atomically through a staged 0600 file, and the spawner publishes them only once its own child answered with a valid proof. A client that cannot authenticate before the startup lock retries under it, after any concurrent spawner has published. --- src/hai_agents_local/runtime/runtime.py | 17 ++++++++++------- src/hai_agents_local/runtime/state.py | 21 ++++++++++++++------- tests/test_runtime_placement.py | 15 +++++++++++++++ 3 files changed, 39 insertions(+), 14 deletions(-) diff --git a/src/hai_agents_local/runtime/runtime.py b/src/hai_agents_local/runtime/runtime.py index 515fdeb..b78a86b 100644 --- a/src/hai_agents_local/runtime/runtime.py +++ b/src/hai_agents_local/runtime/runtime.py @@ -149,7 +149,12 @@ def ensure_started( resolved_port = port if port is not None else int(os.environ.get(PORT_ENV, "").strip() or DEFAULT_PORT) base_url = f"http://{LOOPBACK_HOST}:{resolved_port}" - attached = cls._attach(base_url=base_url, cache_dir=resolved_cache) + try: + attached = cls._attach(base_url=base_url, cache_dir=resolved_cache) + except LocalRuntimeError: + # A concurrent spawner publishes its token only after its child is proven; decide once it has. + with _startup_lock(resolved_cache, resolved_port, timeout_s): + attached = cls._attach(base_url=base_url, cache_dir=resolved_cache) if attached is not None: attached.require_recipe(required_recipe) return attached @@ -177,14 +182,9 @@ def ensure_started( (spawn_env or {}).get(AUTH_TOKEN_ENV, os.environ.get(AUTH_TOKEN_ENV, "") if inherit_env else "").strip() ) token = explicit_token or secrets.token_urlsafe(32) - # Publish the token before the health wait so a client racing our probe can authenticate. - token_file = ( - None - if explicit_token - else write_owner_only(token_file_path(resolved_port, cache_dir=resolved_cache), token) - ) log_path = runtime_log_path(resolved_port, cache_dir=resolved_cache) proc = None + token_file = None try: proc = spawn( cmd, @@ -201,6 +201,9 @@ def ensure_started( ) if status != 200 or proc.poll() is not None: raise RuntimeUnhealthyError("spawned runtime failed authenticated readiness probe") + # Published only once the child proved it owns the port, so another runtime's file is never replaced. + if not explicit_token: + token_file = write_owner_only(token_file_path(resolved_port, cache_dir=resolved_cache), token) pid_file = write_owner_only(pid_file_path(resolved_port, cache_dir=resolved_cache), str(proc.pid)) except BaseException: # Covers KeyboardInterrupt mid-spawn: never leak the child or its token file. diff --git a/src/hai_agents_local/runtime/state.py b/src/hai_agents_local/runtime/state.py index ddf80b1..2f3fc68 100644 --- a/src/hai_agents_local/runtime/state.py +++ b/src/hai_agents_local/runtime/state.py @@ -11,6 +11,7 @@ import contextlib import os import pathlib +import tempfile import typing CACHE_DIR_ENV = "HAI_AGENT_LOCAL_CACHE_DIR" @@ -51,16 +52,22 @@ def runtime_log_path(port: int, *, cache_dir: typing.Optional[_PathInput] = None def write_owner_only(path: pathlib.Path, content: str) -> pathlib.Path: - """Write `content` to `path` owner-only (0600), refusing a pre-existing symlink at the path.""" + """Atomically publish `content` at `path` owner-only (0600), refusing a pre-existing symlink at the path.""" path.parent.mkdir(parents=True, exist_ok=True) with contextlib.suppress(OSError): path.parent.chmod(0o700) # owner-only state dir; no-op on Windows - flags = os.O_WRONLY | os.O_CREAT | os.O_TRUNC | getattr(os, "O_NOFOLLOW", 0) - fd = os.open(path, flags, 0o600) - with os.fdopen(fd, "w", encoding="utf-8") as fh: - if os.name == "posix": - os.fchmod(fd, 0o600) # enforce owner-only even if the file pre-existed - fh.write(content) + if path.is_symlink(): + raise OSError(f"refusing to replace the symlink at {path}") + # mkstemp creates the staging file 0600 with O_EXCL, so readers only ever see complete content. + fd, staged = tempfile.mkstemp(dir=path.parent, prefix=f".{path.name}.") + try: + with os.fdopen(fd, "w", encoding="utf-8") as fh: + fh.write(content) + os.replace(staged, path) + except BaseException: + with contextlib.suppress(OSError): + os.unlink(staged) + raise return path diff --git a/tests/test_runtime_placement.py b/tests/test_runtime_placement.py index 98b3d12..8c9890d 100644 --- a/tests/test_runtime_placement.py +++ b/tests/test_runtime_placement.py @@ -311,6 +311,21 @@ def cleanup(): assert token_file.read_text() == "replacement-token" +def test_spawner_never_overwrites_the_live_runtime_token(tmp_path, monkeypatch, runtime_server): + import sys + + token_file = write_owner_only(token_file_path(runtime_server.port, cache_dir=tmp_path), runtime_server.token) + monkeypatch.setattr(LocalRuntime, "_attach", lambda **kwargs: None) + with pytest.raises(LocalRuntimeError, match="is not the runtime"): + LocalRuntime.ensure_started( + command=[sys.executable, "-c", "import time; time.sleep(30)"], + cache_dir=tmp_path, + port=runtime_server.port, + timeout_s=5, + ) + assert token_file.read_text() == runtime_server.token + + @pytest.mark.asyncio @pytest.mark.parametrize("asynchronous", [False, True]) async def test_attached_runtime_with_another_recipe_is_rejected_and_left_running(asynchronous): From 8e9b38f7cd02548a11d394202e7f11d7d0e4b73c Mon Sep 17 00:00:00 2001 From: abonneth Date: Thu, 1 Oct 2026 18:23:54 +0200 Subject: [PATCH 17/35] fix(local): stop a bridge cleanly when its session's channel closes The runtime answers 410 once a session ended and its command channel closed. The bridge raised it as a crash, so every normal session end logged an error and fired the loss handler, which cancelled the session again. --- src/hai_agents_local/bridge.py | 5 ++++- src/hai_agents_local/errors.py | 4 ++++ src/hai_agents_local/transport.py | 6 +++++- tests/test_local.py | 20 ++++++++++++++++++++ 4 files changed, 33 insertions(+), 2 deletions(-) diff --git a/src/hai_agents_local/bridge.py b/src/hai_agents_local/bridge.py index 96f37fa..e625dd7 100644 --- a/src/hai_agents_local/bridge.py +++ b/src/hai_agents_local/bridge.py @@ -17,7 +17,7 @@ import httpx from .config import default_base_url -from .errors import RateLimitedError, SessionNotFoundError +from .errors import ChannelClosedError, RateLimitedError, SessionNotFoundError from .runtime import identity from .transport import Command, CommandExchange, Json, deserialize_args, serialize_result @@ -210,6 +210,9 @@ async def _poll_loop(self, exchange: CommandExchange) -> None: ): # Instant empty polls are paced so a misbehaving server cannot cause a busy loop. break + except ChannelClosedError: + logger.info("channel %s closed; the session ended", self.session_id) + return except SessionNotFoundError: # Channel was garbage-collected server-side; recreate on the next iteration so # rate limits and transient errors during recreation hit the handlers below. diff --git a/src/hai_agents_local/errors.py b/src/hai_agents_local/errors.py index 85293b0..1bbac28 100644 --- a/src/hai_agents_local/errors.py +++ b/src/hai_agents_local/errors.py @@ -9,6 +9,10 @@ class SessionNotFoundError(Exception): """The command channel disappeared server-side and must be recreated.""" +class ChannelClosedError(Exception): + """The session ended and the platform closed its command channel for good.""" + + class RateLimitedError(Exception): """The platform asked the bridge to back off polling.""" diff --git a/src/hai_agents_local/transport.py b/src/hai_agents_local/transport.py index 5ee29fa..63fb981 100644 --- a/src/hai_agents_local/transport.py +++ b/src/hai_agents_local/transport.py @@ -12,7 +12,7 @@ import httpx from pydantic import BaseModel, ConfigDict, Field, ValidationError -from .errors import AuthError, RateLimitedError, SessionNotFoundError +from .errors import AuthError, ChannelClosedError, RateLimitedError, SessionNotFoundError logger = logging.getLogger(__name__) @@ -89,6 +89,8 @@ async def ensure_channel(self, session_id: str) -> None: raise RateLimitedError(_retry_after(resp)) if resp.status_code == HTTPStatus.CONFLICT: return + if resp.status_code == HTTPStatus.GONE: + raise ChannelClosedError(f"channel {session_id!r} is closed") resp.raise_for_status() async def fetch_commands( @@ -107,6 +109,8 @@ async def fetch_commands( return None case HTTPStatus.NOT_FOUND: raise SessionNotFoundError(f"channel {session_id!r} not found") + case HTTPStatus.GONE: + raise ChannelClosedError(f"channel {session_id!r} is closed") case HTTPStatus.UNAUTHORIZED | HTTPStatus.FORBIDDEN: raise AuthError(f"auth error ({resp.status_code})") case HTTPStatus.TOO_MANY_REQUESTS: diff --git a/tests/test_local.py b/tests/test_local.py index b28702e..4c168bd 100644 --- a/tests/test_local.py +++ b/tests/test_local.py @@ -1,5 +1,6 @@ import asyncio import json +import logging import sys import threading import types @@ -881,6 +882,25 @@ async def fetch_commands(self, session_id: str, **kwargs: Any) -> None: assert manager._runners[bridge.session_id].thread.is_alive() manager.stop([bridge.session_id]) + def test_closed_channel_is_a_clean_stop(self, manager, monkeypatch, caplog): + original = httpx.AsyncClient + + def respond(request: httpx.Request) -> httpx.Response: + return httpx.Response(410 if request.url.path.startswith("/api/v1/commands/") else 200, json={}) + + monkeypatch.setattr( + httpx, "AsyncClient", lambda **kwargs: original(transport=httpx.MockTransport(respond), **kwargs) + ) + bridge = FakeBridge(api_key="k", base_url="http://runtime.test") + crashed = threading.Event() + bridge.on_crash = crashed.set + manager.ensure([bridge]) + runner = manager._runners[bridge.session_id] + runner.thread.join(5.0) + assert not runner.thread.is_alive() and runner.error is None + assert not crashed.is_set() + assert not [record for record in caplog.records if record.levelno >= logging.ERROR] + def test_crash_after_ready_fires_on_crash(self, manager): crashed = threading.Event() From f738f1047348330774cdda4bd2ab1cd50976a3fb Mon Sep 17 00:00:00 2001 From: abonneth Date: Thu, 1 Oct 2026 18:26:17 +0200 Subject: [PATCH 18/35] fix(local): keep cancel_session's signature and close only live sessions cancel_session overrides renamed the generated 'id' parameter, so cancel_session(id=...) raised TypeError. They now match the generated signature. close() cancelled every session that ever started a bridge, including long finished ones; once the server forgot them the cancels returned 404 and close() raised. It now cancels only sessions whose bridges are still serving, and treats 404 or 409 as already stopped. --- src/hai_agents_local/manager.py | 9 ++++++ src/hai_agents_local/sessions.py | 51 ++++++++++++++++++++++---------- tests/test_local.py | 31 +++++++++++++++++++ 3 files changed, 76 insertions(+), 15 deletions(-) diff --git a/src/hai_agents_local/manager.py b/src/hai_agents_local/manager.py index 044aeef..61c83cf 100644 --- a/src/hai_agents_local/manager.py +++ b/src/hai_agents_local/manager.py @@ -126,6 +126,11 @@ def stop_all(self) -> None: session_ids = list(self._runners) self.stop(session_ids) + def serving(self, session_ids: Sequence[str]) -> list[str]: + """The given bridges that are still running.""" + with self._lock: + return [sid for sid in session_ids if sid in self._runners and self._runners[sid].thread.is_alive()] + class _Runner: def __init__(self, bridge: LocalBridge) -> None: @@ -183,6 +188,10 @@ def ensure_bridges(bridges: Sequence[LocalBridge]) -> list[str]: return _default_manager.ensure(bridges) +def serving_bridges(session_ids: Sequence[str]) -> list[str]: + return _default_manager.serving(session_ids) + + def stop_bridges(session_ids: Sequence[str] | None = None) -> None: if session_ids is None: _default_manager.stop_all() diff --git a/src/hai_agents_local/sessions.py b/src/hai_agents_local/sessions.py index 43befbb..9efdc5f 100644 --- a/src/hai_agents_local/sessions.py +++ b/src/hai_agents_local/sessions.py @@ -15,12 +15,14 @@ import httpx from hai_agents.base_client import BaseClient +from hai_agents.core.api_error import ApiError +from hai_agents.core.request_options import RequestOptions from hai_agents.sessions.client import AsyncSessionsClient, SessionsClient from .bridge import LocalBridge, TokenSource from .config import auto_bridges_enabled from .killswitch import StopWatcher -from .manager import ensure_bridges, stop_bridges +from .manager import ensure_bridges, serving_bridges, stop_bridges from .routing import localize_agent if typing.TYPE_CHECKING: @@ -209,6 +211,17 @@ def _cancel_sessions_at_exit() -> None: atexit.register(_cancel_sessions_at_exit) +# The session already ended or is gone: nothing left to stop. +STOPPED_CANCEL_STATUSES = frozenset({404, 409}) + + +def _live_sessions(owned_bridges: typing.Dict[str, typing.List[str]]) -> typing.List[str]: + """Forget sessions whose bridges all stopped, since each such stop already ended or cancelled its session.""" + for session_id, bridge_ids in list(owned_bridges.items()): + if not serving_bridges(bridge_ids): + del owned_bridges[session_id] + return list(owned_bridges) + class LocalSessionsClient(SessionsClient): def __init__( @@ -222,26 +235,30 @@ def __init__( def close(self) -> None: failures = [] - for session_id in list(self._owned_bridges): + for session_id in _live_sessions(self._owned_bridges): try: self.cancel_session(session_id) + except ApiError as error: + if error.status_code not in STOPPED_CANCEL_STATUSES: + failures.append(error) + else: + _deregister_exit_cancel(session_id) except Exception as error: failures.append(error) if failures: raise RuntimeError("Could not confirm all client-owned sessions stopped") from failures[0] - def cancel_session(self, session_id: str, **kwargs: typing.Any) -> typing.Any: + def cancel_session(self, id: str, *, request_options: typing.Optional[RequestOptions] = None) -> None: # Stop local execution even if the remote cancellation cannot be delivered. - owned = self._owned_bridges.get(str(session_id), []) + owned = self._owned_bridges.get(str(id), []) try: if owned: stop_bridges(owned) - self._owned_bridges.pop(str(session_id), None) + self._owned_bridges.pop(str(id), None) finally: - response = super().cancel_session(session_id, **kwargs) + super().cancel_session(id, request_options=request_options) # Keep the exit retry registered until cancellation is confirmed. - _deregister_exit_cancel(str(session_id)) - return response + _deregister_exit_cancel(str(id)) @functools.wraps(SessionsClient.create_session) def create_session(self, **kwargs: typing.Any) -> typing.Any: @@ -283,24 +300,28 @@ def __init__( async def aclose(self) -> None: failures = [] - for session_id in list(self._owned_bridges): + for session_id in _live_sessions(self._owned_bridges): try: await self.cancel_session(session_id) + except ApiError as error: + if error.status_code not in STOPPED_CANCEL_STATUSES: + failures.append(error) + else: + _deregister_exit_cancel(session_id) except Exception as error: failures.append(error) if failures: raise RuntimeError("Could not confirm all client-owned sessions stopped") from failures[0] - async def cancel_session(self, session_id: str, **kwargs: typing.Any) -> typing.Any: - owned = self._owned_bridges.get(str(session_id), []) + async def cancel_session(self, id: str, *, request_options: typing.Optional[RequestOptions] = None) -> None: + owned = self._owned_bridges.get(str(id), []) try: if owned: await asyncio.to_thread(stop_bridges, owned) - self._owned_bridges.pop(str(session_id), None) + self._owned_bridges.pop(str(id), None) finally: - response = await super().cancel_session(session_id, **kwargs) - _deregister_exit_cancel(str(session_id)) - return response + await super().cancel_session(id, request_options=request_options) + _deregister_exit_cancel(str(id)) @functools.wraps(AsyncSessionsClient.create_session) async def create_session(self, **kwargs: typing.Any) -> typing.Any: diff --git a/tests/test_local.py b/tests/test_local.py index 4c168bd..b0523c6 100644 --- a/tests/test_local.py +++ b/tests/test_local.py @@ -1063,3 +1063,34 @@ async def async_cancel(*args, **kwargs): await client.aclose() else: client.close() + + +@pytest.mark.asyncio +@pytest.mark.parametrize("asynchronous", [False, True]) +async def test_close_cancels_only_sessions_still_served(monkeypatch, asynchronous): + from hai_agents import AsyncClient + + cancelled = [] + + def respond(request: httpx.Request) -> httpx.Response: + session_id = request.url.path.rsplit("/", 1)[1] + cancelled.append(session_id) + return httpx.Response(404 if session_id == "evicted" else 204) + + monkeypatch.setattr("hai_agents_local.sessions.stop_bridges", lambda ids: None) + monkeypatch.setattr( + "hai_agents_local.sessions.serving_bridges", lambda ids: [i for i in ids if i != "ended-bridge"] + ) + transport = httpx.MockTransport(respond) + http = httpx.AsyncClient(transport=transport) if asynchronous else httpx.Client(transport=transport) + client = (AsyncClient if asynchronous else Client)(api_key=API_KEY, base_url="http://api.test", httpx_client=http) + sessions = client.sessions + sessions._owned_bridges = {"live": ["live-bridge"], "ended": ["ended-bridge"], "evicted": ["evicted-bridge"]} + if asynchronous: + await sessions.cancel_session(id="unbridged") + await sessions.aclose() + else: + sessions.cancel_session(id="unbridged") + sessions.close() + assert cancelled == ["unbridged", "live", "evicted"] + assert sessions._owned_bridges == {} From d02c2de1b276a86c2c39a9cec08c8db40e1391c0 Mon Sep 17 00:00:00 2001 From: abonneth Date: Thu, 1 Oct 2026 18:30:28 +0200 Subject: [PATCH 19/35] fix(sdk): keep a started runtime alive while other clients use it Closing the client that started the runtime terminated it for every attached client. close() now stops the runtime only when no session other than this client's own is still active; shutdown() still stops it unconditionally. --- README.md | 4 +- src/hai_agents/client.py | 10 ++-- src/hai_agents_local/runtime/runtime.py | 21 +++++--- src/hai_agents_local/sessions.py | 8 +++ tests/test_runtime_placement.py | 69 ++++++++++++++----------- 5 files changed, 69 insertions(+), 43 deletions(-) diff --git a/README.md b/README.md index 2e573c1..dbe8518 100644 --- a/README.md +++ b/README.md @@ -111,8 +111,8 @@ does not by itself imply local inference: from `hai_agents_local.runtime`, selects a model endpoint for a newly started local runtime. Hosted agents do not currently accept that override. -Closing the client shuts down a runtime it started; an attached runtime remains -owned by its caller. `cancel()` ends the agent session. For a cloud workstation, +Closing the client shuts down a runtime it started, unless sessions from other +clients are still active there. An attached runtime remains owned by its caller. `cancel()` ends the agent session. For a cloud workstation, an explicit `session_id` attaches to a caller-owned environment, which the caller must eventually release. In the candidate shared recipe, automatically provisioned cloud workstations survive agent cancellation and expire through the environment diff --git a/src/hai_agents/client.py b/src/hai_agents/client.py index ae7ae55..09b497c 100644 --- a/src/hai_agents/client.py +++ b/src/hai_agents/client.py @@ -61,7 +61,7 @@ def local( return client def close(self) -> None: - """Stop sessions this client bridged; a local client also releases its runtime and connections.""" + """Stop sessions this client bridged; a local client also stops its runtime once no other client uses it.""" try: if self._sessions is not None: self._sessions.close() @@ -69,7 +69,7 @@ def close(self) -> None: if self.local_runtime is not None: try: if self._owns_runtime: - self.local_runtime.shutdown() + self.local_runtime.shutdown_if_idle(getattr(self._sessions, "own_session_ids", ())) finally: self._client_wrapper.httpx_client.httpx_client.close() @@ -168,7 +168,7 @@ async def local( return client async def aclose(self) -> None: - """Stop sessions this client bridged; a local client also releases its runtime and connections.""" + """Stop sessions this client bridged; a local client also stops its runtime once no other client uses it.""" try: if self._sessions is not None: await self._sessions.aclose() @@ -176,7 +176,9 @@ async def aclose(self) -> None: if self.local_runtime is not None: try: if self._owns_runtime: - await asyncio.to_thread(self.local_runtime.shutdown) + await asyncio.to_thread( + self.local_runtime.shutdown_if_idle, getattr(self._sessions, "own_session_ids", ()) + ) finally: await self._client_wrapper.httpx_client.httpx_client.aclose() diff --git a/src/hai_agents_local/runtime/runtime.py b/src/hai_agents_local/runtime/runtime.py index b78a86b..ec79dbb 100644 --- a/src/hai_agents_local/runtime/runtime.py +++ b/src/hai_agents_local/runtime/runtime.py @@ -57,6 +57,7 @@ PORT_ENV = "HAI_AGENT_RUNTIME_PORT" AUTH_TOKEN_ENV = "HAI_AGENT_RUNTIME_API_TOKEN" CLIENT_TIMEOUT_S = 60.0 +IDLE_PROBE_PAGE_SIZE = 50 _PathInput = typing.Union[str, "os.PathLike[str]"] @@ -408,13 +409,21 @@ def shutdown(self) -> None: "awaiting_tool_results", ) - def shutdown_if_idle(self) -> bool: - """Stop the owned runtime only when it hosts no active sessions; True when it was stopped.""" + def shutdown_if_idle(self, ignore: typing.Collection[str] = ()) -> bool: + """Stop the owned runtime unless it hosts an active session outside ``ignore``; True when it was stopped.""" with self.http_client() as http: - client = BaseClient(base_url=self.base_url, api_key=self.api_key, httpx_client=http) - page = client.sessions.list_sessions(status=list(self.ACTIVE_SESSION_STATUSES), size=1) - if page.items: - return False + sessions = BaseClient(base_url=self.base_url, api_key=self.api_key, httpx_client=http).sessions + page, seen = 1, 0 + while True: + listed = sessions.list_sessions( + status=list(self.ACTIVE_SESSION_STATUSES), page=page, size=IDLE_PROBE_PAGE_SIZE + ) + if any(item.id not in ignore for item in listed.items): + return False + seen += len(listed.items) + if not listed.items or seen >= listed.total: + break + page += 1 self.shutdown() return True diff --git a/src/hai_agents_local/sessions.py b/src/hai_agents_local/sessions.py index 9efdc5f..25e352d 100644 --- a/src/hai_agents_local/sessions.py +++ b/src/hai_agents_local/sessions.py @@ -232,6 +232,8 @@ def __init__( self._auto_bridges = auto_bridges self._cancel_remote: RemoteCancel = functools.partial(_cancel_remote_session, client_wrapper, runtime) self._owned_bridges: typing.Dict[str, typing.List[str]] = {} + # Sessions this client created on a local runtime; they never keep that runtime alive past close(). + self.own_session_ids: typing.Set[str] = set() def close(self) -> None: failures = [] @@ -283,6 +285,8 @@ def create_session(self, **kwargs: typing.Any) -> typing.Any: if stop_watcher is not None and not stop_watcher.active: # A stop was filed while bridges or the session were starting; apply it now. _panic_stop() + if self._runtime is not None: + self.own_session_ids.add(str(session.id)) if started: self._owned_bridges[str(session.id)] = started return session @@ -297,6 +301,8 @@ def __init__( self._auto_bridges = auto_bridges self._cancel_remote: RemoteCancel = functools.partial(_cancel_remote_session, client_wrapper, runtime) self._owned_bridges: typing.Dict[str, typing.List[str]] = {} + # Sessions this client created on a local runtime; they never keep that runtime alive past close(). + self.own_session_ids: typing.Set[str] = set() async def aclose(self) -> None: failures = [] @@ -350,6 +356,8 @@ async def create_session(self, **kwargs: typing.Any) -> typing.Any: if stop_watcher is not None and not stop_watcher.active: # A stop was filed while bridges or the session were starting; apply it now. await asyncio.to_thread(_panic_stop) + if self._runtime is not None: + self.own_session_ids.add(str(session.id)) if started: self._owned_bridges[str(session.id)] = started return session diff --git a/tests/test_runtime_placement.py b/tests/test_runtime_placement.py index 8c9890d..c4c92fb 100644 --- a/tests/test_runtime_placement.py +++ b/tests/test_runtime_placement.py @@ -107,6 +107,21 @@ def shutdown(self): self.stopped = True +def _owned_runtime(server, cache_dir, monkeypatch, stopped): + runtime = LocalRuntime( + base_url=f"http://127.0.0.1:{server.port}", + api_key=server.token, + pid=123, + version=None, + log_path=None, + owned=True, + cache_dir=cache_dir, + port=server.port, + ) + monkeypatch.setattr(runtime, "shutdown", lambda: stopped.append(True)) + return runtime + + @pytest.mark.asyncio @pytest.mark.parametrize("asynchronous", [False, True]) async def test_local_client_without_auto_bridges_leaves_execution_to_the_product(monkeypatch, asynchronous): @@ -190,18 +205,34 @@ def test_local_attach_rejects_remote_or_credential_urls(tmp_path, monkeypatch, u @pytest.mark.asyncio @pytest.mark.parametrize("asynchronous", [False, True]) -@pytest.mark.parametrize("borrowed", [False, True]) -async def test_client_close_releases_only_owned_runtime(monkeypatch, asynchronous, borrowed): - runtime = FakeRuntime() +@pytest.mark.parametrize("runtime_use", ["borrowed", "idle", "shared"]) +async def test_close_stops_only_an_owned_runtime_no_other_client_uses( + tmp_path, monkeypatch, runtime_server, asynchronous, runtime_use +): + from types import SimpleNamespace + + from hai_agents.sessions.client import AsyncSessionsClient, SessionsClient + + stopped = [] + runtime_server.active = ["mine", "another-client-run"] if runtime_use == "shared" else ["mine"] + runtime = _owned_runtime(runtime_server, tmp_path, monkeypatch, stopped) monkeypatch.setattr(LocalRuntime, "ensure_started", lambda **options: runtime) - options = {"runtime": runtime} if borrowed else {} + monkeypatch.setattr(SessionsClient, "create_session", lambda self, **kwargs: SimpleNamespace(id="mine")) + + async def async_create(self, **kwargs): + return SimpleNamespace(id="mine") + + monkeypatch.setattr(AsyncSessionsClient, "create_session", async_create) + options = {"runtime": runtime} if runtime_use == "borrowed" else {} if asynchronous: - client = await AsyncClient.local(**options) + client = await AsyncClient.local(auto_bridges=False, **options) + await client.sessions.create_session(agent="h/agent", messages="hi") await client.aclose() else: - client = Client.local(**options) + client = Client.local(auto_bridges=False, **options) + client.sessions.create_session(agent="h/agent", messages="hi") client.close() - assert runtime.stopped is (not borrowed) + assert stopped == ([True] if runtime_use == "idle" else []) assert client._client_wrapper.httpx_client.httpx_client.is_closed @@ -222,30 +253,6 @@ def resolve(**kwargs): LocalRuntime.ensure_started(cache_dir=tmp_path, port=18795) -def _owned_runtime(server, cache_dir, monkeypatch, stopped): - runtime = LocalRuntime( - base_url=f"http://127.0.0.1:{server.port}", - api_key=server.token, - pid=123, - version=None, - log_path=None, - owned=True, - cache_dir=cache_dir, - port=server.port, - ) - monkeypatch.setattr(runtime, "shutdown", lambda: stopped.append(True)) - return runtime - - -@pytest.mark.parametrize("active", [[], ["other-client-run"]]) -def test_idle_probe_stops_only_an_unused_runtime(tmp_path, monkeypatch, runtime_server, active): - stopped = [] - runtime_server.active = list(active) - runtime = _owned_runtime(runtime_server, tmp_path, monkeypatch, stopped) - assert runtime.shutdown_if_idle() is (not active) - assert stopped == ([] if active else [True]) - - @pytest.mark.asyncio async def test_async_local_startup_keeps_loop_responsive_and_cancellation_cleans_child(monkeypatch): import asyncio From fb3f3d329faebddfc98ef2da4a4237f5dafca666 Mon Sep 17 00:00:00 2001 From: abonneth Date: Thu, 1 Oct 2026 18:35:45 +0200 Subject: [PATCH 20/35] chore(sdk): drop internal markers and em dashes from local runtime docs --- README.md | 4 ++-- src/hai_agents_local/runtime/manifest.py | 4 +--- src/hai_agents_local/runtime/runtime.py | 2 +- tests/test_runtime_state.py | 4 ++-- 4 files changed, 6 insertions(+), 8 deletions(-) diff --git a/README.md b/README.md index dbe8518..9445d1d 100644 --- a/README.md +++ b/README.md @@ -69,8 +69,8 @@ print(result.answer) ## Candidate local-agent support -The placement branch adds `Client.local()` and `await AsyncClient.local()`, which -start a local agent runtime or, with `runtime=...`, use an already prepared one. +`Client.local()` and `await AsyncClient.local()` start a local agent runtime or, +with `runtime=...`, use an already prepared one. Agent placement and environment placement are separate: the client selects where the agent runs; each agent environment selects `host="user_device"` or `host="cloud"`. `Client()` continues to use the hosted Agents API. diff --git a/src/hai_agents_local/runtime/manifest.py b/src/hai_agents_local/runtime/manifest.py index e1f782c..1677d2f 100644 --- a/src/hai_agents_local/runtime/manifest.py +++ b/src/hai_agents_local/runtime/manifest.py @@ -12,7 +12,7 @@ import sys import typing -# SHIP-GATE: pin a release that serves the shared recipe before merge +# TODO: pin a runtime release that serves the shared recipe, with its artifact hashes below. PINNED_RUNTIME_VERSION = "0.1.8" RUNTIME_CDN_BASE = "https://assets.hcompanyprod.fr/hai-agent-runtime" # Guard value: published manifest entries must never use it (every download would fail verification). @@ -34,12 +34,10 @@ def _artifact(filename: str, sha256: str) -> RuntimeArtifact: MANIFEST: typing.Dict[str, RuntimeArtifact] = { "darwin-arm64": _artifact( "hai-agent-runtime-darwin-arm64.zip", - # SHIP-GATE: pin a release that serves the shared recipe before merge "1aed0055898116732aee031dc4a1235782b2909ee51e0367e2d50bb3be6671c9", ), "windows-x86_64": _artifact( "hai-agent-runtime-windows-x86_64.zip", - # SHIP-GATE: pin a release that serves the shared recipe before merge "4e6b2bcd42af2bb6b22197fcde947327497f5c62fd60d48bc9037730d80dc691", ), } diff --git a/src/hai_agents_local/runtime/runtime.py b/src/hai_agents_local/runtime/runtime.py index ec79dbb..f2d91fc 100644 --- a/src/hai_agents_local/runtime/runtime.py +++ b/src/hai_agents_local/runtime/runtime.py @@ -322,7 +322,7 @@ def _child_env( Inheriting os.environ passes the model-gateway HAI_API_KEY / HAI_BASE_URL through to the binary (without them local sessions cannot run inference) and forwards caller flags such as HAI_AGENT_RUNTIME_MODEL/FAKE/FAST/RUNS_DIR. inherit_env=False takes spawn_env as the - complete base environment instead — for callers that must *remove* inherited keys, which an + complete base environment instead, for callers that must *remove* inherited keys, which an overlay cannot express (e.g. stripping HAI_API_KEY for self-hosted base URLs). The generated local bearer and the cloud HAI_API_KEY are different credentials: the token below is the only local bearer, and the cloud key is never used to authenticate against the local diff --git a/tests/test_runtime_state.py b/tests/test_runtime_state.py index 523c504..6b7d508 100644 --- a/tests/test_runtime_state.py +++ b/tests/test_runtime_state.py @@ -1,8 +1,8 @@ """Behavioural tests for the runtime pid file's on-disk hardening. -The pid file's contents drive a privileged operation: ``holo stop --force`` reads it and +The pid file's contents drive a privileged operation: a forced stop reads it and ``os.killpg(..., SIGKILL)`` the pid it finds. It therefore earns the same protections as the bearer -token — owner-only permissions and a refusal to follow a symlink planted at its path — so it can never +token (owner-only permissions and a refusal to replace a symlink planted at its path), so it can never be steered into killing an arbitrary process group. """ From 1d6640bddd6bdd2daa5f33a8a56ba342cd07e2d0 Mon Sep 17 00:00:00 2001 From: abonneth Date: Thu, 1 Oct 2026 18:43:37 +0200 Subject: [PATCH 21/35] fix(local): give the runtime time to release environments before SIGKILL --- src/hai_agents_local/runtime/process.py | 6 ++++-- 1 file changed, 4 insertions(+), 2 deletions(-) diff --git a/src/hai_agents_local/runtime/process.py b/src/hai_agents_local/runtime/process.py index 650ff41..62b89c7 100644 --- a/src/hai_agents_local/runtime/process.py +++ b/src/hai_agents_local/runtime/process.py @@ -21,7 +21,9 @@ LOOPBACK_HOST = "127.0.0.1" SPAWN_TIMEOUT_S = 45.0 HEALTH_POLL_INTERVAL_S = 0.25 -TERM_GRACE_S = 2.0 +# Exceeds the runtime's own shutdown teardown (10 s), which releases cloud environments. +TERM_GRACE_S = 15.0 +KILL_WAIT_S = 2.0 LOG_TAIL_CHARS = 4000 @@ -155,6 +157,6 @@ def terminate(proc: subprocess.Popen) -> None: pass if _signal(proc, force=True): try: - proc.wait(timeout=TERM_GRACE_S) + proc.wait(timeout=KILL_WAIT_S) except subprocess.TimeoutExpired: logger.warning("hai-agent-runtime (pid %d) did not exit after forced kill", proc.pid) From c2c9a9642ae686af4076be7e6bb5de66af2f5a92 Mon Sep 17 00:00:00 2001 From: abonneth Date: Thu, 1 Oct 2026 18:46:42 +0200 Subject: [PATCH 22/35] ci: block SDK releases whose pinned runtime lacks the shared recipe --- .github/workflows/publish.yml | 27 +++++++++++++++++++++++- src/hai_agents_local/runtime/identity.py | 3 ++- 2 files changed, 28 insertions(+), 2 deletions(-) diff --git a/.github/workflows/publish.yml b/.github/workflows/publish.yml index d3db56a..4dcc3f9 100644 --- a/.github/workflows/publish.yml +++ b/.github/workflows/publish.yml @@ -37,8 +37,33 @@ jobs: name: dist path: dist/ - publish: + runtime-pin: needs: build + runs-on: macos-14 + timeout-minutes: 10 + steps: + - name: Install uv + uses: astral-sh/setup-uv@v10.1.0 + + - uses: actions/download-artifact@v8 + with: + name: dist + path: dist/ + + - name: Pinned runtime downloads, verifies and serves the shared recipe + env: + HAI_AGENT_RUNTIME_FAKE: "1" + run: | + uv venv -q && uv pip install -q dist/*.whl + .venv/bin/python -c ' + import tempfile + from hai_agents_local.runtime import LocalRuntime + LocalRuntime.ensure_started(required_recipe="shared", cache_dir=tempfile.mkdtemp(), port=18799, + spawn_env={"HAI_AGENT_RUNTIME_RECIPE": "shared"}).shutdown() + ' + + publish: + needs: [build, runtime-pin] runs-on: ubuntu-latest timeout-minutes: 10 environment: release diff --git a/src/hai_agents_local/runtime/identity.py b/src/hai_agents_local/runtime/identity.py index 5aa6b0c..c55736a 100644 --- a/src/hai_agents_local/runtime/identity.py +++ b/src/hai_agents_local/runtime/identity.py @@ -26,7 +26,8 @@ def verify(token: str, request: httpx.Request, response: httpx.Response) -> None received = response.headers.get(PROOF_HEADER, "") if not sent or not hmac.compare_digest(received.encode(), expected.encode()): raise LocalRuntimeError( - f"{request.url.scheme}://{request.url.netloc.decode()} is not the runtime this client started or attached to" + f"{request.url.scheme}://{request.url.netloc.decode()} did not prove it is the runtime this client " + "started or attached to: another process holds the port, or the runtime predates identity proofs" ) From a42b972ab5b3ed923d475d91b1a3d9244151606c Mon Sep 17 00:00:00 2001 From: abonneth Date: Thu, 1 Oct 2026 18:48:03 +0200 Subject: [PATCH 23/35] test: match the runtime identity error --- tests/test_runtime_placement.py | 8 ++++---- 1 file changed, 4 insertions(+), 4 deletions(-) diff --git a/tests/test_runtime_placement.py b/tests/test_runtime_placement.py index c4c92fb..fba5916 100644 --- a/tests/test_runtime_placement.py +++ b/tests/test_runtime_placement.py @@ -166,7 +166,7 @@ def test_only_the_runtime_holding_the_token_ever_receives_it(tmp_path, runtime_s write_owner_only(token_file_path(runtime_server.port, cache_dir=tmp_path), "local-token") runtime_server.proof_token = proof_token if proof_token != "local-token": - with pytest.raises(LocalRuntimeError, match="not the runtime"): + with pytest.raises(LocalRuntimeError, match="did not prove"): LocalRuntime.attach(port=runtime_server.port, cache_dir=tmp_path) assert runtime_server.requests and not any("authorization" in seen for seen in runtime_server.requests) return @@ -178,7 +178,7 @@ def test_only_the_runtime_holding_the_token_ever_receives_it(tmp_path, runtime_s with Client.local(runtime=attached) as client: assert client.sessions.list_sessions().items == [] runtime_server.proof_token = "squatter-token" - with pytest.raises(LocalRuntimeError, match="not the runtime"): + with pytest.raises(LocalRuntimeError, match="did not prove"): client.sessions.list_sessions() @@ -192,7 +192,7 @@ async def test_bridge_never_serves_an_unproven_runtime(runtime_server): agent, api_key="local-token", base_url=f"http://127.0.0.1:{runtime_server.port}", verify_runtime=True ) bridge.create_driver = lambda: pytest.fail("a driver started for an unproven runtime") - with pytest.raises(LocalRuntimeError, match="not the runtime"): + with pytest.raises(LocalRuntimeError, match="did not prove"): await bridge.run() @@ -323,7 +323,7 @@ def test_spawner_never_overwrites_the_live_runtime_token(tmp_path, monkeypatch, token_file = write_owner_only(token_file_path(runtime_server.port, cache_dir=tmp_path), runtime_server.token) monkeypatch.setattr(LocalRuntime, "_attach", lambda **kwargs: None) - with pytest.raises(LocalRuntimeError, match="is not the runtime"): + with pytest.raises(LocalRuntimeError, match="did not prove"): LocalRuntime.ensure_started( command=[sys.executable, "-c", "import time; time.sleep(30)"], cache_dir=tmp_path, From c6297707e682e228690da432ad0b64e3f2292cc7 Mon Sep 17 00:00:00 2001 From: abonneth Date: Thu, 1 Oct 2026 19:27:50 +0200 Subject: [PATCH 24/35] test: skip desktop stop test until hai-drivers ships DesktopCommandRunner --- tests/test_desktop_stop.py | 8 +++++++- 1 file changed, 7 insertions(+), 1 deletion(-) diff --git a/tests/test_desktop_stop.py b/tests/test_desktop_stop.py index 679557f..7ff9fc1 100644 --- a/tests/test_desktop_stop.py +++ b/tests/test_desktop_stop.py @@ -8,8 +8,14 @@ from types import SimpleNamespace import pytest + +pytest.importorskip("hai_drivers.desktop.utils") +try: + from hai_drivers.desktop.utils import DesktopCommandRunner +except ImportError: + pytest.skip("installed hai-drivers lacks DesktopCommandRunner", allow_module_level=True) + from hai_drivers.desktop.scaled import ScaledDesktopDriver -from hai_drivers.desktop.utils import DesktopCommandRunner from hai_agents_local.desktop import PyautoguiDesktopBridge from hai_agents_local.transport import Command From bd285d209be9466a6334ac39dd047e0e05513db9 Mon Sep 17 00:00:00 2001 From: abonneth Date: Thu, 1 Oct 2026 19:27:50 +0200 Subject: [PATCH 25/35] fix(local): treat a closed channel on result post as a clean stop --- src/hai_agents_local/transport.py | 12 ++++++++---- tests/test_local.py | 9 +++++++-- 2 files changed, 15 insertions(+), 6 deletions(-) diff --git a/src/hai_agents_local/transport.py b/src/hai_agents_local/transport.py index 63fb981..84df979 100644 --- a/src/hai_agents_local/transport.py +++ b/src/hai_agents_local/transport.py @@ -89,8 +89,7 @@ async def ensure_channel(self, session_id: str) -> None: raise RateLimitedError(_retry_after(resp)) if resp.status_code == HTTPStatus.CONFLICT: return - if resp.status_code == HTTPStatus.GONE: - raise ChannelClosedError(f"channel {session_id!r} is closed") + _raise_if_gone(resp, session_id) resp.raise_for_status() async def fetch_commands( @@ -104,13 +103,12 @@ async def fetch_commands( url = f"{self._base}/api/v1/commands/{session_id}/commands" for attempt in range(max_retries + 1): resp = await self._client.get(url, params={"wait_for_seconds": wait_for_seconds}, timeout=read_timeout) + _raise_if_gone(resp, session_id) match resp.status_code: case HTTPStatus.NO_CONTENT: return None case HTTPStatus.NOT_FOUND: raise SessionNotFoundError(f"channel {session_id!r} not found") - case HTTPStatus.GONE: - raise ChannelClosedError(f"channel {session_id!r} is closed") case HTTPStatus.UNAUTHORIZED | HTTPStatus.FORBIDDEN: raise AuthError(f"auth error ({resp.status_code})") case HTTPStatus.TOO_MANY_REQUESTS: @@ -133,9 +131,15 @@ async def post_result( if resp.status_code == HTTPStatus.CONFLICT: # Another delivery of the same command_uid already landed; the result is recorded. return + _raise_if_gone(resp, command_id) resp.raise_for_status() +def _raise_if_gone(resp: httpx.Response, channel: str) -> None: + if resp.status_code == HTTPStatus.GONE: + raise ChannelClosedError(f"channel for {channel!r} is closed") + + def _retry_after(resp: httpx.Response) -> float: try: return max(0.0, float(resp.headers.get("Retry-After", ""))) diff --git a/tests/test_local.py b/tests/test_local.py index b0523c6..78b6afb 100644 --- a/tests/test_local.py +++ b/tests/test_local.py @@ -882,11 +882,16 @@ async def fetch_commands(self, session_id: str, **kwargs: Any) -> None: assert manager._runners[bridge.session_id].thread.is_alive() manager.stop([bridge.session_id]) - def test_closed_channel_is_a_clean_stop(self, manager, monkeypatch, caplog): + @pytest.mark.parametrize("closed_on", ["/commands", "/result"]) + def test_closed_channel_is_a_clean_stop(self, manager, monkeypatch, caplog, closed_on): original = httpx.AsyncClient + command = {"id": "c1", "command_uid": "u1", "name": "noop", "args": {}} def respond(request: httpx.Request) -> httpx.Response: - return httpx.Response(410 if request.url.path.startswith("/api/v1/commands/") else 200, json={}) + path = request.url.path + if path.endswith(closed_on): + return httpx.Response(410, json={}) + return httpx.Response(200, json=[command] if path.endswith("/commands") else {}) monkeypatch.setattr( httpx, "AsyncClient", lambda **kwargs: original(transport=httpx.MockTransport(respond), **kwargs) From b677898932cbef9e5e000b71f94e6f12f03e5dd4 Mon Sep 17 00:00:00 2001 From: abonneth Date: Thu, 1 Oct 2026 19:27:51 +0200 Subject: [PATCH 26/35] fix(local): serialize idle shutdown with startup and cover a full spawn in lock waits --- src/hai_agents_local/runtime/runtime.py | 36 ++++++++++++++++++++----- tests/test_runtime_placement.py | 35 ++++++++++++++++++++++++ 2 files changed, 64 insertions(+), 7 deletions(-) diff --git a/src/hai_agents_local/runtime/runtime.py b/src/hai_agents_local/runtime/runtime.py index f2d91fc..46de55f 100644 --- a/src/hai_agents_local/runtime/runtime.py +++ b/src/hai_agents_local/runtime/runtime.py @@ -29,8 +29,10 @@ from .install import DOWNLOAD_SHA256_ENV, DOWNLOAD_URL_ENV, install_runtime, installed_binary, pinned_artifact from .manifest import PINNED_RUNTIME_VERSION from .process import ( + KILL_WAIT_S, LOOPBACK_HOST, SPAWN_TIMEOUT_S, + TERM_GRACE_S, probe_health, responds, spawn, @@ -58,6 +60,8 @@ AUTH_TOKEN_ENV = "HAI_AGENT_RUNTIME_API_TOKEN" CLIENT_TIMEOUT_S = 60.0 IDLE_PROBE_PAGE_SIZE = 50 +# Beyond the holder's health budget: its authenticated probe, then a failed child's graceful stop. +STARTUP_LOCK_GRACE_S = TERM_GRACE_S + KILL_WAIT_S + 5.0 _PathInput = typing.Union[str, "os.PathLike[str]"] @@ -150,11 +154,12 @@ def ensure_started( resolved_port = port if port is not None else int(os.environ.get(PORT_ENV, "").strip() or DEFAULT_PORT) base_url = f"http://{LOOPBACK_HOST}:{resolved_port}" + lock_wait_s = timeout_s + STARTUP_LOCK_GRACE_S try: attached = cls._attach(base_url=base_url, cache_dir=resolved_cache) except LocalRuntimeError: # A concurrent spawner publishes its token only after its child is proven; decide once it has. - with _startup_lock(resolved_cache, resolved_port, timeout_s): + with _startup_lock(resolved_cache, resolved_port, lock_wait_s): attached = cls._attach(base_url=base_url, cache_dir=resolved_cache) if attached is not None: attached.require_recipe(required_recipe) @@ -171,7 +176,7 @@ def ensure_started( binary_path=binary_path, version=version, cache_dir=resolved_cache, download=download ) ) - with _startup_lock(resolved_cache, resolved_port, timeout_s): + with _startup_lock(resolved_cache, resolved_port, lock_wait_s): attached = cls._attach(base_url=base_url, cache_dir=resolved_cache) if attached is not None: attached.require_recipe(required_recipe) @@ -411,6 +416,15 @@ def shutdown(self) -> None: def shutdown_if_idle(self, ignore: typing.Collection[str] = ()) -> bool: """Stop the owned runtime unless it hosts an active session outside ``ignore``; True when it was stopped.""" + # Serialized with spawns and locked attaches only: a client that attached without the lock + # can still start a session between the listing and the stop. + with _startup_lock(self._cache_dir, self._port, SPAWN_TIMEOUT_S + STARTUP_LOCK_GRACE_S): + if self._hosts_active_session(ignore): + return False + self.shutdown() + return True + + def _hosts_active_session(self, ignore: typing.Collection[str]) -> bool: with self.http_client() as http: sessions = BaseClient(base_url=self.base_url, api_key=self.api_key, httpx_client=http).sessions page, seen = 1, 0 @@ -419,13 +433,11 @@ def shutdown_if_idle(self, ignore: typing.Collection[str] = ()) -> bool: status=list(self.ACTIVE_SESSION_STATUSES), page=page, size=IDLE_PROBE_PAGE_SIZE ) if any(item.id not in ignore for item in listed.items): - return False + return True seen += len(listed.items) if not listed.items or seen >= listed.total: - break + return False page += 1 - self.shutdown() - return True def force_kill(self) -> None: """Stop only the process this manager spawned; never trust a saved PID to claim ownership.""" @@ -436,7 +448,7 @@ def force_kill(self) -> None: def _cleanup_state_files(self) -> None: # Serialize compare-and-unlink with publication of a replacement runtime's state. - with _startup_lock(self._cache_dir, self._port, SPAWN_TIMEOUT_S): + with _startup_lock(self._cache_dir, self._port, SPAWN_TIMEOUT_S + STARTUP_LOCK_GRACE_S): if self._token_file is not None: unlink_if_content(self._token_file, self.api_key) self._token_file = None @@ -445,9 +457,17 @@ def _cleanup_state_files(self) -> None: self._pid_file = None +_held_startup_locks = threading.local() + + @contextlib.contextmanager def _startup_lock(cache_dir: pathlib.Path, port: int, timeout_s: float): path = cache_dir / "state" / f"startup-{port}.lock" + # Reentrant per thread: a second open file description would block on this thread's own flock. + held: typing.Set[pathlib.Path] = _held_startup_locks.__dict__.setdefault("paths", set()) + if path in held: + yield + return path.parent.mkdir(parents=True, exist_ok=True, mode=0o700) fd = os.open(path, os.O_RDWR | os.O_CREAT | getattr(os, "O_NOFOLLOW", 0), 0o600) with os.fdopen(fd, "r+b") as handle: @@ -468,9 +488,11 @@ def _startup_lock(cache_dir: pathlib.Path, port: int, timeout_s: float): if time.monotonic() >= deadline: raise LocalRuntimeError("timed out waiting for local runtime startup lock") time.sleep(0.05) + held.add(path) try: yield finally: + held.discard(path) if os.name == "posix": fcntl.flock(handle, fcntl.LOCK_UN) else: diff --git a/tests/test_runtime_placement.py b/tests/test_runtime_placement.py index fba5916..1a293cd 100644 --- a/tests/test_runtime_placement.py +++ b/tests/test_runtime_placement.py @@ -318,6 +318,41 @@ def cleanup(): assert token_file.read_text() == "replacement-token" +def test_idle_shutdown_waits_for_a_concurrent_startup(tmp_path, runtime_server): + import subprocess + import sys + from concurrent.futures import ThreadPoolExecutor + + from hai_agents_local.runtime import runtime as module + + proc = subprocess.Popen([sys.executable, "-c", "import time; time.sleep(30)"], start_new_session=True) + token_file = write_owner_only(token_file_path(runtime_server.port, cache_dir=tmp_path), runtime_server.token) + runtime = LocalRuntime( + base_url=f"http://127.0.0.1:{runtime_server.port}", + api_key=runtime_server.token, + pid=proc.pid, + version=None, + log_path=None, + owned=True, + cache_dir=tmp_path, + port=runtime_server.port, + proc=proc, + token_file=token_file, + ) + try: + with ThreadPoolExecutor() as pool: + with module._startup_lock(tmp_path, runtime_server.port, 1): + stopping = pool.submit(runtime.shutdown_if_idle) + with pytest.raises(TimeoutError): + stopping.result(timeout=0.2) + assert proc.poll() is None and not runtime_server.requests + assert stopping.result(timeout=10) + assert proc.poll() is not None and not token_file.exists() + finally: + proc.kill() + proc.wait() + + def test_spawner_never_overwrites_the_live_runtime_token(tmp_path, monkeypatch, runtime_server): import sys From 434d07ab0fe3e1a14ed3b16a7db38576c58c40d8 Mon Sep 17 00:00:00 2001 From: abonneth Date: Thu, 1 Oct 2026 19:41:26 +0200 Subject: [PATCH 27/35] fix(local): keep the create error when stopping its bridges fails --- src/hai_agents_local/sessions.py | 12 ++++++++++-- tests/test_local.py | 7 ++++++- 2 files changed, 16 insertions(+), 3 deletions(-) diff --git a/src/hai_agents_local/sessions.py b/src/hai_agents_local/sessions.py index 25e352d..289abc9 100644 --- a/src/hai_agents_local/sessions.py +++ b/src/hai_agents_local/sessions.py @@ -43,6 +43,14 @@ def _apply_runaway_budgets(kwargs: typing.Dict[str, typing.Any]) -> None: kwargs.setdefault("max_time_s", DEFAULT_LOCAL_MAX_TIME_S) +def _stop_bridges_keeping_error(session_ids: typing.Sequence[str]) -> None: + """Stop bridges inside an except block; a stop failure is logged so the handled error still propagates.""" + try: + stop_bridges(session_ids) + except Exception: + logger.warning("could not confirm local bridges stopped after a failed session create", exc_info=True) + + def _token_source(client_wrapper: typing.Any) -> TokenSource: return getattr(client_wrapper, "_async_token", None) or client_wrapper._get_api_key @@ -273,7 +281,7 @@ def create_session(self, **kwargs: typing.Any) -> typing.Any: try: session = super().create_session(**kwargs) except BaseException: - stop_bridges(started) + _stop_bridges_keeping_error(started) raise if bridges: cancel = _cancel_action(self._cancel_remote, bridges, session) @@ -343,7 +351,7 @@ async def create_session(self, **kwargs: typing.Any) -> typing.Any: try: session = await super().create_session(**kwargs) except BaseException: - await asyncio.to_thread(stop_bridges, started) + await asyncio.to_thread(_stop_bridges_keeping_error, started) raise if bridges: cancel = _cancel_action(self._cancel_remote, bridges, session) diff --git a/tests/test_local.py b/tests/test_local.py index 78b6afb..35249be 100644 --- a/tests/test_local.py +++ b/tests/test_local.py @@ -193,8 +193,13 @@ def test_named_agent_is_passed_through_without_bridges(self, monkeypatch): def test_create_session_failure_stops_newly_started_bridges(self, monkeypatch): monkeypatch.setenv(AUTO_BRIDGE_ENV_VAR, "1") stopped: list = [] + + def stop_stuck(ids): + stopped.extend(ids) + raise TimeoutError("bridge did not stop") + monkeypatch.setattr("hai_agents_local.sessions.ensure_bridges", lambda bridges: ["new-sid"]) - monkeypatch.setattr("hai_agents_local.sessions.stop_bridges", stopped.extend) + monkeypatch.setattr("hai_agents_local.sessions.stop_bridges", stop_stuck) monkeypatch.setattr( SessionsClient, "create_session", lambda self, **kw: (_ for _ in ()).throw(RuntimeError("api down")) ) From 292d46045d6f9e13a199a49070adac685c4c7be6 Mon Sep 17 00:00:00 2001 From: abonneth Date: Thu, 1 Oct 2026 19:41:26 +0200 Subject: [PATCH 28/35] fix(local): count queued sessions as active in idle shutdown --- src/hai_agents_local/runtime/runtime.py | 1 + 1 file changed, 1 insertion(+) diff --git a/src/hai_agents_local/runtime/runtime.py b/src/hai_agents_local/runtime/runtime.py index 46de55f..233b42a 100644 --- a/src/hai_agents_local/runtime/runtime.py +++ b/src/hai_agents_local/runtime/runtime.py @@ -407,6 +407,7 @@ def shutdown(self) -> None: # Statuses that mean the runtime still holds live session state a shutdown would destroy. # ("idle" sessions await user input but keep runtime state.) ACTIVE_SESSION_STATUSES: typing.ClassVar[typing.Tuple[str, ...]] = ( + "queued", "pending", "running", "paused", From 688453bbd90a277ee923b718e5b1659ced6f7285 Mon Sep 17 00:00:00 2001 From: abonneth Date: Fri, 2 Oct 2026 15:16:27 +0200 Subject: [PATCH 29/35] fix(runtime): stop an owned runtime whose idle probe fails --- src/hai_agents_local/runtime/runtime.py | 7 ++++-- tests/test_runtime_placement.py | 30 +++++++++++++++++++++++++ 2 files changed, 35 insertions(+), 2 deletions(-) diff --git a/src/hai_agents_local/runtime/runtime.py b/src/hai_agents_local/runtime/runtime.py index 233b42a..6f79404 100644 --- a/src/hai_agents_local/runtime/runtime.py +++ b/src/hai_agents_local/runtime/runtime.py @@ -420,8 +420,11 @@ def shutdown_if_idle(self, ignore: typing.Collection[str] = ()) -> bool: # Serialized with spawns and locked attaches only: a client that attached without the lock # can still start a session between the listing and the stop. with _startup_lock(self._cache_dir, self._port, SPAWN_TIMEOUT_S + STARTUP_LOCK_GRACE_S): - if self._hosts_active_session(ignore): - return False + try: + if self._hosts_active_session(ignore): + return False + except Exception: + logger.warning("idle probe failed on %s; stopping the owned runtime", self.base_url, exc_info=True) self.shutdown() return True diff --git a/tests/test_runtime_placement.py b/tests/test_runtime_placement.py index 1a293cd..4607dee 100644 --- a/tests/test_runtime_placement.py +++ b/tests/test_runtime_placement.py @@ -353,6 +353,36 @@ def test_idle_shutdown_waits_for_a_concurrent_startup(tmp_path, runtime_server): proc.wait() +def test_idle_shutdown_stops_an_owned_runtime_that_cannot_answer(tmp_path): + import socket + import subprocess + import sys + + with socket.socket() as probe: + probe.bind(("127.0.0.1", 0)) + port = probe.getsockname()[1] + proc = subprocess.Popen([sys.executable, "-c", "import time; time.sleep(30)"], start_new_session=True) + token_file = write_owner_only(token_file_path(port, cache_dir=tmp_path), "token") + runtime = LocalRuntime( + base_url=f"http://127.0.0.1:{port}", + api_key="token", + pid=proc.pid, + version=None, + log_path=None, + owned=True, + cache_dir=tmp_path, + port=port, + proc=proc, + token_file=token_file, + ) + try: + assert runtime.shutdown_if_idle() + assert proc.poll() is not None and not token_file.exists() + finally: + proc.kill() + proc.wait() + + def test_spawner_never_overwrites_the_live_runtime_token(tmp_path, monkeypatch, runtime_server): import sys From 5c66e7c757ffa5a41f4c44dc605eb7b0cd932e72 Mon Sep 17 00:00:00 2001 From: abonneth Date: Fri, 2 Oct 2026 17:03:37 +0200 Subject: [PATCH 30/35] refactor(runtime): drop unused LocalRuntime.force_kill and health --- src/hai_agents_local/runtime/runtime.py | 14 -------------- src/hai_agents_local/runtime/state.py | 10 ++-------- tests/test_runtime_placement.py | 2 +- 3 files changed, 3 insertions(+), 23 deletions(-) diff --git a/src/hai_agents_local/runtime/runtime.py b/src/hai_agents_local/runtime/runtime.py index 6f79404..10b7dbe 100644 --- a/src/hai_agents_local/runtime/runtime.py +++ b/src/hai_agents_local/runtime/runtime.py @@ -390,13 +390,6 @@ def async_http_client(self, timeout: typing.Optional[float] = None) -> httpx.Asy self.api_key, timeout=CLIENT_TIMEOUT_S if timeout is None else timeout, follow_redirects=True ) - def health(self) -> typing.Dict[str, typing.Any]: - """The /health JSON body; raises RuntimeUnhealthyError when the runtime is not answering.""" - payload = probe_health(self.base_url, self.api_key) - if payload is None: - raise RuntimeUnhealthyError(f"hai-agent-runtime at {self.base_url} is not answering /health") - return payload - def shutdown(self) -> None: """Gracefully stop the runtime this LocalRuntime spawned (SIGTERM group, grace, SIGKILL group).""" if not self.owned or self._proc is None: @@ -443,13 +436,6 @@ def _hosts_active_session(self, ignore: typing.Collection[str]) -> bool: return False page += 1 - def force_kill(self) -> None: - """Stop only the process this manager spawned; never trust a saved PID to claim ownership.""" - if not self.owned or self._proc is None: - raise LocalRuntimeError("cannot force-kill a borrowed runtime from a persisted PID") - terminate(self._proc) - self._cleanup_state_files() - def _cleanup_state_files(self) -> None: # Serialize compare-and-unlink with publication of a replacement runtime's state. with _startup_lock(self._cache_dir, self._port, SPAWN_TIMEOUT_S + STARTUP_LOCK_GRACE_S): diff --git a/src/hai_agents_local/runtime/state.py b/src/hai_agents_local/runtime/state.py index 2f3fc68..6c23eac 100644 --- a/src/hai_agents_local/runtime/state.py +++ b/src/hai_agents_local/runtime/state.py @@ -1,10 +1,4 @@ -"""Owner-only on-disk discovery state for locally spawned hai-agent-runtime processes. - -A spawner persists the generated bearer token and the runtime pid under the SDK -cache dir so a second process can attach (token) or force-kill (pid) without any -IPC. Both files drive privileged actions, so they are 0600 from the first byte -and refuse pre-planted symlinks. -""" +"""Owner-only (0600, symlink-refusing) token and pid files that let other local processes find a spawned runtime.""" from __future__ import annotations @@ -42,7 +36,7 @@ def token_file_path(port: int, *, cache_dir: typing.Optional[_PathInput] = None) def pid_file_path(port: int, *, cache_dir: typing.Optional[_PathInput] = None) -> pathlib.Path: - """Where a spawner publishes the runtime pid so force_kill() works from another process.""" + """Where a spawner publishes the runtime pid for out-of-process stop tools.""" return state_dir(cache_dir) / f"agent-pid-{port}" diff --git a/tests/test_runtime_placement.py b/tests/test_runtime_placement.py index 4607dee..12a1883 100644 --- a/tests/test_runtime_placement.py +++ b/tests/test_runtime_placement.py @@ -174,7 +174,7 @@ def test_only_the_runtime_holding_the_token_ever_receives_it(tmp_path, runtime_s with pytest.raises(BinaryIncompatibleError): attached.require_recipe("desktop") with pytest.raises(LocalRuntimeError): - attached.force_kill() + attached.shutdown() with Client.local(runtime=attached) as client: assert client.sessions.list_sessions().items == [] runtime_server.proof_token = "squatter-token" From 871349e3e8963ccf92d02107108564f07b967a95 Mon Sep 17 00:00:00 2001 From: abonneth Date: Fri, 2 Oct 2026 17:04:29 +0200 Subject: [PATCH 31/35] docs: trim local-agent README to usage --- README.md | 56 ++++++++----------------------------------------------- 1 file changed, 8 insertions(+), 48 deletions(-) diff --git a/README.md b/README.md index 62106ed..11064ef 100644 --- a/README.md +++ b/README.md @@ -67,25 +67,17 @@ print(result.answer) `result` is a `SessionRunResult`: `id`, `status`, `answer`, the accumulated `events`, and `final_changes`. -## Candidate local-agent support +## Local agents -`Client.local()` and `await AsyncClient.local()` start a local agent runtime or, -with `runtime=...`, use an already prepared one. -Agent placement and environment placement are separate: the client -selects where the agent runs; each agent environment selects `host="user_device"` -or `host="cloud"`. `Client()` continues to use the hosted Agents API. - -For development, install the candidate SDK and use a prepared HAI source checkout: +`Client.local()` and `await AsyncClient.local()` run the agent on this machine through a local agent runtime, +started on demand or passed in with `runtime=...`. Each environment picks `host="user_device"` or `host="cloud"`. +`Client()` keeps using the hosted Agents API. Closing the client stops a runtime it started, unless other clients still +have active sessions there. ```python from hai_agents import Client -with Client.local( - local_options={ - "command": ["/path/to/hai/.venv/bin/python", "-m", "hai_agent_runtime"], - "download": False, - }, -) as client: +with Client.local(local_options={"binary_path": "/path/to/hai-agent-runtime"}) as client: session = client.sessions.create_session( agent={ "name": "local-example", @@ -102,40 +94,8 @@ with Client.local( # Poll or steer the session here, before leaving the client context. ``` -The source runtime must support the `shared` recipe. The current pinned binary -0.1.8 does not; candidate source or an explicitly supplied compatible -`local_options["binary_path"]` is required until a compatible release is pinned. -The runtime inherits `HAI_API_KEY` for hosted inference. Local agent execution -does not by itself imply local inference: -`Client.local(inference=Inference.self_hosted(url, model=...))`, with `Inference` -from `hai_agents_local.runtime`, selects a model endpoint for a newly started -local runtime. Hosted agents do not currently accept that override. - -Closing the client shuts down a runtime it started, unless sessions from other -clients are still active there. An attached runtime remains owned by its caller. `cancel()` ends the agent session. For a cloud workstation, -an explicit `session_id` attaches to a caller-owned environment, which the caller -must eventually release. In the candidate shared recipe, automatically provisioned -cloud workstations survive agent cancellation and expire through the environment -manager after 30 minutes without commands (the runner's fixed deadline still -applies). Reattach using the `RunnerSessionEvent` ID in a new agent session; -cancellation does not revive the old agent session. The existing environment -manager API can release the workstation earlier. A manager that cannot confirm -the requested expiry is rejected and the newly created runner is deleted. This -is temporary compute retention, not a durable-storage or pause/resume guarantee. -The candidate source runtime accepts base64 -message attachments and exposes files shared by the agent through -`sessions.get_session_resource(id, "local", key)`. Download them before closing -the runtime: local shared resources expire with the retained session. The limits -are 50 MiB per file, 64 MiB and 128 shared files per session. - -## Runtime release maintenance - -The SDK runtime manifest is maintained independently of schema generation. After -publishing and verifying a compatible runtime, run `scripts/bump_runtime.py` with -`--version` and one `--sha PLATFORM=SHA256` for every platform already in the -manifest. Partial updates are rejected so a new URL cannot retain an old digest. -Schema generation never touches it. The release workflow still opens -legacy CLI pin PRs; retarget it only after the SDK/CLI migration has shipped. +Inference stays hosted (`HAI_API_KEY`) unless you pass `inference=Inference.self_hosted(url, model=...)` +(`from hai_agents_local.runtime import Inference`). ## How a session works From 59a1d3b893948cf89cd5bdea5f8a092ec2225f83 Mon Sep 17 00:00:00 2001 From: abonneth Date: Fri, 2 Oct 2026 17:05:55 +0200 Subject: [PATCH 32/35] refactor(local): set verify_runtime once where local sessions localize --- src/hai_agents_local/bridge.py | 6 ++---- src/hai_agents_local/browser.py | 5 +---- src/hai_agents_local/desktop.py | 5 +---- src/hai_agents_local/routing.py | 20 +++++++------------- src/hai_agents_local/sessions.py | 14 ++++++++------ src/hai_agents_local/workstation.py | 5 +---- tests/test_runtime_placement.py | 12 +++++++----- 7 files changed, 27 insertions(+), 40 deletions(-) diff --git a/src/hai_agents_local/bridge.py b/src/hai_agents_local/bridge.py index e625dd7..9416a3d 100644 --- a/src/hai_agents_local/bridge.py +++ b/src/hai_agents_local/bridge.py @@ -62,6 +62,8 @@ class LocalBridge(ABC, Generic[DriverT]): environment_kind: ClassVar[str] startup_hint: ClassVar[str | None] = None """Appended to the manager's not-ready timeout error; names the common cause of a hung startup.""" + verify_runtime: bool = False + """Reject responses not HMAC-proven with api_key; requires api_key to be the local runtime's token string.""" def __init__( self, @@ -70,12 +72,9 @@ def __init__( api_key: TokenSource, base_url: str | None = None, session_id: str | None = None, - verify_runtime: bool = False, ) -> None: if not api_key: raise ValueError("api_key is required") - if verify_runtime and not isinstance(api_key, str): - raise ValueError("verify_runtime needs api_key to be the local runtime's token string") if session_id is not None: try: uuid.UUID(session_id) @@ -85,7 +84,6 @@ def __init__( self.api_key = api_key self.base_url = base_url or default_base_url() self.session_id = session_id or str(uuid.uuid4()) - self.verify_runtime = verify_runtime self.ready = threading.Event() self.on_crash: Callable[[], None] | None = None self._driver: DriverT | None = None diff --git a/src/hai_agents_local/browser.py b/src/hai_agents_local/browser.py index ca43a1a..e522588 100644 --- a/src/hai_agents_local/browser.py +++ b/src/hai_agents_local/browser.py @@ -54,11 +54,8 @@ def __init__( api_key: TokenSource, base_url: str | None = None, session_id: str | None = None, - verify_runtime: bool = False, ) -> None: - super().__init__( - environment_id, api_key=api_key, base_url=base_url, session_id=session_id, verify_runtime=verify_runtime - ) + super().__init__(environment_id, api_key=api_key, base_url=base_url, session_id=session_id) self.debugging_port = debugging_port def create_driver(self) -> SeleniumWebDriver: diff --git a/src/hai_agents_local/desktop.py b/src/hai_agents_local/desktop.py index fc4b1a8..0339b98 100644 --- a/src/hai_agents_local/desktop.py +++ b/src/hai_agents_local/desktop.py @@ -62,15 +62,12 @@ def __init__( api_key: TokenSource, base_url: str | None = None, session_id: str | None = None, - verify_runtime: bool = False, max_width: int | None = DEFAULT_MAX_WIDTH, max_height: int | None = None, image_format: ImageFormat | None = DEFAULT_IMAGE_FORMAT, quality: int = DEFAULT_QUALITY, ) -> None: - super().__init__( - environment_id, api_key=api_key, base_url=base_url, session_id=session_id, verify_runtime=verify_runtime - ) + super().__init__(environment_id, api_key=api_key, base_url=base_url, session_id=session_id) self.max_width = max_width self.max_height = max_height self.image_format = image_format diff --git a/src/hai_agents_local/routing.py b/src/hai_agents_local/routing.py index 9056f1e..4c6f8f3 100644 --- a/src/hai_agents_local/routing.py +++ b/src/hai_agents_local/routing.py @@ -25,42 +25,36 @@ def localize_agent( - agent: AgentLike, *, api_key: TokenSource, base_url: str | None = None, verify_runtime: bool = False + agent: AgentLike, *, api_key: TokenSource, base_url: str | None = None ) -> tuple[AgentLike, list[LocalBridge]]: """Copy of the agent where every unclaimed user_device environment is stamped with the session id of a freshly built bridge, plus those bridges. Environments that already carry a session_id are assumed to be served elsewhere and left alone, as are string agent references.""" bridges: list[LocalBridge] = [] - return _localize_agent(agent, bridges, api_key, base_url, verify_runtime), bridges + return _localize_agent(agent, bridges, api_key, base_url), bridges def _localize_agent( - agent: AgentLike, bridges: list[LocalBridge], api_key: TokenSource, base_url: str | None, verify_runtime: bool + agent: AgentLike, bridges: list[LocalBridge], api_key: TokenSource, base_url: str | None ) -> AgentLike: if isinstance(agent, str): return agent changes: dict[str, Any] = {} environments = _read(agent, "environments") if isinstance(environments, (list, tuple)): - localized_envs = [ - _localize_environment(env, bridges, api_key, base_url, verify_runtime) for env in environments - ] + localized_envs = [_localize_environment(env, bridges, api_key, base_url) for env in environments] if _any_replaced(localized_envs, environments): changes["environments"] = localized_envs subagents = _read(agent, "subagents") if isinstance(subagents, (list, tuple)): - localized_subs = [_localize_agent(sub, bridges, api_key, base_url, verify_runtime) for sub in subagents] + localized_subs = [_localize_agent(sub, bridges, api_key, base_url) for sub in subagents] if _any_replaced(localized_subs, subagents): changes["subagents"] = localized_subs return _replace(agent, **changes) if changes else agent def _localize_environment( - env: EnvironmentLike, - bridges: list[LocalBridge], - api_key: TokenSource, - base_url: str | None, - verify_runtime: bool, + env: EnvironmentLike, bridges: list[LocalBridge], api_key: TokenSource, base_url: str | None ) -> EnvironmentLike: kind = _local_kind(env) if kind is None or _read(env, "session_id"): @@ -71,7 +65,7 @@ def _localize_environment( f"serve one local {kind}; give the extra environments an explicit session_id and serve each " "from its own machine with `hai local browser|desktop|workstation --session-id `" ) - bridge = BRIDGE_TYPES[kind](_read(env, "id"), api_key=api_key, base_url=base_url, verify_runtime=verify_runtime) + bridge = BRIDGE_TYPES[kind](_read(env, "id"), api_key=api_key, base_url=base_url) bridges.append(bridge) return _replace(env, kind=kind, session_id=bridge.session_id) diff --git a/src/hai_agents_local/sessions.py b/src/hai_agents_local/sessions.py index 289abc9..d79656b 100644 --- a/src/hai_agents_local/sessions.py +++ b/src/hai_agents_local/sessions.py @@ -94,12 +94,14 @@ def _localize( if agent is None or isinstance(agent, str) or not auto_bridges_enabled(): return [] _warn_if_overrides_target_user_device(kwargs) - credentials: typing.Dict[str, typing.Any] = ( - {"api_key": runtime.api_key, "base_url": runtime.base_url, "verify_runtime": True} - if runtime is not None - else {"api_key": _token_source(client_wrapper), "base_url": client_wrapper.get_base_url()} - ) - localized, bridges = localize_agent(agent, **credentials) + if runtime is None: + localized, bridges = localize_agent( + agent, api_key=_token_source(client_wrapper), base_url=client_wrapper.get_base_url() + ) + else: + localized, bridges = localize_agent(agent, api_key=runtime.api_key, base_url=runtime.base_url) + for bridge in bridges: + bridge.verify_runtime = True kwargs["agent"] = localized return bridges diff --git a/src/hai_agents_local/workstation.py b/src/hai_agents_local/workstation.py index ed079a9..f2724be 100644 --- a/src/hai_agents_local/workstation.py +++ b/src/hai_agents_local/workstation.py @@ -36,11 +36,8 @@ def __init__( api_key: TokenSource, base_url: str | None = None, session_id: str | None = None, - verify_runtime: bool = False, ) -> None: - super().__init__( - environment_id, api_key=api_key, base_url=base_url, session_id=session_id, verify_runtime=verify_runtime - ) + super().__init__(environment_id, api_key=api_key, base_url=base_url, session_id=session_id) self.workspace = Path(workspace).expanduser() if workspace else Path.home() / "hai" / self.session_id def preflight(self) -> None: diff --git a/tests/test_runtime_placement.py b/tests/test_runtime_placement.py index 12a1883..69af01b 100644 --- a/tests/test_runtime_placement.py +++ b/tests/test_runtime_placement.py @@ -183,14 +183,16 @@ def test_only_the_runtime_holding_the_token_ever_receives_it(tmp_path, runtime_s @pytest.mark.asyncio -async def test_bridge_never_serves_an_unproven_runtime(runtime_server): - from hai_agents_local.routing import localize_agent +async def test_bridge_never_serves_an_unproven_runtime(monkeypatch, runtime_server): + from types import SimpleNamespace + + from hai_agents_local.sessions import _localize + monkeypatch.delenv("HAI_AUTO_BRIDGE", raising=False) runtime_server.proof_token = "squatter-token" + runtime = SimpleNamespace(api_key="local-token", base_url=f"http://127.0.0.1:{runtime_server.port}") agent = {"environments": [{"id": "workstation", "kind": "workstation", "host": "user_device"}]} - _, [bridge] = localize_agent( - agent, api_key="local-token", base_url=f"http://127.0.0.1:{runtime_server.port}", verify_runtime=True - ) + [bridge] = _localize(None, runtime, {"agent": agent}) bridge.create_driver = lambda: pytest.fail("a driver started for an unproven runtime") with pytest.raises(LocalRuntimeError, match="did not prove"): await bridge.run() From 3b15d981293be74f1d8ae6498e7b5d41d975970c Mon Sep 17 00:00:00 2001 From: abonneth Date: Fri, 2 Oct 2026 17:08:23 +0200 Subject: [PATCH 33/35] refactor(runtime): keep the runtime pin in pin.json; bump script loads, validates, dumps --- scripts/bump_runtime.py | 89 +++++++----------------- src/hai_agents_local/runtime/manifest.py | 29 +++----- src/hai_agents_local/runtime/pin.json | 7 ++ tests/test_bump_runtime.py | 38 +++++----- 4 files changed, 66 insertions(+), 97 deletions(-) create mode 100644 src/hai_agents_local/runtime/pin.json diff --git a/scripts/bump_runtime.py b/scripts/bump_runtime.py index 6cb0e78..21e5e2e 100644 --- a/scripts/bump_runtime.py +++ b/scripts/bump_runtime.py @@ -1,74 +1,39 @@ -"""Rewrite the pinned hai-agent-runtime version + per-platform sha256 in the SDK runtime manifest.""" +"""Pin a hai-agent-runtime release: its version plus a fresh sha256 for every platform in the pin file.""" from __future__ import annotations import argparse +import json import re import sys -from dataclasses import dataclass from pathlib import Path -# Stdlib only on purpose: this runs as `python scripts/bump_runtime.py` in a -# checkout with no dependencies installed, so a third-party import would break -# the bump step. -RUNTIME_INSTALL = Path(__file__).parents[1] / "src" / "hai_agents_local" / "runtime" / "manifest.py" -_SHA256_RE = re.compile(r"^[0-9a-f]{64}$") -_MANIFEST_FILENAME_RE = re.compile(r'"hai-agent-runtime-([^".]+)\.zip"') - - -@dataclass(frozen=True) -class RuntimeBump: - version: str - shas: dict[str, str] # platform key (e.g. darwin-arm64) -> sha256 hex - - def __post_init__(self) -> None: - if not re.fullmatch(r"[0-9]+\.[0-9]+\.[0-9]+(?:[-+][0-9A-Za-z.-]+)?", self.version): - raise ValueError("runtime version must be a release version, e.g. 0.1.13") - for platform, sha in self.shas.items(): - if not _SHA256_RE.fullmatch(sha) or sha == "0" * 64: - raise ValueError(f"{platform}: {sha!r} is not a lowercase 64-char sha256") - - -def _filename_for(platform: str) -> str: - return f"hai-agent-runtime-{platform}.zip" - - -def _manifest_platforms(source: str) -> set[str]: - """Platform keys that have a published artifact literal in `source`.""" - return set(_MANIFEST_FILENAME_RE.findall(source)) - - -def apply_bump(source: str, bump: RuntimeBump) -> str: - """Return `source` with PINNED_RUNTIME_VERSION and the manifest digests replaced; raises if any anchor is missing.""" - published = _manifest_platforms(source) - extra = bump.shas.keys() - published +# Stdlib only: runs in a bare checkout with no dependencies installed. +PIN_FILE = Path(__file__).parents[1] / "src" / "hai_agents_local" / "runtime" / "pin.json" +_VERSION_RE = re.compile(r"[0-9]+\.[0-9]+\.[0-9]+(?:[-+][0-9A-Za-z.-]+)?") +_SHA256_RE = re.compile(r"[0-9a-f]{64}") +_PLACEHOLDER_SHA256 = "0" * 64 + + +def apply_bump(pin: dict, version: str, shas: dict[str, str]) -> dict: + """`pin` moved to `version`; raises unless `shas` holds a real digest for exactly the pinned platforms.""" + if not _VERSION_RE.fullmatch(version): + raise ValueError("runtime version must be a release version, e.g. 0.1.13") + for platform, sha in shas.items(): + if not _SHA256_RE.fullmatch(sha) or sha == _PLACEHOLDER_SHA256: + raise ValueError(f"{platform}: {sha!r} is not a lowercase 64-char sha256") + published = set(pin["sha256"]) + extra = shas.keys() - published if extra: raise ValueError(f"no manifest entry for platform(s): {sorted(extra)}") - # The version is a single literal feeding every derived URL, so any published - # platform left without a fresh sha would keep a stale digest at the new - # version's URL and fail verification on download. Refuse the partial bump. - missing = published - bump.shas.keys() + # A platform left on its old digest would fail verification at the new version's URL. + missing = published - shas.keys() if missing: raise ValueError(f"missing sha for published platform(s): {sorted(missing)}") - - updated, count = re.subn( - r'PINNED_RUNTIME_VERSION = "[^"]*"', - f'PINNED_RUNTIME_VERSION = "{bump.version}"', - source, - ) - if count != 1: - raise ValueError(f"expected exactly one PINNED_RUNTIME_VERSION assignment, found {count}") - - for platform, sha in bump.shas.items(): - filename = _filename_for(platform) - pattern = re.compile(rf'("{re.escape(filename)}",\s*(?:#[^\n]*\n\s*)*")[0-9a-fA-F]{{64}}(")') - updated, count = pattern.subn(rf"\g<1>{sha}\g<2>", updated) - if count != 1: - raise ValueError(f"expected exactly one sha256 literal for {filename}, found {count}") - return updated + return {**pin, "version": version, "sha256": {platform: shas[platform] for platform in pin["sha256"]}} -def _parse_args(argv: list[str]) -> RuntimeBump: +def _parse_args(argv: list[str]) -> tuple[str, dict[str, str]]: parser = argparse.ArgumentParser(description=__doc__) parser.add_argument("--version", required=True) parser.add_argument( @@ -87,14 +52,14 @@ def _parse_args(argv: list[str]) -> RuntimeBump: if platform in shas: parser.error(f"duplicate --sha for {platform}") shas[platform] = sha.lower() - return RuntimeBump(version=args.version, shas=shas) + return args.version, shas def main(argv: list[str]) -> int: - bump = _parse_args(argv) - source = RUNTIME_INSTALL.read_text() - RUNTIME_INSTALL.write_text(apply_bump(source, bump)) - print(f"bumped runtime to {bump.version} ({', '.join(sorted(bump.shas))})") + version, shas = _parse_args(argv) + pin = apply_bump(json.loads(PIN_FILE.read_text(encoding="utf-8")), version, shas) + PIN_FILE.write_text(json.dumps(pin, indent=2) + "\n", encoding="utf-8") + print(f"bumped runtime to {version} ({', '.join(sorted(shas))})") return 0 diff --git a/src/hai_agents_local/runtime/manifest.py b/src/hai_agents_local/runtime/manifest.py index e4f2d72..acab3dc 100644 --- a/src/hai_agents_local/runtime/manifest.py +++ b/src/hai_agents_local/runtime/manifest.py @@ -1,19 +1,14 @@ -"""Pinned hai-agent-runtime version and per-platform artifact digests. - -This module is the SDK's single runtime pin: a runtime release bumps -PINNED_RUNTIME_VERSION and MANIFEST here with scripts/bump_runtime.py. Artifacts live under an -immutable version-scoped CDN prefix, so an edge can never serve stale bytes. -""" +"""Pinned hai-agent-runtime artifacts, loaded from pin.json (updated by scripts/bump_runtime.py).""" from __future__ import annotations import dataclasses +import json +import pathlib import platform import sys import typing -# TODO: pin a runtime release that serves the shared recipe, with its artifact hashes below. -PINNED_RUNTIME_VERSION = "0.1.8" RUNTIME_CDN_BASE = "https://assets.hcompanyprod.fr/hai-agent-runtime" # Guard value: published manifest entries must never use it (every download would fail verification). PLACEHOLDER_SHA256 = "0" * 64 @@ -26,20 +21,16 @@ class RuntimeArtifact: sha256: str -def _artifact(filename: str, sha256: str) -> RuntimeArtifact: - """A published release file resolved to its pinned, version-scoped CDN URL.""" - return RuntimeArtifact(url=f"{RUNTIME_CDN_BASE}/{PINNED_RUNTIME_VERSION}/{filename}", sha256=sha256) +_PIN = json.loads(pathlib.Path(__file__).with_name("pin.json").read_text(encoding="utf-8")) +# TODO: pin a runtime release that serves the shared recipe. +PINNED_RUNTIME_VERSION: str = _PIN["version"] MANIFEST: typing.Dict[str, RuntimeArtifact] = { - "darwin-arm64": _artifact( - "hai-agent-runtime-darwin-arm64.zip", - "1aed0055898116732aee031dc4a1235782b2909ee51e0367e2d50bb3be6671c9", - ), - "windows-x86_64": _artifact( - "hai-agent-runtime-windows-x86_64.zip", - "4e6b2bcd42af2bb6b22197fcde947327497f5c62fd60d48bc9037730d80dc691", - ), + platform_name: RuntimeArtifact( + url=f"{RUNTIME_CDN_BASE}/{PINNED_RUNTIME_VERSION}/hai-agent-runtime-{platform_name}.zip", sha256=sha256 + ) + for platform_name, sha256 in _PIN["sha256"].items() } UNIMPLEMENTED_PLATFORMS: typing.Dict[str, str] = { diff --git a/src/hai_agents_local/runtime/pin.json b/src/hai_agents_local/runtime/pin.json new file mode 100644 index 0000000..7e022d5 --- /dev/null +++ b/src/hai_agents_local/runtime/pin.json @@ -0,0 +1,7 @@ +{ + "version": "0.1.8", + "sha256": { + "darwin-arm64": "1aed0055898116732aee031dc4a1235782b2909ee51e0367e2d50bb3be6671c9", + "windows-x86_64": "4e6b2bcd42af2bb6b22197fcde947327497f5c62fd60d48bc9037730d80dc691" + } +} diff --git a/tests/test_bump_runtime.py b/tests/test_bump_runtime.py index 476735b..a635696 100644 --- a/tests/test_bump_runtime.py +++ b/tests/test_bump_runtime.py @@ -9,30 +9,36 @@ import pytest ROOT = Path(__file__).resolve().parents[1] +RUNTIME = Path("src") / "hai_agents_local" / "runtime" -@pytest.mark.parametrize("complete", [True, False]) -def test_release_pin_update_is_complete_or_leaves_manifest_unchanged(tmp_path, complete): - script = tmp_path / "scripts" / "bump_runtime.py" - manifest = tmp_path / "src" / "hai_agents_local" / "runtime" / "manifest.py" - script.parent.mkdir(parents=True) - manifest.parent.mkdir(parents=True) - shutil.copyfile(ROOT / "scripts" / "bump_runtime.py", script) - shutil.copyfile(ROOT / "src" / "hai_agents_local" / "runtime" / "manifest.py", manifest) - before = manifest.read_bytes() +@pytest.mark.parametrize("case", ["complete", "partial", "unknown-platform", "placeholder-sha"]) +def test_release_pin_update_is_complete_or_leaves_pin_unchanged(tmp_path, case): + for relative in (Path("scripts") / "bump_runtime.py", RUNTIME / "manifest.py", RUNTIME / "pin.json"): + (tmp_path / relative).parent.mkdir(parents=True, exist_ok=True) + shutil.copyfile(ROOT / relative, tmp_path / relative) + pin, manifest = tmp_path / RUNTIME / "pin.json", tmp_path / RUNTIME / "manifest.py" + before = pin.read_bytes() original = runpy.run_path(str(manifest))["MANIFEST"] shas = {platform: f"{index:064x}" for index, platform in enumerate(original, start=1)} - args = [sys.executable, str(script), "--version", "9.8.7"] - for platform, sha in list(shas.items())[: len(shas) if complete else -1]: + if case == "partial": + shas.popitem() + elif case == "unknown-platform": + shas["plan9-mips"] = "f" * 64 + elif case == "placeholder-sha": + shas[next(iter(shas))] = "0" * 64 + args = [sys.executable, str(tmp_path / "scripts" / "bump_runtime.py"), "--version", "9.8.7"] + for platform, sha in shas.items(): args.extend(["--sha", f"{platform}={sha}"]) result = subprocess.run(args, capture_output=True, text=True) - if not complete: + if case != "complete": assert result.returncode != 0 - assert manifest.read_bytes() == before + assert pin.read_bytes() == before return assert result.returncode == 0, result.stderr - updated = runpy.run_path(str(manifest))["MANIFEST"] - assert set(updated) == set(original) - for platform, artifact in updated.items(): + updated = runpy.run_path(str(manifest)) + assert updated["PINNED_RUNTIME_VERSION"] == "9.8.7" + assert set(updated["MANIFEST"]) == set(original) + for platform, artifact in updated["MANIFEST"].items(): assert artifact.sha256 == shas[platform] assert artifact.url.endswith(f"/9.8.7/hai-agent-runtime-{platform}.zip") From e2b0410a11705ac7633be7d574093e06188dae1e Mon Sep 17 00:00:00 2001 From: abonneth Date: Fri, 2 Oct 2026 17:10:18 +0200 Subject: [PATCH 34/35] refactor(local): share sync/async session bookkeeping and close-failure handling --- src/hai_agents_local/sessions.py | 102 +++++++++++++++---------------- 1 file changed, 51 insertions(+), 51 deletions(-) diff --git a/src/hai_agents_local/sessions.py b/src/hai_agents_local/sessions.py index d79656b..c0b49d7 100644 --- a/src/hai_agents_local/sessions.py +++ b/src/hai_agents_local/sessions.py @@ -5,6 +5,7 @@ import asyncio import atexit import concurrent.futures +import contextlib import functools import inspect import json @@ -225,15 +226,32 @@ def _cancel_sessions_at_exit() -> None: STOPPED_CANCEL_STATUSES = frozenset({404, 409}) -def _live_sessions(owned_bridges: typing.Dict[str, typing.List[str]]) -> typing.List[str]: - """Forget sessions whose bridges all stopped, since each such stop already ended or cancelled its session.""" - for session_id, bridge_ids in list(owned_bridges.items()): - if not serving_bridges(bridge_ids): - del owned_bridges[session_id] - return list(owned_bridges) +class _CloseFailures: + """Cancel errors collected while closing; a session that already ended is not a failure.""" + def __init__(self) -> None: + self._errors: typing.List[Exception] = [] + + @contextlib.contextmanager + def cancelling(self, session_id: str) -> typing.Iterator[None]: + try: + yield + except ApiError as error: + if error.status_code not in STOPPED_CANCEL_STATUSES: + self._errors.append(error) + else: + _deregister_exit_cancel(session_id) + except Exception as error: + self._errors.append(error) + + def raise_any(self) -> None: + if self._errors: + raise RuntimeError("Could not confirm all client-owned sessions stopped") from self._errors[0] + + +class _LocalSessionsState: + """Bridge and session bookkeeping shared by the sync and async local sessions clients.""" -class LocalSessionsClient(SessionsClient): def __init__( self, *, client_wrapper: typing.Any, runtime: typing.Optional[LocalRuntime] = None, auto_bridges: bool = True ) -> None: @@ -245,20 +263,27 @@ def __init__( # Sessions this client created on a local runtime; they never keep that runtime alive past close(). self.own_session_ids: typing.Set[str] = set() + def _live_sessions(self) -> typing.List[str]: + """Forget sessions whose bridges all stopped, since each such stop already ended or cancelled its session.""" + for session_id, bridge_ids in list(self._owned_bridges.items()): + if not serving_bridges(bridge_ids): + del self._owned_bridges[session_id] + return list(self._owned_bridges) + + def _track(self, session: typing.Any, started: typing.List[str]) -> None: + if self._runtime is not None: + self.own_session_ids.add(str(session.id)) + if started: + self._owned_bridges[str(session.id)] = started + + +class LocalSessionsClient(_LocalSessionsState, SessionsClient): def close(self) -> None: - failures = [] - for session_id in _live_sessions(self._owned_bridges): - try: + failures = _CloseFailures() + for session_id in self._live_sessions(): + with failures.cancelling(session_id): self.cancel_session(session_id) - except ApiError as error: - if error.status_code not in STOPPED_CANCEL_STATUSES: - failures.append(error) - else: - _deregister_exit_cancel(session_id) - except Exception as error: - failures.append(error) - if failures: - raise RuntimeError("Could not confirm all client-owned sessions stopped") from failures[0] + failures.raise_any() def cancel_session(self, id: str, *, request_options: typing.Optional[RequestOptions] = None) -> None: # Stop local execution even if the remote cancellation cannot be delivered. @@ -295,39 +320,17 @@ def create_session(self, **kwargs: typing.Any) -> typing.Any: if stop_watcher is not None and not stop_watcher.active: # A stop was filed while bridges or the session were starting; apply it now. _panic_stop() - if self._runtime is not None: - self.own_session_ids.add(str(session.id)) - if started: - self._owned_bridges[str(session.id)] = started + self._track(session, started) return session -class LocalAsyncSessionsClient(AsyncSessionsClient): - def __init__( - self, *, client_wrapper: typing.Any, runtime: typing.Optional[LocalRuntime] = None, auto_bridges: bool = True - ) -> None: - super().__init__(client_wrapper=client_wrapper) - self._runtime = runtime - self._auto_bridges = auto_bridges - self._cancel_remote: RemoteCancel = functools.partial(_cancel_remote_session, client_wrapper, runtime) - self._owned_bridges: typing.Dict[str, typing.List[str]] = {} - # Sessions this client created on a local runtime; they never keep that runtime alive past close(). - self.own_session_ids: typing.Set[str] = set() - +class LocalAsyncSessionsClient(_LocalSessionsState, AsyncSessionsClient): async def aclose(self) -> None: - failures = [] - for session_id in _live_sessions(self._owned_bridges): - try: + failures = _CloseFailures() + for session_id in self._live_sessions(): + with failures.cancelling(session_id): await self.cancel_session(session_id) - except ApiError as error: - if error.status_code not in STOPPED_CANCEL_STATUSES: - failures.append(error) - else: - _deregister_exit_cancel(session_id) - except Exception as error: - failures.append(error) - if failures: - raise RuntimeError("Could not confirm all client-owned sessions stopped") from failures[0] + failures.raise_any() async def cancel_session(self, id: str, *, request_options: typing.Optional[RequestOptions] = None) -> None: owned = self._owned_bridges.get(str(id), []) @@ -366,8 +369,5 @@ async def create_session(self, **kwargs: typing.Any) -> typing.Any: if stop_watcher is not None and not stop_watcher.active: # A stop was filed while bridges or the session were starting; apply it now. await asyncio.to_thread(_panic_stop) - if self._runtime is not None: - self.own_session_ids.add(str(session.id)) - if started: - self._owned_bridges[str(session.id)] = started + self._track(session, started) return session From 9202c2fbfe0cce30f2cbbfa0238b5b3461223cb2 Mon Sep 17 00:00:00 2001 From: abonneth Date: Fri, 2 Oct 2026 17:12:18 +0200 Subject: [PATCH 35/35] style(runtime): one-line docstrings, drop narrating comments --- src/hai_agents_local/runtime/__init__.py | 5 +---- src/hai_agents_local/runtime/install.py | 6 +----- src/hai_agents_local/runtime/process.py | 10 ++-------- src/hai_agents_local/runtime/runtime.py | 21 ++++----------------- src/hai_agents_local/runtime/state.py | 1 - tests/test_runtime_state.py | 10 +--------- 6 files changed, 9 insertions(+), 44 deletions(-) diff --git a/src/hai_agents_local/runtime/__init__.py b/src/hai_agents_local/runtime/__init__.py index 10207f6..65a16b9 100644 --- a/src/hai_agents_local/runtime/__init__.py +++ b/src/hai_agents_local/runtime/__init__.py @@ -1,7 +1,4 @@ -"""Local agent runtime management: install, find, start, attach to and verify a hai-agent-runtime binary. - -Imported lazily by ``Client.local`` so remote-only users pay nothing for it. -""" +"""Local agent runtime management, imported lazily by ``Client.local`` so remote-only users pay nothing for it.""" from .acquire import acquire_runtime, acquire_runtime_async from .errors import ( diff --git a/src/hai_agents_local/runtime/install.py b/src/hai_agents_local/runtime/install.py index d928013..4f4953b 100644 --- a/src/hai_agents_local/runtime/install.py +++ b/src/hai_agents_local/runtime/install.py @@ -1,8 +1,4 @@ -"""Verified download and atomic install of the hai-agent-runtime binary. - -The SDK is a library, so consent is the caller's ``download=True`` and progress -is plain logging. -""" +"""Verified download and atomic install of the hai-agent-runtime binary; consent is the caller's ``download=True``.""" from __future__ import annotations diff --git a/src/hai_agents_local/runtime/process.py b/src/hai_agents_local/runtime/process.py index 62b89c7..6228302 100644 --- a/src/hai_agents_local/runtime/process.py +++ b/src/hai_agents_local/runtime/process.py @@ -53,12 +53,7 @@ def probe_health(base_url: str, token: str) -> typing.Optional[typing.Dict[str, def spawn(cmd: typing.List[str], *, env: typing.Dict[str, str], log_path: pathlib.Path) -> subprocess.Popen: - """Start the runtime in its own process group with stderr to `log_path`. - - stderr goes to a file, not a pipe: nobody drains a pipe after spawn, so the - buffer would fill and block. Own process group so we can reap grandchildren - (e.g. desktop helpers) the binary may spawn. - """ + """Start the runtime in its own process group (to reap grandchildren); stderr to a file, as nobody drains a pipe.""" log_path.parent.mkdir(parents=True, exist_ok=True) logger.info("spawning hai-agent-runtime: %s (stderr -> %s)", " ".join(cmd), log_path) with log_path.open("wb") as log_file: # child inherits the fd; the parent handle can close right away @@ -136,8 +131,7 @@ def _signal(proc: subprocess.Popen, *, force: bool) -> bool: if os.name == "posix": return _killpg_posix(proc.pid, signal.SIGKILL if force else signal.SIGTERM) try: - # Windows has no portable graceful process-group signal. /F is required - # to ensure an executable launched through a .cmd shim cannot outlive it. + # No portable graceful group signal on Windows; /F keeps a .cmd-shim child from outliving it. subprocess.run(["taskkill", "/F", "/T", "/PID", str(proc.pid)], check=True, capture_output=True) except (OSError, subprocess.CalledProcessError): return False diff --git a/src/hai_agents_local/runtime/runtime.py b/src/hai_agents_local/runtime/runtime.py index 10b7dbe..77eee80 100644 --- a/src/hai_agents_local/runtime/runtime.py +++ b/src/hai_agents_local/runtime/runtime.py @@ -167,8 +167,7 @@ def ensure_started( if command is not None and (not command or binary_path is not None): raise ValueError("command must be nonempty and cannot be combined with binary_path") - # Downloads are atomically installed and can outlast the process health budget. - # Keep them outside the port lock; recheck attachment before spawning. + # Downloads can outlast the lock budget: resolve outside it, then recheck attachment under it. cmd = ( list(command) if command is not None @@ -322,17 +321,7 @@ def _child_env( spawn_env: typing.Optional[typing.Dict[str, str]], inherit_env: bool, ) -> typing.Dict[str, str]: - """Child env: inherited-plus-overlay by default, caller-verbatim with inherit_env=False. - - Inheriting os.environ passes the model-gateway HAI_API_KEY / HAI_BASE_URL through to the - binary (without them local sessions cannot run inference) and forwards caller flags such as - HAI_AGENT_RUNTIME_MODEL/FAKE/FAST/RUNS_DIR. inherit_env=False takes spawn_env as the - complete base environment instead, for callers that must *remove* inherited keys, which an - overlay cannot express (e.g. stripping HAI_API_KEY for self-hosted base URLs). The - generated local bearer and the cloud HAI_API_KEY are different credentials: the token below - is the only local bearer, and the cloud key is never used to authenticate against the local - runtime. Port and token are set last in both modes so caller input never clobbers them. - """ + """os.environ plus spawn_env, or spawn_env verbatim to drop inherited keys; port and token always win.""" env = {**os.environ, **(spawn_env or {})} if inherit_env else dict(spawn_env or {}) env[PORT_ENV] = str(port) env[AUTH_TOKEN_ENV] = token @@ -397,8 +386,7 @@ def shutdown(self) -> None: terminate(self._proc) self._cleanup_state_files() - # Statuses that mean the runtime still holds live session state a shutdown would destroy. - # ("idle" sessions await user input but keep runtime state.) + # Statuses whose runtime state a shutdown would destroy ("idle" awaits user input but keeps state). ACTIVE_SESSION_STATUSES: typing.ClassVar[typing.Tuple[str, ...]] = ( "queued", "pending", @@ -410,8 +398,7 @@ def shutdown(self) -> None: def shutdown_if_idle(self, ignore: typing.Collection[str] = ()) -> bool: """Stop the owned runtime unless it hosts an active session outside ``ignore``; True when it was stopped.""" - # Serialized with spawns and locked attaches only: a client that attached without the lock - # can still start a session between the listing and the stop. + # Only spawns and locked attaches are serialized; an unlocked attacher can still race the listing. with _startup_lock(self._cache_dir, self._port, SPAWN_TIMEOUT_S + STARTUP_LOCK_GRACE_S): try: if self._hosts_active_session(ignore): diff --git a/src/hai_agents_local/runtime/state.py b/src/hai_agents_local/runtime/state.py index 6c23eac..175258b 100644 --- a/src/hai_agents_local/runtime/state.py +++ b/src/hai_agents_local/runtime/state.py @@ -10,7 +10,6 @@ CACHE_DIR_ENV = "HAI_AGENT_LOCAL_CACHE_DIR" DEFAULT_CACHE_DIR = pathlib.Path.home() / ".hai" / "agent-runtime" -# The shared well-known local runtime port. DEFAULT_PORT = 18795 _PathInput = typing.Union[str, "os.PathLike[str]"] diff --git a/tests/test_runtime_state.py b/tests/test_runtime_state.py index 6b7d508..8d61f3f 100644 --- a/tests/test_runtime_state.py +++ b/tests/test_runtime_state.py @@ -1,10 +1,4 @@ -"""Behavioural tests for the runtime pid file's on-disk hardening. - -The pid file's contents drive a privileged operation: a forced stop reads it and -``os.killpg(..., SIGKILL)`` the pid it finds. It therefore earns the same protections as the bearer -token (owner-only permissions and a refusal to replace a symlink planted at its path), so it can never -be steered into killing an arbitrary process group. -""" +"""A forced stop SIGKILLs the pid file's process group, so the file is owner-only and refuses planted symlinks.""" from __future__ import annotations @@ -29,8 +23,6 @@ def test_pid_file_written_owner_only(tmp_path: Path, monkeypatch: pytest.MonkeyP @pytest.mark.skipif(sys.platform == "win32", reason="O_NOFOLLOW is POSIX-only") def test_pid_file_write_refuses_symlink_at_path(tmp_path: Path, monkeypatch: pytest.MonkeyPatch) -> None: - # An attacker pre-plants a symlink at the pid path pointing at a file they want clobbered (and, - # later, read back by `holo stop --force`). The hardened write must refuse to follow it. victim = tmp_path / "victim" victim.write_text("untouched", encoding="utf-8") pid_path = pid_file_path(4242, cache_dir=tmp_path)