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/README.md b/README.md index c2bb478..9755de3 100644 --- a/README.md +++ b/README.md @@ -67,6 +67,36 @@ print(result.answer) `result` is a `SessionRunResult`: `id`, `status`, `answer`, the accumulated `events`, and `final_changes`. +## Local agents + +`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={"binary_path": "/path/to/hai-agent-runtime"}) 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. +``` + +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 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..21e5e2e --- /dev/null +++ b/scripts/bump_runtime.py @@ -0,0 +1,67 @@ +"""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 pathlib import Path + +# 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)}") + # 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)}") + return {**pin, "version": version, "sha256": {platform: shas[platform] for platform in pin["sha256"]}} + + +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( + "--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 args.version, shas + + +def main(argv: list[str]) -> int: + 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 + + +if __name__ == "__main__": + raise SystemExit(main(sys.argv[1:])) diff --git a/src/hai_agents/client.py b/src/hai_agents/client.py index 82e5354..09b497c 100644 --- a/src/hai_agents/client.py +++ b/src/hai_agents/client.py @@ -7,6 +7,7 @@ from __future__ import annotations +import asyncio import typing import typing_extensions @@ -27,8 +28,57 @@ 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): + local_runtime: typing.Optional[LocalRuntime] = None + _owns_runtime = False + _auto_bridges = True + + @classmethod + def local( + cls, + *, + runtime: typing.Optional[LocalRuntime] = None, + inference: typing.Optional[Inference] = None, + local_options: typing.Optional[typing.Dict[str, typing.Any]] = None, + 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 + + runtime, owned = acquire_runtime(runtime, inference=inference, local_options=local_options) + try: + client = cls(base_url=runtime.base_url, api_key=runtime.api_key, httpx_client=runtime.http_client(timeout)) + except BaseException: + if owned: + runtime.shutdown() + raise + client.local_runtime, client._owns_runtime, client._auto_bridges = runtime, owned, auto_bridges + return client + + def close(self) -> None: + """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() + finally: + if self.local_runtime is not None: + try: + if self._owns_runtime: + self.local_runtime.shutdown_if_idle(getattr(self._sessions, "own_session_ids", ())) + finally: + 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, *, @@ -81,11 +131,63 @@ def sessions(self) -> SessionsClient: 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): + 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, + 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 + + runtime, owned = await acquire_runtime_async(runtime, inference=inference, local_options=local_options) + try: + client = cls( + base_url=runtime.base_url, api_key=runtime.api_key, httpx_client=runtime.async_http_client(timeout) + ) + except BaseException: + 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: + """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() + finally: + if self.local_runtime is not None: + try: + if self._owns_runtime: + 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() + + async def __aenter__(self) -> AsyncClient: + return self + + async def __aexit__(self, *exc: typing.Any) -> None: + await self.aclose() + async def run_session( self, *, @@ -138,5 +240,7 @@ def sessions(self) -> AsyncSessionsClient: 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/bridge.py b/src/hai_agents_local/bridge.py index 3f4ea4b..9416a3d 100644 --- a/src/hai_agents_local/bridge.py +++ b/src/hai_agents_local/bridge.py @@ -17,7 +17,8 @@ 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 logger = logging.getLogger(__name__) @@ -61,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, @@ -110,13 +113,23 @@ 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. 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): @@ -195,6 +208,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. @@ -256,7 +272,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/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/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/desktop.py b/src/hai_agents_local/desktop.py index e0b1ac0..0339b98 100644 --- a/src/hai_agents_local/desktop.py +++ b/src/hai_agents_local/desktop.py @@ -1,6 +1,8 @@ from __future__ import annotations +import asyncio import sys +import threading from typing import TYPE_CHECKING, Literal from .bridge import LocalBridge, TokenSource @@ -28,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)") @@ -98,6 +101,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/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/manager.py b/src/hai_agents_local/manager.py index b10a388..61c83cf 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 @@ -53,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 "" @@ -71,7 +78,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 @@ -97,16 +107,29 @@ 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) + + 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: @@ -152,6 +175,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. @@ -163,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/runtime/__init__.py b/src/hai_agents_local/runtime/__init__.py new file mode 100644 index 0000000..65a16b9 --- /dev/null +++ b/src/hai_agents_local/runtime/__init__.py @@ -0,0 +1,26 @@ +"""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 ( + BinaryIncompatibleError, + BinaryNotFoundError, + DownloadVerificationError, + LocalRuntimeError, + 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/runtime/errors.py b/src/hai_agents_local/runtime/errors.py new file mode 100644 index 0000000..970aa39 --- /dev/null +++ b/src/hai_agents_local/runtime/errors.py @@ -0,0 +1,27 @@ +"""Error types for hai_agents_local.runtime.""" + +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/runtime/identity.py b/src/hai_agents_local/runtime/identity.py new file mode 100644 index 0000000..c55736a --- /dev/null +++ b/src/hai_agents_local/runtime/identity.py @@ -0,0 +1,61 @@ +"""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()} did not prove it is the runtime this client " + "started or attached to: another process holds the port, or the runtime predates identity proofs" + ) + + +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/inference.py b/src/hai_agents_local/runtime/inference.py new file mode 100644 index 0000000..c1eac1a --- /dev/null +++ b/src/hai_agents_local/runtime/inference.py @@ -0,0 +1,37 @@ +"""Inference placement for a local agent runtime, independent of 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/runtime/install.py b/src/hai_agents_local/runtime/install.py new file mode 100644 index 0000000..4f4953b --- /dev/null +++ b/src/hai_agents_local/runtime/install.py @@ -0,0 +1,179 @@ +"""Verified download and atomic install of the hai-agent-runtime binary; consent is the caller's ``download=True``.""" + +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/runtime/manifest.py b/src/hai_agents_local/runtime/manifest.py new file mode 100644 index 0000000..acab3dc --- /dev/null +++ b/src/hai_agents_local/runtime/manifest.py @@ -0,0 +1,56 @@ +"""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 + +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 + + +_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] = { + 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] = { + "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/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/src/hai_agents_local/runtime/process.py b/src/hai_agents_local/runtime/process.py new file mode 100644 index 0000000..6228302 --- /dev/null +++ b/src/hai_agents_local/runtime/process.py @@ -0,0 +1,156 @@ +"""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 . import identity +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 +# 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 + + +def responds(base_url: str) -> bool: + """Whether any HTTP server answers at `base_url`, proven or not.""" + try: + 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: + 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 (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 + 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, + *, + token: str, + timeout_s: float, + log_path: pathlib.Path, + cancel_event: typing.Optional[threading.Event] = None, +) -> typing.Dict[str, typing.Any]: + """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, token) + 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: + # 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 + 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=KILL_WAIT_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/runtime.py b/src/hai_agents_local/runtime/runtime.py new file mode 100644 index 0000000..77eee80 --- /dev/null +++ b/src/hai_agents_local/runtime/runtime.py @@ -0,0 +1,477 @@ +"""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 hai_agents.base_client import BaseClient + +from . import identity +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 ( + KILL_WAIT_S, + LOOPBACK_HOST, + SPAWN_TIMEOUT_S, + TERM_GRACE_S, + probe_health, + responds, + 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" +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]"] + + +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 + + +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.""" + + 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) + 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, lock_wait_s): + 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 can outlast the lock budget: resolve outside it, then recheck attachment under it. + 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, lock_wait_s): + attached = cls._attach(base_url=base_url, cache_dir=resolved_cache) + if attached is not None: + attached.require_recipe(required_recipe) + return attached + + 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) + log_path = runtime_log_path(resolved_port, cache_dir=resolved_cache) + proc = None + token_file = 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, 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 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. + 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, 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" + ) + + @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) + 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: + status = _authenticated_probe(base_url, token) + except httpx.HTTPError as exc: + raise LocalRuntimeError("runtime attachment failed authenticated session probe") from exc + if status != 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]: + """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 + 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 http_client(self, timeout: typing.Optional[float] = None) -> httpx.Client: + """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: + """``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 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 whose runtime state a shutdown would destroy ("idle" awaits user input but keeps state). + ACTIVE_SESSION_STATUSES: typing.ClassVar[typing.Tuple[str, ...]] = ( + "queued", + "pending", + "running", + "paused", + "idle", + "awaiting_tool_results", + ) + + 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.""" + # 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): + 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 + + 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 + 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 True + seen += len(listed.items) + if not listed.items or seen >= listed.total: + return False + page += 1 + + 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): + 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 + + +_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: + 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) + held.add(path) + try: + yield + finally: + held.discard(path) + 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/runtime/state.py b/src/hai_agents_local/runtime/state.py new file mode 100644 index 0000000..175258b --- /dev/null +++ b/src/hai_agents_local/runtime/state.py @@ -0,0 +1,90 @@ +"""Owner-only (0600, symlink-refusing) token and pid files that let other local processes find a spawned runtime.""" + +from __future__ import annotations + +import contextlib +import os +import pathlib +import tempfile +import typing + +CACHE_DIR_ENV = "HAI_AGENT_LOCAL_CACHE_DIR" +DEFAULT_CACHE_DIR = pathlib.Path.home() / ".hai" / "agent-runtime" +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 for out-of-process stop tools.""" + 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: + """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 + 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 + + +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 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/sessions.py b/src/hai_agents_local/sessions.py index d8dccf5..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 @@ -12,14 +13,22 @@ import threading import typing +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: + from .runtime import LocalRuntime + logger = logging.getLogger(__name__) # Runaway guards for sessions whose caller passes no budget: a local session left running (its @@ -27,6 +36,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: @@ -34,6 +44,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 @@ -65,7 +83,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 @@ -75,9 +95,14 @@ 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() - ) + 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 @@ -111,8 +136,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) @@ -123,7 +151,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: @@ -132,12 +160,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 @@ -156,11 +187,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) @@ -191,12 +222,84 @@ 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}) + + +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.""" + + 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() + + 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 = _CloseFailures() + for session_id in self._live_sessions(): + with failures.cancelling(session_id): + self.cancel_session(session_id) + 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. + owned = self._owned_bridges.get(str(id), []) + try: + if owned: + stop_bridges(owned) + self._owned_bridges.pop(str(id), None) + finally: + super().cancel_session(id, request_options=request_options) + # Keep the exit retry registered until cancellation is confirmed. + _deregister_exit_cancel(str(id)) -class LocalSessionsClient(SessionsClient): @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(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 @@ -205,45 +308,66 @@ 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(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() + self._track(session, started) return session -class LocalAsyncSessionsClient(AsyncSessionsClient): +class LocalAsyncSessionsClient(_LocalSessionsState, AsyncSessionsClient): + async def aclose(self) -> None: + failures = _CloseFailures() + for session_id in self._live_sessions(): + with failures.cancelling(session_id): + await self.cancel_session(session_id) + failures.raise_any() + + 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(id), None) + finally: + 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: - wrapper = self._raw_client._client_wrapper - bridges = _localize(wrapper, kwargs) + 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. + 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) 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(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) + self._track(session, started) return session diff --git a/src/hai_agents_local/transport.py b/src/hai_agents_local/transport.py index 5ee29fa..84df979 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,7 @@ async def ensure_channel(self, session_id: str) -> None: raise RateLimitedError(_retry_after(resp)) if resp.status_code == HTTPStatus.CONFLICT: return + _raise_if_gone(resp, session_id) resp.raise_for_status() async def fetch_commands( @@ -102,6 +103,7 @@ 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 @@ -129,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/src/hai_agents_local/workstation.py b/src/hai_agents_local/workstation.py index aca6579..1a446bf 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 @@ -63,6 +64,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_bump_runtime.py b/tests/test_bump_runtime.py new file mode 100644 index 0000000..a635696 --- /dev/null +++ b/tests/test_bump_runtime.py @@ -0,0 +1,44 @@ +"""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] +RUNTIME = Path("src") / "hai_agents_local" / "runtime" + + +@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)} + 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 case != "complete": + assert result.returncode != 0 + assert pin.read_bytes() == before + return + assert result.returncode == 0, result.stderr + 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") diff --git a/tests/test_desktop_stop.py b/tests/test_desktop_stop.py new file mode 100644 index 0000000..7ff9fc1 --- /dev/null +++ b/tests/test_desktop_stop.py @@ -0,0 +1,82 @@ +"""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 + +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_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 diff --git a/tests/test_local.py b/tests/test_local.py index 6d02de4..31fddc0 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 @@ -192,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")) ) @@ -475,12 +481,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 +528,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): @@ -725,13 +770,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: @@ -761,7 +817,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") @@ -770,6 +827,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() @@ -822,6 +887,30 @@ 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]) + @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: + 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) + ) + 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() @@ -869,3 +958,149 @@ 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 + + +@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() + + +@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 == {} diff --git a/tests/test_runtime_placement.py b/tests/test_runtime_placement.py new file mode 100644 index 0000000..69af01b --- /dev/null +++ b/tests/test_runtime_placement.py @@ -0,0 +1,416 @@ +"""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 + +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 + + +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" + 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 + + +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): + 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): + 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() + + +@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="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 + attached = LocalRuntime.attach(port=runtime_server.port, cache_dir=tmp_path) + with pytest.raises(BinaryIncompatibleError): + attached.require_recipe("desktop") + with pytest.raises(LocalRuntimeError): + attached.shutdown() + with Client.local(runtime=attached) as client: + assert client.sessions.list_sessions().items == [] + runtime_server.proof_token = "squatter-token" + with pytest.raises(LocalRuntimeError, match="did not prove"): + client.sessions.list_sessions() + + +@pytest.mark.asyncio +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(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() + + +@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("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) + 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(auto_bridges=False, **options) + await client.sessions.create_session(agent="h/agent", messages="hi") + await client.aclose() + else: + client = Client.local(auto_bridges=False, **options) + client.sessions.create_session(agent="h/agent", messages="hi") + client.close() + assert stopped == ([True] if runtime_use == "idle" else []) + 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.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) + + 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.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.runtime 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" + + +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_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 + + 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="did not prove"): + 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): + class Runtime(FakeRuntime): + def require_recipe(self, recipe): + raise BinaryIncompatibleError("recipe changed") + + runtime = Runtime() + with pytest.raises(BinaryIncompatibleError, match="recipe changed"): + if asynchronous: + await AsyncClient.local(runtime=runtime) + else: + Client.local(runtime=runtime) + assert not runtime.stopped diff --git a/tests/test_runtime_state.py b/tests/test_runtime_state.py new file mode 100644 index 0000000..8d61f3f --- /dev/null +++ b/tests/test_runtime_state.py @@ -0,0 +1,35 @@ +"""A forced stop SIGKILLs the pid file's process group, so the file is owner-only and refuses planted symlinks.""" + +from __future__ import annotations + +import os +import stat +import sys +from pathlib import Path + +import pytest + +from hai_agents_local.runtime.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: + 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"