From cd54df4fe52341448ce1dffaf921833389d85134 Mon Sep 17 00:00:00 2001 From: Jaideep <67646710+jdchawla29@users.noreply.github.com> Date: Sat, 22 Aug 2026 15:03:50 -0700 Subject: [PATCH 1/7] Bind hosted inference inside isolated workspaces --- README.md | 6 + hud/agents/claude/sdk/agent.py | 27 ++- hud/agents/cli.py | 10 + hud/agents/codex/agent.py | 30 ++- hud/agents/tests/test_claude_cli_agent.py | 19 ++ hud/agents/tests/test_codex_cli_agent.py | 17 ++ hud/clients/__init__.py | 3 +- hud/clients/client.py | 44 +++- hud/clients/tests/test_connect.py | 59 ++++- hud/environment/env.py | 73 ++++++ hud/environment/platform_inference.py | 281 ++++++++++++++++++++++ hud/environment/server.py | 40 ++- hud/environment/tests/test_workspace.py | 90 +++++++ hud/environment/workspace.py | 16 ++ hud/eval/runtime/__init__.py | 2 + hud/eval/runtime/core.py | 10 + 16 files changed, 708 insertions(+), 19 deletions(-) create mode 100644 hud/environment/platform_inference.py diff --git a/README.md b/README.md index 4a06edc15..b1bb099f5 100644 --- a/README.md +++ b/README.md @@ -156,6 +156,12 @@ A **capability** is a connection the environment exposes; a **harness** attaches From the [platform UI](https://hud.ai) you can run batches, compare models on the same taskset, and inspect every trace. +Hosted Claude Code and Codex harnesses reach platform inference through an +environment-owned, workspace-local endpoint. The endpoint is available only to +`bwrap` workspaces with network isolation; the workspace receives an opaque +per-session key, while platform credentials and trace attribution stay outside +its environment and manifest. + → [Run & deploy](https://docs.hud.ai/v6/reference/runtime) ## Train on rewards diff --git a/hud/agents/claude/sdk/agent.py b/hud/agents/claude/sdk/agent.py index 76fd85c8e..69a556e75 100644 --- a/hud/agents/claude/sdk/agent.py +++ b/hud/agents/claude/sdk/agent.py @@ -21,6 +21,7 @@ WINDOWS_SHELLS, powershell, powershell_quote, + require_platform_isolation, resolve_executable, run_jsonl, ) @@ -34,6 +35,7 @@ if TYPE_CHECKING: from hud.capabilities import SSHClient + from hud.environment.platform_inference import InferenceBinding from hud.eval.run import Run logger = logging.getLogger(__name__) @@ -64,6 +66,7 @@ def __init__(self, config: ClaudeCLIConfig | None = None) -> None: async def __call__(self, run: Run) -> None: mcp_servers: dict[str, dict[str, Any]] = {} ssh = cast("SSHClient", await run.client.open("ssh")) + require_platform_isolation(ssh, run.client.inference) manifest = run.client.manifest assert manifest is not None bindings = manifest.bindings @@ -111,6 +114,7 @@ async def __call__(self, run: Run) -> None: mcp_servers=mcp_servers, prompt=run.prompt_text, executable=executable, + inference=run.client.inference, ) async def _exec( @@ -122,6 +126,7 @@ async def _exec( mcp_servers: dict[str, dict[str, Any]], prompt: str, executable: str = "claude", + inference: InferenceBinding | None = None, ) -> None: mcp_config_path = await self._write_mcp_config(ssh, mcp_servers) input_text = ( @@ -145,6 +150,7 @@ async def _exec( shell=shell, mcp_config_path=mcp_config_path, executable=executable, + inference=inference, ) if shell in WINDOWS_SHELLS: await ssh.write_text(RUN_SCRIPT_PATH, f"@echo off\r\n{command}\r\n") @@ -173,20 +179,26 @@ async def _exec( except (OSError, asyncssh.Error): logger.warning("Failed to remove Claude CLI runtime files") - def _build_env_vars(self) -> dict[str, str]: + def _build_env_vars(self, inference: InferenceBinding | None = None) -> dict[str, str]: env: dict[str, str] = {} use_hud_gateway = self.config.use_hud_gateway if use_hud_gateway is None: - use_hud_gateway = settings.api_key is not None + use_hud_gateway = inference is not None or settings.api_key is not None if use_hud_gateway: - if not settings.api_key: + if inference is not None: + base_url = inference.base_url + api_key = inference.api_key + elif settings.api_key: + base_url = settings.hud_gateway_url + api_key = settings.api_key + else: raise ValueError("HUD_API_KEY is required for HUD gateway routing") - env["ANTHROPIC_BASE_URL"] = settings.hud_gateway_url - env["ANTHROPIC_API_KEY"] = settings.api_key + env["ANTHROPIC_BASE_URL"] = base_url + env["ANTHROPIC_API_KEY"] = api_key env["CLAUDE_CODE_DISABLE_EXPERIMENTAL_BETAS"] = "1" env["DISABLE_AUTO_COMPACT"] = "1" - if trace_id := get_current_trace_id(): + if inference is None and (trace_id := get_current_trace_id()): env["ANTHROPIC_CUSTOM_HEADERS"] = f"Trace-Id: {trace_id}" elif settings.anthropic_api_key: env["ANTHROPIC_API_KEY"] = settings.anthropic_api_key @@ -227,8 +239,9 @@ def _build_cli_command( shell: str, mcp_config_path: str | None = None, executable: str = "claude", + inference: InferenceBinding | None = None, ) -> str: - env_vars = self._build_env_vars() + env_vars = self._build_env_vars(inference) is_win = shell in WINDOWS_SHELLS base_args: list[str] = [ executable, diff --git a/hud/agents/cli.py b/hud/agents/cli.py index d626c024e..3f7850627 100644 --- a/hud/agents/cli.py +++ b/hud/agents/cli.py @@ -14,12 +14,22 @@ from collections.abc import Callable from hud.capabilities import SSHClient + from hud.environment.platform_inference import InferenceBinding from hud.eval.runtime import RuntimeConfig WINDOWS_SHELLS = ("cmd", "powershell") PROCESS_CLOSE_TIMEOUT_S = 5.0 +def require_platform_isolation(ssh: SSHClient, binding: InferenceBinding | None) -> None: + """Refuse a platform credential binding when the remote shell is not isolated.""" + if binding is not None and ssh.capability.params.get("isolation") != "bwrap": + raise RuntimeError( + "platform inference requires a bwrap-isolated workspace; refusing to expose " + "the workspace-local binding to an unisolated SSH session" + ) + + async def resolve_executable( ssh: SSHClient, command: str, diff --git a/hud/agents/codex/agent.py b/hud/agents/codex/agent.py index ed4f6946b..c54fbe367 100644 --- a/hud/agents/codex/agent.py +++ b/hud/agents/codex/agent.py @@ -14,6 +14,7 @@ WINDOWS_SHELLS, powershell, powershell_quote, + require_platform_isolation, resolve_executable, run_jsonl, ) @@ -25,6 +26,7 @@ if TYPE_CHECKING: from hud.capabilities import SSHClient + from hud.environment.platform_inference import InferenceBinding from hud.eval.run import Run logger = logging.getLogger(__name__) @@ -208,7 +210,12 @@ def record_tool(self, item: dict[str, Any], started_at: str, ended_at: str) -> N ) -def codex_command(config: CodexCLIConfig, shell: str, executable: str = "codex") -> str: +def codex_command( + config: CodexCLIConfig, + shell: str, + executable: str = "codex", + inference: InferenceBinding | None = None, +) -> str: env: dict[str, str] = {} args = [ executable, @@ -226,21 +233,27 @@ def codex_command(config: CodexCLIConfig, shell: str, executable: str = "codex") use_hud_gateway = config.use_hud_gateway if use_hud_gateway is None: - use_hud_gateway = settings.api_key is not None + use_hud_gateway = inference is not None or settings.api_key is not None if use_hud_gateway: - if not settings.api_key: + if inference is not None: + base_url = inference.base_url + api_key = inference.api_key + elif settings.api_key: + base_url = settings.hud_gateway_url + api_key = settings.api_key + else: raise ValueError("HUD_API_KEY is required for HUD gateway routing") - env["HUD_API_KEY"] = settings.api_key + env["HUD_API_KEY"] = api_key overrides = { "model_provider": "hud", "model_providers.hud.name": "HUD", - "model_providers.hud.base_url": settings.hud_gateway_url, + "model_providers.hud.base_url": base_url, "model_providers.hud.env_key": "HUD_API_KEY", "model_providers.hud.wire_api": "responses", } for key, value in overrides.items(): args.extend(["-c", f"{key}={json.dumps(value)}"]) - if trace_id := get_current_trace_id(): + if inference is None and (trace_id := get_current_trace_id()): args.extend( [ "-c", @@ -290,8 +303,9 @@ async def run_codex( shell: str, prompt: str, executable: str = "codex", + inference: InferenceBinding | None = None, ) -> None: - command = codex_command(config, shell, executable) + command = codex_command(config, shell, executable, inference=inference) logger.info("SSH exec codex CLI (%d chars)", len(command)) events = CodexEvents(run, model=config.model, started_at=now_iso()) returncode, stderr = await run_jsonl(ssh, command, events.consume, input_text=prompt) @@ -309,6 +323,7 @@ def __init__(self, config: CodexCLIConfig | None = None) -> None: async def __call__(self, run: Run) -> None: ssh = cast("SSHClient", await run.client.open("ssh")) + require_platform_isolation(ssh, run.client.inference) executable = await resolve_executable( ssh, "codex", @@ -322,6 +337,7 @@ async def __call__(self, run: Run) -> None: shell=ssh.capability.params.get("shell", "bash"), prompt=run.prompt_text, executable=executable, + inference=run.client.inference, ) diff --git a/hud/agents/tests/test_claude_cli_agent.py b/hud/agents/tests/test_claude_cli_agent.py index b6528b24b..73a79fa09 100644 --- a/hud/agents/tests/test_claude_cli_agent.py +++ b/hud/agents/tests/test_claude_cli_agent.py @@ -30,6 +30,7 @@ from hud.agents.types import AgentStep, ClaudeCLIConfig, ToolStep from hud.capabilities import Capability, SSHClient from hud.capabilities.rfb import WebPScreenshotEncoding +from hud.environment.platform_inference import InferenceBinding from hud.settings import settings from hud.telemetry.context import set_trace_context from hud.types import MCPToolResult @@ -68,6 +69,20 @@ def test_command_follows_explicit_gateway_routing(monkeypatch: pytest.MonkeyPatc assert "ANTHROPIC_MODEL=claude-sonnet-5" in provider +def test_command_prefers_environment_inference_binding() -> None: + binding = InferenceBinding( + base_url="http://127.0.0.1:49123/p/opaque", + api_key="workspace-key", + ) + + gateway = claude_command(ClaudeCLIConfig(use_hud_gateway=True), "bash", inference=binding) + + assert "ANTHROPIC_BASE_URL=http://127.0.0.1:49123/p/opaque" in gateway + assert "ANTHROPIC_API_KEY=workspace-key" in gateway + assert "HUD_API_KEY" not in gateway + assert "Trace-Id" not in gateway + + def test_windows_command_encodes_environment_and_arguments( monkeypatch: pytest.MonkeyPatch, ) -> None: @@ -458,6 +473,7 @@ async def test_manifest_mcp_capability_is_written_for_remote_claude( ssh = SSHClient(shell, cast("Any", object())) class Client: + inference = None manifest = SimpleNamespace(bindings=[shell, mcp]) async def open(self, ref: str) -> SSHClient: @@ -499,6 +515,7 @@ async def test_remote_claude_passes_screenshot_encoding_to_computer_mcp( bridge_active = False class Client: + inference = None manifest = SimpleNamespace(bindings=[shell, screen]) async def open(self, ref: str) -> SSHClient: @@ -583,6 +600,7 @@ async def test_remote_claude_preserves_multiple_rfb_bindings( bridged: list[str] = [] class Client: + inference = None manifest = SimpleNamespace(bindings=[shell, *screens]) async def open(self, ref: str) -> SSHClient: @@ -870,6 +888,7 @@ async def test_concurrent_runs_keep_their_ssh_state_isolated( class Client: def __init__(self, shell: Capability, ssh: SSHClient) -> None: + self.inference = None self.manifest = SimpleNamespace(bindings=[shell]) self.ssh = ssh diff --git a/hud/agents/tests/test_codex_cli_agent.py b/hud/agents/tests/test_codex_cli_agent.py index 9819c1827..8eb5d503c 100644 --- a/hud/agents/tests/test_codex_cli_agent.py +++ b/hud/agents/tests/test_codex_cli_agent.py @@ -18,6 +18,7 @@ from hud.agents.tests.cli_fakes import fake_run as _fake_run from hud.agents.types import AgentStep, CodexCLIConfig, ToolStep from hud.capabilities import Capability, SSHClient +from hud.environment.platform_inference import InferenceBinding from hud.eval.runtime import RuntimeConfig, RuntimeResources from hud.settings import settings from hud.telemetry.context import set_trace_context @@ -104,6 +105,19 @@ def test_command_follows_explicit_gateway_routing(monkeypatch: pytest.MonkeyPatc assert command.endswith(" -") +def test_command_prefers_environment_inference_binding() -> None: + binding = InferenceBinding( + base_url="http://127.0.0.1:49123/p/opaque", + api_key="workspace-key", + ) + + command = codex_command(CodexCLIConfig(use_hud_gateway=True), "bash", inference=binding) + + assert "HUD_API_KEY=workspace-key" in command + assert 'model_providers.hud.base_url="http://127.0.0.1:49123/p/opaque"' in command + assert "Trace-Id" not in command + + def test_windows_command_encodes_environment_and_arguments( monkeypatch: pytest.MonkeyPatch, ) -> None: @@ -276,6 +290,8 @@ async def test_agent_opens_ssh_and_uses_workspace_prompt(monkeypatch: pytest.Mon ssh = _FakeSSH(_FakeProcess(_STREAM_JSON), shell="powershell") class Client: + inference = None + async def open(self, ref: str) -> _FakeSSH: assert ref == "ssh" return ssh @@ -294,6 +310,7 @@ async def open(self, ref: str) -> _FakeSSH: shell="powershell", prompt="Fix it", executable="codex", + inference=None, ) diff --git a/hud/clients/__init__.py b/hud/clients/__init__.py index 7c670788c..746f827c9 100644 --- a/hud/clients/__init__.py +++ b/hud/clients/__init__.py @@ -2,11 +2,12 @@ from __future__ import annotations -from .client import HudClient, HudProtocolError, Manifest, ServerInfo, connect +from .client import HudClient, HudProtocolError, InferenceBinding, Manifest, ServerInfo, connect __all__ = [ "HudClient", "HudProtocolError", + "InferenceBinding", "Manifest", "ServerInfo", "connect", diff --git a/hud/clients/client.py b/hud/clients/client.py index 755c6f542..6c7443637 100644 --- a/hud/clients/client.py +++ b/hud/clients/client.py @@ -27,6 +27,7 @@ RFBClient, SSHClient, ) +from hud.environment.platform_inference import InferenceBinding from hud.environment.utils import read_frame, send_frame, splice if TYPE_CHECKING: @@ -124,6 +125,7 @@ def __init__( self._opened: dict[str, CapabilityClient] = {} self._forwarders: list[asyncio.Server] = [] self._tunnels: set[asyncio.Task[None]] = set() + self.inference: InferenceBinding | None = None # ─── lifecycle ──────────────────────────────────────────────────── @@ -322,6 +324,29 @@ async def grade(self, payload: dict[str, Any]) -> dict[str, Any]: async def cancel(self) -> None: await self._call("tasks.cancel", {}) + async def bind_inference( + self, + *, + upstream_url: str, + token: str, + trace_id: str | None = None, + ) -> InferenceBinding: + """Ask the environment server to expose scoped inference inside its workspace.""" + result = await self._call( + "platform.inference.bind", + { + "upstream_url": upstream_url, + "token": token, + "trace_id": trace_id, + }, + ) + base_url = result.get("base_url") + api_key = result.get("api_key") + if not isinstance(base_url, str) or not isinstance(api_key, str): + raise HudProtocolError(-32603, "platform inference binding was malformed") + self.inference = InferenceBinding(base_url=base_url, api_key=api_key) + return self.inference + # ─── JSON-RPC plumbing ──────────────────────────────────────────── async def _call( @@ -447,6 +472,16 @@ async def connect(runtime: Runtime, *, ready_timeout: float = 240.0) -> AsyncIte parts.port or 0, ready_timeout=_runtime_ready_timeout(runtime, ready_timeout), ) + try: + if runtime.inference is not None: + await client.bind_inference( + upstream_url=runtime.inference.upstream_url, + token=runtime.inference.token, + trace_id=runtime.inference.trace_id, + ) + except BaseException: + await client.close() + raise owner = asyncio.current_task() assert owner is not None heartbeat_error: Exception | None = None @@ -485,4 +520,11 @@ async def heartbeat() -> None: raise -__all__ = ["HudClient", "HudProtocolError", "Manifest", "ServerInfo", "connect"] +__all__ = [ + "HudClient", + "HudProtocolError", + "InferenceBinding", + "Manifest", + "ServerInfo", + "connect", +] diff --git a/hud/clients/tests/test_connect.py b/hud/clients/tests/test_connect.py index a5f921d4c..813b47f34 100644 --- a/hud/clients/tests/test_connect.py +++ b/hud/clients/tests/test_connect.py @@ -20,11 +20,68 @@ from hud.capabilities import Capability, CapabilityClient from hud.clients import connect from hud.environment.utils import read_frame, send_frame -from hud.eval.runtime import Runtime +from hud.eval.runtime import Runtime, RuntimeInference HELLO_RESULT = {"session_id": "s-1", "env": {"name": "stub", "version": "1.0"}, "bindings": []} +async def test_connect_binds_runtime_inference_after_hello() -> None: + requests: list[dict[str, object]] = [] + + async def handler(reader: asyncio.StreamReader, writer: asyncio.StreamWriter) -> None: + try: + hello = await read_frame(reader) + assert hello is not None + requests.append(hello) + await send_frame(writer, {"jsonrpc": "2.0", "id": hello["id"], "result": HELLO_RESULT}) + bind = await read_frame(reader) + assert bind is not None + requests.append(bind) + await send_frame( + writer, + { + "jsonrpc": "2.0", + "id": bind["id"], + "result": { + "base_url": "http://127.0.0.1:49123/p/opaque", + "api_key": "workspace-key", + }, + }, + ) + await read_frame(reader) + finally: + writer.close() + + server = await asyncio.start_server(handler, "127.0.0.1", 0) + port = server.sockets[0].getsockname()[1] + runtime = Runtime( + f"tcp://127.0.0.1:{port}", + inference=RuntimeInference( + upstream_url="https://inference.hud.so", + token="scoped-runtime-token", + trace_id="trace-1", + ), + ) + assert "scoped-runtime-token" not in repr(runtime) + try: + async with connect(runtime) as client: + assert client.inference is not None + assert client.inference.base_url == "http://127.0.0.1:49123/p/opaque" + assert client.inference.api_key == "workspace-key" + finally: + server.close() + await server.wait_closed() + + assert [request["method"] for request in requests] == ["hello", "platform.inference.bind"] + params = requests[1]["params"] + assert isinstance(params, dict) + assert params == { + "upstream_url": "https://inference.hud.so", + "token": "scoped-runtime-token", + "trace_id": "trace-1", + } + + async def test_open_retries_transient_capability_connection_failures( monkeypatch: pytest.MonkeyPatch, ) -> None: diff --git a/hud/environment/env.py b/hud/environment/env.py index 86023f3da..b0985ec56 100644 --- a/hud/environment/env.py +++ b/hud/environment/env.py @@ -10,6 +10,7 @@ import contextlib import functools import inspect +import secrets from contextvars import ContextVar from typing import TYPE_CHECKING, Any, Generic, ParamSpec, Protocol, TypeVar, cast @@ -17,6 +18,8 @@ from hud.capabilities import Capability +from .egress import Peer +from .platform_inference import InferenceBinding, PlatformInferenceProxy from .workspace import Workspace if TYPE_CHECKING: @@ -162,6 +165,10 @@ def __init__( self._on_stop: list[Callable[[], Awaitable[None]]] = [] # Per task-session end (cancel / bye / post-grade cleanup). self._on_task_teardown: list[Callable[[], Awaitable[None]]] = [] + self._workspaces: list[Workspace] = [] + self._platform_inference = PlatformInferenceProxy() + self._platform_peer_port: int | None = None + self._platform_peer_name: str | None = None # ─── task registration ─────────────────────────────────────────── @@ -285,6 +292,7 @@ def workspace( track_files = settings.file_tracking_enabled ws = Workspace(root, track_files=track_files, **kwargs) + self._workspaces.append(ws) @self.initialize async def _up() -> None: @@ -349,5 +357,70 @@ async def stop(self) -> None: for hook in reversed(self._on_stop): with contextlib.suppress(Exception): await hook() + if self._platform_peer_name is not None: + for workspace in self._workspaces: + workspace.remove_peer(self._platform_peer_name) + self._platform_inference.stop() + self._platform_peer_name = None + self._platform_peer_port = None self._started = False self._hooks_done = False + + def _bind_platform_peer(self) -> None: + if self._platform_peer_port is not None or not self._workspaces: + return + unavailable = { + port + for workspace in self._workspaces + for port in (*workspace.ports, *(peer.port for peer in workspace.peers)) + } + for _ in range(100): + port = 20_000 + secrets.randbelow(40_000) + if port not in unavailable: + break + else: + raise RuntimeError("could not allocate a workspace-local platform service port") + peer_name = "platform-" + secrets.token_hex(8) + peer = Peer( + peer_name, + port, + target=self._platform_inference.address, + ) + for workspace in self._workspaces: + workspace.add_peer(peer, first=True) + self._platform_peer_name = peer_name + self._platform_peer_port = port + + def bind_platform_inference( + self, + session_id: str, + *, + upstream_url: str, + token: str, + trace_id: str | None, + ) -> InferenceBinding: + """Bind one control session to inference through every bounded workspace.""" + if not self._started or not self._workspaces: + raise RuntimeError("environment has no workspace available for platform inference") + if any(not workspace.bwrap_available for workspace in self._workspaces): + raise RuntimeError("platform inference requires bwrap-isolated workspaces") + if any(not workspace.owns_netns for workspace in self._workspaces): + raise RuntimeError("platform inference requires network-isolated workspaces") + if self._platform_peer_port is None: + self._platform_inference.start() + try: + self._bind_platform_peer() + except BaseException: + self._platform_inference.stop() + raise + assert self._platform_peer_port is not None + return self._platform_inference.register( + session_id, + upstream_url=upstream_url, + token=token, + trace_id=trace_id, + workspace_url=f"http://127.0.0.1:{self._platform_peer_port}", + ) + + def unbind_platform_inference(self, session_id: str) -> None: + self._platform_inference.unregister(session_id) diff --git a/hud/environment/platform_inference.py b/hud/environment/platform_inference.py new file mode 100644 index 000000000..f57ed66ab --- /dev/null +++ b/hud/environment/platform_inference.py @@ -0,0 +1,281 @@ +"""Platform inference made available inside bounded workspaces.""" + +from __future__ import annotations + +import hmac +import http.client +import secrets +import threading +import urllib.parse +from dataclasses import dataclass +from http.server import BaseHTTPRequestHandler, ThreadingHTTPServer +from typing import TYPE_CHECKING, Any + +from .egress import _HOP_BY_HOP, _field, _Unrelayable + +if TYPE_CHECKING: + from collections.abc import Mapping + + +_CREDENTIAL_HEADERS = frozenset( + {"authorization", "hud-api-key", "hud-runtime-token", "x-api-key", "x-goog-api-key"} +) +_TRACE_HEADERS = frozenset({"trace-id", "x-trace-id", "x-hud-trace-id"}) + + +@dataclass(frozen=True, slots=True) +class InferenceBinding: + """Workspace-local connection details for one platform inference lease.""" + + base_url: str + api_key: str + + +@dataclass(frozen=True, slots=True) +class _Lease: + upstream: urllib.parse.SplitResult + upstream_token: str + trace_id: str | None + client_key: str + + +def _request_key(headers: Mapping[str, str]) -> str | None: + for name in ("hud-api-key", "x-api-key", "x-goog-api-key"): + if value := headers.get(name): + return value + authorization = headers.get("authorization", "") + scheme, _, value = authorization.partition(" ") + return value if scheme.lower() == "bearer" and value else None + + +class _InferenceHandler(BaseHTTPRequestHandler): + protocol_version = "HTTP/1.1" + proxy: PlatformInferenceProxy + + def log_message(self, format: str, *args: Any) -> None: + """Requests and opaque lease paths must not enter environment logs.""" + + def _fail(self, status: int, reason: str) -> None: + self.send_response(status) + self.send_header("X-Proxy-Error", reason) + self.send_header("Content-Length", "0") + self.end_headers() + + def _request_body(self) -> bytes | None: + transfer = self.headers.get("Transfer-Encoding") + length = self.headers.get("Content-Length") + if transfer is None: + if length is None: + return None + size = int(length) + if size < 0: + raise ValueError + body = self.rfile.read(size) + if len(body) != size: + raise ValueError + return body + if length is not None or transfer.strip().lower() != "chunked": + raise ValueError + + chunks: list[bytes] = [] + while True: + line = self.rfile.readline(65537) + if len(line) > 65536 or not line.endswith(b"\r\n"): + raise ValueError + size_text = line[:-2].split(b";", 1)[0].strip() + if not size_text or any(byte not in b"0123456789abcdefABCDEF" for byte in size_text): + raise ValueError + size = int(size_text, 16) + if size == 0: + while True: + trailer = self.rfile.readline(65537) + if len(trailer) > 65536 or not trailer.endswith(b"\r\n"): + raise ValueError + if trailer == b"\r\n": + return b"".join(chunks) + chunk = self.rfile.read(size) + if len(chunk) != size or self.rfile.read(2) != b"\r\n": + raise ValueError + chunks.append(chunk) + + def _forward(self) -> None: + parts = urllib.parse.urlsplit(self.path) + lease, upstream_path = self.proxy.resolve(parts.path) + if lease is None: + self._fail(404, "unknown-lease") + return + supplied = _request_key({key.lower(): value for key, value in self.headers.items()}) + if supplied is None or not hmac.compare_digest(supplied, lease.client_key): + self._fail(401, "invalid-lease-key") + return + try: + body = self._request_body() + except (ValueError, OverflowError): + self.close_connection = True + self._fail(400, "invalid-request-body") + return + + headers = { + key: value + for key, value in self.headers.items() + if key.lower() not in _HOP_BY_HOP | _CREDENTIAL_HEADERS | _TRACE_HEADERS | {"host"} + } + headers["Hud-Runtime-Token"] = lease.upstream_token + if lease.trace_id is not None: + headers["Trace-Id"] = lease.trace_id + + upstream = lease.upstream + base = upstream.path.rstrip("/") + path = f"{base}/{upstream_path.lstrip('/')}" + if parts.query: + path = f"{path}?{parts.query}" + host = upstream.hostname + assert host is not None + connection: http.client.HTTPConnection + if upstream.scheme == "https": + connection = http.client.HTTPSConnection(host, upstream.port, timeout=300) + else: + connection = http.client.HTTPConnection(host, upstream.port, timeout=300) + response_started = False + try: + connection.request(self.command, path, body=body, headers=headers) + response = connection.getresponse() + relayed = [ + _field(key, value) + for key, value in response.getheaders() + if key.lower() not in _HOP_BY_HOP and key.lower() != "content-length" + ] + length = response.getheader("Content-Length") + framed = length is not None and length.strip().isdigit() + _field("Reason", response.reason or "") + response_started = True + self.send_response(response.status, response.reason) + for key, value in relayed: + self.send_header(key, value) + if framed: + assert length is not None + self.send_header("Content-Length", length.strip()) + else: + self.send_header("Connection", "close") + self.close_connection = True + self.end_headers() + while chunk := response.read(65536): + self.wfile.write(chunk) + self.wfile.flush() + except _Unrelayable: + self._fail(502, "unrelayable-upstream-header") + except (OSError, http.client.HTTPException): + if response_started: + self.close_connection = True + else: + self._fail(502, "upstream-failure") + finally: + connection.close() + + do_GET = _forward + do_HEAD = _forward + do_POST = _forward + do_PUT = _forward + do_DELETE = _forward + do_PATCH = _forward + do_OPTIONS = _forward + + +class PlatformInferenceProxy: + """Environment-owned reverse proxy with per-control-session credentials.""" + + def __init__(self) -> None: + self._server: ThreadingHTTPServer | None = None + self._thread: threading.Thread | None = None + self._leases: dict[str, tuple[str, _Lease]] = {} + self._lock = threading.Lock() + + @property + def address(self) -> tuple[str, int]: + if self._server is None: + raise RuntimeError("platform inference proxy is not started") + host, port = self._server.server_address[:2] + return str(host), int(port) + + def start(self) -> None: + if self._server is not None: + return + handler = type("_ScopedInferenceHandler", (_InferenceHandler,), {"proxy": self}) + server = ThreadingHTTPServer(("127.0.0.1", 0), handler) + server.daemon_threads = True + self._server = server + self._thread = threading.Thread(target=server.serve_forever, daemon=True) + self._thread.start() + + def register( + self, + session_id: str, + *, + upstream_url: str, + token: str, + trace_id: str | None, + workspace_url: str, + ) -> InferenceBinding: + upstream = urllib.parse.urlsplit(upstream_url) + if ( + upstream.scheme not in {"http", "https"} + or upstream.hostname is None + or upstream.username is not None + or upstream.password is not None + or upstream.query + or upstream.fragment + ): + raise ValueError("platform inference upstream must be an HTTP(S) base URL") + if not token: + raise ValueError("platform inference token must not be empty") + with self._lock: + existing = self._leases.get(session_id) + if existing is not None: + route, lease = existing + requested = (upstream, token, trace_id) + current = (lease.upstream, lease.upstream_token, lease.trace_id) + if requested != current: + raise RuntimeError("platform inference is already bound for this session") + return InferenceBinding(f"{workspace_url}/{route}", lease.client_key) + route = "p/" + secrets.token_urlsafe(18) + lease = _Lease( + upstream=upstream, + upstream_token=token, + trace_id=trace_id, + client_key=secrets.token_urlsafe(32), + ) + self._leases[session_id] = (route, lease) + return InferenceBinding(f"{workspace_url}/{route}", lease.client_key) + + def resolve(self, path: str) -> tuple[_Lease | None, str]: + stripped = path.lstrip("/") + prefix, separator, remainder = stripped.partition("/") + if prefix != "p" or not separator: + return None, "" + token, separator, upstream_path = remainder.partition("/") + if not token or not separator: + return None, "" + route = f"p/{token}" + with self._lock: + for stored_route, lease in self._leases.values(): + if hmac.compare_digest(route, stored_route): + return lease, "/" + upstream_path + return None, "" + + def unregister(self, session_id: str) -> None: + with self._lock: + self._leases.pop(session_id, None) + + def stop(self) -> None: + server, self._server = self._server, None + if server is not None: + server.shutdown() + server.server_close() + thread, self._thread = self._thread, None + if thread is not None: + thread.join(timeout=5) + with self._lock: + self._leases.clear() + + +__all__ = ["InferenceBinding", "PlatformInferenceProxy"] diff --git a/hud/environment/server.py b/hud/environment/server.py index 36ad47818..44a947cd3 100644 --- a/hud/environment/server.py +++ b/hud/environment/server.py @@ -234,7 +234,7 @@ def __init__(self, env: Environment) -> None: self._live: set[str] = set() async def start(self, session_id: str, task_id: str, args: dict[str, Any]) -> dict[str, Any]: - await self.cancel(session_id) + await self._cancel_runner(session_id) runner = TaskRunner(self.env.tasks[task_id], args) self._runners[session_id] = runner try: @@ -255,6 +255,7 @@ async def grade(self, session_id: str, payload: dict[str, Any]) -> dict[str, Any return await runner.grade(payload) finally: current_session_id.reset(token) + self.env.unbind_platform_inference(claim_sid) def _adopt_parked(self) -> tuple[str, TaskRunner]: """Claim the parked session iff unambiguous — the blind-reconnect grade path.""" @@ -269,7 +270,7 @@ def _adopt_parked(self) -> tuple[str, TaskRunner]: sid = parked[0] return sid, self._runners.pop(sid) - async def cancel(self, session_id: str) -> None: + async def _cancel_runner(self, session_id: str) -> None: runner = self._runners.pop(session_id, None) if runner is None: return @@ -280,6 +281,12 @@ async def cancel(self, session_id: str) -> None: finally: current_session_id.reset(token) + async def cancel(self, session_id: str) -> None: + try: + await self._cancel_runner(session_id) + finally: + self.env.unbind_platform_inference(session_id) + async def cancel_all(self) -> None: """Tear down every suspended/live task (server shutdown).""" for session_id in list(self._runners): @@ -353,6 +360,35 @@ async def error_to(msg_id: int | None, code: int, message: str) -> None: {"tasks": [t.manifest_entry() for t in env.tasks.values()]}, ) + elif method == "platform.inference.bind": + upstream_url = params.get("upstream_url") + token = params.get("token") + trace_id = params.get("trace_id") + if not isinstance(upstream_url, str) or not isinstance(token, str): + await error_to( + msg_id, + -32602, + "platform.inference.bind: upstream_url and token must be strings", + ) + continue + if trace_id is not None and not isinstance(trace_id, str): + await error_to( + msg_id, + -32602, + "platform.inference.bind: trace_id must be a string", + ) + continue + binding = env.bind_platform_inference( + session_id, + upstream_url=upstream_url, + token=token, + trace_id=trace_id, + ) + await reply_to( + msg_id, + {"base_url": binding.base_url, "api_key": binding.api_key}, + ) + elif method == "tasks.start": task_id = params.get("id") if not isinstance(task_id, str): diff --git a/hud/environment/tests/test_workspace.py b/hud/environment/tests/test_workspace.py index 5b857ab59..9d6f461b7 100644 --- a/hud/environment/tests/test_workspace.py +++ b/hud/environment/tests/test_workspace.py @@ -1136,6 +1136,96 @@ def test_a_peer_answers_at_the_address_the_task_expects() -> None: bind_addresses([Peer("db", 5432), Peer("db", 5432)]) +def test_platform_inference_proxy_replaces_workspace_credentials() -> None: + import http.client + from http.server import BaseHTTPRequestHandler, HTTPServer + + from hud.environment.platform_inference import PlatformInferenceProxy + + received: dict[str, object] = {} + + class Upstream(BaseHTTPRequestHandler): + def log_message(self, format: str, *args: Any) -> None: + pass + + def do_POST(self) -> None: + length = int(self.headers.get("Content-Length", "0")) + received.update( + path=self.path, + body=self.rfile.read(length), + headers={key.lower(): value for key, value in self.headers.items()}, + ) + body = b'{"ok":true}' + self.send_response(200) + self.send_header("Content-Type", "application/json") + self.send_header("Content-Length", str(len(body))) + self.end_headers() + self.wfile.write(body) + + upstream = HTTPServer(("127.0.0.1", 0), Upstream) + upstream_thread = threading.Thread(target=upstream.serve_forever, daemon=True) + upstream_thread.start() + proxy = PlatformInferenceProxy() + proxy.start() + host, port = proxy.address + binding = proxy.register( + "session", + upstream_url=f"http://127.0.0.1:{upstream.server_port}/gateway", + token="scoped-runtime-token", + trace_id="trace-1", + workspace_url=f"http://{host}:{port}", + ) + try: + route = urllib.parse.urlsplit(binding.base_url) + assert route.hostname is not None + connection = http.client.HTTPConnection(route.hostname, route.port, timeout=5) + connection.request( + "POST", + f"{route.path}/v1/messages?beta=1", + body=b'{"model":"claude"}', + headers={ + "X-Api-Key": binding.api_key, + "Trace-Id": "workspace-chosen", + "Content-Type": "application/json", + }, + ) + response = connection.getresponse() + assert response.status == 200 + assert response.read() == b'{"ok":true}' + connection.close() + + assert received["path"] == "/gateway/v1/messages?beta=1" + assert received["body"] == b'{"model":"claude"}' + headers = cast("dict[str, str]", received["headers"]) + assert headers["hud-runtime-token"] == "scoped-runtime-token" + assert headers["trace-id"] == "trace-1" + assert headers.get("x-api-key") is None + + denied = http.client.HTTPConnection(route.hostname, route.port, timeout=5) + denied.request("POST", f"{route.path}/v1/messages", headers={"X-Api-Key": "wrong"}) + denied_response = denied.getresponse() + assert denied_response.status == 401 + denied_response.read() + denied.close() + + proxy.unregister("session") + gone = http.client.HTTPConnection(route.hostname, route.port, timeout=5) + gone.request( + "POST", + f"{route.path}/v1/messages", + headers={"X-Api-Key": binding.api_key}, + ) + gone_response = gone.getresponse() + assert gone_response.status == 404 + gone_response.read() + gone.close() + finally: + proxy.stop() + upstream.shutdown() + upstream.server_close() + upstream_thread.join(5) + + def test_workspace_names_are_added_to_the_substrates_hosts_rather_than_replacing_it() -> None: """Dropping the substrate's entries would cost the workspace localhost.""" from hud.environment.egress import Peer, hosts_text diff --git a/hud/environment/workspace.py b/hud/environment/workspace.py index 8dc7784ae..131889301 100644 --- a/hud/environment/workspace.py +++ b/hud/environment/workspace.py @@ -633,6 +633,22 @@ def owns_netns(self) -> bool: """ return not self.network or self.allowed_hosts is not None + def add_peer(self, peer: Peer, *, first: bool = False) -> None: + """Add a substrate service before the workspace accepts sessions.""" + if self._sandbox is not None: + raise RuntimeError("workspace peers must be bound before its sandbox starts") + self.peers = (peer, *self.peers) if first else (*self.peers, peer) + if self._hosts_path is not None: + self._hosts_path = self._write_hosts() + + def remove_peer(self, name: str) -> None: + """Remove a substrate service after the workspace has stopped.""" + if self._sandbox is not None: + raise RuntimeError("workspace peers must be unbound after its sandbox stops") + self.peers = tuple(peer for peer in self.peers if peer.name != name) + if self._hosts_path is not None: + self._hosts_path = self._write_hosts() + def _setpriv(self) -> str | None: """Absolute path to ``setpriv``, resolved via the *server's* PATH. diff --git a/hud/eval/runtime/__init__.py b/hud/eval/runtime/__init__.py index 46d8ab207..d9427454b 100644 --- a/hud/eval/runtime/__init__.py +++ b/hud/eval/runtime/__init__.py @@ -6,6 +6,7 @@ Runtime, RuntimeConfig, RuntimeGPU, + RuntimeInference, RuntimeLimits, RuntimeResources, RuntimeTPU, @@ -30,6 +31,7 @@ "Runtime", "RuntimeConfig", "RuntimeGPU", + "RuntimeInference", "RuntimeLimits", "RuntimeResources", "RuntimeTPU", diff --git a/hud/eval/runtime/core.py b/hud/eval/runtime/core.py index fc178cc7e..fa358c33d 100644 --- a/hud/eval/runtime/core.py +++ b/hud/eval/runtime/core.py @@ -170,6 +170,7 @@ class Runtime: url: str params: dict[str, Any] = field(default_factory=dict) config: RuntimeConfig | None = None + inference: RuntimeInference | None = field(default=None, repr=False, compare=False) def __call__(self, task: Task) -> AbstractAsyncContextManager[Runtime]: return nullcontext(self) @@ -185,6 +186,15 @@ async def restore_session(self, session_id: str, source: Path) -> None: validate_session_id(session_id) +@dataclass(frozen=True, slots=True) +class RuntimeInference: + """Controller-only material the environment binds as a workspace-local peer.""" + + upstream_url: str + token: str = field(repr=False) + trace_id: str | None = None + + class Shared: """Lease provider: at most ``width`` rollouts share each task placement. From 196d4435a46d7486a28256dd0ba915f1f7c0e951 Mon Sep 17 00:00:00 2001 From: Jaideep <67646710+jdchawla29@users.noreply.github.com> Date: Sat, 22 Aug 2026 18:43:01 -0700 Subject: [PATCH 2/7] fix(environment): clear capabilities before nested bubblewrap --- hud/environment/capexec.py | 25 ++++++++ hud/environment/namespace.py | 18 +++++- hud/environment/tests/test_workspace.py | 79 +++++++++++++++++++++++++ 3 files changed, 121 insertions(+), 1 deletion(-) create mode 100644 hud/environment/capexec.py diff --git a/hud/environment/capexec.py b/hud/environment/capexec.py new file mode 100644 index 000000000..6ad3d3602 --- /dev/null +++ b/hud/environment/capexec.py @@ -0,0 +1,25 @@ +"""Execute a trusted process without inherited ambient capabilities.""" + +from __future__ import annotations + +import ctypes +import os +import sys +from typing import NoReturn + +_PR_CAP_AMBIENT = 47 +_PR_CAP_AMBIENT_CLEAR_ALL = 4 + + +def exec_without_ambient_capabilities(argv: list[str]) -> NoReturn: + if not argv: + raise ValueError("command required") + libc = ctypes.CDLL(None, use_errno=True) + if libc.prctl(_PR_CAP_AMBIENT, _PR_CAP_AMBIENT_CLEAR_ALL, 0, 0, 0) != 0: + error = ctypes.get_errno() + raise OSError(error, os.strerror(error)) + os.execvp(argv[0], argv) # noqa: S606 - replace this trusted trampoline process + + +if __name__ == "__main__": + exec_without_ambient_capabilities(sys.argv[1:]) diff --git a/hud/environment/namespace.py b/hud/environment/namespace.py index 569978abc..8018df60d 100644 --- a/hud/environment/namespace.py +++ b/hud/environment/namespace.py @@ -324,6 +324,20 @@ def __init__( self.session_used = False self.forwarders: list[asyncio.AbstractServer] = [] + def _prepare_bwrap(self, argv: list[str]) -> list[str]: + argv = list(argv) + try: + index = argv.index(self.bwrap) + except ValueError: + return argv + argv[index:index] = [ + sys.executable, + "-I", + "-S", + str(Path(__file__).with_name("capexec.py")), + ] + return argv + async def serve(self) -> None: server: asyncssh.SSHAcceptor | None = None try: @@ -408,6 +422,7 @@ async def _start_holder(self) -> tuple[ProcessGroup, int]: argv = list(self.holder_argv) index = argv.index(self.bwrap) + 1 argv[index:index] = ["--info-fd", str(write_fd)] + argv = self._prepare_bwrap(argv) if self.map_identities: block_read, block_write = os.pipe() os.set_inheritable(block_read, True) @@ -518,6 +533,7 @@ async def _spawn( if held is None: raise RuntimeError(f"workspace {scope} holder is not running") _, holder_pid = held + request_argv = self._prepare_bwrap(request["argv"]) argv = [ shutil.which("nsenter") or "/usr/bin/nsenter", "--target", @@ -533,7 +549,7 @@ async def _spawn( ), "--", *command_prefix, - *request["argv"], + *request_argv, ] process: ProcessGroup | None = None if channel.term_type: diff --git a/hud/environment/tests/test_workspace.py b/hud/environment/tests/test_workspace.py index 9d6f461b7..638bcc5c7 100644 --- a/hud/environment/tests/test_workspace.py +++ b/hud/environment/tests/test_workspace.py @@ -28,6 +28,7 @@ import pytest from hud.capabilities import SSHClient +from hud.environment import capexec as capexec_mod from hud.environment import namespace as namespace_mod from hud.environment import workspace as workspace_mod from hud.environment.egress import Peer, _field, _UnixServer, _Unrelayable @@ -933,6 +934,84 @@ async def test_namespace_host_only_terminates_a_used_session_holder( holder.wait.assert_not_awaited() +def test_namespace_host_drops_ambient_capabilities_before_bwrap(tmp_path: Path) -> None: + host = namespace_mod._NamespaceHost( + tmp_path / "namespace.sock", + setup_loopback=False, + holder_argv=[], + bwrap="/usr/bin/bwrap", + launcher_depth=0, + map_identities=False, + ports=frozenset(), + ) + assert namespace_mod.__file__ is not None + + assert host._prepare_bwrap(["/usr/bin/bwrap", "--unshare-user"]) == [ + sys.executable, + "-I", + "-S", + str(Path(namespace_mod.__file__).with_name("capexec.py")), + "/usr/bin/bwrap", + "--unshare-user", + ] + assert host._prepare_bwrap( + ["/usr/bin/unshare", "--pid", "/usr/bin/bwrap", "--unshare-user"] + ) == [ + "/usr/bin/unshare", + "--pid", + sys.executable, + "-I", + "-S", + str(Path(namespace_mod.__file__).with_name("capexec.py")), + "/usr/bin/bwrap", + "--unshare-user", + ] + assert host._prepare_bwrap(["/usr/bin/bridge"]) == ["/usr/bin/bridge"] + + +def test_exec_without_ambient_capabilities_clears_caps_before_exec( + monkeypatch: pytest.MonkeyPatch, +) -> None: + prctl = Mock(return_value=0) + monkeypatch.setattr( + capexec_mod.ctypes, + "CDLL", + Mock(return_value=SimpleNamespace(prctl=prctl)), + ) + execvp = Mock(side_effect=RuntimeError("exec called")) + monkeypatch.setattr(os, "execvp", execvp) + + with pytest.raises(RuntimeError, match="exec called"): + capexec_mod.exec_without_ambient_capabilities(["bwrap", "--unshare-user"]) + + prctl.assert_called_once_with( + capexec_mod._PR_CAP_AMBIENT, + capexec_mod._PR_CAP_AMBIENT_CLEAR_ALL, + 0, + 0, + 0, + ) + execvp.assert_called_once_with("bwrap", ["bwrap", "--unshare-user"]) + + +def test_exec_without_ambient_capabilities_fails_closed( + monkeypatch: pytest.MonkeyPatch, +) -> None: + monkeypatch.setattr( + capexec_mod.ctypes, + "CDLL", + Mock(return_value=SimpleNamespace(prctl=Mock(return_value=-1))), + ) + monkeypatch.setattr(capexec_mod.ctypes, "get_errno", Mock(return_value=1)) + execvp = Mock() + monkeypatch.setattr(os, "execvp", execvp) + + with pytest.raises(PermissionError): + capexec_mod.exec_without_ambient_capabilities(["bwrap"]) + + execvp.assert_not_called() + + @pytest.mark.asyncio async def test_run_can_use_a_fresh_no_network_sandbox( tmp_path: Path, monkeypatch: pytest.MonkeyPatch From 26d4d188f37f4f141c9d0e48be1926d92edc7b7e Mon Sep 17 00:00:00 2001 From: Jaideep <67646710+jdchawla29@users.noreply.github.com> Date: Sat, 22 Aug 2026 19:03:27 -0700 Subject: [PATCH 3/7] Revert "fix(environment): clear capabilities before nested bubblewrap" This reverts commit b2dc21d24f2b4303ec177d8c92803d11307c0133. --- hud/environment/capexec.py | 25 -------- hud/environment/namespace.py | 18 +----- hud/environment/tests/test_workspace.py | 79 ------------------------- 3 files changed, 1 insertion(+), 121 deletions(-) delete mode 100644 hud/environment/capexec.py diff --git a/hud/environment/capexec.py b/hud/environment/capexec.py deleted file mode 100644 index 6ad3d3602..000000000 --- a/hud/environment/capexec.py +++ /dev/null @@ -1,25 +0,0 @@ -"""Execute a trusted process without inherited ambient capabilities.""" - -from __future__ import annotations - -import ctypes -import os -import sys -from typing import NoReturn - -_PR_CAP_AMBIENT = 47 -_PR_CAP_AMBIENT_CLEAR_ALL = 4 - - -def exec_without_ambient_capabilities(argv: list[str]) -> NoReturn: - if not argv: - raise ValueError("command required") - libc = ctypes.CDLL(None, use_errno=True) - if libc.prctl(_PR_CAP_AMBIENT, _PR_CAP_AMBIENT_CLEAR_ALL, 0, 0, 0) != 0: - error = ctypes.get_errno() - raise OSError(error, os.strerror(error)) - os.execvp(argv[0], argv) # noqa: S606 - replace this trusted trampoline process - - -if __name__ == "__main__": - exec_without_ambient_capabilities(sys.argv[1:]) diff --git a/hud/environment/namespace.py b/hud/environment/namespace.py index 8018df60d..569978abc 100644 --- a/hud/environment/namespace.py +++ b/hud/environment/namespace.py @@ -324,20 +324,6 @@ def __init__( self.session_used = False self.forwarders: list[asyncio.AbstractServer] = [] - def _prepare_bwrap(self, argv: list[str]) -> list[str]: - argv = list(argv) - try: - index = argv.index(self.bwrap) - except ValueError: - return argv - argv[index:index] = [ - sys.executable, - "-I", - "-S", - str(Path(__file__).with_name("capexec.py")), - ] - return argv - async def serve(self) -> None: server: asyncssh.SSHAcceptor | None = None try: @@ -422,7 +408,6 @@ async def _start_holder(self) -> tuple[ProcessGroup, int]: argv = list(self.holder_argv) index = argv.index(self.bwrap) + 1 argv[index:index] = ["--info-fd", str(write_fd)] - argv = self._prepare_bwrap(argv) if self.map_identities: block_read, block_write = os.pipe() os.set_inheritable(block_read, True) @@ -533,7 +518,6 @@ async def _spawn( if held is None: raise RuntimeError(f"workspace {scope} holder is not running") _, holder_pid = held - request_argv = self._prepare_bwrap(request["argv"]) argv = [ shutil.which("nsenter") or "/usr/bin/nsenter", "--target", @@ -549,7 +533,7 @@ async def _spawn( ), "--", *command_prefix, - *request_argv, + *request["argv"], ] process: ProcessGroup | None = None if channel.term_type: diff --git a/hud/environment/tests/test_workspace.py b/hud/environment/tests/test_workspace.py index 638bcc5c7..9d6f461b7 100644 --- a/hud/environment/tests/test_workspace.py +++ b/hud/environment/tests/test_workspace.py @@ -28,7 +28,6 @@ import pytest from hud.capabilities import SSHClient -from hud.environment import capexec as capexec_mod from hud.environment import namespace as namespace_mod from hud.environment import workspace as workspace_mod from hud.environment.egress import Peer, _field, _UnixServer, _Unrelayable @@ -934,84 +933,6 @@ async def test_namespace_host_only_terminates_a_used_session_holder( holder.wait.assert_not_awaited() -def test_namespace_host_drops_ambient_capabilities_before_bwrap(tmp_path: Path) -> None: - host = namespace_mod._NamespaceHost( - tmp_path / "namespace.sock", - setup_loopback=False, - holder_argv=[], - bwrap="/usr/bin/bwrap", - launcher_depth=0, - map_identities=False, - ports=frozenset(), - ) - assert namespace_mod.__file__ is not None - - assert host._prepare_bwrap(["/usr/bin/bwrap", "--unshare-user"]) == [ - sys.executable, - "-I", - "-S", - str(Path(namespace_mod.__file__).with_name("capexec.py")), - "/usr/bin/bwrap", - "--unshare-user", - ] - assert host._prepare_bwrap( - ["/usr/bin/unshare", "--pid", "/usr/bin/bwrap", "--unshare-user"] - ) == [ - "/usr/bin/unshare", - "--pid", - sys.executable, - "-I", - "-S", - str(Path(namespace_mod.__file__).with_name("capexec.py")), - "/usr/bin/bwrap", - "--unshare-user", - ] - assert host._prepare_bwrap(["/usr/bin/bridge"]) == ["/usr/bin/bridge"] - - -def test_exec_without_ambient_capabilities_clears_caps_before_exec( - monkeypatch: pytest.MonkeyPatch, -) -> None: - prctl = Mock(return_value=0) - monkeypatch.setattr( - capexec_mod.ctypes, - "CDLL", - Mock(return_value=SimpleNamespace(prctl=prctl)), - ) - execvp = Mock(side_effect=RuntimeError("exec called")) - monkeypatch.setattr(os, "execvp", execvp) - - with pytest.raises(RuntimeError, match="exec called"): - capexec_mod.exec_without_ambient_capabilities(["bwrap", "--unshare-user"]) - - prctl.assert_called_once_with( - capexec_mod._PR_CAP_AMBIENT, - capexec_mod._PR_CAP_AMBIENT_CLEAR_ALL, - 0, - 0, - 0, - ) - execvp.assert_called_once_with("bwrap", ["bwrap", "--unshare-user"]) - - -def test_exec_without_ambient_capabilities_fails_closed( - monkeypatch: pytest.MonkeyPatch, -) -> None: - monkeypatch.setattr( - capexec_mod.ctypes, - "CDLL", - Mock(return_value=SimpleNamespace(prctl=Mock(return_value=-1))), - ) - monkeypatch.setattr(capexec_mod.ctypes, "get_errno", Mock(return_value=1)) - execvp = Mock() - monkeypatch.setattr(os, "execvp", execvp) - - with pytest.raises(PermissionError): - capexec_mod.exec_without_ambient_capabilities(["bwrap"]) - - execvp.assert_not_called() - - @pytest.mark.asyncio async def test_run_can_use_a_fresh_no_network_sandbox( tmp_path: Path, monkeypatch: pytest.MonkeyPatch From e0d19dcbb51f1857190c8a7ba9f9ff03db69266f Mon Sep 17 00:00:00 2001 From: Jaideep <67646710+jdchawla29@users.noreply.github.com> Date: Sat, 22 Aug 2026 20:02:15 -0700 Subject: [PATCH 4/7] fix(runtime): keep managed agents outside Harbor harness --- hud/agents/claude/sdk/agent.py | 4 ++-- hud/agents/codex/agent.py | 4 ++-- hud/agents/tests/test_codex_cli_agent.py | 4 ++-- hud/integrations/harbor/env.py | 1 - hud/integrations/harbor/tests/test_contract.py | 1 - 5 files changed, 6 insertions(+), 8 deletions(-) diff --git a/hud/agents/claude/sdk/agent.py b/hud/agents/claude/sdk/agent.py index 69a556e75..ba16c7b83 100644 --- a/hud/agents/claude/sdk/agent.py +++ b/hud/agents/claude/sdk/agent.py @@ -45,8 +45,8 @@ RUN_SCRIPT_PATH = ".hud_run.bat" _MANAGED_CLAUDE_PATHS = { - "linux-x64": "/media/hud/bin/claude/linux-x64/claude", - "linux-x64-musl": "/media/hud/bin/claude/linux-x64-musl/claude", + "linux-x64": "/usr/local/lib/agents/claude/linux-x64/claude", + "linux-x64-musl": "/usr/local/lib/agents/claude/linux-x64-musl/claude", } diff --git a/hud/agents/codex/agent.py b/hud/agents/codex/agent.py index c54fbe367..5cbdbf702 100644 --- a/hud/agents/codex/agent.py +++ b/hud/agents/codex/agent.py @@ -32,8 +32,8 @@ logger = logging.getLogger(__name__) _MANAGED_CODEX_PATHS = { - "linux-x64": "/media/hud/bin/codex/bin/codex", - "linux-x64-musl": "/media/hud/bin/codex/bin/codex", + "linux-x64": "/usr/local/lib/agents/codex/bin/codex", + "linux-x64-musl": "/usr/local/lib/agents/codex/bin/codex", } diff --git a/hud/agents/tests/test_codex_cli_agent.py b/hud/agents/tests/test_codex_cli_agent.py index 8eb5d503c..24999f261 100644 --- a/hud/agents/tests/test_codex_cli_agent.py +++ b/hud/agents/tests/test_codex_cli_agent.py @@ -328,11 +328,11 @@ async def test_executable_resolution_prefers_matching_managed_bundle() -> None: executable = await resolve_executable( cast("Any", ssh), "codex", - {"linux-x64": "/media/hud/bin/codex/bin/codex"}, + {"linux-x64": "/usr/local/lib/agents/codex/bin/codex"}, RuntimeConfig(resources=RuntimeResources(os="linux")), ) - assert executable == "/media/hud/bin/codex/bin/codex" + assert executable == "/usr/local/lib/agents/codex/bin/codex" assert ssh.run.await_count == 2 diff --git a/hud/integrations/harbor/env.py b/hud/integrations/harbor/env.py index 40f2fb973..df8137e83 100644 --- a/hud/integrations/harbor/env.py +++ b/hud/integrations/harbor/env.py @@ -260,7 +260,6 @@ def home(user_id: int | None, *, root: Path | None = None) -> str | None: ) agent_mounts = ( *harness_mounts, - Mount("ro", src=str(ROOT / "bin"), dst=str(ROOT / "bin")), Mount("tmpfs", dst=str(TESTS)), Mount("tmpfs", dst=str(VERIFIER_LOGS)), Mount("ro", src="/dev/null", dst=str(AGENT_ANSWER)), diff --git a/hud/integrations/harbor/tests/test_contract.py b/hud/integrations/harbor/tests/test_contract.py index 531bf4c55..58124a04f 100644 --- a/hud/integrations/harbor/tests/test_contract.py +++ b/hud/integrations/harbor/tests/test_contract.py @@ -130,7 +130,6 @@ def test_adapt_packages_an_image_task_as_a_compose_project(tmp_path: Path) -> No served = (context / "env.py").read_text(encoding="utf-8") assert f'Environment("{context.name}")' in served assert 'Environment(CONFIG["name"])' not in served - assert 'Mount("ro", src=str(ROOT / "bin"), dst=str(ROOT / "bin"))' in served project_root = context / "compose-project" assert _tree_snapshot(project_root / "environment") == authored_environment payload = project_root / "hud" From 76fccb9f5e3b9a3d9f7d41f4faa033bf308dd5af Mon Sep 17 00:00:00 2001 From: Jaideep <67646710+jdchawla29@users.noreply.github.com> Date: Sun, 23 Aug 2026 01:00:09 -0700 Subject: [PATCH 5/7] refactor(runtime): route scoped inference directly --- hud/__init__.py | 2 + hud/agents/claude/sdk/agent.py | 13 +- hud/agents/cli.py | 10 - hud/agents/codex/agent.py | 11 +- hud/agents/tests/test_claude_cli_agent.py | 37 ++- hud/agents/tests/test_codex_cli_agent.py | 20 +- hud/clients/__init__.py | 3 +- hud/clients/client.py | 68 ++---- hud/clients/tests/test_connect.py | 48 +--- hud/environment/__init__.py | 3 +- hud/environment/egress.py | 40 +++ hud/environment/env.py | 138 +++++------ hud/environment/platform_inference.py | 281 ---------------------- hud/environment/server.py | 50 ++-- hud/environment/tests/test_workspace.py | 102 ++------ hud/environment/workspace.py | 4 +- hud/eval/__init__.py | 3 +- hud/eval/run.py | 40 ++- hud/eval/runtime/__init__.py | 2 - hud/eval/runtime/core.py | 10 - 20 files changed, 269 insertions(+), 616 deletions(-) delete mode 100644 hud/environment/platform_inference.py diff --git a/hud/__init__.py b/hud/__init__.py index fdcdc7b0f..90ff1adf8 100644 --- a/hud/__init__.py +++ b/hud/__init__.py @@ -16,6 +16,7 @@ Grade, HostedRuntime, HUDRuntime, + InferenceAccess, Job, LocalRuntime, Run, @@ -43,6 +44,7 @@ "Grade", "HUDRuntime", "HostedRuntime", + "InferenceAccess", "Job", "LocalRuntime", "Run", diff --git a/hud/agents/claude/sdk/agent.py b/hud/agents/claude/sdk/agent.py index ba16c7b83..3b77c1db4 100644 --- a/hud/agents/claude/sdk/agent.py +++ b/hud/agents/claude/sdk/agent.py @@ -21,7 +21,6 @@ WINDOWS_SHELLS, powershell, powershell_quote, - require_platform_isolation, resolve_executable, run_jsonl, ) @@ -35,8 +34,7 @@ if TYPE_CHECKING: from hud.capabilities import SSHClient - from hud.environment.platform_inference import InferenceBinding - from hud.eval.run import Run + from hud.eval.run import InferenceAccess, Run logger = logging.getLogger(__name__) @@ -66,7 +64,6 @@ def __init__(self, config: ClaudeCLIConfig | None = None) -> None: async def __call__(self, run: Run) -> None: mcp_servers: dict[str, dict[str, Any]] = {} ssh = cast("SSHClient", await run.client.open("ssh")) - require_platform_isolation(ssh, run.client.inference) manifest = run.client.manifest assert manifest is not None bindings = manifest.bindings @@ -114,7 +111,7 @@ async def __call__(self, run: Run) -> None: mcp_servers=mcp_servers, prompt=run.prompt_text, executable=executable, - inference=run.client.inference, + inference=run.inference, ) async def _exec( @@ -126,7 +123,7 @@ async def _exec( mcp_servers: dict[str, dict[str, Any]], prompt: str, executable: str = "claude", - inference: InferenceBinding | None = None, + inference: InferenceAccess | None = None, ) -> None: mcp_config_path = await self._write_mcp_config(ssh, mcp_servers) input_text = ( @@ -179,7 +176,7 @@ async def _exec( except (OSError, asyncssh.Error): logger.warning("Failed to remove Claude CLI runtime files") - def _build_env_vars(self, inference: InferenceBinding | None = None) -> dict[str, str]: + def _build_env_vars(self, inference: InferenceAccess | None = None) -> dict[str, str]: env: dict[str, str] = {} use_hud_gateway = self.config.use_hud_gateway if use_hud_gateway is None: @@ -239,7 +236,7 @@ def _build_cli_command( shell: str, mcp_config_path: str | None = None, executable: str = "claude", - inference: InferenceBinding | None = None, + inference: InferenceAccess | None = None, ) -> str: env_vars = self._build_env_vars(inference) is_win = shell in WINDOWS_SHELLS diff --git a/hud/agents/cli.py b/hud/agents/cli.py index 3f7850627..d626c024e 100644 --- a/hud/agents/cli.py +++ b/hud/agents/cli.py @@ -14,22 +14,12 @@ from collections.abc import Callable from hud.capabilities import SSHClient - from hud.environment.platform_inference import InferenceBinding from hud.eval.runtime import RuntimeConfig WINDOWS_SHELLS = ("cmd", "powershell") PROCESS_CLOSE_TIMEOUT_S = 5.0 -def require_platform_isolation(ssh: SSHClient, binding: InferenceBinding | None) -> None: - """Refuse a platform credential binding when the remote shell is not isolated.""" - if binding is not None and ssh.capability.params.get("isolation") != "bwrap": - raise RuntimeError( - "platform inference requires a bwrap-isolated workspace; refusing to expose " - "the workspace-local binding to an unisolated SSH session" - ) - - async def resolve_executable( ssh: SSHClient, command: str, diff --git a/hud/agents/codex/agent.py b/hud/agents/codex/agent.py index 5cbdbf702..cd5e0e127 100644 --- a/hud/agents/codex/agent.py +++ b/hud/agents/codex/agent.py @@ -14,7 +14,6 @@ WINDOWS_SHELLS, powershell, powershell_quote, - require_platform_isolation, resolve_executable, run_jsonl, ) @@ -26,8 +25,7 @@ if TYPE_CHECKING: from hud.capabilities import SSHClient - from hud.environment.platform_inference import InferenceBinding - from hud.eval.run import Run + from hud.eval.run import InferenceAccess, Run logger = logging.getLogger(__name__) @@ -214,7 +212,7 @@ def codex_command( config: CodexCLIConfig, shell: str, executable: str = "codex", - inference: InferenceBinding | None = None, + inference: InferenceAccess | None = None, ) -> str: env: dict[str, str] = {} args = [ @@ -303,7 +301,7 @@ async def run_codex( shell: str, prompt: str, executable: str = "codex", - inference: InferenceBinding | None = None, + inference: InferenceAccess | None = None, ) -> None: command = codex_command(config, shell, executable, inference=inference) logger.info("SSH exec codex CLI (%d chars)", len(command)) @@ -323,7 +321,6 @@ def __init__(self, config: CodexCLIConfig | None = None) -> None: async def __call__(self, run: Run) -> None: ssh = cast("SSHClient", await run.client.open("ssh")) - require_platform_isolation(ssh, run.client.inference) executable = await resolve_executable( ssh, "codex", @@ -337,7 +334,7 @@ async def __call__(self, run: Run) -> None: shell=ssh.capability.params.get("shell", "bash"), prompt=run.prompt_text, executable=executable, - inference=run.client.inference, + inference=run.inference, ) diff --git a/hud/agents/tests/test_claude_cli_agent.py b/hud/agents/tests/test_claude_cli_agent.py index 73a79fa09..af20348be 100644 --- a/hud/agents/tests/test_claude_cli_agent.py +++ b/hud/agents/tests/test_claude_cli_agent.py @@ -30,7 +30,7 @@ from hud.agents.types import AgentStep, ClaudeCLIConfig, ToolStep from hud.capabilities import Capability, SSHClient from hud.capabilities.rfb import WebPScreenshotEncoding -from hud.environment.platform_inference import InferenceBinding +from hud.eval import InferenceAccess from hud.settings import settings from hud.telemetry.context import set_trace_context from hud.types import MCPToolResult @@ -69,16 +69,19 @@ def test_command_follows_explicit_gateway_routing(monkeypatch: pytest.MonkeyPatc assert "ANTHROPIC_MODEL=claude-sonnet-5" in provider -def test_command_prefers_environment_inference_binding() -> None: - binding = InferenceBinding( - base_url="http://127.0.0.1:49123/p/opaque", - api_key="workspace-key", +def test_command_prefers_rollout_inference_access() -> None: + inference = InferenceAccess( + base_url="https://inference.hud.so", + api_key="scoped-runtime-token", ) - gateway = claude_command(ClaudeCLIConfig(use_hud_gateway=True), "bash", inference=binding) + gateway = ClaudeCLIAgent(ClaudeCLIConfig(use_hud_gateway=True))._build_cli_command( + shell="bash", + inference=inference, + ) - assert "ANTHROPIC_BASE_URL=http://127.0.0.1:49123/p/opaque" in gateway - assert "ANTHROPIC_API_KEY=workspace-key" in gateway + assert "ANTHROPIC_BASE_URL=https://inference.hud.so" in gateway + assert "ANTHROPIC_API_KEY=scoped-runtime-token" in gateway assert "HUD_API_KEY" not in gateway assert "Trace-Id" not in gateway @@ -487,7 +490,9 @@ async def open(self, ref: str) -> SSHClient: await agent( cast( "Any", - SimpleNamespace(client=Client(), prompt_text="call the tool", runtime_config=None), + SimpleNamespace( + client=Client(), prompt_text="call the tool", runtime_config=None, inference=None + ), ) ) @@ -559,7 +564,9 @@ async def execute(*_args: Any, **_kwargs: Any) -> None: await agent( cast( "Any", - SimpleNamespace(client=Client(), prompt_text="use the computer", runtime_config=None), + SimpleNamespace( + client=Client(), prompt_text="use the computer", runtime_config=None, inference=None + ), ) ) @@ -651,6 +658,7 @@ async def execute(*_args: Any, **kwargs: Any) -> None: client=Client(), prompt_text="use both screens", runtime_config=None, + inference=None, ), ) ) @@ -916,9 +924,14 @@ async def execute( agent = ClaudeCLIAgent() monkeypatch.setattr(agent, "_exec", execute) - run_a = SimpleNamespace(client=Client(shell_a, ssh_a), prompt_text="first", runtime_config=None) + run_a = SimpleNamespace( + client=Client(shell_a, ssh_a), prompt_text="first", runtime_config=None, inference=None + ) run_b = SimpleNamespace( - client=Client(shell_b, ssh_b), prompt_text="second", runtime_config=None + client=Client(shell_b, ssh_b), + prompt_text="second", + runtime_config=None, + inference=None, ) first = asyncio.create_task(agent(cast("Any", run_a))) diff --git a/hud/agents/tests/test_codex_cli_agent.py b/hud/agents/tests/test_codex_cli_agent.py index 24999f261..872b9f700 100644 --- a/hud/agents/tests/test_codex_cli_agent.py +++ b/hud/agents/tests/test_codex_cli_agent.py @@ -18,7 +18,7 @@ from hud.agents.tests.cli_fakes import fake_run as _fake_run from hud.agents.types import AgentStep, CodexCLIConfig, ToolStep from hud.capabilities import Capability, SSHClient -from hud.environment.platform_inference import InferenceBinding +from hud.eval import InferenceAccess from hud.eval.runtime import RuntimeConfig, RuntimeResources from hud.settings import settings from hud.telemetry.context import set_trace_context @@ -105,16 +105,16 @@ def test_command_follows_explicit_gateway_routing(monkeypatch: pytest.MonkeyPatc assert command.endswith(" -") -def test_command_prefers_environment_inference_binding() -> None: - binding = InferenceBinding( - base_url="http://127.0.0.1:49123/p/opaque", - api_key="workspace-key", +def test_command_prefers_rollout_inference_access() -> None: + inference = InferenceAccess( + base_url="https://inference.hud.so", + api_key="scoped-runtime-token", ) - command = codex_command(CodexCLIConfig(use_hud_gateway=True), "bash", inference=binding) + command = codex_command(CodexCLIConfig(use_hud_gateway=True), "bash", inference=inference) - assert "HUD_API_KEY=workspace-key" in command - assert 'model_providers.hud.base_url="http://127.0.0.1:49123/p/opaque"' in command + assert "HUD_API_KEY=scoped-runtime-token" in command + assert 'model_providers.hud.base_url="https://inference.hud.so"' in command assert "Trace-Id" not in command @@ -299,7 +299,9 @@ async def open(self, ref: str) -> _FakeSSH: agent = CodexCLIAgent() execute = AsyncMock() monkeypatch.setattr("hud.agents.codex.agent.run_codex", execute) - run = SimpleNamespace(client=Client(), prompt_text="Fix it", runtime_config=None) + run = SimpleNamespace( + client=Client(), prompt_text="Fix it", runtime_config=None, inference=None + ) await agent(cast("Any", run)) diff --git a/hud/clients/__init__.py b/hud/clients/__init__.py index 746f827c9..7c670788c 100644 --- a/hud/clients/__init__.py +++ b/hud/clients/__init__.py @@ -2,12 +2,11 @@ from __future__ import annotations -from .client import HudClient, HudProtocolError, InferenceBinding, Manifest, ServerInfo, connect +from .client import HudClient, HudProtocolError, Manifest, ServerInfo, connect __all__ = [ "HudClient", "HudProtocolError", - "InferenceBinding", "Manifest", "ServerInfo", "connect", diff --git a/hud/clients/client.py b/hud/clients/client.py index 6c7443637..9ee17a92a 100644 --- a/hud/clients/client.py +++ b/hud/clients/client.py @@ -27,12 +27,12 @@ RFBClient, SSHClient, ) -from hud.environment.platform_inference import InferenceBinding from hud.environment.utils import read_frame, send_frame, splice if TYPE_CHECKING: - from collections.abc import AsyncIterator + from collections.abc import AsyncIterator, Sequence + from hud.environment.egress import WorkspaceRoute from hud.eval.runtime import Runtime LOGGER = logging.getLogger("hud.clients") @@ -125,7 +125,6 @@ def __init__( self._opened: dict[str, CapabilityClient] = {} self._forwarders: list[asyncio.Server] = [] self._tunnels: set[asyncio.Task[None]] = set() - self.inference: InferenceBinding | None = None # ─── lifecycle ──────────────────────────────────────────────────── @@ -157,14 +156,23 @@ def abort(self) -> None: # ─── handshake ──────────────────────────────────────────────────── - async def hello(self, session_id: str | None = None) -> Manifest: + async def hello( + self, + session_id: str | None = None, + *, + workspace_routes: Sequence[WorkspaceRoute] = (), + ) -> Manifest: """Send ``hello``; cache and return the parsed ``Manifest``. ``session_id`` resumes that parked session on the env — its suspended task, e.g. one a prior connection started — instead of minting a fresh session. """ - params: dict[str, Any] = {} if session_id is None else {"session_id": session_id} + params: dict[str, Any] = {} + if workspace_routes: + params["workspace_routes"] = [route.to_wire() for route in workspace_routes] + if session_id is not None: + params["session_id"] = session_id result = await self._call("hello", params) env = result["env"] bindings = [Capability.from_manifest(binding) for binding in result["bindings"]] @@ -324,29 +332,6 @@ async def grade(self, payload: dict[str, Any]) -> dict[str, Any]: async def cancel(self) -> None: await self._call("tasks.cancel", {}) - async def bind_inference( - self, - *, - upstream_url: str, - token: str, - trace_id: str | None = None, - ) -> InferenceBinding: - """Ask the environment server to expose scoped inference inside its workspace.""" - result = await self._call( - "platform.inference.bind", - { - "upstream_url": upstream_url, - "token": token, - "trace_id": trace_id, - }, - ) - base_url = result.get("base_url") - api_key = result.get("api_key") - if not isinstance(base_url, str) or not isinstance(api_key, str): - raise HudProtocolError(-32603, "platform inference binding was malformed") - self.inference = InferenceBinding(base_url=base_url, api_key=api_key) - return self.inference - # ─── JSON-RPC plumbing ──────────────────────────────────────────── async def _call( @@ -398,6 +383,7 @@ async def _connect_ready( port: int, *, ready_timeout: float, + workspace_routes: Sequence[WorkspaceRoute], interval: float = 0.5, ) -> HudClient: """Connect and complete ``hello``, retrying until the env is ready. @@ -421,7 +407,7 @@ async def _connect_ready( client = HudClient(reader, writer, endpoint=(host, port)) try: - await client.hello() + await client.hello(workspace_routes=workspace_routes) except asyncio.CancelledError: client.abort() raise @@ -454,7 +440,12 @@ def _runtime_ready_timeout(runtime: Runtime, default: float) -> float: @asynccontextmanager -async def connect(runtime: Runtime, *, ready_timeout: float = 240.0) -> AsyncIterator[HudClient]: +async def connect( + runtime: Runtime, + *, + ready_timeout: float = 240.0, + workspace_routes: Sequence[WorkspaceRoute] = (), +) -> AsyncIterator[HudClient]: """Connect a :class:`HudClient` to a provisioned substrate's control channel. Takes the :class:`~hud.eval.runtime.Runtime` a provider yielded (or @@ -471,17 +462,8 @@ async def connect(runtime: Runtime, *, ready_timeout: float = 240.0) -> AsyncIte parts.hostname or "127.0.0.1", parts.port or 0, ready_timeout=_runtime_ready_timeout(runtime, ready_timeout), + workspace_routes=workspace_routes, ) - try: - if runtime.inference is not None: - await client.bind_inference( - upstream_url=runtime.inference.upstream_url, - token=runtime.inference.token, - trace_id=runtime.inference.trace_id, - ) - except BaseException: - await client.close() - raise owner = asyncio.current_task() assert owner is not None heartbeat_error: Exception | None = None @@ -492,9 +474,12 @@ async def heartbeat() -> None: await asyncio.sleep(_CONTROL_HEARTBEAT_INTERVAL_SECONDS) assert client.manifest is not None try: + params: dict[str, Any] = {"session_id": client.manifest.session_id} + if workspace_routes: + params["workspace_routes"] = [route.to_wire() for route in workspace_routes] await client._call( "hello", - {"session_id": client.manifest.session_id}, + params, reply_timeout=_CONTROL_HEARTBEAT_TIMEOUT_SECONDS, ) except HudProtocolError as exc: @@ -523,7 +508,6 @@ async def heartbeat() -> None: __all__ = [ "HudClient", "HudProtocolError", - "InferenceBinding", "Manifest", "ServerInfo", "connect", diff --git a/hud/clients/tests/test_connect.py b/hud/clients/tests/test_connect.py index 813b47f34..c5d62ea47 100644 --- a/hud/clients/tests/test_connect.py +++ b/hud/clients/tests/test_connect.py @@ -19,13 +19,14 @@ import hud.clients.client as client_module from hud.capabilities import Capability, CapabilityClient from hud.clients import connect +from hud.environment import WorkspaceRoute from hud.environment.utils import read_frame, send_frame -from hud.eval.runtime import Runtime, RuntimeInference +from hud.eval.runtime import Runtime HELLO_RESULT = {"session_id": "s-1", "env": {"name": "stub", "version": "1.0"}, "bindings": []} -async def test_connect_binds_runtime_inference_after_hello() -> None: +async def test_connect_sends_workspace_routes_in_hello() -> None: requests: list[dict[str, object]] = [] async def handler(reader: asyncio.StreamReader, writer: asyncio.StreamWriter) -> None: @@ -34,52 +35,25 @@ async def handler(reader: asyncio.StreamReader, writer: asyncio.StreamWriter) -> assert hello is not None requests.append(hello) await send_frame(writer, {"jsonrpc": "2.0", "id": hello["id"], "result": HELLO_RESULT}) - bind = await read_frame(reader) - assert bind is not None - requests.append(bind) - await send_frame( - writer, - { - "jsonrpc": "2.0", - "id": bind["id"], - "result": { - "base_url": "http://127.0.0.1:49123/p/opaque", - "api_key": "workspace-key", - }, - }, - ) await read_frame(reader) finally: writer.close() server = await asyncio.start_server(handler, "127.0.0.1", 0) port = server.sockets[0].getsockname()[1] - runtime = Runtime( - f"tcp://127.0.0.1:{port}", - inference=RuntimeInference( - upstream_url="https://inference.hud.so", - token="scoped-runtime-token", - trace_id="trace-1", - ), - ) - assert "scoped-runtime-token" not in repr(runtime) + runtime = Runtime(f"tcp://127.0.0.1:{port}") + route = WorkspaceRoute("ssh", "inference.hud.so", 443) try: - async with connect(runtime) as client: - assert client.inference is not None - assert client.inference.base_url == "http://127.0.0.1:49123/p/opaque" - assert client.inference.api_key == "workspace-key" + async with connect(runtime, workspace_routes=(route,)): + pass finally: server.close() await server.wait_closed() - assert [request["method"] for request in requests] == ["hello", "platform.inference.bind"] - params = requests[1]["params"] + assert [request["method"] for request in requests] == ["hello"] + params = requests[0]["params"] assert isinstance(params, dict) - assert params == { - "upstream_url": "https://inference.hud.so", - "token": "scoped-runtime-token", - "trace_id": "trace-1", - } + assert params == {"workspace_routes": [route.to_wire()]} async def test_open_retries_transient_capability_connection_failures( @@ -345,11 +319,13 @@ async def fake_connect_ready( port: int, *, ready_timeout: float, + workspace_routes: tuple[WorkspaceRoute, ...], interval: float = 0.5, ) -> _FakeClient: seen["host"] = host seen["port"] = port seen["ready_timeout"] = ready_timeout + assert workspace_routes == () seen["interval"] = interval return _FakeClient() diff --git a/hud/environment/__init__.py b/hud/environment/__init__.py index 1beb9bfb9..acbbb96b5 100644 --- a/hud/environment/__init__.py +++ b/hud/environment/__init__.py @@ -23,7 +23,7 @@ from hud.utils.modules import iter_modules from .arguments import DataFileArg, DataFileRef, DataFilesArg, GradingArg, PromptArg -from .egress import Peer +from .egress import Peer, WorkspaceRoute from .env import Answer, Environment from .workspace import DEFAULT_SYSTEM_MOUNTS, Mount, MountKind, Workspace @@ -100,5 +100,6 @@ def load_environment( "Peer", "PromptArg", "Workspace", + "WorkspaceRoute", "load_environment", ] diff --git a/hud/environment/egress.py b/hud/environment/egress.py index 37ddc3732..3b81ac3db 100644 --- a/hud/environment/egress.py +++ b/hud/environment/egress.py @@ -185,6 +185,45 @@ def address(self) -> tuple[str, int]: return self.target or ("127.0.0.1", self.port) +@dataclass(frozen=True, slots=True) +class WorkspaceRoute: + """A controller-provided host route exposed through one workspace capability.""" + + capability: str + host: str + port: int + + def __post_init__(self) -> None: + if not self.capability or self.capability.strip() != self.capability: + raise ValueError("workspace route capability must not be empty or padded") + if not self.host or self.host.strip() != self.host or any(c.isspace() for c in self.host): + raise ValueError("workspace route host must be a hostname without whitespace") + try: + ipaddress.ip_address(self.host) + except ValueError: + pass + else: + raise ValueError("workspace routes require a hostname, not an IP address") + if not 1 <= self.port <= 65535: + raise ValueError("workspace route port must be between 1 and 65535") + + def to_wire(self) -> dict[str, str | int]: + return {"capability": self.capability, "host": self.host, "port": self.port} + + @classmethod + def from_wire(cls, value: object) -> WorkspaceRoute: + if not isinstance(value, dict): + raise ValueError("workspace routes must be objects") + capability = value.get("capability") + host = value.get("host") + port = value.get("port") + if not isinstance(capability, str) or not isinstance(host, str): + raise ValueError("workspace route capability and host must be strings") + if isinstance(port, bool) or not isinstance(port, int): + raise ValueError("workspace route port must be an integer") + return cls(capability=capability, host=host, port=port) + + def bind_addresses( peers: Sequence[Peer], *, @@ -665,6 +704,7 @@ def stop(self) -> None: "VISITOR_PORT", "Egress", "Peer", + "WorkspaceRoute", "bind_addresses", "hosts_text", "permitted", diff --git a/hud/environment/env.py b/hud/environment/env.py index b0985ec56..90364a8c8 100644 --- a/hud/environment/env.py +++ b/hud/environment/env.py @@ -10,7 +10,6 @@ import contextlib import functools import inspect -import secrets from contextvars import ContextVar from typing import TYPE_CHECKING, Any, Generic, ParamSpec, Protocol, TypeVar, cast @@ -18,8 +17,7 @@ from hud.capabilities import Capability -from .egress import Peer -from .platform_inference import InferenceBinding, PlatformInferenceProxy +from .egress import Peer, WorkspaceRoute from .workspace import Workspace if TYPE_CHECKING: @@ -165,10 +163,8 @@ def __init__( self._on_stop: list[Callable[[], Awaitable[None]]] = [] # Per task-session end (cancel / bye / post-grade cleanup). self._on_task_teardown: list[Callable[[], Awaitable[None]]] = [] - self._workspaces: list[Workspace] = [] - self._platform_inference = PlatformInferenceProxy() - self._platform_peer_port: int | None = None - self._platform_peer_name: str | None = None + self._workspaces: dict[str, Workspace] = {} + self._workspace_routes: dict[WorkspaceRoute, tuple[Workspace, Peer | None]] = {} # ─── task registration ─────────────────────────────────────────── @@ -291,8 +287,10 @@ def workspace( from hud.settings import settings track_files = settings.file_tracking_enabled + if name in self._workspaces: + raise ValueError(f"workspace capability {name!r} is already attached") ws = Workspace(root, track_files=track_files, **kwargs) - self._workspaces.append(ws) + self._workspaces[name] = ws @self.initialize async def _up() -> None: @@ -357,70 +355,66 @@ async def stop(self) -> None: for hook in reversed(self._on_stop): with contextlib.suppress(Exception): await hook() - if self._platform_peer_name is not None: - for workspace in self._workspaces: - workspace.remove_peer(self._platform_peer_name) - self._platform_inference.stop() - self._platform_peer_name = None - self._platform_peer_port = None + for workspace, peer in reversed(self._workspace_routes.values()): + if peer is not None: + workspace.remove_peer(peer) + self._workspace_routes.clear() self._started = False self._hooks_done = False - def _bind_platform_peer(self) -> None: - if self._platform_peer_port is not None or not self._workspaces: - return - unavailable = { - port - for workspace in self._workspaces - for port in (*workspace.ports, *(peer.port for peer in workspace.peers)) - } - for _ in range(100): - port = 20_000 + secrets.randbelow(40_000) - if port not in unavailable: - break - else: - raise RuntimeError("could not allocate a workspace-local platform service port") - peer_name = "platform-" + secrets.token_hex(8) - peer = Peer( - peer_name, - port, - target=self._platform_inference.address, - ) - for workspace in self._workspaces: - workspace.add_peer(peer, first=True) - self._platform_peer_name = peer_name - self._platform_peer_port = port - - def bind_platform_inference( - self, - session_id: str, - *, - upstream_url: str, - token: str, - trace_id: str | None, - ) -> InferenceBinding: - """Bind one control session to inference through every bounded workspace.""" - if not self._started or not self._workspaces: - raise RuntimeError("environment has no workspace available for platform inference") - if any(not workspace.bwrap_available for workspace in self._workspaces): - raise RuntimeError("platform inference requires bwrap-isolated workspaces") - if any(not workspace.owns_netns for workspace in self._workspaces): - raise RuntimeError("platform inference requires network-isolated workspaces") - if self._platform_peer_port is None: - self._platform_inference.start() - try: - self._bind_platform_peer() - except BaseException: - self._platform_inference.stop() - raise - assert self._platform_peer_port is not None - return self._platform_inference.register( - session_id, - upstream_url=upstream_url, - token=token, - trace_id=trace_id, - workspace_url=f"http://127.0.0.1:{self._platform_peer_port}", - ) - - def unbind_platform_inference(self, session_id: str) -> None: - self._platform_inference.unregister(session_id) + def bind_workspace_routes(self, routes: Sequence[WorkspaceRoute]) -> None: + """Install controller routes before a workspace starts its sandbox.""" + if not self._started: + raise RuntimeError("environment must be started before workspace routes are bound") + + planned: list[tuple[WorkspaceRoute, Workspace, Peer | None]] = [] + for route in dict.fromkeys(routes): + if route in self._workspace_routes: + continue + workspace = self._workspaces.get(route.capability) + if workspace is None and route.capability in {"ssh", "ssh/2"}: + if len(self._workspaces) > 1: + names = ", ".join(sorted(self._workspaces)) + raise RuntimeError( + f"workspace capability {route.capability!r} is ambiguous: {names}" + ) + workspace = next(iter(self._workspaces.values()), None) + if workspace is None: + raise RuntimeError(f"workspace capability {route.capability!r} does not exist") + if not workspace.bwrap_available or not workspace.owns_netns: + raise RuntimeError( + f"workspace route for {route.capability!r} requires an isolated network" + ) + matching = [ + peer + for peer in workspace.peers + if peer.name == route.host and peer.port == route.port + ] + if matching: + if any(peer.address != (route.host, route.port) for peer in matching): + raise RuntimeError( + f"workspace route {route.host}:{route.port} conflicts with an authored peer" + ) + planned.append((route, workspace, None)) + continue + planned.append( + ( + route, + workspace, + Peer(route.host, route.port, target=(route.host, route.port)), + ) + ) + + bound: list[tuple[WorkspaceRoute, Workspace, Peer | None]] = [] + try: + for route, workspace, peer in planned: + if peer is not None: + workspace.add_peer(peer, first=True) + self._workspace_routes[route] = (workspace, peer) + bound.append((route, workspace, peer)) + except BaseException: + for route, workspace, peer in reversed(bound): + if peer is not None: + workspace.remove_peer(peer) + self._workspace_routes.pop(route, None) + raise diff --git a/hud/environment/platform_inference.py b/hud/environment/platform_inference.py deleted file mode 100644 index f57ed66ab..000000000 --- a/hud/environment/platform_inference.py +++ /dev/null @@ -1,281 +0,0 @@ -"""Platform inference made available inside bounded workspaces.""" - -from __future__ import annotations - -import hmac -import http.client -import secrets -import threading -import urllib.parse -from dataclasses import dataclass -from http.server import BaseHTTPRequestHandler, ThreadingHTTPServer -from typing import TYPE_CHECKING, Any - -from .egress import _HOP_BY_HOP, _field, _Unrelayable - -if TYPE_CHECKING: - from collections.abc import Mapping - - -_CREDENTIAL_HEADERS = frozenset( - {"authorization", "hud-api-key", "hud-runtime-token", "x-api-key", "x-goog-api-key"} -) -_TRACE_HEADERS = frozenset({"trace-id", "x-trace-id", "x-hud-trace-id"}) - - -@dataclass(frozen=True, slots=True) -class InferenceBinding: - """Workspace-local connection details for one platform inference lease.""" - - base_url: str - api_key: str - - -@dataclass(frozen=True, slots=True) -class _Lease: - upstream: urllib.parse.SplitResult - upstream_token: str - trace_id: str | None - client_key: str - - -def _request_key(headers: Mapping[str, str]) -> str | None: - for name in ("hud-api-key", "x-api-key", "x-goog-api-key"): - if value := headers.get(name): - return value - authorization = headers.get("authorization", "") - scheme, _, value = authorization.partition(" ") - return value if scheme.lower() == "bearer" and value else None - - -class _InferenceHandler(BaseHTTPRequestHandler): - protocol_version = "HTTP/1.1" - proxy: PlatformInferenceProxy - - def log_message(self, format: str, *args: Any) -> None: - """Requests and opaque lease paths must not enter environment logs.""" - - def _fail(self, status: int, reason: str) -> None: - self.send_response(status) - self.send_header("X-Proxy-Error", reason) - self.send_header("Content-Length", "0") - self.end_headers() - - def _request_body(self) -> bytes | None: - transfer = self.headers.get("Transfer-Encoding") - length = self.headers.get("Content-Length") - if transfer is None: - if length is None: - return None - size = int(length) - if size < 0: - raise ValueError - body = self.rfile.read(size) - if len(body) != size: - raise ValueError - return body - if length is not None or transfer.strip().lower() != "chunked": - raise ValueError - - chunks: list[bytes] = [] - while True: - line = self.rfile.readline(65537) - if len(line) > 65536 or not line.endswith(b"\r\n"): - raise ValueError - size_text = line[:-2].split(b";", 1)[0].strip() - if not size_text or any(byte not in b"0123456789abcdefABCDEF" for byte in size_text): - raise ValueError - size = int(size_text, 16) - if size == 0: - while True: - trailer = self.rfile.readline(65537) - if len(trailer) > 65536 or not trailer.endswith(b"\r\n"): - raise ValueError - if trailer == b"\r\n": - return b"".join(chunks) - chunk = self.rfile.read(size) - if len(chunk) != size or self.rfile.read(2) != b"\r\n": - raise ValueError - chunks.append(chunk) - - def _forward(self) -> None: - parts = urllib.parse.urlsplit(self.path) - lease, upstream_path = self.proxy.resolve(parts.path) - if lease is None: - self._fail(404, "unknown-lease") - return - supplied = _request_key({key.lower(): value for key, value in self.headers.items()}) - if supplied is None or not hmac.compare_digest(supplied, lease.client_key): - self._fail(401, "invalid-lease-key") - return - try: - body = self._request_body() - except (ValueError, OverflowError): - self.close_connection = True - self._fail(400, "invalid-request-body") - return - - headers = { - key: value - for key, value in self.headers.items() - if key.lower() not in _HOP_BY_HOP | _CREDENTIAL_HEADERS | _TRACE_HEADERS | {"host"} - } - headers["Hud-Runtime-Token"] = lease.upstream_token - if lease.trace_id is not None: - headers["Trace-Id"] = lease.trace_id - - upstream = lease.upstream - base = upstream.path.rstrip("/") - path = f"{base}/{upstream_path.lstrip('/')}" - if parts.query: - path = f"{path}?{parts.query}" - host = upstream.hostname - assert host is not None - connection: http.client.HTTPConnection - if upstream.scheme == "https": - connection = http.client.HTTPSConnection(host, upstream.port, timeout=300) - else: - connection = http.client.HTTPConnection(host, upstream.port, timeout=300) - response_started = False - try: - connection.request(self.command, path, body=body, headers=headers) - response = connection.getresponse() - relayed = [ - _field(key, value) - for key, value in response.getheaders() - if key.lower() not in _HOP_BY_HOP and key.lower() != "content-length" - ] - length = response.getheader("Content-Length") - framed = length is not None and length.strip().isdigit() - _field("Reason", response.reason or "") - response_started = True - self.send_response(response.status, response.reason) - for key, value in relayed: - self.send_header(key, value) - if framed: - assert length is not None - self.send_header("Content-Length", length.strip()) - else: - self.send_header("Connection", "close") - self.close_connection = True - self.end_headers() - while chunk := response.read(65536): - self.wfile.write(chunk) - self.wfile.flush() - except _Unrelayable: - self._fail(502, "unrelayable-upstream-header") - except (OSError, http.client.HTTPException): - if response_started: - self.close_connection = True - else: - self._fail(502, "upstream-failure") - finally: - connection.close() - - do_GET = _forward - do_HEAD = _forward - do_POST = _forward - do_PUT = _forward - do_DELETE = _forward - do_PATCH = _forward - do_OPTIONS = _forward - - -class PlatformInferenceProxy: - """Environment-owned reverse proxy with per-control-session credentials.""" - - def __init__(self) -> None: - self._server: ThreadingHTTPServer | None = None - self._thread: threading.Thread | None = None - self._leases: dict[str, tuple[str, _Lease]] = {} - self._lock = threading.Lock() - - @property - def address(self) -> tuple[str, int]: - if self._server is None: - raise RuntimeError("platform inference proxy is not started") - host, port = self._server.server_address[:2] - return str(host), int(port) - - def start(self) -> None: - if self._server is not None: - return - handler = type("_ScopedInferenceHandler", (_InferenceHandler,), {"proxy": self}) - server = ThreadingHTTPServer(("127.0.0.1", 0), handler) - server.daemon_threads = True - self._server = server - self._thread = threading.Thread(target=server.serve_forever, daemon=True) - self._thread.start() - - def register( - self, - session_id: str, - *, - upstream_url: str, - token: str, - trace_id: str | None, - workspace_url: str, - ) -> InferenceBinding: - upstream = urllib.parse.urlsplit(upstream_url) - if ( - upstream.scheme not in {"http", "https"} - or upstream.hostname is None - or upstream.username is not None - or upstream.password is not None - or upstream.query - or upstream.fragment - ): - raise ValueError("platform inference upstream must be an HTTP(S) base URL") - if not token: - raise ValueError("platform inference token must not be empty") - with self._lock: - existing = self._leases.get(session_id) - if existing is not None: - route, lease = existing - requested = (upstream, token, trace_id) - current = (lease.upstream, lease.upstream_token, lease.trace_id) - if requested != current: - raise RuntimeError("platform inference is already bound for this session") - return InferenceBinding(f"{workspace_url}/{route}", lease.client_key) - route = "p/" + secrets.token_urlsafe(18) - lease = _Lease( - upstream=upstream, - upstream_token=token, - trace_id=trace_id, - client_key=secrets.token_urlsafe(32), - ) - self._leases[session_id] = (route, lease) - return InferenceBinding(f"{workspace_url}/{route}", lease.client_key) - - def resolve(self, path: str) -> tuple[_Lease | None, str]: - stripped = path.lstrip("/") - prefix, separator, remainder = stripped.partition("/") - if prefix != "p" or not separator: - return None, "" - token, separator, upstream_path = remainder.partition("/") - if not token or not separator: - return None, "" - route = f"p/{token}" - with self._lock: - for stored_route, lease in self._leases.values(): - if hmac.compare_digest(route, stored_route): - return lease, "/" + upstream_path - return None, "" - - def unregister(self, session_id: str) -> None: - with self._lock: - self._leases.pop(session_id, None) - - def stop(self) -> None: - server, self._server = self._server, None - if server is not None: - server.shutdown() - server.server_close() - thread, self._thread = self._thread, None - if thread is not None: - thread.join(timeout=5) - with self._lock: - self._leases.clear() - - -__all__ = ["InferenceBinding", "PlatformInferenceProxy"] diff --git a/hud/environment/server.py b/hud/environment/server.py index 44a947cd3..28d8b15aa 100644 --- a/hud/environment/server.py +++ b/hud/environment/server.py @@ -28,6 +28,7 @@ from hud.graders.results import EvaluationResult +from .egress import WorkspaceRoute from .env import Answer, current_session_id from .utils import error, read_frame, reply, send_frame, splice @@ -255,7 +256,6 @@ async def grade(self, session_id: str, payload: dict[str, Any]) -> dict[str, Any return await runner.grade(payload) finally: current_session_id.reset(token) - self.env.unbind_platform_inference(claim_sid) def _adopt_parked(self) -> tuple[str, TaskRunner]: """Claim the parked session iff unambiguous — the blind-reconnect grade path.""" @@ -282,10 +282,7 @@ async def _cancel_runner(self, session_id: str) -> None: current_session_id.reset(token) async def cancel(self, session_id: str) -> None: - try: - await self._cancel_runner(session_id) - finally: - self.env.unbind_platform_inference(session_id) + await self._cancel_runner(session_id) async def cancel_all(self) -> None: """Tear down every suspended/live task (server shutdown).""" @@ -341,6 +338,20 @@ async def error_to(msg_id: int | None, code: int, message: str) -> None: self._live.add(session_id) current_session_id.reset(session_token) session_token = current_session_id.set(session_id) + raw_routes = params.get("workspace_routes", []) + if not isinstance(raw_routes, list): + await error_to( + msg_id, -32602, "hello: 'workspace_routes' must be a list" + ) + continue + try: + workspace_routes = [ + WorkspaceRoute.from_wire(route) for route in raw_routes + ] + except ValueError as exc: + await error_to(msg_id, -32602, f"hello: {exc}") + continue + env.bind_workspace_routes(workspace_routes) # env.start() ran before serving, so hook-published # capabilities (e.g. a workspace's ssh address) are # already concrete here. @@ -360,35 +371,6 @@ async def error_to(msg_id: int | None, code: int, message: str) -> None: {"tasks": [t.manifest_entry() for t in env.tasks.values()]}, ) - elif method == "platform.inference.bind": - upstream_url = params.get("upstream_url") - token = params.get("token") - trace_id = params.get("trace_id") - if not isinstance(upstream_url, str) or not isinstance(token, str): - await error_to( - msg_id, - -32602, - "platform.inference.bind: upstream_url and token must be strings", - ) - continue - if trace_id is not None and not isinstance(trace_id, str): - await error_to( - msg_id, - -32602, - "platform.inference.bind: trace_id must be a string", - ) - continue - binding = env.bind_platform_inference( - session_id, - upstream_url=upstream_url, - token=token, - trace_id=trace_id, - ) - await reply_to( - msg_id, - {"base_url": binding.base_url, "api_key": binding.api_key}, - ) - elif method == "tasks.start": task_id = params.get("id") if not isinstance(task_id, str): diff --git a/hud/environment/tests/test_workspace.py b/hud/environment/tests/test_workspace.py index 9d6f461b7..451873137 100644 --- a/hud/environment/tests/test_workspace.py +++ b/hud/environment/tests/test_workspace.py @@ -30,7 +30,7 @@ from hud.capabilities import SSHClient from hud.environment import namespace as namespace_mod from hud.environment import workspace as workspace_mod -from hud.environment.egress import Peer, _field, _UnixServer, _Unrelayable +from hud.environment.egress import Peer, WorkspaceRoute, _field, _UnixServer, _Unrelayable from hud.environment.workspace import Bubblewrap, Mount, Workspace from hud.utils.process import ProcessGroup, ProcessResult @@ -1136,94 +1136,26 @@ def test_a_peer_answers_at_the_address_the_task_expects() -> None: bind_addresses([Peer("db", 5432), Peer("db", 5432)]) -def test_platform_inference_proxy_replaces_workspace_credentials() -> None: - import http.client - from http.server import BaseHTTPRequestHandler, HTTPServer +async def test_workspace_route_is_bound_once_and_removed_on_stop(tmp_path: Path) -> None: + from hud.environment import Environment - from hud.environment.platform_inference import PlatformInferenceProxy + env = Environment() + workspace = env.workspace(tmp_path / "root", track_files=False) + workspace._bwrap = cast("Any", object()) + env._started = True + route = WorkspaceRoute("ssh", "inference.hud.so", 443) - received: dict[str, object] = {} + env.bind_workspace_routes([route, route]) + env.bind_workspace_routes([route]) - class Upstream(BaseHTTPRequestHandler): - def log_message(self, format: str, *args: Any) -> None: - pass + assert workspace.peers == (Peer("inference.hud.so", 443, target=("inference.hud.so", 443)),) + await env.stop() + assert workspace.peers == () - def do_POST(self) -> None: - length = int(self.headers.get("Content-Length", "0")) - received.update( - path=self.path, - body=self.rfile.read(length), - headers={key.lower(): value for key, value in self.headers.items()}, - ) - body = b'{"ok":true}' - self.send_response(200) - self.send_header("Content-Type", "application/json") - self.send_header("Content-Length", str(len(body))) - self.end_headers() - self.wfile.write(body) - - upstream = HTTPServer(("127.0.0.1", 0), Upstream) - upstream_thread = threading.Thread(target=upstream.serve_forever, daemon=True) - upstream_thread.start() - proxy = PlatformInferenceProxy() - proxy.start() - host, port = proxy.address - binding = proxy.register( - "session", - upstream_url=f"http://127.0.0.1:{upstream.server_port}/gateway", - token="scoped-runtime-token", - trace_id="trace-1", - workspace_url=f"http://{host}:{port}", - ) - try: - route = urllib.parse.urlsplit(binding.base_url) - assert route.hostname is not None - connection = http.client.HTTPConnection(route.hostname, route.port, timeout=5) - connection.request( - "POST", - f"{route.path}/v1/messages?beta=1", - body=b'{"model":"claude"}', - headers={ - "X-Api-Key": binding.api_key, - "Trace-Id": "workspace-chosen", - "Content-Type": "application/json", - }, - ) - response = connection.getresponse() - assert response.status == 200 - assert response.read() == b'{"ok":true}' - connection.close() - - assert received["path"] == "/gateway/v1/messages?beta=1" - assert received["body"] == b'{"model":"claude"}' - headers = cast("dict[str, str]", received["headers"]) - assert headers["hud-runtime-token"] == "scoped-runtime-token" - assert headers["trace-id"] == "trace-1" - assert headers.get("x-api-key") is None - - denied = http.client.HTTPConnection(route.hostname, route.port, timeout=5) - denied.request("POST", f"{route.path}/v1/messages", headers={"X-Api-Key": "wrong"}) - denied_response = denied.getresponse() - assert denied_response.status == 401 - denied_response.read() - denied.close() - - proxy.unregister("session") - gone = http.client.HTTPConnection(route.hostname, route.port, timeout=5) - gone.request( - "POST", - f"{route.path}/v1/messages", - headers={"X-Api-Key": binding.api_key}, - ) - gone_response = gone.getresponse() - assert gone_response.status == 404 - gone_response.read() - gone.close() - finally: - proxy.stop() - upstream.shutdown() - upstream.server_close() - upstream_thread.join(5) + +def test_workspace_route_rejects_ip_literals() -> None: + with pytest.raises(ValueError, match="hostname"): + WorkspaceRoute("ssh", "127.0.0.1", 443) def test_workspace_names_are_added_to_the_substrates_hosts_rather_than_replacing_it() -> None: diff --git a/hud/environment/workspace.py b/hud/environment/workspace.py index 131889301..c0c93d028 100644 --- a/hud/environment/workspace.py +++ b/hud/environment/workspace.py @@ -641,11 +641,11 @@ def add_peer(self, peer: Peer, *, first: bool = False) -> None: if self._hosts_path is not None: self._hosts_path = self._write_hosts() - def remove_peer(self, name: str) -> None: + def remove_peer(self, peer: Peer) -> None: """Remove a substrate service after the workspace has stopped.""" if self._sandbox is not None: raise RuntimeError("workspace peers must be unbound after its sandbox stops") - self.peers = tuple(peer for peer in self.peers if peer.name != name) + self.peers = tuple(candidate for candidate in self.peers if candidate != peer) if self._hosts_path is not None: self._hosts_path = self._write_hosts() diff --git a/hud/eval/__init__.py b/hud/eval/__init__.py index 0bce06d61..7f67c6e94 100644 --- a/hud/eval/__init__.py +++ b/hud/eval/__init__.py @@ -32,7 +32,7 @@ from .chat import Chat from .job import Job -from .run import Grade, Run, rollout +from .run import Grade, InferenceAccess, Run, rollout from .runtime import ( ComposeProject, DaytonaRuntime, @@ -63,6 +63,7 @@ "Grade", "HUDRuntime", "HostedRuntime", + "InferenceAccess", "Job", "LocalRuntime", "ModalRuntime", diff --git a/hud/eval/run.py b/hud/eval/run.py index 839d106e9..d8ab985c9 100644 --- a/hud/eval/run.py +++ b/hud/eval/run.py @@ -27,10 +27,12 @@ import uuid from dataclasses import dataclass, field from typing import TYPE_CHECKING, Any, Literal, Self, cast +from urllib.parse import urlsplit import mcp.types as mcp_types from hud.clients import HudProtocolError, connect +from hud.environment import WorkspaceRoute from hud.graders.results import SubScore from hud.telemetry.context import set_trace_context from hud.types import Step, TaskCall, Trace @@ -52,6 +54,35 @@ logger = logging.getLogger("hud.eval.run") +@dataclass(frozen=True, slots=True) +class InferenceAccess: + """Scoped inference access for an agent running inside a workspace.""" + + base_url: str + api_key: str = field(repr=False) + workspace: str = "ssh" + + def __post_init__(self) -> None: + parts = urlsplit(self.base_url) + if parts.scheme not in {"http", "https"} or parts.hostname is None: + raise ValueError("inference base_url must be an HTTP(S) URL with a hostname") + if parts.username is not None or parts.password is not None: + raise ValueError("inference base_url must not contain credentials") + if parts.query or parts.fragment: + raise ValueError("inference base_url must not contain a query or fragment") + if not self.api_key: + raise ValueError("inference api_key must not be empty") + + def workspace_route(self) -> WorkspaceRoute: + parts = urlsplit(self.base_url) + assert parts.hostname is not None + return WorkspaceRoute( + capability=self.workspace, + host=parts.hostname, + port=parts.port or (443 if parts.scheme == "https" else 80), + ) + + def validate_rollout_timeouts( task: Task, agent: Agent, @@ -200,12 +231,14 @@ def __init__( *, best_effort_grade: bool = False, runtime_config: RuntimeConfig | None = None, + inference: InferenceAccess | None = None, ) -> None: self._client = client self._task_id = task_id self._args = args self._best_effort_grade = best_effort_grade self.runtime_config = runtime_config + self.inference = inference #: The task's opening prompt as ``tasks.start`` returned it: plain #: text, or a list of message dicts (``{"role", "content"}``) for #: chat-style / multi-turn prompts. Agents consume the normalized @@ -445,6 +478,7 @@ async def rollout( group_id: str | None = None, trace_id: str | None = None, rollout_timeout: float | None = None, + inference: InferenceAccess | None = None, ) -> Run: """Drive one task to a graded :class:`Run` here, against ``runtime``'s channel. @@ -535,7 +569,8 @@ async def close_actor() -> None: scope.push_async_callback(close_actor) addr = await actor.enter_async_context(runtime(task)) _phase = "starting task" - async with connect(addr) as actor_client: + workspace_routes = (inference.workspace_route(),) if inference is not None else () + async with connect(addr, workspace_routes=workspace_routes) as actor_client: client = actor_client live = Run( actor_client, @@ -543,6 +578,7 @@ async def close_actor() -> None: task.args, best_effort_grade=task.verifier is not None, runtime_config=addr.config or actor_runtime_config, + inference=inference, ) live._runtime = addr.url # the placement record for the receipt async with live: # start on enter; complete on exit @@ -684,4 +720,4 @@ def _consume_task_result(task: asyncio.Future[Any]) -> None: task.result() -__all__ = ["Grade", "Run", "rollout"] +__all__ = ["Grade", "InferenceAccess", "Run", "rollout"] diff --git a/hud/eval/runtime/__init__.py b/hud/eval/runtime/__init__.py index d9427454b..46d8ab207 100644 --- a/hud/eval/runtime/__init__.py +++ b/hud/eval/runtime/__init__.py @@ -6,7 +6,6 @@ Runtime, RuntimeConfig, RuntimeGPU, - RuntimeInference, RuntimeLimits, RuntimeResources, RuntimeTPU, @@ -31,7 +30,6 @@ "Runtime", "RuntimeConfig", "RuntimeGPU", - "RuntimeInference", "RuntimeLimits", "RuntimeResources", "RuntimeTPU", diff --git a/hud/eval/runtime/core.py b/hud/eval/runtime/core.py index fa358c33d..fc178cc7e 100644 --- a/hud/eval/runtime/core.py +++ b/hud/eval/runtime/core.py @@ -170,7 +170,6 @@ class Runtime: url: str params: dict[str, Any] = field(default_factory=dict) config: RuntimeConfig | None = None - inference: RuntimeInference | None = field(default=None, repr=False, compare=False) def __call__(self, task: Task) -> AbstractAsyncContextManager[Runtime]: return nullcontext(self) @@ -186,15 +185,6 @@ async def restore_session(self, session_id: str, source: Path) -> None: validate_session_id(session_id) -@dataclass(frozen=True, slots=True) -class RuntimeInference: - """Controller-only material the environment binds as a workspace-local peer.""" - - upstream_url: str - token: str = field(repr=False) - trace_id: str | None = None - - class Shared: """Lease provider: at most ``width`` rollouts share each task placement. From 4f442431f6f02a91e7877bae089c284cf73503e7 Mon Sep 17 00:00:00 2001 From: Jaideep <67646710+jdchawla29@users.noreply.github.com> Date: Sun, 23 Aug 2026 01:42:22 -0700 Subject: [PATCH 6/7] refactor(runtime): separate inference connection from reachability --- hud/__init__.py | 4 +- hud/agents/claude/sdk/agent.py | 10 +-- hud/agents/codex/agent.py | 48 ++++++++------ hud/agents/tests/test_claude_cli_agent.py | 17 +++-- hud/agents/tests/test_codex_cli_agent.py | 29 ++++++-- hud/agents/types.py | 6 +- hud/clients/tests/test_connect.py | 13 ++++ hud/environment/egress.py | 14 ++++ hud/eval/__init__.py | 4 +- hud/eval/run.py | 81 +++++++++++------------ hud/eval/tests/test_rollout.py | 33 ++++++++- hud/tests/test_init_module.py | 1 + 12 files changed, 178 insertions(+), 82 deletions(-) diff --git a/hud/__init__.py b/hud/__init__.py index 90ff1adf8..6cde008be 100644 --- a/hud/__init__.py +++ b/hud/__init__.py @@ -16,7 +16,7 @@ Grade, HostedRuntime, HUDRuntime, - InferenceAccess, + InferenceConnection, Job, LocalRuntime, Run, @@ -44,7 +44,7 @@ "Grade", "HUDRuntime", "HostedRuntime", - "InferenceAccess", + "InferenceConnection", "Job", "LocalRuntime", "Run", diff --git a/hud/agents/claude/sdk/agent.py b/hud/agents/claude/sdk/agent.py index 3b77c1db4..abd1336f7 100644 --- a/hud/agents/claude/sdk/agent.py +++ b/hud/agents/claude/sdk/agent.py @@ -34,7 +34,7 @@ if TYPE_CHECKING: from hud.capabilities import SSHClient - from hud.eval.run import InferenceAccess, Run + from hud.eval.run import InferenceConnection, Run logger = logging.getLogger(__name__) @@ -123,7 +123,7 @@ async def _exec( mcp_servers: dict[str, dict[str, Any]], prompt: str, executable: str = "claude", - inference: InferenceAccess | None = None, + inference: InferenceConnection | None = None, ) -> None: mcp_config_path = await self._write_mcp_config(ssh, mcp_servers) input_text = ( @@ -176,7 +176,7 @@ async def _exec( except (OSError, asyncssh.Error): logger.warning("Failed to remove Claude CLI runtime files") - def _build_env_vars(self, inference: InferenceAccess | None = None) -> dict[str, str]: + def _build_env_vars(self, inference: InferenceConnection | None = None) -> dict[str, str]: env: dict[str, str] = {} use_hud_gateway = self.config.use_hud_gateway if use_hud_gateway is None: @@ -185,7 +185,7 @@ def _build_env_vars(self, inference: InferenceAccess | None = None) -> dict[str, if use_hud_gateway: if inference is not None: base_url = inference.base_url - api_key = inference.api_key + api_key = inference.credential elif settings.api_key: base_url = settings.hud_gateway_url api_key = settings.api_key @@ -236,7 +236,7 @@ def _build_cli_command( shell: str, mcp_config_path: str | None = None, executable: str = "claude", - inference: InferenceAccess | None = None, + inference: InferenceConnection | None = None, ) -> str: env_vars = self._build_env_vars(inference) is_win = shell in WINDOWS_SHELLS diff --git a/hud/agents/codex/agent.py b/hud/agents/codex/agent.py index cd5e0e127..b9a4019f6 100644 --- a/hud/agents/codex/agent.py +++ b/hud/agents/codex/agent.py @@ -25,7 +25,7 @@ if TYPE_CHECKING: from hud.capabilities import SSHClient - from hud.eval.run import InferenceAccess, Run + from hud.eval.run import InferenceConnection, Run logger = logging.getLogger(__name__) @@ -212,7 +212,7 @@ def codex_command( config: CodexCLIConfig, shell: str, executable: str = "codex", - inference: InferenceAccess | None = None, + inference: InferenceConnection | None = None, ) -> str: env: dict[str, str] = {} args = [ @@ -235,18 +235,20 @@ def codex_command( if use_hud_gateway: if inference is not None: base_url = inference.base_url - api_key = inference.api_key + credential = inference.credential + credential_env = "HUD_RUNTIME_INFERENCE_TOKEN" elif settings.api_key: base_url = settings.hud_gateway_url - api_key = settings.api_key + credential = settings.api_key + credential_env = "HUD_API_KEY" else: raise ValueError("HUD_API_KEY is required for HUD gateway routing") - env["HUD_API_KEY"] = api_key + env[credential_env] = credential overrides = { "model_provider": "hud", "model_providers.hud.name": "HUD", "model_providers.hud.base_url": base_url, - "model_providers.hud.env_key": "HUD_API_KEY", + "model_providers.hud.env_key": credential_env, "model_providers.hud.wire_api": "responses", } for key, value in overrides.items(): @@ -262,35 +264,43 @@ def codex_command( env["CODEX_API_KEY"] = settings.openai_api_key args.append("-") + isolate_home = bool(env) if shell in WINDOWS_SHELLS: - script = ";".join( - [ + invocation = ( + f"& {powershell_quote(executable)} " + f"{' '.join(powershell_quote(arg) for arg in args[1:])}; " + "$hudExitCode=$LASTEXITCODE" + ) + statements = [ + *(f"$env:{key}={powershell_quote(value)}" for key, value in env.items()), + invocation, + "exit $hudExitCode", + ] + if isolate_home: + statements = [ "$codexHome=Join-Path ([System.IO.Path]::GetTempPath()) " "('hud-codex-' + [System.Guid]::NewGuid())", "New-Item -ItemType Directory -Force -Path $codexHome | Out-Null", "$env:CODEX_HOME=$codexHome", *(f"$env:{key}={powershell_quote(value)}" for key, value in env.items()), - f"try {{ & {powershell_quote(executable)} " - f"{' '.join(powershell_quote(arg) for arg in args[1:])}; " - "$hudExitCode=$LASTEXITCODE } finally { Remove-Item -Recurse -Force " + f"try {{ {invocation} }} finally {{ Remove-Item -Recurse -Force " "$codexHome }", "exit $hudExitCode", ] - ) - return powershell(script) + return powershell(";".join(statements)) command = " ".join(shlex.quote(arg) for arg in args) env_prefix = " ".join(f"{key}={shlex.quote(value)}" for key, value in env.items()) invocation = f"{env_prefix} {command}" if env_prefix else command - return "; ".join( - [ + statements = ['export PATH="$HOME/.local/bin:$PATH"', invocation] + if isolate_home: + statements = [ 'codex_home=$(mktemp -d "${TMPDIR:-/tmp}/hud-codex.XXXXXX") || exit 1', "trap 'rm -rf -- \"$codex_home\"' EXIT", 'export CODEX_HOME="$codex_home"', - 'export PATH="$HOME/.local/bin:$PATH"', - invocation, + *statements, ] - ) + return "; ".join(statements) async def run_codex( @@ -301,7 +311,7 @@ async def run_codex( shell: str, prompt: str, executable: str = "codex", - inference: InferenceAccess | None = None, + inference: InferenceConnection | None = None, ) -> None: command = codex_command(config, shell, executable, inference=inference) logger.info("SSH exec codex CLI (%d chars)", len(command)) diff --git a/hud/agents/tests/test_claude_cli_agent.py b/hud/agents/tests/test_claude_cli_agent.py index af20348be..07a793d73 100644 --- a/hud/agents/tests/test_claude_cli_agent.py +++ b/hud/agents/tests/test_claude_cli_agent.py @@ -30,7 +30,7 @@ from hud.agents.types import AgentStep, ClaudeCLIConfig, ToolStep from hud.capabilities import Capability, SSHClient from hud.capabilities.rfb import WebPScreenshotEncoding -from hud.eval import InferenceAccess +from hud.eval import InferenceConnection from hud.settings import settings from hud.telemetry.context import set_trace_context from hud.types import MCPToolResult @@ -69,10 +69,10 @@ def test_command_follows_explicit_gateway_routing(monkeypatch: pytest.MonkeyPatc assert "ANTHROPIC_MODEL=claude-sonnet-5" in provider -def test_command_prefers_rollout_inference_access() -> None: - inference = InferenceAccess( +def test_command_prefers_rollout_inference_connection() -> None: + inference = InferenceConnection( base_url="https://inference.hud.so", - api_key="scoped-runtime-token", + credential="scoped-runtime-token", ) gateway = ClaudeCLIAgent(ClaudeCLIConfig(use_hud_gateway=True))._build_cli_command( @@ -84,6 +84,15 @@ def test_command_prefers_rollout_inference_access() -> None: assert "ANTHROPIC_API_KEY=scoped-runtime-token" in gateway assert "HUD_API_KEY" not in gateway assert "Trace-Id" not in gateway + for name in ( + "ANTHROPIC_MODEL", + "ANTHROPIC_SMALL_FAST_MODEL", + "ANTHROPIC_DEFAULT_SONNET_MODEL", + "ANTHROPIC_DEFAULT_OPUS_MODEL", + "ANTHROPIC_DEFAULT_HAIKU_MODEL", + "CLAUDE_CODE_SUBAGENT_MODEL", + ): + assert f"{name}=claude-sonnet-5" in gateway def test_windows_command_encodes_environment_and_arguments( diff --git a/hud/agents/tests/test_codex_cli_agent.py b/hud/agents/tests/test_codex_cli_agent.py index 872b9f700..3caeb021d 100644 --- a/hud/agents/tests/test_codex_cli_agent.py +++ b/hud/agents/tests/test_codex_cli_agent.py @@ -18,7 +18,7 @@ from hud.agents.tests.cli_fakes import fake_run as _fake_run from hud.agents.types import AgentStep, CodexCLIConfig, ToolStep from hud.capabilities import Capability, SSHClient -from hud.eval import InferenceAccess +from hud.eval import InferenceConnection from hud.eval.runtime import RuntimeConfig, RuntimeResources from hud.settings import settings from hud.telemetry.context import set_trace_context @@ -105,19 +105,38 @@ def test_command_follows_explicit_gateway_routing(monkeypatch: pytest.MonkeyPatc assert command.endswith(" -") -def test_command_prefers_rollout_inference_access() -> None: - inference = InferenceAccess( +def test_command_prefers_rollout_inference_connection() -> None: + inference = InferenceConnection( base_url="https://inference.hud.so", - api_key="scoped-runtime-token", + credential="scoped-runtime-token", ) command = codex_command(CodexCLIConfig(use_hud_gateway=True), "bash", inference=inference) - assert "HUD_API_KEY=scoped-runtime-token" in command + assert "HUD_RUNTIME_INFERENCE_TOKEN=scoped-runtime-token" in command + assert 'model_providers.hud.env_key="HUD_RUNTIME_INFERENCE_TOKEN"' in command + assert "HUD_API_KEY" not in command assert 'model_providers.hud.base_url="https://inference.hud.so"' in command assert "Trace-Id" not in command +@pytest.mark.parametrize("shell", ["bash", "powershell"]) +def test_command_preserves_ambient_codex_login_without_explicit_credentials(shell: str) -> None: + command = codex_command(CodexCLIConfig(use_hud_gateway=False), shell) + script = ( + base64.b64decode(command.rsplit(" ", 1)[1]).decode("utf-16-le") + if shell == "powershell" + else command + ) + + assert "CODEX_HOME" not in script + assert "CODEX_API_KEY" not in script + assert "HUD_API_KEY" not in script + assert "HUD_RUNTIME_INFERENCE_TOKEN" not in script + assert "mktemp" not in script + assert "codex exec" in script or "& 'codex' 'exec'" in script + + def test_windows_command_encodes_environment_and_arguments( monkeypatch: pytest.MonkeyPatch, ) -> None: diff --git a/hud/agents/types.py b/hud/agents/types.py index 4dbf3e3bd..be1956597 100644 --- a/hud/agents/types.py +++ b/hud/agents/types.py @@ -182,7 +182,11 @@ class ClaudeCLIConfig(AgentConfig): class CodexCLIConfig(AgentConfig): - """Configuration for CodexCLIAgent (runs ``codex exec`` over SSH).""" + """Configuration for CodexCLIAgent (runs ``codex exec`` over SSH). + + Without an explicit inference connection or API key, the agent leaves + ``CODEX_HOME`` unchanged so a login in that execution environment can apply. + """ model_name: str = "Codex CLI" model: str = Field(default="gpt-5.6-sol", validation_alias=_model_alias) diff --git a/hud/clients/tests/test_connect.py b/hud/clients/tests/test_connect.py index c5d62ea47..555e78f1c 100644 --- a/hud/clients/tests/test_connect.py +++ b/hud/clients/tests/test_connect.py @@ -26,6 +26,19 @@ HELLO_RESULT = {"session_id": "s-1", "env": {"name": "stub", "version": "1.0"}, "bindings": []} +def test_workspace_route_from_url_extracts_transport_address() -> None: + assert WorkspaceRoute.from_url("ssh", "https://inference.hud.so/v1") == WorkspaceRoute( + "ssh", + "inference.hud.so", + 443, + ) + assert WorkspaceRoute.from_url("shell", "http://gateway.test:8080") == WorkspaceRoute( + "shell", + "gateway.test", + 8080, + ) + + async def test_connect_sends_workspace_routes_in_hello() -> None: requests: list[dict[str, object]] = [] diff --git a/hud/environment/egress.py b/hud/environment/egress.py index 3b81ac3db..44cd93488 100644 --- a/hud/environment/egress.py +++ b/hud/environment/egress.py @@ -210,6 +210,20 @@ def __post_init__(self) -> None: def to_wire(self) -> dict[str, str | int]: return {"capability": self.capability, "host": self.host, "port": self.port} + @classmethod + def from_url(cls, capability: str, url: str) -> WorkspaceRoute: + """Build a host route for one HTTP(S) endpoint.""" + parts = urllib.parse.urlsplit(url) + if parts.scheme not in {"http", "https"} or parts.hostname is None: + raise ValueError("workspace route URL must be HTTP(S) with a hostname") + if parts.username is not None or parts.password is not None: + raise ValueError("workspace route URL must not contain credentials") + return cls( + capability=capability, + host=parts.hostname, + port=parts.port or (443 if parts.scheme == "https" else 80), + ) + @classmethod def from_wire(cls, value: object) -> WorkspaceRoute: if not isinstance(value, dict): diff --git a/hud/eval/__init__.py b/hud/eval/__init__.py index 7f67c6e94..7411f9105 100644 --- a/hud/eval/__init__.py +++ b/hud/eval/__init__.py @@ -32,7 +32,7 @@ from .chat import Chat from .job import Job -from .run import Grade, InferenceAccess, Run, rollout +from .run import Grade, InferenceConnection, Run, rollout from .runtime import ( ComposeProject, DaytonaRuntime, @@ -63,7 +63,7 @@ "Grade", "HUDRuntime", "HostedRuntime", - "InferenceAccess", + "InferenceConnection", "Job", "LocalRuntime", "ModalRuntime", diff --git a/hud/eval/run.py b/hud/eval/run.py index d8ab985c9..1be0b8b1e 100644 --- a/hud/eval/run.py +++ b/hud/eval/run.py @@ -32,7 +32,6 @@ import mcp.types as mcp_types from hud.clients import HudProtocolError, connect -from hud.environment import WorkspaceRoute from hud.graders.results import SubScore from hud.telemetry.context import set_trace_context from hud.types import Step, TaskCall, Trace @@ -42,10 +41,12 @@ from .job import job_enter, trace_enter, trace_exit if TYPE_CHECKING: + from collections.abc import Sequence from types import TracebackType from hud.agents.base import Agent from hud.clients.client import HudClient + from hud.environment import WorkspaceRoute from .runtime import Provider from .runtime.core import RuntimeConfig @@ -55,12 +56,11 @@ @dataclass(frozen=True, slots=True) -class InferenceAccess: - """Scoped inference access for an agent running inside a workspace.""" +class InferenceConnection: + """Execution-scoped inference connection exposed to a live agent.""" base_url: str - api_key: str = field(repr=False) - workspace: str = "ssh" + credential: str = field(repr=False) def __post_init__(self) -> None: parts = urlsplit(self.base_url) @@ -70,17 +70,8 @@ def __post_init__(self) -> None: raise ValueError("inference base_url must not contain credentials") if parts.query or parts.fragment: raise ValueError("inference base_url must not contain a query or fragment") - if not self.api_key: - raise ValueError("inference api_key must not be empty") - - def workspace_route(self) -> WorkspaceRoute: - parts = urlsplit(self.base_url) - assert parts.hostname is not None - return WorkspaceRoute( - capability=self.workspace, - host=parts.hostname, - port=parts.port or (443 if parts.scheme == "https" else 80), - ) + if not self.credential: + raise ValueError("inference credential must not be empty") def validate_rollout_timeouts( @@ -231,7 +222,7 @@ def __init__( *, best_effort_grade: bool = False, runtime_config: RuntimeConfig | None = None, - inference: InferenceAccess | None = None, + inference: InferenceConnection | None = None, ) -> None: self._client = client self._task_id = task_id @@ -478,7 +469,8 @@ async def rollout( group_id: str | None = None, trace_id: str | None = None, rollout_timeout: float | None = None, - inference: InferenceAccess | None = None, + inference: InferenceConnection | None = None, + workspace_routes: Sequence[WorkspaceRoute] = (), ) -> Run: """Drive one task to a graded :class:`Run` here, against ``runtime``'s channel. @@ -569,7 +561,6 @@ async def close_actor() -> None: scope.push_async_callback(close_actor) addr = await actor.enter_async_context(runtime(task)) _phase = "starting task" - workspace_routes = (inference.workspace_route(),) if inference is not None else () async with connect(addr, workspace_routes=workspace_routes) as actor_client: client = actor_client live = Run( @@ -585,29 +576,32 @@ async def close_actor() -> None: run = live # bound only once live: an earlier failure synthesizes _phase = "agent loop" try: - async with file_tracking_observer(actor_client): - if agent_timeout is None: - await agent(run) - else: - deadline = asyncio.timeout(agent_timeout) - try: - async with deadline: - await agent(run) - except TimeoutError: - if not deadline.expired(): - raise - detail = f"agent timed out after {agent_timeout:g}s" - logger.warning(detail) - run.trace.status = "error" - run.trace.stop_reason = "timeout" - run.record(Step(source="system", error=detail)) - except Exception as exc: - if task.verifier is None: - raise - detail = "".join(traceback.format_exception_only(exc)).strip() - logger.warning("rollout failed mid-run (%s): %s", _phase, detail) - run.trace.status = "error" - run.record(Step(source="system", error=f"[{_phase}] {detail}")) + try: + async with file_tracking_observer(actor_client): + if agent_timeout is None: + await agent(run) + else: + deadline = asyncio.timeout(agent_timeout) + try: + async with deadline: + await agent(run) + except TimeoutError: + if not deadline.expired(): + raise + detail = f"agent timed out after {agent_timeout:g}s" + logger.warning(detail) + run.trace.status = "error" + run.trace.stop_reason = "timeout" + run.record(Step(source="system", error=detail)) + except Exception as exc: + if task.verifier is None: + raise + detail = "".join(traceback.format_exception_only(exc)).strip() + logger.warning("rollout failed mid-run (%s): %s", _phase, detail) + run.trace.status = "error" + run.record(Step(source="system", error=f"[{_phase}] {detail}")) + finally: + run.inference = None _phase = "grading" if verifier is not None: @@ -707,6 +701,7 @@ async def close_actor() -> None: run.trace.status = "error" run.record(Step(source="system", error=f"[{_phase}] {detail}")) assert run is not None # the body bound it, or the handler synthesized it + run.inference = None run.trace.trace_id = trace_id run.job_id = job_id run.group_id = group_id @@ -720,4 +715,4 @@ def _consume_task_result(task: asyncio.Future[Any]) -> None: task.result() -__all__ = ["Grade", "InferenceAccess", "Run", "rollout"] +__all__ = ["Grade", "InferenceConnection", "Run", "rollout"] diff --git a/hud/eval/tests/test_rollout.py b/hud/eval/tests/test_rollout.py index 8210e2078..0b031b6ec 100644 --- a/hud/eval/tests/test_rollout.py +++ b/hud/eval/tests/test_rollout.py @@ -32,7 +32,15 @@ from hud.agents.openai_compatible import OpenAIChatAgent from hud.agents.types import OpenAIChatConfig from hud.environment import Answer, Environment -from hud.eval import Job, LocalRuntime, Runtime, SubprocessRuntime, Task, Taskset +from hud.eval import ( + InferenceConnection, + Job, + LocalRuntime, + Runtime, + SubprocessRuntime, + Task, + Taskset, +) from hud.eval.run import Run, rollout if TYPE_CHECKING: @@ -166,6 +174,29 @@ async def test_rollout_returns_graded_run_with_trace_id(env_file: Path) -> None: assert run.runtime.startswith("tcp://127.0.0.1:") +async def test_inference_connection_exists_only_during_agent_execution(env_file: Path) -> None: + connection = InferenceConnection( + base_url="https://inference.hud.so", + credential="scoped-runtime-token", + ) + observed: list[InferenceConnection | None] = [] + + class InspectingAgent(Agent): + async def __call__(self, run: Run) -> None: + observed.append(run.inference) + run.trace.content = _solve_add(run.prompt_text) + + run = await rollout( + _add_task(2, 3), + InspectingAgent(), + runtime=SubprocessRuntime(env_file), + inference=connection, + ) + + assert observed == [connection] + assert run.inference is None + + async def test_verifier_task_replaces_the_actor_grade_in_the_same_runtime() -> None: env = Environment("reviewed") completed: list[str] = [] diff --git a/hud/tests/test_init_module.py b/hud/tests/test_init_module.py index 62bd04723..0281d756d 100644 --- a/hud/tests/test_init_module.py +++ b/hud/tests/test_init_module.py @@ -26,6 +26,7 @@ def test_all_exports(self): "Job", "HUDRuntime", "HostedRuntime", + "InferenceConnection", "Run", "Runtime", "RuntimeConfig", From f01e5bcf7d66bec641386f3aa4637befa495dec9 Mon Sep 17 00:00:00 2001 From: Jaideep <67646710+jdchawla29@users.noreply.github.com> Date: Thu, 27 Aug 2026 13:30:11 -0700 Subject: [PATCH 7/7] style(agents): format Codex cleanup command --- hud/agents/codex/agent.py | 3 +-- 1 file changed, 1 insertion(+), 2 deletions(-) diff --git a/hud/agents/codex/agent.py b/hud/agents/codex/agent.py index b9a4019f6..345443c1f 100644 --- a/hud/agents/codex/agent.py +++ b/hud/agents/codex/agent.py @@ -283,8 +283,7 @@ def codex_command( "New-Item -ItemType Directory -Force -Path $codexHome | Out-Null", "$env:CODEX_HOME=$codexHome", *(f"$env:{key}={powershell_quote(value)}" for key, value in env.items()), - f"try {{ {invocation} }} finally {{ Remove-Item -Recurse -Force " - "$codexHome }", + f"try {{ {invocation} }} finally {{ Remove-Item -Recurse -Force $codexHome }}", "exit $hudExitCode", ] return powershell(";".join(statements))