diff --git a/libs/python/cua-sandbox/README.md b/libs/python/cua-sandbox/README.md index 5ee5e6d824..899a1c8741 100644 --- a/libs/python/cua-sandbox/README.md +++ b/libs/python/cua-sandbox/README.md @@ -77,3 +77,27 @@ async with Localhost.connect() as host: await host.shell.run("echo hello") await host.screenshot() ``` + + +## Cloud sandbox + +Fleet is the OAuth cloud backend. Configure OAuth credentials once; Fleet uses `https://run.cua.ai` by default and can be overridden with `configure(fleet_base_url=...)` or `CUA_FLEET_BASE_URL`. The legacy API-key VM API continues to use `https://api.cua.ai`. Cloud images must use a registry reference; `expose()` declares additional Fleet services. + +Fleet does not support snapshots or custom disks, and currently supports only `us-east-1`. `await sb.tunnel.forward(3000)` returns the authenticated Fleet service URL for an exposed port; it does not open a local SSH tunnel. + +```python +import os + +import cua_sandbox as cua +from cua_sandbox import Image, Sandbox + +cua.configure( + client_id=os.environ["CUA_CLIENT_ID"], + client_secret=os.environ["CUA_CLIENT_SECRET"], +) + +async with Sandbox.ephemeral( + Image.from_registry("registry.example/desktop-workspace@sha256:...").expose(3000) +) as sb: + await sb.shell.run("uname -a") +``` diff --git a/libs/python/cua-sandbox/cua_sandbox/_config.py b/libs/python/cua-sandbox/cua_sandbox/_config.py index 0f9872a559..14baa298b0 100644 --- a/libs/python/cua-sandbox/cua_sandbox/_config.py +++ b/libs/python/cua-sandbox/cua_sandbox/_config.py @@ -1,7 +1,4 @@ -"""Global configuration — configure(api_key, base_url). - -Auth priority: per-call > configure() > ~/.cua/credentials > CUA_API_KEY env var. -""" +"""Global configuration for cloud sandbox provisioning.""" from __future__ import annotations @@ -9,11 +6,19 @@ from dataclasses import dataclass from typing import Optional +_DEFAULT_BASE_URL = "https://api.cua.ai" +_DEFAULT_FLEET_BASE_URL = "https://run.cua.ai" +_DEFAULT_TOKEN_URL = "https://auth.cua.ai/realms/cyclops-cs/protocol/openid-connect/token" + @dataclass class _Config: api_key: Optional[str] = None - base_url: str = "https://api.cua.ai" + base_url: str = _DEFAULT_BASE_URL + fleet_base_url: str = _DEFAULT_FLEET_BASE_URL + token_url: str = _DEFAULT_TOKEN_URL + client_id: Optional[str] = None + client_secret: Optional[str] = None _global_config = _Config() @@ -23,29 +28,39 @@ def configure( *, api_key: Optional[str] = None, base_url: Optional[str] = None, + fleet_base_url: Optional[str] = None, + token_url: Optional[str] = None, + client_id: Optional[str] = None, + client_secret: Optional[str] = None, ) -> None: - """Set global configuration for the CUA SDK. + """Set global configuration for cloud sandboxes. - Args: - api_key: API key for cloud sandboxes. - base_url: Base URL for the CUA cloud API. + API-key cloud operations use ``base_url``. Fleet uses ``fleet_base_url`` and + OAuth client credentials. """ if api_key is not None: _global_config.api_key = api_key if base_url is not None: _global_config.base_url = base_url + if fleet_base_url is not None: + _global_config.fleet_base_url = fleet_base_url + if token_url is not None: + _global_config.token_url = token_url + if client_id is not None: + _global_config.client_id = client_id + if client_secret is not None: + _global_config.client_secret = client_secret def get_api_key(override: Optional[str] = None) -> Optional[str]: - """Resolve API key with priority: override > configure() > credentials file > env.""" + """Resolve a legacy API key with per-call configuration taking priority.""" if override: return override if _global_config.api_key: return _global_config.api_key - # Try credentials file - cred = _read_credentials_key() - if cred: - return cred + credential = _read_credentials_key() + if credential: + return credential return os.environ.get("CUA_API_KEY") @@ -53,12 +68,29 @@ def get_base_url() -> str: return os.environ.get("CUA_BASE_URL") or _global_config.base_url +def get_fleet_base_url() -> str: + """Return the Fleet API endpoint without changing legacy VM API routing.""" + return os.environ.get("CUA_FLEET_BASE_URL") or _global_config.fleet_base_url + + +def get_token_url() -> str: + return os.environ.get("CUA_TOKEN_URL") or _global_config.token_url + + +def get_client_id(override: Optional[str] = None) -> Optional[str]: + return override or _global_config.client_id or os.environ.get("CUA_CLIENT_ID") + + +def get_client_secret(override: Optional[str] = None) -> Optional[str]: + return override or _global_config.client_secret or os.environ.get("CUA_CLIENT_SECRET") + + def _read_credentials_key() -> Optional[str]: - """Read API key from ~/.cua/credentials if it exists.""" - cred_path = os.path.join(os.path.expanduser("~"), ".cua", "credentials") + """Read a legacy API key from ``~/.cua/credentials`` when present.""" + credential_path = os.path.join(os.path.expanduser("~"), ".cua", "credentials") try: - with open(cred_path) as f: - for line in f: + with open(credential_path) as credential_file: + for line in credential_file: line = line.strip() if line.startswith("api_key="): return line[len("api_key=") :] diff --git a/libs/python/cua-sandbox/cua_sandbox/interfaces/tunnel.py b/libs/python/cua-sandbox/cua_sandbox/interfaces/tunnel.py index 7d7f27a136..e3f66b8b5d 100644 --- a/libs/python/cua-sandbox/cua_sandbox/interfaces/tunnel.py +++ b/libs/python/cua-sandbox/cua_sandbox/interfaces/tunnel.py @@ -28,15 +28,18 @@ class TunnelInfo: """A single forwarded port.""" - def __init__(self, host: str, port: int, sandbox_port: int) -> None: + def __init__( + self, host: str, port: int, sandbox_port: int, *, url: Optional[str] = None + ) -> None: self.host = host self.port = port # host-side port self.sandbox_port = sandbox_port # original port inside sandbox + self._url = url self._closer: Optional[object] = None # Callable[[TunnelInfo], Coroutine] @property def url(self) -> str: - return f"http://{self.host}:{self.port}" + return self._url or f"http://{self.host}:{self.port}" async def close(self) -> None: """Close this tunnel (no-op if already closed or inside a context manager).""" diff --git a/libs/python/cua-sandbox/cua_sandbox/sandbox.py b/libs/python/cua-sandbox/cua_sandbox/sandbox.py index 76043b17fa..5335b18fb9 100644 --- a/libs/python/cua-sandbox/cua_sandbox/sandbox.py +++ b/libs/python/cua-sandbox/cua_sandbox/sandbox.py @@ -56,6 +56,7 @@ def record_event(event_name: str, properties: dict | None = None) -> None: pass +from cua_sandbox._config import get_client_id, get_client_secret from cua_sandbox.image import Image from cua_sandbox.interfaces import ( Apps, @@ -72,6 +73,7 @@ def record_event(event_name: str, properties: dict | None = None) -> None: ) from cua_sandbox.transport.base import Transport from cua_sandbox.transport.cloud import CloudTransport +from cua_sandbox.transport.fleet_cloud import FleetCloudTransport from cua_sandbox.transport.http import HTTPTransport from cua_sandbox.transport.websocket import WebSocketTransport @@ -272,7 +274,7 @@ def __init__( async def _connect(self) -> None: await self._transport.connect() # Update name from transport (e.g. CloudTransport resolves name after creating a VM) - if self.name is None and isinstance(self._transport, CloudTransport): + if self.name is None and isinstance(self._transport, (CloudTransport, FleetCloudTransport)): self.name = self._transport.name async def disconnect(self) -> None: @@ -294,7 +296,7 @@ async def snapshot(self, name: str | None = None, stateful: bool = False) -> "Im """ from cua_sandbox.transport.cloud import CloudTransport - if not isinstance(self._transport, CloudTransport): + if not isinstance(self._transport, (CloudTransport, FleetCloudTransport)): raise NotImplementedError("Snapshots are only supported for cloud sandboxes") image_desc = await self._transport.create_snapshot(name=name, stateful=stateful) @@ -333,7 +335,7 @@ async def destroy(self) -> None: await self._transport.disconnect() except Exception: logger.warning("Failed to disconnect transport for sandbox %r", self.name) - if isinstance(self._transport, CloudTransport): + if isinstance(self._transport, (CloudTransport, FleetCloudTransport)): try: await self._transport.delete_vm() except Exception: @@ -420,15 +422,16 @@ async def create( Args: image: Image to run (e.g. ``Image.desktop("ubuntu")``). name: Optional name to assign to the sandbox. - api_key: CUA API key for cloud sandboxes. + api_key: Legacy CUA API key. Providing one uses the legacy VM API; + Fleet cloud sandboxes use OAuth client credentials instead. local: Use a local runtime instead of cloud. runtime: Explicit runtime backend (DockerRuntime, QEMURuntime, etc.). cpu: Number of CPUs for the cloud sandbox. memory_mb: Memory in MB for the cloud sandbox. - disk_gb: Disk size in GB for the cloud sandbox. - region: Cloud region (default ``"us-east-1"``). - time_to_start: Max seconds to wait for the VM to become reachable - (default 600). Only applies to cloud sandboxes. + disk_gb: Unsupported by Fleet; API-key legacy cloud only. + region: Fleet currently supports only ``"us-east-1"``. + time_to_start: Max seconds to wait for Fleet provisioning and + service readiness (default 600). request_timeout: Default HTTP request timeout in seconds for commands sent to the computer-server (default 30, cloud only). Individual commands with a server-side timeout automatically @@ -540,15 +543,16 @@ async def ephemeral( Args: image: Image to run (e.g. ``Image.desktop("ubuntu")``). name: Optional name to assign to the sandbox. - api_key: CUA API key for cloud sandboxes. + api_key: Legacy CUA API key. Providing one uses the legacy VM API; + Fleet cloud sandboxes use OAuth client credentials instead. local: Use a local runtime instead of cloud. runtime: Explicit runtime backend (DockerRuntime, QEMURuntime, etc.). cpu: Number of CPUs for the cloud sandbox. memory_mb: Memory in MB for the cloud sandbox. - disk_gb: Disk size in GB for the cloud sandbox. - region: Cloud region (default ``"us-east-1"``). - time_to_start: Max seconds to wait for the VM to become reachable - (default 600). Only applies to cloud sandboxes. + disk_gb: Unsupported by Fleet; API-key legacy cloud only. + region: Fleet currently supports only ``"us-east-1"``. + time_to_start: Max seconds to wait for Fleet provisioning and + service readiness (default 600). request_timeout: Default HTTP request timeout in seconds for commands sent to the computer-server (default 30, cloud only). Individual commands with a server-side timeout automatically @@ -690,24 +694,45 @@ async def _list_android(): ) return results + @staticmethod + def _uses_fleet(api_key: Optional[str]) -> bool: + """Choose Fleet only for OAuth-configured calls without an explicit API key.""" + return api_key is None and bool(get_client_id() and get_client_secret()) + @classmethod async def _list_cloud(cls, *, api_key: Optional[str] = None) -> "list[SandboxInfo]": - from cua_sandbox.transport.cloud import cloud_list_vms + if not cls._uses_fleet(api_key): + from cua_sandbox.transport.cloud import cloud_list_vms - vms = await cloud_list_vms(api_key=api_key) - results = [] - for vm in vms: - raw_status = vm.get("status", "unknown") - results.append( + vms = await cloud_list_vms(api_key=api_key) + return [ SandboxInfo( name=vm.get("name", ""), - status=raw_status, + status=vm.get("status", "unknown"), source="cloud", os_type=vm.get("os_type") or vm.get("os"), created_at=vm.get("created_at"), ) - ) - return results + for vm in vms + ] + + pools = await FleetCloudTransport.list_sandboxes() + return [cls._fleet_sandbox_info(pool) for pool in pools] + + @staticmethod + def _fleet_sandbox_info(pool: dict[str, Any]) -> SandboxInfo: + metadata = pool.get("metadata") or {} + spec = pool.get("spec") or {} + status = pool.get("status") or {} + replicas = spec.get("replicas", 1) + available = status.get("availableCount", 0) + state = "suspended" if replicas == 0 else "running" if available else "provisioning" + return SandboxInfo( + name=metadata.get("name", ""), + status=state, + source="fleet", + created_at=metadata.get("creationTimestamp"), + ) @classmethod async def get_info( @@ -747,16 +772,18 @@ async def get_info( ), ) raise ValueError(f"Local sandbox '{name}' not found.") - from cua_sandbox.transport.cloud import cloud_get_vm - - vm = await cloud_get_vm(name, api_key=api_key) - return SandboxInfo( - name=vm.get("name", name), - status=vm.get("status", "unknown"), - source="cloud", - os_type=vm.get("os_type") or vm.get("os"), - created_at=vm.get("created_at"), - ) + if not cls._uses_fleet(api_key): + from cua_sandbox.transport.cloud import cloud_get_vm + + vm = await cloud_get_vm(name, api_key=api_key) + return SandboxInfo( + name=vm.get("name", name), + status=vm.get("status", "unknown"), + source="cloud", + os_type=vm.get("os_type") or vm.get("os"), + created_at=vm.get("created_at"), + ) + return cls._fleet_sandbox_info(await FleetCloudTransport.get_sandbox_info(name)) @classmethod async def suspend( @@ -781,9 +808,12 @@ async def suspend( if local: await cls._suspend_local(name) return - from cua_sandbox.transport.cloud import cloud_vm_action + if not cls._uses_fleet(api_key): + from cua_sandbox.transport.cloud import cloud_vm_action - await cloud_vm_action(name, "stop", api_key=api_key) + await cloud_vm_action(name, "stop", api_key=api_key) + return + await FleetCloudTransport.suspend_sandbox(name) @classmethod async def _suspend_local(cls, name: str) -> None: @@ -832,10 +862,13 @@ async def resume( """ if local: return await cls._resume_local(name) - from cua_sandbox.transport.cloud import cloud_vm_action + if not cls._uses_fleet(api_key): + from cua_sandbox.transport.cloud import cloud_vm_action - await cloud_vm_action(name, "run", api_key=api_key) - # Connect to the now-running cloud sandbox + await cloud_vm_action(name, "run", api_key=api_key) + else: + await FleetCloudTransport.resume_sandbox(name) + # Connect to the now-running cloud sandbox. sb = await cls._create(name=name, ephemeral=False, api_key=api_key) return sb @@ -849,14 +882,12 @@ async def _resume_local(cls, name: str) -> "Sandbox": raise ValueError(f"No local sandbox named '{name}' found in state files.") runtime_type = state.get("runtime_type") if runtime_type == "lume": - from cua_sandbox.image import Image from cua_sandbox.runtime.lume import LumeRuntime image = Image.from_dict(state["image"]) rt = LumeRuntime() rt_info = await rt.resume(image, name) elif runtime_type == "qemu-baremetal": - from cua_sandbox.image import Image from cua_sandbox.runtime.qemu import QEMUBaremetalRuntime image = Image.from_dict(state["image"]) @@ -910,9 +941,12 @@ async def restart( if local: await cls._suspend_local(name) return await cls._resume_local(name) - from cua_sandbox.transport.cloud import cloud_vm_action + if not cls._uses_fleet(api_key): + from cua_sandbox.transport.cloud import cloud_vm_action - await cloud_vm_action(name, "restart", api_key=api_key) + await cloud_vm_action(name, "restart", api_key=api_key) + else: + await FleetCloudTransport.restart_sandbox(name) sb = await cls._create(name=name, ephemeral=False, api_key=api_key) return sb @@ -937,9 +971,12 @@ async def delete( if local: await cls._delete_local(name) return - from cua_sandbox.transport.cloud import cloud_vm_action + if not cls._uses_fleet(api_key): + from cua_sandbox.transport.cloud import cloud_vm_action - await cloud_vm_action(name, "delete", api_key=api_key) + await cloud_vm_action(name, "delete", api_key=api_key) + return + await FleetCloudTransport.delete_sandbox(name) @classmethod async def _delete_local(cls, name: str) -> None: @@ -996,7 +1033,7 @@ async def _create( ephemeral = bool(image) rt_info = None - if image and image.kind is None and image._registry: + if image and image.kind is None and image._registry and local: from cua_sandbox.registry.resolve import resolve_image_kind image = resolve_image_kind(image) @@ -1044,11 +1081,10 @@ async def _create( runtime = _auto_runtime(image) if image and not runtime and not local: # image without runtime and not local → cloud creation - if not any([ws_url, http_url]): - transport = CloudTransport( - name=name, - api_key=api_key, + if not any([ws_url, http_url]) and not api_key: + transport = FleetCloudTransport( image=image, + name=name or _random_name(), cpu=cpu, memory_mb=memory_mb, disk_gb=disk_gb, @@ -1080,6 +1116,23 @@ async def _create( sb, image=image, local=False, ephemeral=bool(ephemeral), t_start=_t_start ) return sb + if api_key and not any([ws_url, http_url]): + transport = _make_transport( + api_key=api_key, + name=name, + cpu=cpu, + memory_mb=memory_mb, + disk_gb=disk_gb, + region=region, + ) + sb = cls( + transport, name=name, _ephemeral=ephemeral, _telemetry_enabled=telemetry_enabled + ) + await sb._connect() + _record_sandbox_create( + sb, image=image, local=False, ephemeral=bool(ephemeral), t_start=_t_start + ) + return sb runtime = _auto_runtime(image) if image and runtime: sb_name = name or _random_name() @@ -1153,17 +1206,27 @@ async def _create( container_name=container_name, ) else: - transport = _make_transport( - ws_url=ws_url, - http_url=http_url, - api_key=api_key, - container_name=container_name, - name=name, - cpu=cpu, - memory_mb=memory_mb, - disk_gb=disk_gb, - region=region, - ) + if name and cls._uses_fleet(api_key) and not ws_url and not http_url: + transport = FleetCloudTransport( + image=None, + name=name, + cpu=cpu, + memory_mb=memory_mb, + disk_gb=disk_gb, + region=region, + ) + else: + transport = _make_transport( + ws_url=ws_url, + http_url=http_url, + api_key=api_key, + container_name=container_name, + name=name, + cpu=cpu, + memory_mb=memory_mb, + disk_gb=disk_gb, + region=region, + ) # Write persistent state for local (non-ephemeral) sandboxes if not ephemeral and rt_info and local: from cua_sandbox import sandbox_state diff --git a/libs/python/cua-sandbox/cua_sandbox/transport/__init__.py b/libs/python/cua-sandbox/cua_sandbox/transport/__init__.py index 8339174d02..6dc01eaa15 100644 --- a/libs/python/cua-sandbox/cua_sandbox/transport/__init__.py +++ b/libs/python/cua-sandbox/cua_sandbox/transport/__init__.py @@ -1,6 +1,7 @@ from cua_sandbox.transport.adb import ADBTransport from cua_sandbox.transport.base import Transport from cua_sandbox.transport.cloud import CloudTransport +from cua_sandbox.transport.fleet import FleetTransport from cua_sandbox.transport.http import HTTPTransport from cua_sandbox.transport.local import LocalTransport from cua_sandbox.transport.osworld import OSWorldTransport @@ -16,6 +17,7 @@ "WebSocketTransport", "HTTPTransport", "CloudTransport", + "FleetTransport", "QMPTransport", "OSWorldTransport", "ADBTransport", diff --git a/libs/python/cua-sandbox/cua_sandbox/transport/computer_server.py b/libs/python/cua-sandbox/cua_sandbox/transport/computer_server.py new file mode 100644 index 0000000000..2571cc4b2c --- /dev/null +++ b/libs/python/cua-sandbox/cua_sandbox/transport/computer_server.py @@ -0,0 +1,53 @@ +"""Shared response handling for computer-server HTTP transports.""" + +from __future__ import annotations + +import base64 +import json +from typing import Any, Dict + + +def parse_command_response(text: str) -> Dict[str, Any]: + """Extract and validate the first JSON data frame from an SSE response.""" + for line in text.splitlines(): + if not line.startswith("data: "): + continue + payload = json.loads(line[6:]) + if isinstance(payload, dict) and not payload.get("success", True): + if "error" in payload: + raise RuntimeError(f"Remote error: {payload['error']}") + parts = [] + return_code = payload.get("return_code") + if return_code is not None: + parts.append(f"return_code={return_code}") + stderr = (payload.get("stderr") or "").strip() + if stderr: + parts.append(f"stderr={stderr!r}") + stdout = (payload.get("stdout") or "").strip() + if stdout and not stderr: + parts.append(f"stdout={stdout!r}") + detail = ", ".join(parts) or "no detail" + raise RuntimeError(f"Remote error: {detail}") + return payload + raise RuntimeError(f"No SSE data frame in response: {text[:200]}") + + +def decode_screenshot_response(payload: Dict[str, Any]) -> bytes: + """Decode computer-server's accepted screenshot response shapes.""" + encoded = payload.get("image_data", payload.get("base64_image", payload.get("result", ""))) + if isinstance(encoded, dict): + encoded = encoded.get("image_data", encoded.get("base64_image", encoded.get("base64", ""))) + return base64.b64decode(encoded) + + +def normalize_screen_size(payload: Dict[str, Any]) -> Dict[str, int]: + """Normalize nested computer-server screen-size response shapes.""" + data: Any = payload + if isinstance(data, dict): + data = data.get("size", data.get("result", data)) + if isinstance(data, dict): + width = data.get("width") or data.get("screen_width") or data.get("w") + height = data.get("height") or data.get("screen_height") or data.get("h") + if width is not None and height is not None: + return {"width": int(width), "height": int(height)} + raise KeyError(f"Cannot extract screen size from response: {payload}") diff --git a/libs/python/cua-sandbox/cua_sandbox/transport/cyclops_http_client.py b/libs/python/cua-sandbox/cua_sandbox/transport/cyclops_http_client.py new file mode 100644 index 0000000000..f3c6de6c6e --- /dev/null +++ b/libs/python/cua-sandbox/cua_sandbox/transport/cyclops_http_client.py @@ -0,0 +1,36 @@ +"""httpx bridge for the generated Cyclops SDK callback interface.""" + +from __future__ import annotations + +import httpx +from cyclops_sdk import HttpClient, HttpError, HttpHeader, HttpRequest, HttpResponse + + +class CyclopsHttpClient(HttpClient): + """Execute SDK requests with one shared async httpx client.""" + + def __init__(self, client: httpx.AsyncClient | None = None) -> None: + self._client = client or httpx.AsyncClient(timeout=30.0) + self._owns_client = client is None + + async def execute(self, request: HttpRequest) -> HttpResponse: + try: + response = await self._client.request( + request.method, + request.url, + headers={header.name: header.value for header in request.headers}, + content=request.body, + ) + except httpx.TransportError as error: + raise HttpError.Transport(str(error)) from error + return HttpResponse( + status=response.status_code, + headers=[ + HttpHeader(name=name, value=value) for name, value in response.headers.multi_items() + ], + body=response.content, + ) + + async def aclose(self) -> None: + if self._owns_client: + await self._client.aclose() diff --git a/libs/python/cua-sandbox/cua_sandbox/transport/fleet.py b/libs/python/cua-sandbox/cua_sandbox/transport/fleet.py new file mode 100644 index 0000000000..e6a955d7c7 --- /dev/null +++ b/libs/python/cua-sandbox/cua_sandbox/transport/fleet.py @@ -0,0 +1,137 @@ +"""Computer-server transport routed through Cyclops named services.""" + +from __future__ import annotations + +import asyncio +import json +from typing import Any, Dict, Optional + +import httpx +from cua_sandbox.transport.base import Transport +from cua_sandbox.transport.computer_server import ( + decode_screenshot_response, + normalize_screen_size, + parse_command_response, +) +from cyclops_sdk import HttpHeader, HttpRequest + +_CMD_MAX_RETRIES = 3 +_CMD_RETRY_BACKOFF_S = 0.5 + + +class FleetTransport(Transport): + """Route computer-server requests through ``CyclopsClient.service_request``.""" + + def __init__( + self, + *, + sdk: Any, + bound: Any, + service_name: str = "api", + timeout: float = 30.0, + **_: Any, + ) -> None: + self._sdk = sdk + self._bound = bound + self._service_name = service_name + self._timeout = timeout + self._connected = False + + async def connect(self) -> None: + if self._service_name not in self._bound.services: + raise ValueError(f"Fleet sandbox does not expose service {self._service_name!r}") + self._connected = True + + async def disconnect(self) -> None: + self._connected = False + + async def _request(self, method: str, path: str, *, json_body: Any = None) -> httpx.Response: + assert self._connected, "Transport not connected" + body = None if json_body is None else json.dumps(json_body).encode() + headers = ( + [] if body is None else [HttpHeader(name="content-type", value="application/json")] + ) + result = await self._sdk.service_request( + self._bound, + self._service_name, + path, + HttpRequest( + method=method, url=f"https://service.invalid{path}", headers=headers, body=body + ), + ) + request = httpx.Request(method, f"https://service.invalid{path}") + return httpx.Response( + result.status, + headers={header.name: header.value for header in result.headers}, + content=result.body, + request=request, + ) + + async def _cmd(self, command: str, params: Optional[Dict[str, Any]] = None) -> Dict[str, Any]: + body: Dict[str, Any] = {"command": command} + if params: + body["params"] = params + response = None + for attempt in range(_CMD_MAX_RETRIES): + response = await self._request("POST", "/cmd", json_body=body) + if response.status_code < 500 or attempt == _CMD_MAX_RETRIES - 1: + break + await asyncio.sleep(_CMD_RETRY_BACKOFF_S * (2**attempt)) + assert response is not None + response.raise_for_status() + return parse_command_response(response.text) + + async def send(self, action: str, **params: Any) -> Any: + result = await self._cmd(action, params if params else None) + return result.get("result", result) + + async def screenshot(self, format: str = "png", quality: int = 95) -> bytes: + params = None if format == "png" else {"format": format, "quality": quality} + return decode_screenshot_response(await self._cmd("screenshot", params)) + + async def get_screen_size(self) -> Dict[str, int]: + return normalize_screen_size(await self._cmd("get_screen_size")) + + async def get_environment(self) -> str: + try: + response = await self._request("GET", "/status") + response.raise_for_status() + payload = response.json() + return payload.get("os_type", payload.get("platform", "linux")) + except Exception: + return "linux" + + async def pty_create( + self, + command: Optional[str] = None, + cols: int = 120, + rows: int = 40, + cwd: Optional[str] = None, + envs: Optional[Dict[str, str]] = None, + ) -> Dict[str, Any]: + body: Dict[str, Any] = {"cols": cols, "rows": rows} + if command is not None: + body["command"] = command + if cwd is not None: + body["cwd"] = cwd + if envs is not None: + body["envs"] = envs + response = await self._request("POST", "/pty", json_body=body) + response.raise_for_status() + return response.json() + + async def pty_send(self, pid: int, data: str) -> None: + response = await self._request("POST", f"/pty/{pid}/stdin", json_body={"data": data}) + response.raise_for_status() + + async def pty_kill(self, pid: int) -> bool: + response = await self._request("DELETE", f"/pty/{pid}") + response.raise_for_status() + return bool(response.json().get("killed", True)) + + async def pty_info(self, pid: int) -> Optional[Dict[str, Any]]: + response = await self._request("GET", f"/pty/{pid}") + if response.status_code == 404: + return None + response.raise_for_status() + return response.json() diff --git a/libs/python/cua-sandbox/cua_sandbox/transport/fleet_cloud.py b/libs/python/cua-sandbox/cua_sandbox/transport/fleet_cloud.py new file mode 100644 index 0000000000..633d85cba3 --- /dev/null +++ b/libs/python/cua-sandbox/cua_sandbox/transport/fleet_cloud.py @@ -0,0 +1,387 @@ +"""Fleet-backed implementation of the public cloud sandbox transport.""" + +from __future__ import annotations + +import asyncio +import json +import logging +from typing import TYPE_CHECKING, Any, Optional +from urllib.parse import urlparse + +from cua_sandbox._config import ( + get_client_id, + get_client_secret, + get_fleet_base_url, + get_token_url, +) +from cua_sandbox.image import Image +from cua_sandbox.transport.cyclops_http_client import CyclopsHttpClient +from cua_sandbox.transport.fleet import FleetTransport +from cyclops_sdk import ( + ClaimSpec, + CreateClaimRequest, + CreatePoolRequest, + CyclopsClient, + CyclopsConfiguration, + CyclopsCredentials, + HttpRequest, + PoolSpec, + PoolTemplate, + PreservedJson, + SandboxService, + SandboxTemplateRef, + ServiceProtocol, +) + +if TYPE_CHECKING: + from cua_sandbox.interfaces.tunnel import TunnelInfo + +logger = logging.getLogger(__name__) + + +class _FleetClient: + """Thin async facade over the generated Cyclops SDK.""" + + def __init__(self) -> None: + client_id = get_client_id() + client_secret = get_client_secret() + if not client_id or not client_secret: + raise ValueError( + "Fleet cloud sandboxes require CUA_CLIENT_ID and CUA_CLIENT_SECRET, " + "or cua.configure(client_id=..., client_secret=...)." + ) + self._base_url = get_fleet_base_url().rstrip("/") + self._http_client = CyclopsHttpClient() + configuration = CyclopsConfiguration( + base_url=self._base_url, + token_url=get_token_url(), + credentials=CyclopsCredentials(client_id, client_secret), + pool_poll_interval_ms=2000, + pool_poll_limit=300, + claim_poll_interval_ms=2000, + claim_poll_limit=300, + ) + self._client = CyclopsClient.connect(configuration, self._http_client) + + async def close(self) -> None: + await self._http_client.aclose() + + async def create_pool(self, request: CreatePoolRequest) -> Any: + return await self._client.create_pool(request) + + async def create_claim(self, request: CreateClaimRequest) -> Any: + return await self._client.create_claim(request) + + async def wait_pool(self, pool: Any, timeout: float = 900.0, poll_interval: float = 5.0) -> Any: + deadline = asyncio.get_running_loop().time() + timeout + while True: + pools = await self._client.list_pools(pool.metadata.namespace) + for current_pool in pools: + if current_pool.metadata.name != pool.metadata.name: + continue + if current_pool.status and (current_pool.status.available_count or 0) >= 1: + return current_pool + break + if asyncio.get_running_loop().time() >= deadline: + raise TimeoutError( + f"Timed out waiting for Fleet pool {pool.metadata.name!r} to warm up" + ) + await asyncio.sleep(poll_interval) + + async def wait_claim(self, claim: Any) -> Any: + return await self._client.wait_claim(claim) + + async def delete_claim(self, claim: Any) -> None: + await self._client.delete_claim(claim) + + async def delete_pool(self, pool: Any) -> None: + await self._client.delete_pool(pool) + + async def service_request( + self, sandbox: Any, service: str, path: str, request: HttpRequest + ) -> Any: + return await self._client.service_request(sandbox, service, path, request) + + async def get_pool(self, name: str) -> Any: + for pool in await self.list_pools(): + if pool.metadata.name == name: + return pool + raise LookupError(f"Fleet pool {name!r} was not found") + + async def get_claim(self, pool: Any) -> Any: + expected = f"{pool.metadata.name}-claim" + for claim in await self._client.list_claims(pool.metadata.namespace): + if claim.metadata.name == expected: + return claim + raise LookupError(f"Fleet claim {expected!r} was not found") + + async def list_pools(self) -> list[Any]: + response = await self._http_client.execute( + HttpRequest(method="GET", url=f"{self._base_url}/api/namespaces", headers=[], body=None) + ) + if not 200 <= response.status < 300: + raise RuntimeError(f"Fleet namespace listing failed with HTTP {response.status}") + payload = json.loads(response.body) + items = payload if isinstance(payload, list) else payload.get("items", []) + namespaces = [ + ( + item + if isinstance(item, str) + else item.get("name") or item.get("metadata", {}).get("name") + ) + for item in items + ] + pools = await asyncio.gather( + *(self._client.list_pools(namespace) for namespace in namespaces if namespace) + ) + return [pool for namespace_pools in pools for pool in namespace_pools] + + async def set_pool_replicas(self, pool: Any, replicas: int) -> Any: + pool.spec.replicas = replicas + return await self._client.update_pool(pool) + + async def wait_service_ready( + self, sandbox: Any, service: str, time_to_start: Optional[float] = None + ) -> None: + timeout = time_to_start if time_to_start is not None else 600.0 + deadline = asyncio.get_running_loop().time() + timeout + while True: + response = await self.service_request( + sandbox, + service, + "/status", + HttpRequest( + method="GET", url="https://service.invalid/status", headers=[], body=None + ), + ) + if 200 <= response.status < 500: + return + if asyncio.get_running_loop().time() >= deadline: + raise TimeoutError( + f"Fleet service {service!r} did not become ready within {timeout} seconds" + ) + await asyncio.sleep(2) + + def service_url(self, sandbox: Any, service: str) -> str: + if service not in sandbox.services: + raise ValueError(f"Fleet sandbox does not expose service {service!r}") + return f"{self._base_url}/api/svc/{sandbox.namespace}/{sandbox.name}-{service}/" + + +class FleetCloudTransport(FleetTransport): + """Provision and manage registry-image sandboxes through Fleet.""" + + def __init__( + self, + *, + image: Optional[Image], + name: str, + cpu: Optional[int] = None, + memory_mb: Optional[int] = None, + disk_gb: Optional[int] = None, + region: str = "us-east-1", + time_to_start: Optional[float] = None, + request_timeout: Optional[float] = None, + ) -> None: + if disk_gb is not None: + raise ValueError("disk_gb is not supported by the Fleet cloud transport") + if region != "us-east-1": + raise ValueError("Fleet cloud sandboxes currently support only region='us-east-1'") + self._image = image + self._name = name + self._cpu = cpu + self._memory_mb = memory_mb + self._time_to_start = time_to_start if time_to_start is not None else 600.0 + self._request_timeout = request_timeout or 30.0 + self._provisioned = False + self._owns_resources = image is not None + self._pool: Any = None + self._claim: Any = None + self._sdk: Any = None + + @property + def name(self) -> str: + return self._name + + async def connect(self) -> None: + if not self._provisioned: + if self._sdk is None: + self._sdk = _FleetClient() + try: + if self._pool is None: + if self._image is None: + self._pool = await self._sdk.get_pool(self._name) + else: + self._validate_image(self._image) + self._pool = await self._sdk.create_pool(self._pool_request()) + self._pool = await self._sdk.wait_pool(self._pool) + if self._claim is None: + if self._image is None: + self._claim = await self._sdk.get_claim(self._pool) + else: + self._claim = await self._sdk.create_claim( + CreateClaimRequest( + pool=self._pool, + spec=ClaimSpec( + sandbox_template_ref=SandboxTemplateRef( + name=self._pool.metadata.name + ), + warmpool=None, + bind_deadline=600, + lifecycle=None, + ), + ) + ) + bound = await self._sdk.wait_claim(self._claim) + await self._sdk.wait_service_ready(bound, "server", self._time_to_start) + except BaseException as provisioning_error: + try: + if self._owns_resources: + await self._cleanup_resources() + except BaseException as cleanup_error: + logger.warning( + "Failed to clean up Fleet sandbox %r: %s", self._name, cleanup_error + ) + raise provisioning_error from cleanup_error + raise + FleetTransport.__init__( + self, + sdk=self._sdk, + bound=bound, + service_name="server", + timeout=self._request_timeout, + ) + self._provisioned = True + await FleetTransport.connect(self) + + async def create_snapshot(self, **_: Any) -> dict[str, Any]: + raise NotImplementedError("Snapshots are not supported by the Fleet cloud transport") + + async def forward_tunnel(self, sandbox_port: int | str) -> "TunnelInfo": + if not isinstance(sandbox_port, int): + raise ValueError("Fleet services can only expose numeric TCP ports") + if not self._provisioned: + raise ValueError("Transport not connected") + service = "server" if sandbox_port == 8000 else f"port-{sandbox_port}" + from cua_sandbox.interfaces.tunnel import TunnelInfo + + endpoint = self._sdk.service_url(self._bound, service) + parsed = urlparse(endpoint) + return TunnelInfo( + parsed.hostname or "", + parsed.port or (443 if parsed.scheme == "https" else 80), + sandbox_port, + url=endpoint, + ) + + async def delete_vm(self) -> None: + await self._cleanup_resources() + + async def _cleanup_resources(self) -> None: + if self._claim is not None: + await self._sdk.delete_claim(self._claim) + self._claim = None + if self._pool is not None: + await self._sdk.delete_pool(self._pool) + self._pool = None + self._provisioned = False + + @classmethod + async def list_sandboxes(cls) -> list[Any]: + sdk = _FleetClient() + try: + return await sdk.list_pools() + finally: + await sdk.close() + + @classmethod + async def get_sandbox_info(cls, name: str) -> Any: + sdk = _FleetClient() + try: + return await sdk.get_pool(name) + finally: + await sdk.close() + + @classmethod + async def suspend_sandbox(cls, name: str) -> None: + sdk = _FleetClient() + try: + await sdk.set_pool_replicas(await sdk.get_pool(name), 0) + finally: + await sdk.close() + + @classmethod + async def resume_sandbox(cls, name: str, time_to_start: Optional[float] = None) -> None: + del time_to_start + sdk = _FleetClient() + try: + await sdk.set_pool_replicas(await sdk.get_pool(name), 1) + finally: + await sdk.close() + + @classmethod + async def restart_sandbox(cls, name: str, time_to_start: Optional[float] = None) -> None: + await cls.suspend_sandbox(name) + await cls.resume_sandbox(name, time_to_start) + + @classmethod + async def delete_sandbox(cls, name: str) -> None: + sdk = _FleetClient() + try: + pool = await sdk.get_pool(name) + await sdk.delete_claim(await sdk.get_claim(pool)) + await sdk.delete_pool(pool) + finally: + await sdk.close() + + def _pool_request(self) -> CreatePoolRequest: + assert self._image is not None + services = [SandboxService(name="server", target_port=8000, protocol=ServiceProtocol.TCP)] + services.extend( + SandboxService(name=f"port-{port}", target_port=port, protocol=ServiceProtocol.TCP) + for port in self._image._ports + if port != 8000 + ) + return CreatePoolRequest( + namespace=self._name, + spec=PoolSpec( + replicas=1, + services=services, + template=PoolTemplate( + runtime=None, + runtime_class_name=None, + node_selector=None, + tolerations=None, + command=None, + container_disk_image=self._image._registry, + image_pull_secret="ecr-credentials", + cpu_cores=self._cpu, + memory=None if self._memory_mb is None else f"{self._memory_mb}Mi", + firmware=None, + probes=PreservedJson.from_json( + json.dumps({"readinessProbe": {"tcpSocket": {"port": 8000}}}) + ), + oidc=None, + ), + autoscaling=None, + ), + ) + + @staticmethod + def _service_names(pool: Any) -> list[str]: + return [service.name for service in pool.spec.services or []] or ["server"] + + @staticmethod + def _validate_image(image: Image) -> None: + if not image._registry: + raise ValueError("Fleet cloud sandboxes require Image.from_registry(...)") + if ( + image._layers + or image._env + or image._files + or image._snapshot_source + or image._disk_path + ): + raise ValueError( + "Fleet cloud supports registry images with optional exposed services only" + ) diff --git a/libs/python/cua-sandbox/cua_sandbox/transport/http.py b/libs/python/cua-sandbox/cua_sandbox/transport/http.py index 3d1ce479f9..1489ca284b 100644 --- a/libs/python/cua-sandbox/cua_sandbox/transport/http.py +++ b/libs/python/cua-sandbox/cua_sandbox/transport/http.py @@ -7,13 +7,16 @@ from __future__ import annotations import asyncio -import base64 -import json import logging from typing import Any, Dict, Optional import httpx from cua_sandbox.transport.base import Transport +from cua_sandbox.transport.computer_server import ( + decode_screenshot_response, + normalize_screen_size, + parse_command_response, +) logger = logging.getLogger(__name__) @@ -112,55 +115,7 @@ async def _cmd(self, command: str, params: Optional[Dict[str, Any]] = None) -> D resp.raise_for_status() return self._parse_sse(resp.text) - @staticmethod - def _parse_sse(text: str) -> Dict[str, Any]: - """Extract the first ``data: {...}`` frame from an SSE response. - - The server returns one of two failure shapes when ``success`` is - false: - - 1. Generic handler error — ``{"success": false, "error": ""}`` - 2. Shell-command shape — ``{"success": false, "stdout": "...", - "stderr": "...", "return_code": }`` (no ``error`` key) - - The old code stringified ``payload.get('error', 'unknown')`` for - both shapes, which turned every shell-command failure into - ``Remote error: unknown`` and hid the actual ``stderr`` + exit - code. That masking made ``Command timed out after 10s``, - ``UI hierchary dump failed``, and similar concrete failures - indistinguishable from a genuine internal error — the common - pattern where ``await sb.shell.run(cmd)`` returned non-zero - became a debugging dead end. - - This rewrite preserves the ``error``-key path verbatim and falls - back to a composite ``return_code=...`` / ``stderr=...`` / - ``stdout=...`` string when no ``error`` key is present. - """ - for line in text.splitlines(): - if line.startswith("data: "): - payload = json.loads(line[6:]) - if isinstance(payload, dict) and not payload.get("success", True): - if "error" in payload: - raise RuntimeError(f"Remote error: {payload['error']}") - # Shell-command shape: surface return_code/stderr/stdout. - parts = [] - rc = payload.get("return_code") - if rc is not None: - parts.append(f"return_code={rc}") - stderr = (payload.get("stderr") or "").strip() - if stderr: - parts.append(f"stderr={stderr!r}") - stdout = (payload.get("stdout") or "").strip() - if stdout and not stderr: - # Only surface stdout when there's nothing on - # stderr — saves bloating the message for noisy - # successful-output commands that happened to - # return non-zero. - parts.append(f"stdout={stdout!r}") - detail = ", ".join(parts) or "no detail" - raise RuntimeError(f"Remote error: {detail}") - return payload - raise RuntimeError(f"No SSE data frame in response: {text[:200]}") + _parse_sse = staticmethod(parse_command_response) async def send(self, action: str, **params: Any) -> Any: result = await self._cmd(action, params if params else None) @@ -169,25 +124,11 @@ async def send(self, action: str, **params: Any) -> Any: async def screenshot(self, format: str = "png", quality: int = 95) -> bytes: params = None if format == "png" else {"format": format, "quality": quality} result = await self._cmd("screenshot", params) - # computer-server returns {"success": true, "image_data": "..."} - b64 = result.get("image_data", result.get("base64_image", result.get("result", ""))) - if isinstance(b64, dict): - b64 = b64.get("image_data", b64.get("base64_image", b64.get("base64", ""))) - return base64.b64decode(b64) + return decode_screenshot_response(result) async def get_screen_size(self) -> Dict[str, int]: result = await self._cmd("get_screen_size") - # Flatten nested responses and normalize key names - data = result - if isinstance(data, dict): - # Unwrap nested: {"result": {...}}, {"size": {...}} - data = data.get("size", data.get("result", data)) - if isinstance(data, dict): - w = data.get("width") or data.get("screen_width") or data.get("w") - h = data.get("height") or data.get("screen_height") or data.get("h") - if w is not None and h is not None: - return {"width": int(w), "height": int(h)} - raise KeyError(f"Cannot extract screen size from response: {result}") + return normalize_screen_size(result) # ── PTY over dedicated /pty_* routes ──────────────────────────────── async def pty_create( diff --git a/libs/python/cua-sandbox/hatch_build.py b/libs/python/cua-sandbox/hatch_build.py new file mode 100644 index 0000000000..cd63132064 --- /dev/null +++ b/libs/python/cua-sandbox/hatch_build.py @@ -0,0 +1,40 @@ +import platform +import subprocess +from pathlib import Path + +from hatchling.builders.hooks.plugin.interface import BuildHookInterface + + +class CustomBuildHook(BuildHookInterface): + repository_root = Path(__file__).resolve().parents[3] + native_target_dir = repository_root / "target" / "cyclops-sdk-bindings-native" + + def initialize(self, version, build_data): + mirror_root = self.repository_root / "libs" / "fleet" + subprocess.run( + [ + "cargo", + "build", + "--locked", + "--manifest-path", + str(mirror_root / "Cargo.toml"), + "--package", + "cyclops-sdk", + "--target-dir", + str(self.native_target_dir), + ], + check=True, + ) + + suffix = {"Darwin": ".dylib", "Linux": ".so"}.get(platform.system()) + if suffix is None: + raise RuntimeError(f"unsupported host for Cyclops SDK bindings: {platform.system()}") + + native_library = self.native_target_dir / "debug" / f"libcyclops_sdk{suffix}" + if not native_library.is_file(): + raise RuntimeError(f"native Cyclops SDK library was not built: {native_library}") + + build_data["infer_tag"] = True + force_include = build_data.setdefault("force_include", {}) + force_include[str(mirror_root / "sdk-bindings" / "python" / "cyclops_sdk")] = "cyclops_sdk" + force_include[str(native_library)] = f"cyclops_sdk/{native_library.name}" diff --git a/libs/python/cua-sandbox/pyproject.toml b/libs/python/cua-sandbox/pyproject.toml index 27c5e2b6b0..603685db7e 100644 --- a/libs/python/cua-sandbox/pyproject.toml +++ b/libs/python/cua-sandbox/pyproject.toml @@ -18,7 +18,6 @@ classifiers = [ "Development Status :: 3 - Alpha", "Intended Audience :: Developers", "License :: OSI Approved :: MIT License", - "Operating System :: OS Independent", "Programming Language :: Python :: 3", "Programming Language :: Python :: 3.11", "Programming Language :: Python :: 3.12", @@ -96,3 +95,5 @@ dev = [ [tool.uv.sources] cua-agent = { path = "../agent", editable = true } + +[tool.hatch.build.hooks.custom] diff --git a/libs/python/cua-sandbox/tests/test_computer_server_transport.py b/libs/python/cua-sandbox/tests/test_computer_server_transport.py new file mode 100644 index 0000000000..b14cdcb14b --- /dev/null +++ b/libs/python/cua-sandbox/tests/test_computer_server_transport.py @@ -0,0 +1,46 @@ +import base64 + +import pytest +from cua_sandbox.transport.computer_server import ( + decode_screenshot_response, + normalize_screen_size, + parse_command_response, +) + + +def test_parse_command_response_returns_sse_payload(): + result = parse_command_response( + 'event: result\ndata: {"success": true, "result": {"ok": 1}}\n\n' + ) + + assert result == {"success": True, "result": {"ok": 1}} + + +def test_parse_command_response_raises_remote_error(): + with pytest.raises(RuntimeError, match="return_code=2"): + parse_command_response( + 'data: {"success": false, "return_code": 2, "stderr": "bad command"}\n\n' + ) + + +def test_parse_command_response_requires_data_frame(): + with pytest.raises(RuntimeError, match="No SSE data frame"): + parse_command_response("event: keepalive\n\n") + + +def test_decode_screenshot_response_accepts_nested_payload(): + encoded = base64.b64encode(b"png-data").decode() + + assert decode_screenshot_response({"result": {"image_data": encoded}}) == b"png-data" + + +def test_normalize_screen_size_accepts_nested_aliases(): + assert normalize_screen_size({"result": {"screen_width": "1280", "screen_height": 720}}) == { + "width": 1280, + "height": 720, + } + + +def test_normalize_screen_size_rejects_unknown_shape(): + with pytest.raises(KeyError, match="Cannot extract screen size"): + normalize_screen_size({"result": {"unexpected": True}}) diff --git a/libs/python/cua-sandbox/tests/test_config.py b/libs/python/cua-sandbox/tests/test_config.py index 5dc97fb6ac..fd16db5c3e 100644 --- a/libs/python/cua-sandbox/tests/test_config.py +++ b/libs/python/cua-sandbox/tests/test_config.py @@ -1,12 +1,38 @@ """Unit tests for config and auth modules.""" -from cua_sandbox._config import _global_config, configure, get_api_key, get_base_url +from cua_sandbox._config import ( + _global_config, + configure, + get_api_key, + get_base_url, + get_client_id, + get_client_secret, + get_fleet_base_url, + get_token_url, +) class TestConfig: def setup_method(self): _global_config.api_key = None - _global_config.base_url = "https://api.trycua.com" + _global_config.base_url = "https://api.cua.ai" + _global_config.fleet_base_url = "https://run.cua.ai" + _global_config.token_url = ( + "https://auth.cua.ai/realms/cyclops-cs/protocol/openid-connect/token" + ) + _global_config.client_id = None + _global_config.client_secret = None + + def test_configure_client_credentials_uses_fleet_defaults(self): + configure(client_id="client-id", client_secret="client-secret") + + assert get_client_id() == "client-id" + assert get_client_secret() == "client-secret" + assert get_base_url() == "https://api.cua.ai" + assert get_fleet_base_url() == "https://run.cua.ai" + assert ( + get_token_url() == "https://auth.cua.ai/realms/cyclops-cs/protocol/openid-connect/token" + ) def test_configure_api_key(self): configure(api_key="sk-test-123") diff --git a/libs/python/cua-sandbox/tests/test_cyclops_sdk_packaging.py b/libs/python/cua-sandbox/tests/test_cyclops_sdk_packaging.py new file mode 100644 index 0000000000..0851e7dc02 --- /dev/null +++ b/libs/python/cua-sandbox/tests/test_cyclops_sdk_packaging.py @@ -0,0 +1,76 @@ +import importlib.util +import tempfile +import unittest +from pathlib import Path +from unittest.mock import patch + +PROJECT_ROOT = Path(__file__).resolve().parents[1] +BUILD_HOOK_PATH = PROJECT_ROOT / "hatch_build.py" + + +def load_build_hook(): + spec = importlib.util.spec_from_file_location("cua_sandbox_hatch_build", BUILD_HOOK_PATH) + assert spec and spec.loader + module = importlib.util.module_from_spec(spec) + spec.loader.exec_module(module) + return module.CustomBuildHook + + +class CyclopsSdkPackagingTests(unittest.TestCase): + def test_build_hook_stages_mirrored_cyclops_sdk(self): + CustomBuildHook = load_build_hook() + with tempfile.TemporaryDirectory() as temporary_directory: + temporary_root = Path(temporary_directory) + repository_root = temporary_root / "repository" + mirror_root = repository_root / "libs" / "fleet" + binding_root = mirror_root / "sdk-bindings" / "python" / "cyclops_sdk" + binding_root.mkdir(parents=True) + native_library = temporary_root / "native" / "debug" / "libcyclops_sdk.so" + calls = [] + + def fake_run(command, check): + calls.append((command, check)) + native_library.parent.mkdir(parents=True) + native_library.touch() + + with ( + patch("platform.system", return_value="Linux"), + patch("subprocess.run", fake_run), + patch.object(CustomBuildHook, "repository_root", repository_root), + patch.object(CustomBuildHook, "native_target_dir", temporary_root / "native"), + ): + hook = object.__new__(CustomBuildHook) + build_data = {"force_include": {}} + hook.initialize("0.1.0", build_data) + + self.assertEqual( + calls, + [ + ( + [ + "cargo", + "build", + "--locked", + "--manifest-path", + str(mirror_root / "Cargo.toml"), + "--package", + "cyclops-sdk", + "--target-dir", + str(temporary_root / "native"), + ], + True, + ) + ], + ) + self.assertEqual( + build_data["force_include"], + { + str(binding_root): "cyclops_sdk", + str(native_library): "cyclops_sdk/libcyclops_sdk.so", + }, + ) + self.assertTrue(build_data["infer_tag"]) + + +if __name__ == "__main__": + unittest.main() diff --git a/libs/python/cua-sandbox/tests/test_fleet_cloud_client.py b/libs/python/cua-sandbox/tests/test_fleet_cloud_client.py new file mode 100644 index 0000000000..fced7e9be0 --- /dev/null +++ b/libs/python/cua-sandbox/tests/test_fleet_cloud_client.py @@ -0,0 +1,52 @@ +import pytest +from cua_sandbox.transport.fleet_cloud import _FleetClient + + +@pytest.mark.asyncio +async def test_list_pools_enumerates_namespaces(monkeypatch): + client = _FleetClient.__new__(_FleetClient) + + class Http: + async def execute(self, request): + assert request.url.endswith("/api/namespaces") + return type( + "Response", (), {"status": 200, "body": b'{"items":[{"name":"one"},"two"]}'} + )() + + class SDK: + async def list_pools(self, namespace): + return [namespace] + + client._base_url = "https://fleet.example" + client._http_client = Http() + client._client = SDK() + assert await client.list_pools() == ["one", "two"] + + +@pytest.mark.asyncio +async def test_service_request_delegates_to_generated_client(): + client = _FleetClient.__new__(_FleetClient) + calls = [] + + class SDK: + async def service_request(self, *args): + calls.append(args) + return "response" + + client._client = SDK() + assert await client.service_request("sandbox", "server", "/status", "request") == "response" + assert calls == [("sandbox", "server", "/status", "request")] + + +@pytest.mark.asyncio +async def test_pool_lookup_matches_sdk_resources(): + client = _FleetClient.__new__(_FleetClient) + pool = type("Pool", (), {"metadata": type("Metadata", (), {"name": "demo"})()})() + + async def list_pools(): + return [pool] + + client.list_pools = list_pools + assert await client.get_pool("demo") is pool + with pytest.raises(LookupError): + await client.get_pool("missing") diff --git a/libs/python/cua-sandbox/tests/test_fleet_cloud_transport.py b/libs/python/cua-sandbox/tests/test_fleet_cloud_transport.py new file mode 100644 index 0000000000..f6f79486fa --- /dev/null +++ b/libs/python/cua-sandbox/tests/test_fleet_cloud_transport.py @@ -0,0 +1,75 @@ +import pytest +from cua_sandbox import Image +from cua_sandbox.transport.fleet_cloud import FleetCloudTransport +from cyclops_sdk import Sandbox + + +def test_registry_image_becomes_typed_pool_request(): + request = FleetCloudTransport( + image=Image.from_registry("registry.example/workspace@sha256:abc").expose(3000), + name="demo", + cpu=4, + memory_mb=8192, + )._pool_request() + assert request.namespace == "demo" + assert request.spec.template.container_disk_image == "registry.example/workspace@sha256:abc" + assert request.spec.template.cpu_cores == 4 + assert request.spec.template.memory == "8192Mi" + assert [(service.name, service.target_port) for service in request.spec.services] == [ + ("server", 8000), + ("port-3000", 3000), + ] + + +@pytest.mark.parametrize( + "image", [Image.linux(), Image.from_registry("example:latest").apt_install("curl")] +) +def test_rejects_unsupported_images(image): + with pytest.raises(ValueError): + FleetCloudTransport._validate_image(image) + + +@pytest.mark.asyncio +async def test_cleanup_stops_after_claim_failure(): + transport = FleetCloudTransport(image=Image.from_registry("example:latest"), name="demo") + calls = [] + + class Client: + async def delete_claim(self, claim): + calls.append(("claim", claim)) + raise RuntimeError("claim delete failed") + + async def delete_pool(self, pool): + calls.append(("pool", pool)) + + transport._sdk = Client() + transport._claim = "claim" + transport._pool = "pool" + with pytest.raises(RuntimeError, match="claim delete failed"): + await transport._cleanup_resources() + assert calls == [("claim", "claim")] + + +@pytest.mark.asyncio +async def test_forward_tunnel_uses_named_service_url(): + transport = FleetCloudTransport(image=Image.from_registry("example:latest"), name="demo") + transport._provisioned = True + transport._bound = Sandbox( + namespace="demo", claim="claim", name="sandbox", services=["port-3000"] + ) + + class Client: + def service_url(self, sandbox, service): + assert service == "port-3000" + return "https://run.cua.ai/api/svc/demo/sandbox-port-3000/" + + transport._sdk = Client() + tunnel = await transport.forward_tunnel(3000) + assert tunnel.url == "https://run.cua.ai/api/svc/demo/sandbox-port-3000/" + + +@pytest.mark.asyncio +async def test_snapshot_is_unsupported(): + transport = FleetCloudTransport(image=Image.from_registry("example:latest"), name="demo") + with pytest.raises(NotImplementedError, match="Snapshots"): + await transport.create_snapshot() diff --git a/libs/python/cua-sandbox/tests/test_fleet_transport.py b/libs/python/cua-sandbox/tests/test_fleet_transport.py new file mode 100644 index 0000000000..c5af869e2b --- /dev/null +++ b/libs/python/cua-sandbox/tests/test_fleet_transport.py @@ -0,0 +1,65 @@ +import base64 +import json + +import pytest +from cua_sandbox.transport.fleet import FleetTransport +from cyclops_sdk import HttpResponse, Sandbox + + +class FakeSDK: + def __init__(self, responses): + self.responses = list(responses) + self.calls = [] + + async def service_request(self, sandbox, service, path, request): + self.calls.append((sandbox, service, path, request)) + return self.responses.pop(0) + + +def response(status=200, body=b"{}"): + return HttpResponse(status=status, headers=[], body=body) + + +def sandbox(): + return Sandbox(namespace="demo", claim="claim-demo", name="sandbox-demo", services=["api"]) + + +@pytest.mark.asyncio +async def test_service_request_forwards_command_json(): + sdk = FakeSDK([response(body=b'data: {"success":true,"result":"ok"}\n\n')]) + transport = FleetTransport(sdk=sdk, bound=sandbox()) + await transport.connect() + + assert await transport.send("shell.run", timeout=15) == "ok" + _, service, path, request = sdk.calls[0] + assert (service, path, request.method) == ("api", "/cmd", "POST") + assert json.loads(request.body) == {"command": "shell.run", "params": {"timeout": 15}} + + +@pytest.mark.asyncio +async def test_screenshot_and_pty_use_service_request(): + encoded = base64.b64encode(b"png-data").decode() + sdk = FakeSDK( + [ + response(body=f'data: {{"success":true,"image_data":"{encoded}"}}\n\n'.encode()), + response(body=b'{"pid":42}'), + response(body=b'{"killed":true}'), + ] + ) + transport = FleetTransport(sdk=sdk, bound=sandbox()) + await transport.connect() + + assert await transport.screenshot() == b"png-data" + assert await transport.pty_create(command="bash") == {"pid": 42} + assert await transport.pty_kill(42) is True + assert [call[2] for call in sdk.calls] == ["/cmd", "/pty", "/pty/42"] + + +@pytest.mark.asyncio +async def test_connect_rejects_missing_service(): + transport = FleetTransport( + sdk=FakeSDK([]), + bound=Sandbox(namespace="demo", claim="claim", name="sandbox", services=[]), + ) + with pytest.raises(ValueError, match="does not expose service"): + await transport.connect() diff --git a/libs/python/cua-sandbox/tests/test_vm_cleanup.py b/libs/python/cua-sandbox/tests/test_vm_cleanup.py index e10a2b1cc9..92a1e86ccd 100644 --- a/libs/python/cua-sandbox/tests/test_vm_cleanup.py +++ b/libs/python/cua-sandbox/tests/test_vm_cleanup.py @@ -1,6 +1,6 @@ """Unit tests for VM cleanup on connection failure and destroy() resilience. -These tests mock CloudTransport so they run without a real cloud API. +These tests mock FleetCloudTransport so they run without a real cloud API. They verify that: 1. _create() cleans up a provisioned VM when _connect() fails. 2. destroy() runs every cleanup step independently — a failure in one @@ -15,7 +15,7 @@ import pytest from cua_sandbox.image import Image from cua_sandbox.sandbox import Sandbox -from cua_sandbox.transport.cloud import CloudTransport +from cua_sandbox.transport.fleet_cloud import FleetCloudTransport pytestmark = pytest.mark.asyncio @@ -25,9 +25,9 @@ # --------------------------------------------------------------------------- -def _make_cloud_transport(*, name: str = "test-vm") -> CloudTransport: - """Return a CloudTransport with internal state set as if _create_vm() succeeded.""" - t = CloudTransport.__new__(CloudTransport) +def _make_cloud_transport(*, name: str = "test-vm") -> FleetCloudTransport: + """Return a FleetCloudTransport with internal state set as if _create_vm() succeeded.""" + t = FleetCloudTransport.__new__(FleetCloudTransport) t._name = name t._api_key_override = "sk-fake" t._base_url = "https://api.example.com" @@ -41,7 +41,7 @@ def _make_cloud_transport(*, name: str = "test-vm") -> CloudTransport: return t -def _make_sandbox(transport: CloudTransport, **kwargs) -> Sandbox: +def _make_sandbox(transport: FleetCloudTransport, **kwargs) -> Sandbox: """Return a Sandbox wrapping *transport* without calling _connect().""" return Sandbox( transport, @@ -66,12 +66,12 @@ async def test_delete_vm_called_on_timeout(self): transport.delete_vm = AsyncMock() with patch( - "cua_sandbox.sandbox.CloudTransport", + "cua_sandbox.sandbox.FleetCloudTransport", return_value=transport, ): with pytest.raises(httpx.ReadTimeout): await Sandbox._create( - image=Image.linux("ubuntu", "24.04"), + image=Image.from_registry("registry.example/workspace:latest"), api_key="sk-fake", telemetry_enabled=False, ) @@ -85,12 +85,12 @@ async def test_delete_vm_called_on_generic_exception(self): transport.delete_vm = AsyncMock() with patch( - "cua_sandbox.sandbox.CloudTransport", + "cua_sandbox.sandbox.FleetCloudTransport", return_value=transport, ): with pytest.raises(RuntimeError, match="unexpected"): await Sandbox._create( - image=Image.linux("ubuntu", "24.04"), + image=Image.from_registry("registry.example/workspace:latest"), api_key="sk-fake", telemetry_enabled=False, ) @@ -104,12 +104,12 @@ async def test_original_exception_propagates_even_if_delete_fails(self): transport.delete_vm = AsyncMock(side_effect=httpx.ConnectError("api down")) with patch( - "cua_sandbox.sandbox.CloudTransport", + "cua_sandbox.sandbox.FleetCloudTransport", return_value=transport, ): with pytest.raises(TimeoutError, match="poll timeout"): await Sandbox._create( - image=Image.linux("ubuntu", "24.04"), + image=Image.from_registry("registry.example/workspace:latest"), api_key="sk-fake", telemetry_enabled=False, ) @@ -125,12 +125,12 @@ async def test_no_cleanup_when_vm_not_yet_created(self): transport.delete_vm = AsyncMock() with patch( - "cua_sandbox.sandbox.CloudTransport", + "cua_sandbox.sandbox.FleetCloudTransport", return_value=transport, ): with pytest.raises(ValueError, match="no api key"): await Sandbox._create( - image=Image.linux("ubuntu", "24.04"), + image=Image.from_registry("registry.example/workspace:latest"), api_key="sk-fake", telemetry_enabled=False, ) @@ -144,12 +144,12 @@ async def test_keyboard_interrupt_still_cleans_up(self): transport.delete_vm = AsyncMock() with patch( - "cua_sandbox.sandbox.CloudTransport", + "cua_sandbox.sandbox.FleetCloudTransport", return_value=transport, ): with pytest.raises(KeyboardInterrupt): await Sandbox._create( - image=Image.linux("ubuntu", "24.04"), + image=Image.from_registry("registry.example/workspace:latest"), api_key="sk-fake", telemetry_enabled=False, ) @@ -231,14 +231,14 @@ async def test_destroy_happy_path(self): transport.delete_vm.assert_awaited_once() async def test_non_cloud_transport_skips_delete_vm(self): - """Non-CloudTransport sandboxes should not call delete_vm.""" - transport = AsyncMock() # generic mock, not a CloudTransport instance + """Non-FleetCloudTransport sandboxes should not call delete_vm.""" + transport = AsyncMock() # generic mock, not a FleetCloudTransport instance sb = Sandbox(transport, name="local-vm", _ephemeral=True, _telemetry_enabled=False) await sb.destroy() transport.disconnect.assert_awaited_once() - # delete_vm should not be called since transport is not CloudTransport + # delete_vm should not be called since transport is not FleetCloudTransport assert not hasattr(transport, "delete_vm") or not transport.delete_vm.called @@ -259,18 +259,18 @@ async def test_ephemeral_destroys_on_normal_exit(self): with ( patch.object( - CloudTransport, + FleetCloudTransport, "__init__", lambda self, **kw: None, ), patch.object( - CloudTransport, + FleetCloudTransport, "__new__", lambda cls, **kw: transport, ), ): async with Sandbox.ephemeral( - Image.linux("ubuntu", "24.04"), + Image.from_registry("registry.example/workspace:latest"), api_key="sk-fake", telemetry_enabled=False, ) as sb: @@ -287,19 +287,19 @@ async def test_ephemeral_destroys_on_test_failure(self): with ( patch.object( - CloudTransport, + FleetCloudTransport, "__init__", lambda self, **kw: None, ), patch.object( - CloudTransport, + FleetCloudTransport, "__new__", lambda cls, **kw: transport, ), ): with pytest.raises(AssertionError): async with Sandbox.ephemeral( - Image.linux("ubuntu", "24.04"), + Image.from_registry("registry.example/workspace:latest"), api_key="sk-fake", telemetry_enabled=False, ) as _sb: @@ -315,19 +315,19 @@ async def test_ephemeral_cleans_up_when_create_connect_fails(self): with ( patch.object( - CloudTransport, + FleetCloudTransport, "__init__", lambda self, **kw: None, ), patch.object( - CloudTransport, + FleetCloudTransport, "__new__", lambda cls, **kw: transport, ), ): with pytest.raises(httpx.ReadTimeout): async with Sandbox.ephemeral( - Image.linux("ubuntu", "24.04"), + Image.from_registry("registry.example/workspace:latest"), api_key="sk-fake", telemetry_enabled=False, ) as _sb: