diff --git a/packages/modal-infra/src/sandbox/launch.py b/packages/modal-infra/src/sandbox/launch.py new file mode 100644 index 0000000000..7aee1fbb04 --- /dev/null +++ b/packages/modal-infra/src/sandbox/launch.py @@ -0,0 +1,354 @@ +"""Translate a session launch into Modal image, environment, and resource arguments.""" + +import json +import secrets +import time +from dataclasses import dataclass +from typing import Any + +import modal + +from sandbox_runtime.constants import ( + DOCKER_ENABLED_ENV_VAR, + NOVNC_PORT_ENV_VAR, + SANDBOX_TIMEOUT_ENV_VAR, + VNC_PASSWORD_ENV_VAR, + VNC_PASSWORD_MAX_BYTES, +) +from sandbox_runtime.log_config import get_logger +from sandbox_runtime.types import SandboxStatus + +from ..app import app +from ..app_config import APP_NAME +from ..images.base import base_image +from .launch_policy import ( + docker_allocation_name, + docker_allocation_tags, + docker_base_image, + docker_runtime_env, + launch_kwargs, + parse_launch, +) +from .models import SandboxConfig, SandboxHandle +from .tunnels import SandboxTunnels +from .vcs_env import inject_vcs_env_vars + +_RESERVED_LAUNCH_ENV_VARS = { + "RESTORED_FROM_SNAPSHOT", + "FROM_REPO_IMAGE", + "REPO_IMAGE_SHA", + "IMAGE_BUILD_MODE", + "TERMINAL_ENABLED", + "AGENT_SLACK_NOTIFY_ENABLED", + "SESSION_CONFIG", + VNC_PASSWORD_ENV_VAR, + NOVNC_PORT_ENV_VAR, + DOCKER_ENABLED_ENV_VAR, +} + +log = get_logger("manager") +ACCESS_PASSWORD_READ_TIMEOUT_SECONDS = 30 + + +class RepositoryImageUnavailableError(RuntimeError): + """The selected repository image no longer exists in Modal.""" + + +@dataclass(frozen=True) +class BaseImageSource: + pass + + +@dataclass(frozen=True) +class RepositoryImageSource: + image_id: str + sha: str | None + + +@dataclass(frozen=True) +class SnapshotImageSource: + image_id: str + clone_token: str | None + + +type SandboxImageSource = BaseImageSource | RepositoryImageSource | SnapshotImageSource + + +async def _create_sandbox( + create_kwargs: dict[str, Any], *, repository_image: bool +) -> modal.Sandbox: + """Only a missing repository image at create time is classified as unavailable.""" + try: + return await modal.Sandbox.create.aio( + "python", "-m", "sandbox_runtime.entrypoint", **create_kwargs + ) + except modal.exception.NotFoundError as e: + if repository_image: + raise RepositoryImageUnavailableError("repository image is unavailable") from e + raise + + +def _session_identity(session_config: Any) -> str: + if isinstance(session_config, dict): + session_id = session_config.get("session_id") + elif session_config is not None: + session_id = session_config.session_id + else: + session_id = None + return session_id if isinstance(session_id, str) else "" + + +@dataclass(frozen=True) +class SandboxLaunchSpec: + """Canonical launch configuration paired with one image source variant.""" + + config: SandboxConfig + source: SandboxImageSource + + +class SandboxLauncher: + """Own the common Modal launch path for base, repository, and snapshot images.""" + + @staticmethod + def _generate_code_server_password() -> str: + """Generate a random code-server password.""" + return secrets.token_urlsafe(16) + + @staticmethod + def _generate_vnc_password() -> str: + """Generate a random VNC password.""" + return secrets.token_urlsafe(VNC_PASSWORD_MAX_BYTES)[:VNC_PASSWORD_MAX_BYTES] + + async def launch(self, spec: SandboxLaunchSpec) -> SandboxHandle: + """Launch a Modal sandbox from a normalized create or restore specification.""" + config = spec.config + has_repository = bool(config.repo_owner) + sandbox_id = config.sandbox_id + if not sandbox_id: + sandbox_name = ( + f"{config.repo_owner}-{config.repo_name}" if has_repository else "no-repository" + ) + sandbox_id = f"sandbox-{sandbox_name}-{int(time.time() * 1000)}" + + docker = parse_launch(config.sandbox_backend, config.settings) + env_vars = { + key: value + for key, value in (config.user_env_vars or {}).items() + if key not in _RESERVED_LAUNCH_ENV_VARS + } + env_vars.update( + { + "PYTHONUNBUFFERED": "1", + "SANDBOX_ID": sandbox_id, + "CONTROL_PLANE_URL": config.control_plane_url, + "SANDBOX_AUTH_TOKEN": config.sandbox_auth_token, + SANDBOX_TIMEOUT_ENV_VAR: str(config.timeout_seconds), + "REPO_OWNER": config.repo_owner or "", + "REPO_NAME": config.repo_name or "", + **docker_runtime_env(docker), + } + ) + + clone_token: str | None = None + include_github_cli_aliases = False + snapshot_id: str | None = None + if isinstance(spec.source, BaseImageSource): + image = docker_base_image() if docker.enabled else base_image + elif isinstance(spec.source, RepositoryImageSource): + try: + image = modal.Image.from_id(spec.source.image_id) + except modal.exception.NotFoundError as e: + raise RepositoryImageUnavailableError("repository image is unavailable") from e + env_vars["FROM_REPO_IMAGE"] = "true" + env_vars["REPO_IMAGE_SHA"] = spec.source.sha or "" + else: + image = modal.Image.from_id(spec.source.image_id) + env_vars["RESTORED_FROM_SNAPSHOT"] = "true" + clone_token = spec.source.clone_token + include_github_cli_aliases = True + snapshot_id = spec.source.image_id + + if config.session_config is not None: + env_vars["SESSION_CONFIG"] = ( + json.dumps(config.session_config) + if isinstance(config.session_config, dict) + else config.session_config.model_dump_json() + ) + + inject_vcs_env_vars( + env_vars, + clone_token=clone_token if has_repository else None, + include_github_cli_aliases=include_github_cli_aliases, + ) + + code_server_password: str | None = None + if config.code_server_enabled: + code_server_password = self._generate_code_server_password() + env_vars["CODE_SERVER_PASSWORD"] = code_server_password + + vnc_password: str | None = None + if config.vnc_enabled: + vnc_password = self._generate_vnc_password() + env_vars[VNC_PASSWORD_ENV_VAR] = vnc_password + + if config.agent_slack_notify_enabled: + env_vars["AGENT_SLACK_NOTIFY_ENABLED"] = "true" + + tunnels = SandboxTunnels( + code_server_enabled=config.code_server_enabled, + vnc_enabled=config.vnc_enabled, + settings=config.settings, + ) + env_vars.update(tunnels.environment) + + # A fresh handle avoids Modal caching the ID of a deleted/recreated secret. + llm_secrets = modal.Secret.from_name("llm-api-keys") + await llm_secrets.hydrate.aio() + create_kwargs: dict[str, Any] = { + "image": image, + "app": app, + "secrets": [llm_secrets], + "timeout": config.timeout_seconds, + "workdir": "/workspace", + "env": env_vars, + **launch_kwargs(docker), + } + if tunnels.exposed_ports: + create_kwargs["encrypted_ports"] = tunnels.exposed_ports + + repository_image = isinstance(spec.source, RepositoryImageSource) + if docker.enabled: + sandbox, adopted = await self._launch_docker_sandbox( + session_id=_session_identity(config.session_config), + sandbox_id=sandbox_id, + retire_sandbox_id=config.retire_sandbox_id, + create_kwargs=create_kwargs, + repository_image=repository_image, + launch_deadline_at_ms=config.launch_deadline_at_ms, + ) + if adopted: + passwords = await self._read_access_passwords( + sandbox, + code_server_enabled=config.code_server_enabled, + vnc_enabled=config.vnc_enabled, + ) + code_server_password = passwords.get("CODE_SERVER_PASSWORD") + vnc_password = passwords.get(VNC_PASSWORD_ENV_VAR) + else: + sandbox = await _create_sandbox(create_kwargs, repository_image=repository_image) + modal_object_id = sandbox.object_id + urls = await tunnels.resolve(sandbox, sandbox_id) + + return SandboxHandle( + sandbox_id=sandbox_id, + modal_sandbox=sandbox, + status=SandboxStatus.WARMING, + created_at=time.time(), + snapshot_id=snapshot_id, + modal_object_id=modal_object_id, + code_server_url=urls.code_server_url, + code_server_password=code_server_password, + vnc_url=urls.vnc_url, + vnc_password=vnc_password, + ttyd_url=urls.ttyd_url, + tunnel_urls=urls.tunnel_urls, + sandbox_backend=docker.backend, + ) + + async def _launch_docker_sandbox( + self, + *, + session_id: str, + sandbox_id: str, + retire_sandbox_id: str | None, + create_kwargs: dict[str, Any], + repository_image: bool, + launch_deadline_at_ms: int | None = None, + ) -> tuple[modal.Sandbox, bool]: + """Create a named VM or adopt only the allocation owned by this generation.""" + if retire_sandbox_id: + await self._retire_docker_allocation(session_id, retire_sandbox_id) + name = docker_allocation_name(session_id) + tags = docker_allocation_tags(session_id, sandbox_id) + existing = await self._find_owned_docker_allocation(name, tags) + if existing is None: + if launch_deadline_at_ms is not None and time.time() * 1000 >= launch_deadline_at_ms: + raise RuntimeError("VM launch deadline expired") + try: + sandbox = await _create_sandbox( + {**create_kwargs, "name": name, "tags": tags}, + repository_image=repository_image, + ) + return sandbox, False + except modal.exception.AlreadyExistsError: + existing = await self._find_owned_docker_allocation(name, tags) + if existing is None: + raise + log.info( + "sandbox.docker_allocation_adopted", + sandbox_id=sandbox_id, + modal_object_id=existing.object_id, + ) + return existing, True + + @staticmethod + async def _read_access_passwords( + sandbox: modal.Sandbox, *, code_server_enabled: bool, vnc_enabled: bool + ) -> dict[str, str]: + """Recover enabled service credentials from the owned VM launch environment.""" + keys = [] + if code_server_enabled: + keys.append("CODE_SERVER_PASSWORD") + if vnc_enabled: + keys.append(VNC_PASSWORD_ENV_VAR) + if not keys: + return {} + process = await sandbox.exec.aio( + "python", + "-I", + "-c", + "import json, os, sys; print(json.dumps({k: os.environ.get(k) for k in sys.argv[1:]}))", + *keys, + timeout=ACCESS_PASSWORD_READ_TIMEOUT_SECONDS, + ) + output = await process.stdout.read.aio() + if await process.wait.aio() != 0: + raise RuntimeError("Could not recover adopted sandbox access credentials") + try: + passwords = json.loads(output) + except ValueError: + raise RuntimeError("Could not recover adopted sandbox access credentials") from None + if not isinstance(passwords, dict) or any( + not isinstance(passwords.get(key), str) or not passwords[key] for key in keys + ): + raise RuntimeError("Could not recover adopted sandbox access credentials") + return {key: passwords[key] for key in keys} + + @staticmethod + async def _find_owned_docker_allocation( + name: str, tags: dict[str, str] + ) -> modal.Sandbox | None: + try: + sandbox = await modal.Sandbox.from_name.aio(APP_NAME, name) + except modal.exception.NotFoundError: + return None + if await sandbox.get_tags.aio() != tags: + raise RuntimeError("Docker sandbox allocation ownership mismatch") + return sandbox + + async def _retire_docker_allocation(self, session_id: str, sandbox_id: str) -> None: + """Terminate a prior named VM only when its ownership tags match.""" + name = docker_allocation_name(session_id) + try: + sandbox = await modal.Sandbox.from_name.aio(APP_NAME, name) + except modal.exception.NotFoundError: + return + if await sandbox.get_tags.aio() != docker_allocation_tags(session_id, sandbox_id): + log.warn("sandbox.docker_allocation_retire_mismatch", sandbox_id=sandbox_id) + return + await sandbox.terminate.aio(wait=True) + log.info( + "sandbox.docker_allocation_retired", + sandbox_id=sandbox_id, + modal_object_id=sandbox.object_id, + ) diff --git a/packages/modal-infra/src/sandbox/manager.py b/packages/modal-infra/src/sandbox/manager.py index 30e044eb8e..7a82b3f9f8 100644 --- a/packages/modal-infra/src/sandbox/manager.py +++ b/packages/modal-infra/src/sandbox/manager.py @@ -1,18 +1,6 @@ -""" -Sandbox lifecycle management for Open-Inspect. +"""Provider lifecycle operations for Open-Inspect session sandboxes.""" -This module handles: -- Creating sandboxes from filesystem snapshots -- Taking snapshots for session persistence - -Updated: 2026-01-15 to fix Sandbox.create API -""" - -import asyncio -import json -import secrets import time -from dataclasses import dataclass from typing import Any import modal @@ -38,44 +26,60 @@ from sandbox_runtime.log_config import get_logger from sandbox_runtime.types import SandboxStatus, SessionConfig -from ..app import app from ..app_config import APP_NAME -from ..images.base import base_image +from .launch import ( + ACCESS_PASSWORD_READ_TIMEOUT_SECONDS, + BaseImageSource, + RepositoryImageSource, + RepositoryImageUnavailableError, + SandboxImageSource, + SandboxLauncher, + SandboxLaunchSpec, + SnapshotImageSource, +) from .launch_policy import ( PENDING_VM_REFERENCE_PREFIX, ModalBackend, docker_allocation_name, docker_allocation_tags, - docker_base_image, - docker_runtime_env, - launch_kwargs, - parse_launch, parse_pending_vm_reference, ) -from .vcs_env import inject_vcs_env_vars +from .models import DEFAULT_VNC_ENABLED, SandboxConfig, SandboxHandle +from .tunnels import MAX_TUNNEL_PORTS + +# Preserve the existing public imports after moving their implementations. +__all__ = [ + "ACCESS_PASSWORD_READ_TIMEOUT_SECONDS", + "APP_NAME", + "CODE_SERVER_PORT", + "CODE_SERVER_PORT_ENV_VAR", + "CONTROL_TIMEOUT_SECONDS", + "DEFAULT_SANDBOX_TIMEOUT_SECONDS", + "DEFAULT_VNC_ENABLED", + "DOCKER_ENABLED_ENV_VAR", + "EXPECTED_TUNNEL_PORTS_ENV_VAR", + "MAX_TUNNEL_PORTS", + "NOVNC_PORT", + "NOVNC_PORT_ENV_VAR", + "PENDING_VM_REFERENCE_PREFIX", + "SANDBOX_TIMEOUT_ENV_VAR", + "SNAPSHOT_FILESYSTEM_TIMEOUT_SECONDS", + "TTYD_PROXY_PORT", + "TTYD_PROXY_PORT_ENV_VAR", + "TUNNEL_ENV_FILE_PATH", + "TUNNEL_ENV_SANDBOX_ID_KEY", + "VNC_PASSWORD_ENV_VAR", + "VNC_PASSWORD_MAX_BYTES", + "VNC_PORT", + "RepositoryImageUnavailableError", + "SandboxConfig", + "SandboxHandle", + "SandboxManager", +] log = get_logger("manager") SNAPSHOT_FILESYSTEM_TIMEOUT_SECONDS = 300 -ACCESS_PASSWORD_READ_TIMEOUT_SECONDS = 30 -MAX_TUNNEL_PORTS = 10 -DEFAULT_VNC_ENABLED = False -_RESERVED_LAUNCH_ENV_VARS = { - "RESTORED_FROM_SNAPSHOT", - "FROM_REPO_IMAGE", - "REPO_IMAGE_SHA", - "IMAGE_BUILD_MODE", - "TERMINAL_ENABLED", - "AGENT_SLACK_NOTIFY_ENABLED", - "SESSION_CONFIG", - VNC_PASSWORD_ENV_VAR, - NOVNC_PORT_ENV_VAR, - DOCKER_ENABLED_ENV_VAR, -} - - -class RepositoryImageUnavailableError(RuntimeError): - """The selected repository image no longer exists in Modal.""" class PendingVMReferenceNotVisible(RuntimeError): @@ -90,574 +94,13 @@ def _has_repository(repo_owner: str | None, repo_name: str | None) -> bool: return has_owner -async def _create_sandbox( - create_kwargs: dict[str, Any], *, repository_image: bool -) -> modal.Sandbox: - """The one `Sandbox.create` call; only its own NotFound means the image is gone.""" - try: - return await modal.Sandbox.create.aio( - "python", - "-m", - "sandbox_runtime.entrypoint", - **create_kwargs, - ) - except modal.exception.NotFoundError as e: - if repository_image: - raise RepositoryImageUnavailableError("repository image is unavailable") from e - raise - - -def _session_identity(session_config: SessionConfig | dict[str, Any] | None) -> str: - """The control-plane session id carried in the launch's session config.""" - if isinstance(session_config, dict): - session_id = session_config.get("session_id") - elif session_config is not None: - session_id = session_config.session_id - else: - session_id = None - return session_id if isinstance(session_id, str) else "" - - -@dataclass -class SandboxConfig: - """Configuration for creating a sandbox.""" - - repo_owner: str | None - repo_name: str | None - sandbox_backend: ModalBackend = "modal" - sandbox_id: str | None = None # Expected sandbox ID from control plane - session_config: SessionConfig | dict[str, Any] | None = None - control_plane_url: str = "" - sandbox_auth_token: str = "" - timeout_seconds: int = DEFAULT_SANDBOX_TIMEOUT_SECONDS - user_env_vars: dict[str, str] | None = None # User-provided env vars (repo secrets) - repo_image_id: str | None = None # Pre-built repo image ID from provider - repo_image_sha: str | None = None # Git SHA the repo image was built from - code_server_enabled: bool = False # Whether to start code-server in the sandbox - vnc_enabled: bool = DEFAULT_VNC_ENABLED # Whether to start the browser-accessible VNC desktop - agent_slack_notify_enabled: bool = ( - False # Whether to install the agent-initiated slack-notify tool - ) - settings: dict[str, Any] | None = ( - None # Sandbox settings (tunnelPorts, etc.) from control plane - ) - # A previous generation's sandbox id whose Docker VM may still be running - # after an ambiguous create (the control plane lost the response). Only - # Docker launches act on it; the named allocation is retired if owned. - retire_sandbox_id: str | None = None - launch_deadline_at_ms: int | None = None - - -@dataclass -class SandboxHandle: - """Handle to a sandbox.""" - - sandbox_id: str - modal_sandbox: modal.Sandbox - status: SandboxStatus - created_at: float - snapshot_id: str | None = None - modal_object_id: str | None = None # Modal's internal sandbox ID for API calls - code_server_url: str | None = None - code_server_password: str | None = None - vnc_url: str | None = None - vnc_password: str | None = None - ttyd_url: str | None = None # proxy tunnel URL (not ttyd directly) - tunnel_urls: dict[int, str] | None = None # port -> tunnel URL mapping for extra ports - sandbox_backend: ModalBackend = "modal" - - -@dataclass(frozen=True) -class _BaseImageSource: - pass - - -@dataclass(frozen=True) -class _RepositoryImageSource: - image_id: str - sha: str | None - - -@dataclass(frozen=True) -class _SnapshotImageSource: - image_id: str - clone_token: str | None - - -type _SandboxImageSource = _BaseImageSource | _RepositoryImageSource | _SnapshotImageSource - - -@dataclass(frozen=True) -class _SandboxLaunchSpec: - """Canonical launch configuration paired with one image source variant.""" - - config: SandboxConfig - source: _SandboxImageSource - - class SandboxManager: - """ - Manages sandbox lifecycle for Open-Inspect sessions. + """Normalize create/restore requests and manage existing provider sandboxes. - Responsibilities: - - Create sandboxes from snapshots or fresh images - - Take snapshots for session persistence + Launch translation and networking are owned by provider-local collaborators. + Session readiness and checkpoint/shutdown policy remain in the control plane. """ - @staticmethod - def _generate_code_server_password() -> str: - """Generate a random code-server password.""" - return secrets.token_urlsafe(16) - - @staticmethod - def _generate_vnc_password() -> str: - """Generate a random VNC password.""" - return secrets.token_urlsafe(VNC_PASSWORD_MAX_BYTES)[:VNC_PASSWORD_MAX_BYTES] - - @staticmethod - async def _resolve_tunnels( - sandbox: modal.Sandbox, - sandbox_id: str, - ports: list[int], - retries: int = 3, - backoff: float = 1.0, - ) -> dict[int, str]: - """Resolve tunnel URLs for the given ports from Modal, retrying on failure.""" - resolved: dict[int, str] = {} - for attempt in range(retries): - try: - loop = asyncio.get_running_loop() - tunnels = await loop.run_in_executor(None, sandbox.tunnels) - for port in ports: - if port in tunnels and port not in resolved: - resolved[port] = tunnels[port].url - log.info( - "tunnel.resolved", - sandbox_id=sandbox_id, - port=port, - url=tunnels[port].url, - ) - if len(resolved) == len(ports): - return resolved - except Exception as e: - log.warn( - "tunnel.resolve_error", - sandbox_id=sandbox_id, - attempt=attempt + 1, - retries=retries, - error=type(e).__name__, - exc=e, - ) - if attempt < retries - 1: - await asyncio.sleep(backoff * (attempt + 1)) - return resolved - - @staticmethod - def _validate_ports(raw: list) -> list[int]: - """Validate and sanitize tunnel ports: must be int, 1-65535, max MAX_TUNNEL_PORTS.""" - ports: list[int] = [] - for p in raw: - if isinstance(p, int) and 1 <= p <= 65535: - ports.append(p) - if len(ports) >= MAX_TUNNEL_PORTS: - break - return ports - - @staticmethod - def _resolve_service_ports(settings: dict[str, Any] | None) -> tuple[int, int, int]: - """Return effective (code_server_port, novnc_port, ttyd_proxy_port) from settings. - - Falls back to the service defaults when unset or invalid. The control - plane validates these before they reach here. - """ - s = settings or {} - - def coerce(value: Any, default: int) -> int: - if isinstance(value, int) and not isinstance(value, bool) and 1 <= value <= 65535: - return value - return default - - return ( - coerce(s.get("codeServerPort"), CODE_SERVER_PORT), - coerce(s.get("vncPort"), NOVNC_PORT), - coerce(s.get("terminalPort"), TTYD_PROXY_PORT), - ) - - @staticmethod - def _collect_exposed_ports( - code_server_enabled: bool, - vnc_enabled: bool, - terminal_enabled: bool, - settings: dict[str, Any] | None, - code_server_port: int, - novnc_port: int, - ttyd_proxy_port: int, - ) -> tuple[list[int], list[int]]: - """Return (all_exposed_ports, extra_tunnel_ports) from settings and feature flags.""" - # Raw VNC is localhost-only and must never be exposed, including as a - # user-configured extra tunnel. - reserved: set[int] = {VNC_PORT} - exposed: list[int] = [] - if code_server_enabled: - exposed.append(code_server_port) - reserved.add(code_server_port) - if vnc_enabled: - exposed.append(novnc_port) - reserved.add(novnc_port) - if terminal_enabled: - exposed.append(ttyd_proxy_port) - reserved.add(ttyd_proxy_port) - - raw_ports = (settings or {}).get("tunnelPorts", []) - tunnel_ports = SandboxManager._validate_ports(raw_ports) if raw_ports else [] - # Remove reserved ports from tunnel_ports to avoid duplicates - tunnel_ports = [p for p in tunnel_ports if p not in reserved] - exposed.extend(tunnel_ports) - return exposed, tunnel_ports - - @staticmethod - async def _resolve_and_setup_tunnels( - sandbox: modal.Sandbox, - sandbox_id: str, - code_server_enabled: bool, - vnc_enabled: bool, - terminal_enabled: bool, - extra_ports: list[int], - code_server_port: int, - novnc_port: int, - ttyd_proxy_port: int, - ) -> tuple[str | None, str | None, str | None, dict[int, str] | None]: - """Return (code_server_url, vnc_url, ttyd_url, extra_urls).""" - all_ports: list[int] = [] - if code_server_enabled: - all_ports.append(code_server_port) - if vnc_enabled: - all_ports.append(novnc_port) - if terminal_enabled: - all_ports.append(ttyd_proxy_port) - all_ports.extend(extra_ports) - - if not all_ports: - return None, None, None, None - - resolved = await SandboxManager._resolve_tunnels(sandbox, sandbox_id, all_ports) - - # Only pull a service port out of the resolved map when that service owns - # it. Otherwise a user's own port (e.g. 8080 with code-server disabled) - # would be misrouted to code_server_url and dropped from the tunnel map. - code_server_url = resolved.pop(code_server_port, None) if code_server_enabled else None - vnc_url = resolved.pop(novnc_port, None) if vnc_enabled else None - ttyd_url = resolved.pop(ttyd_proxy_port, None) if terminal_enabled else None - extra_urls = resolved if resolved else None - - if extra_urls: - await SandboxManager._write_tunnel_env_file(sandbox, sandbox_id, extra_urls) - - return code_server_url, vnc_url, ttyd_url, extra_urls - - @staticmethod - async def _write_tunnel_env_file( - sandbox: modal.Sandbox, - sandbox_id: str, - tunnel_urls: dict[int, str], - ) -> None: - """Write tunnel URLs to TUNNEL_ENV_FILE_PATH as a dotenv file. - - The first line tags the file with this sandbox's ID so the supervisor's - stale-file cleanup can tell a fresh write (this write can land before - the entrypoint runs) from a snapshot/image leftover. - - Failures are logged but do not block sandbox creation; URLs are also - returned to the control plane via the SandboxHandle. - """ - lines = [f"{TUNNEL_ENV_SANDBOX_ID_KEY}={sandbox_id}"] - lines += [f"TUNNEL_{port}={url}" for port, url in sorted(tunnel_urls.items())] - content = "\n".join(lines) + "\n" - try: - await sandbox.filesystem.write_text.aio(content, TUNNEL_ENV_FILE_PATH) - log.info( - "tunnel.urls_written", - sandbox_id=sandbox_id, - path=TUNNEL_ENV_FILE_PATH, - ports=list(tunnel_urls.keys()), - ) - except Exception as e: - log.warn( - "tunnel.urls_write_failed", - sandbox_id=sandbox_id, - path=TUNNEL_ENV_FILE_PATH, - exc=e, - ) - - async def _launch_sandbox(self, spec: _SandboxLaunchSpec) -> SandboxHandle: - """Launch a Modal sandbox from a normalized create or restore specification.""" - config = spec.config - has_repository = bool(config.repo_owner) - sandbox_id = config.sandbox_id - if not sandbox_id: - sandbox_name = ( - f"{config.repo_owner}-{config.repo_name}" if has_repository else "no-repository" - ) - sandbox_id = f"sandbox-{sandbox_name}-{int(time.time() * 1000)}" - - docker = parse_launch(config.sandbox_backend, config.settings) - env_vars = { - key: value - for key, value in (config.user_env_vars or {}).items() - if key not in _RESERVED_LAUNCH_ENV_VARS - } - env_vars.update( - { - "PYTHONUNBUFFERED": "1", - "SANDBOX_ID": sandbox_id, - "CONTROL_PLANE_URL": config.control_plane_url, - "SANDBOX_AUTH_TOKEN": config.sandbox_auth_token, - SANDBOX_TIMEOUT_ENV_VAR: str(config.timeout_seconds), - "REPO_OWNER": config.repo_owner or "", - "REPO_NAME": config.repo_name or "", - **docker_runtime_env(docker), - } - ) - - clone_token: str | None = None - include_github_cli_aliases = False - snapshot_id: str | None = None - if isinstance(spec.source, _BaseImageSource): - image = docker_base_image() if docker.enabled else base_image - elif isinstance(spec.source, _RepositoryImageSource): - image = modal.Image.from_id(spec.source.image_id) - env_vars["FROM_REPO_IMAGE"] = "true" - env_vars["REPO_IMAGE_SHA"] = spec.source.sha or "" - else: - image = modal.Image.from_id(spec.source.image_id) - env_vars["RESTORED_FROM_SNAPSHOT"] = "true" - clone_token = spec.source.clone_token - include_github_cli_aliases = True - snapshot_id = spec.source.image_id - - if config.session_config is not None: - env_vars["SESSION_CONFIG"] = ( - json.dumps(config.session_config) - if isinstance(config.session_config, dict) - else config.session_config.model_dump_json() - ) - - inject_vcs_env_vars( - env_vars, - clone_token=clone_token if has_repository else None, - include_github_cli_aliases=include_github_cli_aliases, - ) - - code_server_password: str | None = None - if config.code_server_enabled: - code_server_password = self._generate_code_server_password() - env_vars["CODE_SERVER_PASSWORD"] = code_server_password - - vnc_password: str | None = None - if config.vnc_enabled: - vnc_password = self._generate_vnc_password() - env_vars[VNC_PASSWORD_ENV_VAR] = vnc_password - - terminal_enabled = bool((config.settings or {}).get("terminalEnabled", False)) - if terminal_enabled: - env_vars["TERMINAL_ENABLED"] = "true" - if config.agent_slack_notify_enabled: - env_vars["AGENT_SLACK_NOTIFY_ENABLED"] = "true" - - code_server_port, novnc_port, ttyd_proxy_port = self._resolve_service_ports(config.settings) - if config.code_server_enabled: - env_vars[CODE_SERVER_PORT_ENV_VAR] = str(code_server_port) - if config.vnc_enabled: - env_vars[NOVNC_PORT_ENV_VAR] = str(novnc_port) - if terminal_enabled: - env_vars[TTYD_PROXY_PORT_ENV_VAR] = str(ttyd_proxy_port) - - exposed_ports, tunnel_ports = self._collect_exposed_ports( - config.code_server_enabled, - config.vnc_enabled, - terminal_enabled, - config.settings, - code_server_port, - novnc_port, - ttyd_proxy_port, - ) - if tunnel_ports: - env_vars[EXPECTED_TUNNEL_PORTS_ENV_VAR] = ",".join(str(p) for p in tunnel_ports) - - # from_name handles cache their resolved ID; use a fresh handle on every - # launch so a deleted and recreated secret can be resolved again. - llm_secrets = modal.Secret.from_name("llm-api-keys") - await llm_secrets.hydrate.aio() - - create_kwargs: dict[str, Any] = { - "image": image, - "app": app, - "secrets": [llm_secrets], - "timeout": config.timeout_seconds, - "workdir": "/workspace", - "env": env_vars, - **launch_kwargs(docker), - } - if exposed_ports: - create_kwargs["encrypted_ports"] = exposed_ports - - repository_image = isinstance(spec.source, _RepositoryImageSource) - if docker.enabled: - sandbox, adopted = await self._launch_docker_sandbox( - session_id=_session_identity(config.session_config), - sandbox_id=sandbox_id, - retire_sandbox_id=config.retire_sandbox_id, - create_kwargs=create_kwargs, - repository_image=repository_image, - launch_deadline_at_ms=config.launch_deadline_at_ms, - ) - if adopted: - passwords = await self._read_access_passwords( - sandbox, - code_server_enabled=config.code_server_enabled, - vnc_enabled=config.vnc_enabled, - ) - code_server_password = passwords.get("CODE_SERVER_PASSWORD") - vnc_password = passwords.get(VNC_PASSWORD_ENV_VAR) - else: - sandbox = await _create_sandbox(create_kwargs, repository_image=repository_image) - modal_object_id = sandbox.object_id - ( - code_server_url, - vnc_url, - ttyd_url, - extra_tunnel_urls, - ) = await self._resolve_and_setup_tunnels( - sandbox, - sandbox_id, - config.code_server_enabled, - config.vnc_enabled, - terminal_enabled, - tunnel_ports, - code_server_port, - novnc_port, - ttyd_proxy_port, - ) - - return SandboxHandle( - sandbox_id=sandbox_id, - modal_sandbox=sandbox, - status=SandboxStatus.WARMING, - created_at=time.time(), - snapshot_id=snapshot_id, - modal_object_id=modal_object_id, - code_server_url=code_server_url, - code_server_password=code_server_password, - vnc_url=vnc_url, - vnc_password=vnc_password, - ttyd_url=ttyd_url, - tunnel_urls=extra_tunnel_urls, - sandbox_backend=docker.backend, - ) - - async def _launch_docker_sandbox( - self, - *, - session_id: str, - sandbox_id: str, - retire_sandbox_id: str | None, - create_kwargs: dict[str, Any], - repository_image: bool, - launch_deadline_at_ms: int | None = None, - ) -> tuple[modal.Sandbox, bool]: - """Create a Docker VM under a deterministic name, adopting an existing one. - - VM creation can outlive the control plane's HTTP request. One name per - session serializes generations at Modal even when a predecessor lookup - misses an in-flight create. Only matching generation tags permit adoption. - """ - if retire_sandbox_id: - await self._retire_docker_allocation(session_id, retire_sandbox_id) - name = docker_allocation_name(session_id) - tags = docker_allocation_tags(session_id, sandbox_id) - existing = await self._find_owned_docker_allocation(name, tags) - if existing is None: - if launch_deadline_at_ms is not None and time.time() * 1000 >= launch_deadline_at_ms: - raise RuntimeError("VM launch deadline expired") - try: - sandbox = await _create_sandbox( - {**create_kwargs, "name": name, "tags": tags}, - repository_image=repository_image, - ) - return sandbox, False - except modal.exception.AlreadyExistsError: - existing = await self._find_owned_docker_allocation(name, tags) - if existing is None: - raise - log.info( - "sandbox.docker_allocation_adopted", - sandbox_id=sandbox_id, - modal_object_id=existing.object_id, - ) - return existing, True - - @staticmethod - async def _read_access_passwords( - sandbox: modal.Sandbox, *, code_server_enabled: bool, vnc_enabled: bool - ) -> dict[str, str]: - """Recover only enabled service credentials from the owned VM's launch environment.""" - keys = [] - if code_server_enabled: - keys.append("CODE_SERVER_PASSWORD") - if vnc_enabled: - keys.append(VNC_PASSWORD_ENV_VAR) - if not keys: - return {} - process = await sandbox.exec.aio( - "python", - "-I", - "-c", - "import json, os, sys; print(json.dumps({k: os.environ.get(k) for k in sys.argv[1:]}))", - *keys, - timeout=ACCESS_PASSWORD_READ_TIMEOUT_SECONDS, - ) - output = await process.stdout.read.aio() - if await process.wait.aio() != 0: - raise RuntimeError("Could not recover adopted sandbox access credentials") - try: - passwords = json.loads(output) - except ValueError: - raise RuntimeError("Could not recover adopted sandbox access credentials") from None - if not isinstance(passwords, dict) or any( - not isinstance(passwords.get(key), str) or not passwords[key] for key in keys - ): - raise RuntimeError("Could not recover adopted sandbox access credentials") - return {key: passwords[key] for key in keys} - - @staticmethod - async def _find_owned_docker_allocation( - name: str, tags: dict[str, str] - ) -> modal.Sandbox | None: - try: - sandbox = await modal.Sandbox.from_name.aio(APP_NAME, name) - except modal.exception.NotFoundError: - return None - if await sandbox.get_tags.aio() != tags: - raise RuntimeError("Docker sandbox allocation ownership mismatch") - return sandbox - - async def _retire_docker_allocation(self, session_id: str, sandbox_id: str) -> None: - """Terminate a prior generation's named VM, only when its ownership tags match.""" - name = docker_allocation_name(session_id) - try: - sandbox = await modal.Sandbox.from_name.aio(APP_NAME, name) - except modal.exception.NotFoundError: - return - if await sandbox.get_tags.aio() != docker_allocation_tags(session_id, sandbox_id): - log.warn("sandbox.docker_allocation_retire_mismatch", sandbox_id=sandbox_id) - return - await sandbox.terminate.aio(wait=True) - log.info( - "sandbox.docker_allocation_retired", - sandbox_id=sandbox_id, - modal_object_id=sandbox.object_id, - ) - async def create_sandbox( self, config: SandboxConfig, @@ -667,7 +110,7 @@ async def create_sandbox( Creates from the pre-built repo image when one is provided, otherwise from the base image. Snapshot restores go through - restore_sandbox, not this path. + restore_from_snapshot, not this path. Args: config: Sandbox configuration including repo info and session config @@ -679,14 +122,14 @@ async def create_sandbox( _has_repository(config.repo_owner, config.repo_name) if config.repo_image_id: - source: _SandboxImageSource = _RepositoryImageSource( + source: SandboxImageSource = RepositoryImageSource( image_id=config.repo_image_id, sha=config.repo_image_sha, ) else: - source = _BaseImageSource() + source = BaseImageSource() - handle = await self._launch_sandbox(_SandboxLaunchSpec(config=config, source=source)) + handle = await SandboxLauncher().launch(SandboxLaunchSpec(config=config, source=source)) duration_ms = int((time.time() - start_time) * 1000) log.info( @@ -800,7 +243,9 @@ async def get_sandbox_by_id(self, sandbox_id: str) -> SandboxHandle | None: APP_NAME, docker_allocation_name(identity[0]) ) except modal.exception.NotFoundError: - raise PendingVMReferenceNotVisible("VM launch identity is not yet visible") + raise PendingVMReferenceNotVisible( + "VM launch identity is not yet visible" + ) from None else: try: modal_sandbox = await modal.Sandbox.from_id.aio(sandbox_id) @@ -817,7 +262,7 @@ async def get_sandbox_by_id(self, sandbox_id: str) -> SandboxHandle | None: sandbox_id=sandbox_id, modal_object_id=modal_sandbox.object_id, modal_sandbox=modal_sandbox, - status=SandboxStatus.READY, # Assume ready if we can retrieve it + status=SandboxStatus.READY, created_at=time.time(), ) @@ -874,8 +319,8 @@ async def restore_from_snapshot( # so the gh CLI keeps working on snapshots predating the gh wrapper. # Host scoping remains common with fresh creates. These compatibility # credentials are explicitly requested only by the restore path. - handle = await self._launch_sandbox( - _SandboxLaunchSpec( + handle = await SandboxLauncher().launch( + SandboxLaunchSpec( config=SandboxConfig( repo_owner=repo_owner, repo_name=repo_name, @@ -888,12 +333,12 @@ async def restore_from_snapshot( code_server_enabled=code_server_enabled, vnc_enabled=vnc_enabled, agent_slack_notify_enabled=agent_slack_notify_enabled, - retire_sandbox_id=retire_sandbox_id, settings=settings, + retire_sandbox_id=retire_sandbox_id, sandbox_backend=sandbox_backend, launch_deadline_at_ms=launch_deadline_at_ms, ), - source=_SnapshotImageSource( + source=SnapshotImageSource( image_id=snapshot_image_id, clone_token=clone_token, ), diff --git a/packages/modal-infra/src/sandbox/models.py b/packages/modal-infra/src/sandbox/models.py new file mode 100644 index 0000000000..61125ecaa1 --- /dev/null +++ b/packages/modal-infra/src/sandbox/models.py @@ -0,0 +1,59 @@ +"""Configuration and handles shared by Modal sandbox operations.""" + +from dataclasses import dataclass +from typing import Any + +import modal + +from sandbox_runtime.constants import DEFAULT_SANDBOX_TIMEOUT_SECONDS +from sandbox_runtime.types import SandboxStatus, SessionConfig + +from .launch_policy import ModalBackend + +DEFAULT_VNC_ENABLED = False + + +@dataclass +class SandboxConfig: + """Configuration for creating a sandbox.""" + + repo_owner: str | None + repo_name: str | None + sandbox_id: str | None = None # Expected sandbox ID from control plane + session_config: SessionConfig | dict[str, Any] | None = None + control_plane_url: str = "" + sandbox_auth_token: str = "" + timeout_seconds: int = DEFAULT_SANDBOX_TIMEOUT_SECONDS + user_env_vars: dict[str, str] | None = None # User-provided env vars (repo secrets) + repo_image_id: str | None = None # Pre-built repo image ID from provider + repo_image_sha: str | None = None # Git SHA the repo image was built from + code_server_enabled: bool = False # Whether to start code-server in the sandbox + vnc_enabled: bool = DEFAULT_VNC_ENABLED # Whether to start the browser-accessible VNC desktop + agent_slack_notify_enabled: bool = ( + False # Whether to install the agent-initiated slack-notify tool + ) + settings: dict[str, Any] | None = ( + None # Sandbox settings (tunnelPorts, etc.) from control plane + ) + sandbox_backend: ModalBackend = "modal" + retire_sandbox_id: str | None = None + launch_deadline_at_ms: int | None = None + + +@dataclass +class SandboxHandle: + """Handle to a sandbox.""" + + sandbox_id: str + modal_sandbox: modal.Sandbox + status: SandboxStatus + created_at: float + snapshot_id: str | None = None + modal_object_id: str | None = None # Modal's internal sandbox ID for API calls + code_server_url: str | None = None + code_server_password: str | None = None + vnc_url: str | None = None + vnc_password: str | None = None + ttyd_url: str | None = None # proxy tunnel URL (not ttyd directly) + tunnel_urls: dict[int, str] | None = None # port -> tunnel URL mapping for extra ports + sandbox_backend: ModalBackend = "modal" diff --git a/packages/modal-infra/src/sandbox/tunnels.py b/packages/modal-infra/src/sandbox/tunnels.py new file mode 100644 index 0000000000..17210838aa --- /dev/null +++ b/packages/modal-infra/src/sandbox/tunnels.py @@ -0,0 +1,215 @@ +"""Service port ownership and best-effort Modal tunnel publication.""" + +import asyncio +from typing import Any, NamedTuple + +import modal + +from sandbox_runtime.constants import ( + CODE_SERVER_PORT, + CODE_SERVER_PORT_ENV_VAR, + EXPECTED_TUNNEL_PORTS_ENV_VAR, + NOVNC_PORT, + NOVNC_PORT_ENV_VAR, + TTYD_PROXY_PORT, + TTYD_PROXY_PORT_ENV_VAR, + TUNNEL_ENV_FILE_PATH, + TUNNEL_ENV_SANDBOX_ID_KEY, + VNC_PORT, +) +from sandbox_runtime.log_config import get_logger + +# Preserve the logger name used by existing launch/tunnel dashboards. +log = get_logger("manager") +MAX_TUNNEL_PORTS = 10 +DEFAULT_TUNNEL_RESOLUTION_RETRIES = 3 +DEFAULT_TUNNEL_RESOLUTION_BACKOFF_SECONDS = 1.0 + + +class TunnelUrls(NamedTuple): + """Resolved service URLs and any user-requested tunnels.""" + + code_server_url: str | None = None + vnc_url: str | None = None + ttyd_url: str | None = None + tunnel_urls: dict[int, str] | None = None + + +class SandboxTunnels: + """Keep exposed ports, runtime environment, and URL routing in agreement. + + Service ownership is resolved once for a launch. Disabled service ports + remain available as user tunnels; raw VNC is never an extra tunnel. + """ + + def __init__( + self, + *, + code_server_enabled: bool = False, + vnc_enabled: bool = False, + settings: dict[str, Any] | None = None, + ) -> None: + settings = settings or {} + code_server_port, novnc_port, ttyd_proxy_port = self._resolve_service_ports(settings) + self._code_server_port = code_server_port if code_server_enabled else None + self._novnc_port = novnc_port if vnc_enabled else None + self._ttyd_proxy_port = ( + ttyd_proxy_port if bool(settings.get("terminalEnabled", False)) else None + ) + service_ports = [ + port + for port in (self._code_server_port, self._novnc_port, self._ttyd_proxy_port) + if port is not None + ] + reserved = {VNC_PORT, *service_ports} + raw_ports = settings.get("tunnelPorts", []) + self._extra_ports = ( + [port for port in self._validate_ports(raw_ports) if port not in reserved] + if raw_ports + else [] + ) + self.exposed_ports = service_ports + self._extra_ports + + @property + def environment(self) -> dict[str, str]: + """Runtime settings derived from the same ports used for exposure.""" + env: dict[str, str] = {} + if self._code_server_port is not None: + env[CODE_SERVER_PORT_ENV_VAR] = str(self._code_server_port) + if self._novnc_port is not None: + env[NOVNC_PORT_ENV_VAR] = str(self._novnc_port) + if self._ttyd_proxy_port is not None: + env["TERMINAL_ENABLED"] = "true" + env[TTYD_PROXY_PORT_ENV_VAR] = str(self._ttyd_proxy_port) + if self._extra_ports: + env[EXPECTED_TUNNEL_PORTS_ENV_VAR] = ",".join(str(p) for p in self._extra_ports) + return env + + async def resolve(self, sandbox: modal.Sandbox, sandbox_id: str) -> TunnelUrls: + """Resolve URLs and publish extras; partial resolution/write failures are non-fatal.""" + if not self.exposed_ports: + return TunnelUrls() + + resolved = await self._resolve_tunnels(sandbox, sandbox_id, self.exposed_ports) + # A disabled service does not own its default port: leave it in extras. + code_server_url = ( + resolved.pop(self._code_server_port, None) + if self._code_server_port is not None + else None + ) + vnc_url = resolved.pop(self._novnc_port, None) if self._novnc_port is not None else None + ttyd_url = ( + resolved.pop(self._ttyd_proxy_port, None) if self._ttyd_proxy_port is not None else None + ) + extra_urls = resolved or None + if extra_urls: + await self._write_tunnel_env_file(sandbox, sandbox_id, extra_urls) + return TunnelUrls( + code_server_url=code_server_url, + vnc_url=vnc_url, + ttyd_url=ttyd_url, + tunnel_urls=extra_urls, + ) + + @staticmethod + async def _resolve_tunnels( + sandbox: modal.Sandbox, + sandbox_id: str, + ports: list[int], + retries: int = DEFAULT_TUNNEL_RESOLUTION_RETRIES, + backoff_seconds: float = DEFAULT_TUNNEL_RESOLUTION_BACKOFF_SECONDS, + ) -> dict[int, str]: + """Resolve tunnel URLs for the given ports from Modal, retrying on failure.""" + resolved: dict[int, str] = {} + for attempt in range(retries): + try: + loop = asyncio.get_running_loop() + tunnels = await loop.run_in_executor(None, sandbox.tunnels) + for port in ports: + if port in tunnels and port not in resolved: + resolved[port] = tunnels[port].url + log.info( + "tunnel.resolved", + sandbox_id=sandbox_id, + port=port, + url=tunnels[port].url, + ) + if len(resolved) == len(ports): + return resolved + except Exception as e: + log.warn( + "tunnel.resolve_error", + sandbox_id=sandbox_id, + attempt=attempt + 1, + retries=retries, + error=type(e).__name__, + exc=e, + ) + if attempt < retries - 1: + await asyncio.sleep(backoff_seconds * (attempt + 1)) + return resolved + + @staticmethod + def _validate_ports(raw: list[Any]) -> list[int]: + """Validate and sanitize tunnel ports: must be int, 1-65535, max MAX_TUNNEL_PORTS.""" + ports: list[int] = [] + for p in raw: + if isinstance(p, int) and not isinstance(p, bool) and 1 <= p <= 65535: + ports.append(p) + if len(ports) >= MAX_TUNNEL_PORTS: + break + return ports + + @staticmethod + def _resolve_service_ports(settings: dict[str, Any] | None) -> tuple[int, int, int]: + """Return effective (code_server_port, novnc_port, ttyd_proxy_port) from settings. + + Falls back to the service defaults when unset or invalid. The control + plane validates these before they reach here. + """ + s = settings or {} + + def coerce(value: Any, default: int) -> int: + if isinstance(value, int) and not isinstance(value, bool) and 1 <= value <= 65535: + return value + return default + + return ( + coerce(s.get("codeServerPort"), CODE_SERVER_PORT), + coerce(s.get("vncPort"), NOVNC_PORT), + coerce(s.get("terminalPort"), TTYD_PROXY_PORT), + ) + + @staticmethod + async def _write_tunnel_env_file( + sandbox: modal.Sandbox, + sandbox_id: str, + tunnel_urls: dict[int, str], + ) -> None: + """Write tunnel URLs to TUNNEL_ENV_FILE_PATH as a dotenv file. + + The first line tags the file with this sandbox's ID so the supervisor's + stale-file cleanup can tell a fresh write (this write can land before + the entrypoint runs) from a snapshot/image leftover. + + Failures are logged but do not block sandbox creation; URLs are also + returned to the control plane via the SandboxHandle. + """ + lines = [f"{TUNNEL_ENV_SANDBOX_ID_KEY}={sandbox_id}"] + lines += [f"TUNNEL_{port}={url}" for port, url in sorted(tunnel_urls.items())] + content = "\n".join(lines) + "\n" + try: + await sandbox.filesystem.write_text.aio(content, TUNNEL_ENV_FILE_PATH) + log.info( + "tunnel.urls_written", + sandbox_id=sandbox_id, + path=TUNNEL_ENV_FILE_PATH, + ports=list(tunnel_urls.keys()), + ) + except Exception as e: + log.warn( + "tunnel.urls_write_failed", + sandbox_id=sandbox_id, + path=TUNNEL_ENV_FILE_PATH, + exc=e, + ) diff --git a/packages/modal-infra/tests/test_agent_slack_notify_env.py b/packages/modal-infra/tests/test_agent_slack_notify_env.py index abc20cce92..ede0a3c971 100644 --- a/packages/modal-infra/tests/test_agent_slack_notify_env.py +++ b/packages/modal-infra/tests/test_agent_slack_notify_env.py @@ -5,6 +5,7 @@ import pytest from src.sandbox.manager import SandboxConfig, SandboxManager +from src.sandbox.tunnels import SandboxTunnels, TunnelUrls def _patch_create(monkeypatch, captured: dict) -> None: @@ -21,11 +22,11 @@ class FakeSandbox: fake_create = MagicMock() fake_create.aio = fake_create_aio - monkeypatch.setattr("src.sandbox.manager.modal.Sandbox.create", fake_create) + monkeypatch.setattr("src.sandbox.launch.modal.Sandbox.create", fake_create) monkeypatch.setattr( - SandboxManager, - "_resolve_and_setup_tunnels", - AsyncMock(return_value=(None, None, None, None)), + SandboxTunnels, + "resolve", + AsyncMock(return_value=TunnelUrls(None, None, None, None)), ) @@ -79,7 +80,7 @@ async def test_env_set_when_enabled(self, monkeypatch): class FakeImage: object_id = "img-123" - monkeypatch.setattr("src.sandbox.manager.modal.Image.from_id", lambda *a, **k: FakeImage()) + monkeypatch.setattr("src.sandbox.launch.modal.Image.from_id", lambda *a, **k: FakeImage()) _patch_create(monkeypatch, captured) manager = SandboxManager() diff --git a/packages/modal-infra/tests/test_code_server.py b/packages/modal-infra/tests/test_code_server.py index f950e2a31b..302d8c7675 100644 --- a/packages/modal-infra/tests/test_code_server.py +++ b/packages/modal-infra/tests/test_code_server.py @@ -4,25 +4,27 @@ import pytest +from src.sandbox.launch import SandboxLauncher from src.sandbox.manager import CODE_SERVER_PORT, SandboxConfig, SandboxManager +from src.sandbox.tunnels import SandboxTunnels, TunnelUrls class TestGenerateCodeServerPassword: - """SandboxManager._generate_code_server_password tests.""" + """SandboxLauncher._generate_code_server_password tests.""" def test_returns_nonempty_password(self): - password = SandboxManager._generate_code_server_password() + password = SandboxLauncher._generate_code_server_password() assert len(password) > 0 def test_generates_unique_passwords(self): passwords = set() for _ in range(20): - passwords.add(SandboxManager._generate_code_server_password()) + passwords.add(SandboxLauncher._generate_code_server_password()) assert len(passwords) == 20 class TestResolveCodeServerTunnel: - """SandboxManager._resolve_tunnels tests for code-server port.""" + """SandboxTunnels._resolve_tunnels tests for code-server port.""" @pytest.mark.asyncio async def test_returns_tunnel_url_on_success(self): @@ -32,7 +34,7 @@ async def test_returns_tunnel_url_on_success(self): sandbox = MagicMock() sandbox.tunnels.return_value = {CODE_SERVER_PORT: tunnel} - resolved = await SandboxManager._resolve_tunnels(sandbox, "sb-123", [CODE_SERVER_PORT]) + resolved = await SandboxTunnels._resolve_tunnels(sandbox, "sb-123", [CODE_SERVER_PORT]) assert resolved.get(CODE_SERVER_PORT) == "https://tunnel.example.com" @pytest.mark.asyncio @@ -40,9 +42,9 @@ async def test_returns_empty_on_exception_after_retries(self): sandbox = MagicMock() sandbox.tunnels.side_effect = Exception("tunnel unavailable") - with patch("src.sandbox.manager.asyncio.sleep", new_callable=AsyncMock): - resolved = await SandboxManager._resolve_tunnels( - sandbox, "sb-123", [CODE_SERVER_PORT], retries=2, backoff=0.0 + with patch("src.sandbox.tunnels.asyncio.sleep", new_callable=AsyncMock): + resolved = await SandboxTunnels._resolve_tunnels( + sandbox, "sb-123", [CODE_SERVER_PORT], retries=2, backoff_seconds=0.0 ) assert resolved == {} assert sandbox.tunnels.call_count == 2 @@ -52,9 +54,9 @@ async def test_returns_empty_when_port_missing_after_retries(self): sandbox = MagicMock() sandbox.tunnels.return_value = {} # no entry for CODE_SERVER_PORT - with patch("src.sandbox.manager.asyncio.sleep", new_callable=AsyncMock): - resolved = await SandboxManager._resolve_tunnels( - sandbox, "sb-123", [CODE_SERVER_PORT], retries=2, backoff=0.0 + with patch("src.sandbox.tunnels.asyncio.sleep", new_callable=AsyncMock): + resolved = await SandboxTunnels._resolve_tunnels( + sandbox, "sb-123", [CODE_SERVER_PORT], retries=2, backoff_seconds=0.0 ) assert resolved == {} @@ -69,9 +71,9 @@ async def test_retries_then_succeeds(self): {CODE_SERVER_PORT: tunnel}, ] - with patch("src.sandbox.manager.asyncio.sleep", new_callable=AsyncMock): - resolved = await SandboxManager._resolve_tunnels( - sandbox, "sb-123", [CODE_SERVER_PORT], retries=3, backoff=0.0 + with patch("src.sandbox.tunnels.asyncio.sleep", new_callable=AsyncMock): + resolved = await SandboxTunnels._resolve_tunnels( + sandbox, "sb-123", [CODE_SERVER_PORT], retries=3, backoff_seconds=0.0 ) assert resolved.get(CODE_SERVER_PORT) == "https://tunnel.example.com" assert sandbox.tunnels.call_count == 2 @@ -96,12 +98,12 @@ class FakeSandbox: fake_create = MagicMock() fake_create.aio = fake_create_aio - monkeypatch.setattr("src.sandbox.manager.modal.Sandbox.create", fake_create) + monkeypatch.setattr("src.sandbox.launch.modal.Sandbox.create", fake_create) monkeypatch.setattr( - SandboxManager, - "_resolve_and_setup_tunnels", - AsyncMock(return_value=("https://cs.example.com", None, None, None)), + SandboxTunnels, + "resolve", + AsyncMock(return_value=TunnelUrls("https://cs.example.com", None, None, None)), ) manager = SandboxManager() @@ -140,10 +142,10 @@ class FakeSandbox: fake_create = MagicMock() fake_create.aio = fake_create_aio - monkeypatch.setattr("src.sandbox.manager.modal.Sandbox.create", fake_create) + monkeypatch.setattr("src.sandbox.launch.modal.Sandbox.create", fake_create) - tunnel_mock = AsyncMock(return_value=(None, None, None, None)) - monkeypatch.setattr(SandboxManager, "_resolve_and_setup_tunnels", tunnel_mock) + tunnel_mock = AsyncMock(return_value=TunnelUrls(None, None, None, None)) + monkeypatch.setattr(SandboxTunnels, "resolve", tunnel_mock) manager = SandboxManager() config = SandboxConfig( @@ -187,12 +189,12 @@ class FakeSandbox: fake_create = MagicMock() fake_create.aio = fake_create_aio - monkeypatch.setattr("src.sandbox.manager.modal.Image.from_id", fake_from_id) - monkeypatch.setattr("src.sandbox.manager.modal.Sandbox.create", fake_create) + monkeypatch.setattr("src.sandbox.launch.modal.Image.from_id", fake_from_id) + monkeypatch.setattr("src.sandbox.launch.modal.Sandbox.create", fake_create) monkeypatch.setattr( - SandboxManager, - "_resolve_and_setup_tunnels", - AsyncMock(return_value=("https://cs-restored.example.com", None, None, None)), + SandboxTunnels, + "resolve", + AsyncMock(return_value=TunnelUrls("https://cs-restored.example.com", None, None, None)), ) manager = SandboxManager() @@ -238,10 +240,10 @@ class FakeSandbox: fake_create = MagicMock() fake_create.aio = fake_create_aio - monkeypatch.setattr("src.sandbox.manager.modal.Image.from_id", fake_from_id) - monkeypatch.setattr("src.sandbox.manager.modal.Sandbox.create", fake_create) - tunnel_mock = AsyncMock(return_value=(None, None, None, None)) - monkeypatch.setattr(SandboxManager, "_resolve_and_setup_tunnels", tunnel_mock) + monkeypatch.setattr("src.sandbox.launch.modal.Image.from_id", fake_from_id) + monkeypatch.setattr("src.sandbox.launch.modal.Sandbox.create", fake_create) + tunnel_mock = AsyncMock(return_value=TunnelUrls(None, None, None, None)) + monkeypatch.setattr(SandboxTunnels, "resolve", tunnel_mock) manager = SandboxManager() handle = await manager.restore_from_snapshot( diff --git a/packages/modal-infra/tests/test_llm_secrets.py b/packages/modal-infra/tests/test_llm_secrets.py index daf8935755..a6448ecc7b 100644 --- a/packages/modal-infra/tests/test_llm_secrets.py +++ b/packages/modal-infra/tests/test_llm_secrets.py @@ -26,7 +26,7 @@ class FakeSandbox: return FakeSandbox() fake_create_aio.aio = fake_create_aio - monkeypatch.setattr("src.sandbox.manager.modal.Sandbox.create", fake_create_aio) + monkeypatch.setattr("src.sandbox.launch.modal.Sandbox.create", fake_create_aio) return captured @@ -44,7 +44,7 @@ async def test_restore_attaches_the_deployment_wide_secret( class FakeImage: object_id = "img-llm-secrets" - monkeypatch.setattr("src.sandbox.manager.modal.Image.from_id", lambda *a, **k: FakeImage()) + monkeypatch.setattr("src.sandbox.launch.modal.Image.from_id", lambda *a, **k: FakeImage()) await SandboxManager().restore_from_snapshot( snapshot_image_id="img-abc", diff --git a/packages/modal-infra/tests/test_manager_exports.py b/packages/modal-infra/tests/test_manager_exports.py new file mode 100644 index 0000000000..cddf8c6564 --- /dev/null +++ b/packages/modal-infra/tests/test_manager_exports.py @@ -0,0 +1,45 @@ +"""Compatibility coverage for the manager's pre-refactor constant imports.""" + +import pytest + +from sandbox_runtime import constants +from sandbox_runtime.docker_control import CONTROL_TIMEOUT_SECONDS +from src.app_config import APP_NAME +from src.sandbox import manager +from src.sandbox.launch import ACCESS_PASSWORD_READ_TIMEOUT_SECONDS +from src.sandbox.launch_policy import PENDING_VM_REFERENCE_PREFIX +from src.sandbox.models import DEFAULT_VNC_ENABLED +from src.sandbox.tunnels import MAX_TUNNEL_PORTS + + +@pytest.mark.parametrize( + "name, expected", + [ + ("ACCESS_PASSWORD_READ_TIMEOUT_SECONDS", ACCESS_PASSWORD_READ_TIMEOUT_SECONDS), + ("APP_NAME", APP_NAME), + ("CODE_SERVER_PORT", constants.CODE_SERVER_PORT), + ("CODE_SERVER_PORT_ENV_VAR", constants.CODE_SERVER_PORT_ENV_VAR), + ("DEFAULT_SANDBOX_TIMEOUT_SECONDS", constants.DEFAULT_SANDBOX_TIMEOUT_SECONDS), + ("DEFAULT_VNC_ENABLED", DEFAULT_VNC_ENABLED), + ("DOCKER_ENABLED_ENV_VAR", constants.DOCKER_ENABLED_ENV_VAR), + ("EXPECTED_TUNNEL_PORTS_ENV_VAR", constants.EXPECTED_TUNNEL_PORTS_ENV_VAR), + ("MAX_TUNNEL_PORTS", MAX_TUNNEL_PORTS), + ("NOVNC_PORT", constants.NOVNC_PORT), + ("NOVNC_PORT_ENV_VAR", constants.NOVNC_PORT_ENV_VAR), + ("PENDING_VM_REFERENCE_PREFIX", PENDING_VM_REFERENCE_PREFIX), + ("SANDBOX_TIMEOUT_ENV_VAR", constants.SANDBOX_TIMEOUT_ENV_VAR), + ("SNAPSHOT_FILESYSTEM_TIMEOUT_SECONDS", 300), + ("CONTROL_TIMEOUT_SECONDS", CONTROL_TIMEOUT_SECONDS), + ("TTYD_PROXY_PORT", constants.TTYD_PROXY_PORT), + ("TTYD_PROXY_PORT_ENV_VAR", constants.TTYD_PROXY_PORT_ENV_VAR), + ("TUNNEL_ENV_FILE_PATH", constants.TUNNEL_ENV_FILE_PATH), + ("TUNNEL_ENV_SANDBOX_ID_KEY", constants.TUNNEL_ENV_SANDBOX_ID_KEY), + ("VNC_PASSWORD_ENV_VAR", constants.VNC_PASSWORD_ENV_VAR), + ("VNC_PASSWORD_MAX_BYTES", constants.VNC_PASSWORD_MAX_BYTES), + ("VNC_PORT", constants.VNC_PORT), + ], +) +def test_legacy_manager_constant_exports(name, expected): + """Legacy import names retain their values and are explicitly public.""" + assert getattr(manager, name) == expected + assert name in manager.__all__ diff --git a/packages/modal-infra/tests/test_sandbox_env_vars.py b/packages/modal-infra/tests/test_sandbox_env_vars.py index 131ff0ef1f..be4130a938 100644 --- a/packages/modal-infra/tests/test_sandbox_env_vars.py +++ b/packages/modal-infra/tests/test_sandbox_env_vars.py @@ -8,6 +8,7 @@ VNC_PASSWORD_MAX_BYTES, ) from sandbox_runtime.types import SessionConfig +from src.sandbox.launch import SandboxLauncher from src.sandbox.manager import ( DEFAULT_SANDBOX_TIMEOUT_SECONDS, SandboxConfig, @@ -86,7 +87,7 @@ class FakeSandbox: return FakeSandbox() fake_create_aio.aio = fake_create_aio - monkeypatch.setattr("src.sandbox.manager.modal.Sandbox.create", fake_create_aio) + monkeypatch.setattr("src.sandbox.launch.modal.Sandbox.create", fake_create_aio) manager = SandboxManager() config = SandboxConfig( @@ -132,8 +133,8 @@ class FakeSandbox: return FakeSandbox() fake_create_aio.aio = fake_create_aio - monkeypatch.setattr("src.sandbox.manager.modal.Image.from_id", fake_from_id) - monkeypatch.setattr("src.sandbox.manager.modal.Sandbox.create", fake_create_aio) + monkeypatch.setattr("src.sandbox.launch.modal.Image.from_id", fake_from_id) + monkeypatch.setattr("src.sandbox.launch.modal.Sandbox.create", fake_create_aio) manager = SandboxManager() await manager.restore_from_snapshot( @@ -179,7 +180,7 @@ async def test_create_preserves_managed_provider_env_isolation( monkeypatch, managed_marker, suppressed_api_key ): captured = {} - monkeypatch.setattr("src.sandbox.manager.modal.Sandbox.create", _fake_sandbox_create(captured)) + monkeypatch.setattr("src.sandbox.launch.modal.Sandbox.create", _fake_sandbox_create(captured)) await SandboxManager().create_sandbox( SandboxConfig( @@ -217,7 +218,7 @@ async def test_restore_preserves_managed_provider_env_isolation( def test_generated_vnc_password_respects_protocol_limit(): - assert len(SandboxManager._generate_vnc_password().encode()) == VNC_PASSWORD_MAX_BYTES + assert len(SandboxLauncher._generate_vnc_password().encode()) == VNC_PASSWORD_MAX_BYTES @pytest.mark.asyncio @@ -241,8 +242,8 @@ class FakeSandbox: return FakeSandbox() fake_create_aio.aio = fake_create_aio - monkeypatch.setattr("src.sandbox.manager.modal.Image.from_id", fake_from_id) - monkeypatch.setattr("src.sandbox.manager.modal.Sandbox.create", fake_create_aio) + monkeypatch.setattr("src.sandbox.launch.modal.Image.from_id", fake_from_id) + monkeypatch.setattr("src.sandbox.launch.modal.Sandbox.create", fake_create_aio) manager = SandboxManager() await manager.restore_from_snapshot( @@ -281,8 +282,8 @@ class FakeSandbox: return FakeSandbox() fake_create_aio.aio = fake_create_aio - monkeypatch.setattr("src.sandbox.manager.modal.Image.from_id", fake_from_id) - monkeypatch.setattr("src.sandbox.manager.modal.Sandbox.create", fake_create_aio) + monkeypatch.setattr("src.sandbox.launch.modal.Image.from_id", fake_from_id) + monkeypatch.setattr("src.sandbox.launch.modal.Sandbox.create", fake_create_aio) manager = SandboxManager() await manager.restore_from_snapshot( @@ -327,8 +328,8 @@ class FakeSandbox: return FakeSandbox() fake_create_aio.aio = fake_create_aio - monkeypatch.setattr("src.sandbox.manager.modal.Image.from_id", fake_from_id) - monkeypatch.setattr("src.sandbox.manager.modal.Sandbox.create", fake_create_aio) + monkeypatch.setattr("src.sandbox.launch.modal.Image.from_id", fake_from_id) + monkeypatch.setattr("src.sandbox.launch.modal.Sandbox.create", fake_create_aio) manager = SandboxManager() @@ -369,8 +370,8 @@ def _fake_restore_setup(monkeypatch): class FakeImage: object_id = "img-123" - monkeypatch.setattr("src.sandbox.manager.modal.Image.from_id", lambda *a, **kw: FakeImage()) - monkeypatch.setattr("src.sandbox.manager.modal.Sandbox.create", _fake_sandbox_create(captured)) + monkeypatch.setattr("src.sandbox.launch.modal.Image.from_id", lambda *a, **kw: FakeImage()) + monkeypatch.setattr("src.sandbox.launch.modal.Sandbox.create", _fake_sandbox_create(captured)) return captured @@ -466,7 +467,7 @@ class FakeSandbox: async def test_vcs_env_vars_default_github(monkeypatch): """SCM_PROVIDER unset → github.com defaults, no token in env.""" captured = {} - monkeypatch.setattr("src.sandbox.manager.modal.Sandbox.create", _fake_sandbox_create(captured)) + monkeypatch.setattr("src.sandbox.launch.modal.Sandbox.create", _fake_sandbox_create(captured)) monkeypatch.delenv("SCM_PROVIDER", raising=False) manager = SandboxManager() @@ -488,7 +489,7 @@ async def test_vcs_env_vars_default_github(monkeypatch): async def test_vcs_env_vars_gitlab(monkeypatch): """SCM_PROVIDER=gitlab → gitlab.com + oauth2, no token in env.""" captured = {} - monkeypatch.setattr("src.sandbox.manager.modal.Sandbox.create", _fake_sandbox_create(captured)) + monkeypatch.setattr("src.sandbox.launch.modal.Sandbox.create", _fake_sandbox_create(captured)) monkeypatch.setenv("SCM_PROVIDER", "gitlab") manager = SandboxManager() @@ -508,7 +509,7 @@ async def test_vcs_env_vars_gitlab(monkeypatch): async def test_vcs_env_vars_bitbucket(monkeypatch): """SCM_PROVIDER=bitbucket → bitbucket.org + x-token-auth, no token in env.""" captured = {} - monkeypatch.setattr("src.sandbox.manager.modal.Sandbox.create", _fake_sandbox_create(captured)) + monkeypatch.setattr("src.sandbox.launch.modal.Sandbox.create", _fake_sandbox_create(captured)) monkeypatch.setenv("SCM_PROVIDER", "bitbucket") manager = SandboxManager() @@ -532,8 +533,8 @@ async def test_repo_image_boot_omits_fallback_tokens(monkeypatch): class FakeImage: object_id = "repo-img-1" - monkeypatch.setattr("src.sandbox.manager.modal.Image.from_id", lambda *a, **kw: FakeImage()) - monkeypatch.setattr("src.sandbox.manager.modal.Sandbox.create", _fake_sandbox_create(captured)) + monkeypatch.setattr("src.sandbox.launch.modal.Image.from_id", lambda *a, **kw: FakeImage()) + monkeypatch.setattr("src.sandbox.launch.modal.Sandbox.create", _fake_sandbox_create(captured)) monkeypatch.delenv("SCM_PROVIDER", raising=False) manager = SandboxManager() @@ -561,8 +562,8 @@ async def test_repo_image_boot_preserves_user_github_cli_token(monkeypatch, toke class FakeImage: object_id = "repo-img-1" - monkeypatch.setattr("src.sandbox.manager.modal.Image.from_id", lambda *a, **kw: FakeImage()) - monkeypatch.setattr("src.sandbox.manager.modal.Sandbox.create", _fake_sandbox_create(captured)) + monkeypatch.setattr("src.sandbox.launch.modal.Image.from_id", lambda *a, **kw: FakeImage()) + monkeypatch.setattr("src.sandbox.launch.modal.Sandbox.create", _fake_sandbox_create(captured)) monkeypatch.delenv("SCM_PROVIDER", raising=False) manager = SandboxManager() @@ -591,7 +592,7 @@ async def test_no_repo_sandbox_gets_provider_host_scoping(monkeypatch): fall back to github.com credential-helper behavior. """ captured = {} - monkeypatch.setattr("src.sandbox.manager.modal.Sandbox.create", _fake_sandbox_create(captured)) + monkeypatch.setattr("src.sandbox.launch.modal.Sandbox.create", _fake_sandbox_create(captured)) monkeypatch.setenv("SCM_PROVIDER", "gitlab") manager = SandboxManager() @@ -611,8 +612,8 @@ async def test_restore_no_repo_gets_host_scoping_without_tokens(monkeypatch): class FakeImage: object_id = "img-123" - monkeypatch.setattr("src.sandbox.manager.modal.Image.from_id", lambda *a, **kw: FakeImage()) - monkeypatch.setattr("src.sandbox.manager.modal.Sandbox.create", _fake_sandbox_create(captured)) + monkeypatch.setattr("src.sandbox.launch.modal.Image.from_id", lambda *a, **kw: FakeImage()) + monkeypatch.setattr("src.sandbox.launch.modal.Sandbox.create", _fake_sandbox_create(captured)) monkeypatch.setenv("SCM_PROVIDER", "bitbucket") manager = SandboxManager() @@ -649,8 +650,8 @@ async def test_restore_preserves_vcs_clone_token_for_legacy_snapshots(monkeypatc class FakeImage: object_id = "img-123" - monkeypatch.setattr("src.sandbox.manager.modal.Image.from_id", lambda *a, **kw: FakeImage()) - monkeypatch.setattr("src.sandbox.manager.modal.Sandbox.create", _fake_sandbox_create(captured)) + monkeypatch.setattr("src.sandbox.launch.modal.Image.from_id", lambda *a, **kw: FakeImage()) + monkeypatch.setattr("src.sandbox.launch.modal.Sandbox.create", _fake_sandbox_create(captured)) monkeypatch.setenv("SCM_PROVIDER", "bitbucket") manager = SandboxManager() @@ -685,8 +686,8 @@ async def test_restore_github_includes_gh_cli_aliases(monkeypatch): class FakeImage: object_id = "img-123" - monkeypatch.setattr("src.sandbox.manager.modal.Image.from_id", lambda *a, **kw: FakeImage()) - monkeypatch.setattr("src.sandbox.manager.modal.Sandbox.create", _fake_sandbox_create(captured)) + monkeypatch.setattr("src.sandbox.launch.modal.Image.from_id", lambda *a, **kw: FakeImage()) + monkeypatch.setattr("src.sandbox.launch.modal.Sandbox.create", _fake_sandbox_create(captured)) monkeypatch.delenv("SCM_PROVIDER", raising=False) manager = SandboxManager() @@ -720,8 +721,8 @@ async def test_no_repo_restore_omits_clone_token(monkeypatch): class FakeImage: object_id = "img-123" - monkeypatch.setattr("src.sandbox.manager.modal.Image.from_id", lambda *a, **kw: FakeImage()) - monkeypatch.setattr("src.sandbox.manager.modal.Sandbox.create", _fake_sandbox_create(captured)) + monkeypatch.setattr("src.sandbox.launch.modal.Image.from_id", lambda *a, **kw: FakeImage()) + monkeypatch.setattr("src.sandbox.launch.modal.Sandbox.create", _fake_sandbox_create(captured)) monkeypatch.delenv("SCM_PROVIDER", raising=False) manager = SandboxManager() diff --git a/packages/modal-infra/tests/test_sandbox_launch.py b/packages/modal-infra/tests/test_sandbox_launch.py index 1fb73b3e32..5831ebbf18 100644 --- a/packages/modal-infra/tests/test_sandbox_launch.py +++ b/packages/modal-infra/tests/test_sandbox_launch.py @@ -6,6 +6,7 @@ from unittest.mock import AsyncMock, Mock import pytest +from modal.exception import NotFoundError from sandbox_runtime.constants import ( CODE_SERVER_PORT_ENV_VAR, @@ -13,9 +14,12 @@ EXPECTED_TUNNEL_PORTS_ENV_VAR, NOVNC_PORT_ENV_VAR, TTYD_PROXY_PORT_ENV_VAR, + TUNNEL_ENV_FILE_PATH, + TUNNEL_ENV_SANDBOX_ID_KEY, VNC_PASSWORD_ENV_VAR, ) -from sandbox_runtime.types import SessionConfig +from sandbox_runtime.types import SandboxStatus, SessionConfig +from src.sandbox.launch import SandboxLauncher from src.sandbox.launch_policy import ( DockerImageUnavailableError, InvalidDockerSettingsError, @@ -27,13 +31,25 @@ SandboxConfig, SandboxManager, ) +from src.sandbox.tunnels import SandboxTunnels, TunnelUrls def _fake_create(captured: dict): async def create_aio(*args, **kwargs): captured["command"] = args captured["kwargs"] = kwargs - return SimpleNamespace(object_id="modal-object-1", stdout=None) + return SimpleNamespace( + object_id="modal-object-1", + tunnels=Mock( + return_value={ + 9000: SimpleNamespace(url="https://code.example"), + 9001: SimpleNamespace(url="https://vnc.example"), + 9002: SimpleNamespace(url="https://terminal.example"), + 3000: SimpleNamespace(url="https://app.example"), + } + ), + filesystem=SimpleNamespace(write_text=SimpleNamespace(aio=AsyncMock())), + ) create_aio.aio = create_aio return create_aio @@ -50,27 +66,14 @@ async def test_launch_matrix_preserves_common_and_source_specific_behavior( "repo-image-1": object(), "snapshot-image-1": object(), } - monkeypatch.setattr("src.sandbox.manager.base_image", base_image) - monkeypatch.setattr("src.sandbox.manager.modal.Image.from_id", images.__getitem__) - monkeypatch.setattr("src.sandbox.manager.modal.Sandbox.create", _fake_create(captured)) + monkeypatch.setattr("src.sandbox.launch.base_image", base_image) + monkeypatch.setattr("src.sandbox.launch.modal.Image.from_id", images.__getitem__) + monkeypatch.setattr("src.sandbox.launch.modal.Sandbox.create", _fake_create(captured)) monkeypatch.delenv("SCM_PROVIDER", raising=False) - resolve_tunnels = AsyncMock( - return_value=( - "https://code.example", - "https://vnc.example", - "https://terminal.example", - {3000: "https://app.example"}, - ) - ) - monkeypatch.setattr( - SandboxManager, - "_resolve_and_setup_tunnels", - resolve_tunnels, - ) monkeypatch.setattr( - SandboxManager, "_generate_code_server_password", staticmethod(lambda: "code-password") + SandboxLauncher, "_generate_code_server_password", staticmethod(lambda: "code-password") ) - monkeypatch.setattr(SandboxManager, "_generate_vnc_password", staticmethod(lambda: "vnc-pass")) + monkeypatch.setattr(SandboxLauncher, "_generate_vnc_password", staticmethod(lambda: "vnc-pass")) manager = SandboxManager() settings = { @@ -189,23 +192,17 @@ async def test_launch_matrix_preserves_common_and_source_specific_behavior( assert handle.vnc_password == "vnc-pass" assert handle.ttyd_url == "https://terminal.example" assert handle.tunnel_urls == {3000: "https://app.example"} - resolve_tunnels.assert_awaited_once_with( - handle.modal_sandbox, - "sandbox-1", - True, - True, - True, - [3000], - 9000, - 9001, - 9002, + handle.modal_sandbox.tunnels.assert_called_once_with() + handle.modal_sandbox.filesystem.write_text.aio.assert_awaited_once_with( + f"{TUNNEL_ENV_SANDBOX_ID_KEY}=sandbox-1\nTUNNEL_3000=https://app.example\n", + TUNNEL_ENV_FILE_PATH, ) @pytest.mark.asyncio async def test_repository_image_create_validates_repo_before_image_lookup(monkeypatch): from_id = Mock(side_effect=AssertionError("image lookup should not run")) - monkeypatch.setattr("src.sandbox.manager.modal.Image.from_id", from_id) + monkeypatch.setattr("src.sandbox.launch.modal.Image.from_id", from_id) with pytest.raises(ValueError, match="repo_owner and repo_name must be provided together"): await SandboxManager().create_sandbox( @@ -216,25 +213,147 @@ async def test_repository_image_create_validates_repo_before_image_lookup(monkey @pytest.mark.asyncio -async def test_repository_image_not_found_is_reported_explicitly(monkeypatch, fake_llm_secret): - from modal.exception import NotFoundError +@pytest.mark.parametrize("image_source", ["repository", "snapshot"]) +@pytest.mark.parametrize("failure_stage", ["lookup", "create"]) +@pytest.mark.parametrize("missing", [False, True]) +async def test_launch_preserves_image_error_classification( + monkeypatch, image_source, failure_stage, missing +): + error = NotFoundError("missing image") if missing else RuntimeError("transient failure") + from_id = Mock( + return_value=object(), + side_effect=error if failure_stage == "lookup" else None, + ) + create = AsyncMock(side_effect=error) + monkeypatch.setattr("src.sandbox.launch.modal.Image.from_id", from_id) + monkeypatch.setattr("src.sandbox.launch.modal.Sandbox.create", SimpleNamespace(aio=create)) + expected_error = ( + RepositoryImageUnavailableError if image_source == "repository" and missing else type(error) + ) + + with pytest.raises(expected_error) as raised: + if image_source == "snapshot": + await SandboxManager().restore_from_snapshot( + snapshot_image_id="image-1", + session_config={"repo_owner": "acme", "repo_name": "repo"}, + ) + else: + await SandboxManager().create_sandbox( + SandboxConfig(repo_owner="acme", repo_name="repo", repo_image_id="image-1") + ) + + if expected_error is RepositoryImageUnavailableError: + assert raised.value.__cause__ is error + else: + assert raised.value is error + from_id.assert_called_once_with("image-1") + if failure_stage == "lookup": + create.assert_not_awaited() + else: + # A spawn error must not silently fall back to a different image or retry. + create.assert_awaited_once() + + +@pytest.mark.asyncio +@pytest.mark.parametrize("missing", [False, True]) +async def test_base_image_spawn_errors_propagate_without_retry(monkeypatch, missing): + error = NotFoundError("missing image") if missing else RuntimeError("transient failure") + create = AsyncMock(side_effect=error) + from_id = Mock() + monkeypatch.setattr("src.sandbox.launch.modal.Image.from_id", from_id) + monkeypatch.setattr("src.sandbox.launch.modal.Sandbox.create", SimpleNamespace(aio=create)) + + with pytest.raises(type(error)) as raised: + await SandboxManager().create_sandbox(SandboxConfig(repo_owner=None, repo_name=None)) + + assert raised.value is error + create.assert_awaited_once() + from_id.assert_not_called() + + +@pytest.mark.asyncio +@pytest.mark.parametrize("image_source", ["base", "repository", "snapshot"]) +@pytest.mark.parametrize("failure", ["partial", "unavailable", "write"]) +async def test_launch_returns_handle_despite_tunnel_failures(monkeypatch, image_source, failure): + write_text = AsyncMock(side_effect=OSError("write failed") if failure == "write" else None) + sandbox = SimpleNamespace( + object_id="modal-object-1", + tunnels=Mock( + side_effect=( + [RuntimeError("unavailable")] * 3 + if failure == "unavailable" + else [ + {9000: SimpleNamespace(url="https://code.example")}, + RuntimeError("not ready"), + {3000: SimpleNamespace(url="https://app.example")}, + ] + ) + ), + filesystem=SimpleNamespace(write_text=SimpleNamespace(aio=write_text)), + ) + create = AsyncMock(return_value=sandbox) + monkeypatch.setattr("src.sandbox.launch.modal.Sandbox.create", SimpleNamespace(aio=create)) + monkeypatch.setattr("src.sandbox.launch.modal.Image.from_id", lambda _: object()) + sleep = AsyncMock() + monkeypatch.setattr("src.sandbox.tunnels.asyncio.sleep", sleep) + common = { + "sandbox_id": "sandbox-partial", + "code_server_enabled": True, + "settings": {"codeServerPort": 9000, "tunnelPorts": [3000, 3001]}, + } + manager = SandboxManager() + + if image_source == "snapshot": + handle = await manager.restore_from_snapshot( + snapshot_image_id="image-1", + session_config={"repo_owner": "acme", "repo_name": "repo"}, + **common, + ) + else: + handle = await manager.create_sandbox( + SandboxConfig( + repo_owner="acme", + repo_name="repo", + repo_image_id="image-1" if image_source == "repository" else None, + **common, + ) + ) - monkeypatch.setattr("src.sandbox.manager.modal.Image.from_id", lambda _image_id: object()) + assert handle.status is SandboxStatus.WARMING + assert handle.modal_sandbox is sandbox + assert handle.modal_object_id == "modal-object-1" + assert handle.code_server_password == create.call_args.kwargs["env"]["CODE_SERVER_PASSWORD"] + assert create.call_args.kwargs["encrypted_ports"] == [9000, 3000, 3001] + assert sandbox.tunnels.call_count == 3 + assert [call.args for call in sleep.await_args_list] == [(1.0,), (2.0,)] + create.assert_awaited_once() + if failure == "unavailable": + assert handle.code_server_url is None + assert handle.tunnel_urls is None + write_text.assert_not_awaited() + else: + assert handle.code_server_url == "https://code.example" + assert handle.tunnel_urls == {3000: "https://app.example"} + write_text.assert_awaited_once_with( + f"{TUNNEL_ENV_SANDBOX_ID_KEY}=sandbox-partial\nTUNNEL_3000=https://app.example\n", + TUNNEL_ENV_FILE_PATH, + ) + + +@pytest.mark.asyncio +async def test_repository_image_not_found_is_reported_explicitly(monkeypatch, fake_llm_secret): + monkeypatch.setattr("src.sandbox.launch.modal.Image.from_id", lambda _image_id: object()) async def create_aio(*_args, **_kwargs): fake_llm_secret[0].hydrate.aio.assert_awaited_once_with() raise NotFoundError("image not found") create = SimpleNamespace(aio=AsyncMock(side_effect=create_aio)) - monkeypatch.setattr("src.sandbox.manager.modal.Sandbox.create", create) + monkeypatch.setattr("src.sandbox.launch.modal.Sandbox.create", create) with pytest.raises(RepositoryImageUnavailableError) as exc_info: await SandboxManager().create_sandbox( - SandboxConfig( - repo_owner="acme", - repo_name="repo", - repo_image_id="repo-image-missing", - ) + SandboxConfig(repo_owner="acme", repo_name="repo", repo_image_id="image-1") ) assert isinstance(exc_info.value.__cause__, NotFoundError) @@ -243,17 +362,15 @@ async def create_aio(*_args, **_kwargs): @pytest.mark.asyncio async def test_missing_secret_does_not_mark_repository_image_unavailable(monkeypatch): - from modal.exception import NotFoundError - create = SimpleNamespace(aio=AsyncMock()) - monkeypatch.setattr("src.sandbox.manager.modal.Sandbox.create", create) + monkeypatch.setattr("src.sandbox.launch.modal.Sandbox.create", create) def missing_secret(_name, **_kwargs): secret = Mock() secret.hydrate.aio = AsyncMock(side_effect=NotFoundError("secret not found")) return secret - monkeypatch.setattr("src.sandbox.manager.modal.Secret.from_name", missing_secret) + monkeypatch.setattr("src.sandbox.launch.modal.Secret.from_name", missing_secret) with pytest.raises(NotFoundError, match="secret not found"): await SandboxManager().create_sandbox( @@ -265,10 +382,8 @@ def missing_secret(_name, **_kwargs): @pytest.mark.asyncio async def test_base_image_not_found_is_not_classified_as_repository_image(monkeypatch): - from modal.exception import NotFoundError - create = SimpleNamespace(aio=AsyncMock(side_effect=NotFoundError("base image not found"))) - monkeypatch.setattr("src.sandbox.manager.modal.Sandbox.create", create) + monkeypatch.setattr("src.sandbox.launch.modal.Sandbox.create", create) with pytest.raises(NotFoundError, match="base image not found"): await SandboxManager().create_sandbox(SandboxConfig(repo_owner="acme", repo_name="repo")) @@ -282,13 +397,12 @@ async def test_base_image_not_found_is_not_classified_as_repository_image(monkey def _docker_manager(monkeypatch) -> tuple[SandboxManager, dict, object]: captured: dict = {} docker_image = object() - monkeypatch.setattr("src.sandbox.manager.base_image", object()) + monkeypatch.setattr("src.sandbox.launch.base_image", object()) monkeypatch.setattr("src.images.base.docker_image", docker_image) - monkeypatch.setattr("src.sandbox.manager.modal.Sandbox.create", _fake_create(captured)) + monkeypatch.setattr("src.sandbox.launch.modal.Sandbox.create", _fake_create(captured)) monkeypatch.setattr( - SandboxManager, - "_resolve_and_setup_tunnels", - AsyncMock(return_value=(None, None, None, {})), + "src.sandbox.tunnels.SandboxTunnels.resolve", + AsyncMock(return_value=TunnelUrls()), ) return SandboxManager(), captured, docker_image @@ -311,8 +425,6 @@ def _docker_config(**overrides) -> SandboxConfig: def _not_found(*_args, **_kwargs): - from modal.exception import NotFoundError - raise NotFoundError("no sandbox") @@ -321,9 +433,9 @@ def _not_found(*_args, **_kwargs): async def test_docker_launch_selects_vm_runtime_and_named_allocation(monkeypatch, image_source): manager, captured, docker_image = _docker_manager(monkeypatch) artifact = object() - monkeypatch.setattr("src.sandbox.manager.modal.Image.from_id", lambda _id: artifact) + monkeypatch.setattr("src.sandbox.launch.modal.Image.from_id", lambda _id: artifact) monkeypatch.setattr( - "src.sandbox.manager.modal.Sandbox.from_name", + "src.sandbox.launch.modal.Sandbox.from_name", SimpleNamespace(aio=AsyncMock(side_effect=_not_found)), ) @@ -350,7 +462,6 @@ async def test_docker_launch_selects_vm_runtime_and_named_allocation(monkeypatch assert kwargs["memory"] == 4096 assert kwargs["name"] == docker_allocation_name("session-1") assert kwargs["tags"] == docker_allocation_tags("session-1", "sandbox-acme-repo-1700000000000") - # The trusted signal wins over any user-supplied value. assert kwargs["env"][DOCKER_ENABLED_ENV_VAR] == "true" assert handle.sandbox_backend == "modal-vm" @@ -359,9 +470,9 @@ async def test_docker_launch_selects_vm_runtime_and_named_allocation(monkeypatch @pytest.mark.parametrize("image_source", ["base", "snapshot"]) async def test_expired_vm_launch_cannot_materialize_after_lookup(monkeypatch, image_source): manager, captured, _ = _docker_manager(monkeypatch) - monkeypatch.setattr("src.sandbox.manager.modal.Image.from_id", lambda _id: object()) + monkeypatch.setattr("src.sandbox.launch.modal.Image.from_id", lambda _id: object()) monkeypatch.setattr( - "src.sandbox.manager.modal.Sandbox.from_name", + "src.sandbox.launch.modal.Sandbox.from_name", SimpleNamespace(aio=AsyncMock(side_effect=_not_found)), ) if image_source == "base": @@ -409,7 +520,7 @@ async def test_docker_launch_adopts_an_existing_owned_allocation(monkeypatch): existing.get_tags.aio = existing.get_tags from_name = AsyncMock(return_value=existing) monkeypatch.setattr( - "src.sandbox.manager.modal.Sandbox.from_name", SimpleNamespace(aio=from_name) + "src.sandbox.launch.modal.Sandbox.from_name", SimpleNamespace(aio=from_name) ) handle = await manager.create_sandbox(_docker_config()) @@ -425,19 +536,19 @@ async def test_docker_launch_adopts_an_existing_owned_allocation(monkeypatch): async def test_docker_retry_returns_the_original_access_credentials( monkeypatch, create_race, image_source ): - from modal.exception import AlreadyExistsError, NotFoundError + from modal.exception import AlreadyExistsError manager, captured, _ = _docker_manager(monkeypatch) - monkeypatch.setattr("src.sandbox.manager.modal.Image.from_id", lambda _id: object()) + monkeypatch.setattr("src.sandbox.launch.modal.Image.from_id", lambda _id: object()) monkeypatch.setattr( - SandboxManager, "_generate_code_server_password", Mock(side_effect=["original", "new"]) + SandboxLauncher, "_generate_code_server_password", Mock(side_effect=["original", "new"]) ) monkeypatch.setattr( - SandboxManager, "_generate_vnc_password", Mock(side_effect=["old-vnc", "new-vnc"]) + SandboxLauncher, "_generate_vnc_password", Mock(side_effect=["old-vnc", "new-vnc"]) ) from_name = AsyncMock(side_effect=NotFoundError("not created")) monkeypatch.setattr( - "src.sandbox.manager.modal.Sandbox.from_name", SimpleNamespace(aio=from_name) + "src.sandbox.launch.modal.Sandbox.from_name", SimpleNamespace(aio=from_name) ) async def launch(): @@ -470,7 +581,7 @@ async def launch(): ) from_name.side_effect = [NotFoundError("racing"), existing] if create_race else [existing] create = AsyncMock(side_effect=AlreadyExistsError("already created")) - monkeypatch.setattr("src.sandbox.manager.modal.Sandbox.create", SimpleNamespace(aio=create)) + monkeypatch.setattr("src.sandbox.launch.modal.Sandbox.create", SimpleNamespace(aio=create)) adopted = await launch() @@ -503,7 +614,7 @@ async def test_docker_adoption_fails_if_original_credentials_cannot_be_recovered exec=SimpleNamespace(aio=AsyncMock(return_value=process)), ) monkeypatch.setattr( - "src.sandbox.manager.modal.Sandbox.from_name", + "src.sandbox.launch.modal.Sandbox.from_name", SimpleNamespace(aio=AsyncMock(return_value=existing)), ) @@ -511,7 +622,7 @@ async def test_docker_adoption_fails_if_original_credentials_cannot_be_recovered await manager.create_sandbox(_docker_config(code_server_enabled=True)) assert "kwargs" not in captured - manager._resolve_and_setup_tunnels.assert_not_awaited() + SandboxTunnels.resolve.assert_not_awaited() @pytest.mark.asyncio @@ -523,7 +634,7 @@ async def test_docker_launch_refuses_a_same_named_allocation_it_does_not_own(mon ) foreign.get_tags.aio = foreign.get_tags monkeypatch.setattr( - "src.sandbox.manager.modal.Sandbox.from_name", + "src.sandbox.launch.modal.Sandbox.from_name", SimpleNamespace(aio=AsyncMock(return_value=foreign)), ) @@ -553,7 +664,7 @@ async def from_name(_app, name): _not_found() monkeypatch.setattr( - "src.sandbox.manager.modal.Sandbox.from_name", SimpleNamespace(aio=from_name) + "src.sandbox.launch.modal.Sandbox.from_name", SimpleNamespace(aio=from_name) ) await manager.create_sandbox( @@ -563,7 +674,6 @@ async def from_name(_app, name): prior.terminate.assert_awaited_once_with(wait=True) assert captured["kwargs"]["name"] == docker_allocation_name("session-1") - # A prior allocation with foreign tags is left alone. prior.terminate.reset_mock() prior.get_tags = AsyncMock(return_value={"openinspect_kind": "other"}) prior.get_tags.aio = prior.get_tags @@ -576,9 +686,10 @@ async def from_name(_app, name): @pytest.mark.asyncio async def test_late_predecessor_cannot_materialize_beside_successor(monkeypatch): - from modal.exception import AlreadyExistsError, NotFoundError + from modal.exception import AlreadyExistsError - manager, _, _ = _docker_manager(monkeypatch) + _, _, _ = _docker_manager(monkeypatch) + launcher = SandboxLauncher() predecessor_name = docker_allocation_name("session-1") predecessor = SimpleNamespace( object_id="late-predecessor", @@ -589,17 +700,16 @@ async def test_late_predecessor_cannot_materialize_beside_successor(monkeypatch) lookup = AsyncMock( side_effect=[NotFoundError("still creating"), NotFoundError("still creating"), predecessor] ) - monkeypatch.setattr("src.sandbox.manager.modal.Sandbox.from_name", SimpleNamespace(aio=lookup)) + monkeypatch.setattr("src.sandbox.launch.modal.Sandbox.from_name", SimpleNamespace(aio=lookup)) async def create(kwargs, *, repository_image): - # Provider-side naming wins the race after both client lookups missed it. if kwargs["name"] == predecessor_name: raise AlreadyExistsError("predecessor won the name") return SimpleNamespace(object_id="duplicate-successor") - monkeypatch.setattr("src.sandbox.manager._create_sandbox", create) + monkeypatch.setattr("src.sandbox.launch._create_sandbox", create) with pytest.raises(RuntimeError, match="ownership mismatch"): - await manager._launch_docker_sandbox( + await launcher._launch_docker_sandbox( session_id="session-1", sandbox_id="successor", retire_sandbox_id="prior", @@ -631,10 +741,8 @@ async def terminate(*, wait=False): ), terminate=SimpleNamespace(aio=terminate), ) - from modal.exception import NotFoundError - monkeypatch.setattr( - "src.sandbox.manager.modal.Sandbox.from_name", + "src.sandbox.launch.modal.Sandbox.from_name", SimpleNamespace(aio=AsyncMock(side_effect=[prior, NotFoundError("no successor")])), ) launch = asyncio.create_task( @@ -656,3 +764,50 @@ async def terminate(*, wait=False): finally: launch.cancel() await asyncio.gather(launch, return_exceptions=True) + + +@pytest.mark.asyncio +@pytest.mark.parametrize("restore", [False, True], ids=["create", "restore"]) +@pytest.mark.parametrize( + "ports, expected", + [ + ([True, False], []), + ([True, False, 0, -1, 65536, "3000", 3.5, None, 1, 3000, 65535], [1, 3000, 65535]), + ([True] * 10 + [3000], [3000]), + ], + ids=["booleans-only", "mixed-with-boundary-ports", "booleans-do-not-consume-limit"], +) +async def test_launch_rejects_boolean_tunnel_ports(monkeypatch, restore, ports, expected): + """Invalid extras never reach Modal or the runtime's expected-port list.""" + urls = {port: f"https://port-{port}.example" for port in expected} + sandbox = SimpleNamespace( + object_id="modal-ports", + tunnels=Mock(return_value={port: SimpleNamespace(url=url) for port, url in urls.items()}), + filesystem=SimpleNamespace(write_text=SimpleNamespace(aio=AsyncMock())), + ) + create = AsyncMock(return_value=sandbox) + monkeypatch.setattr("src.sandbox.launch.modal.Sandbox.create", SimpleNamespace(aio=create)) + monkeypatch.setattr("src.sandbox.launch.modal.Image.from_id", lambda _: object()) + manager = SandboxManager() + settings = {"tunnelPorts": ports} + + if restore: + handle = await manager.restore_from_snapshot( + snapshot_image_id="image-1", + session_config={"repo_owner": "acme", "repo_name": "repo"}, + settings=settings, + ) + else: + handle = await manager.create_sandbox( + SandboxConfig(repo_owner="acme", repo_name="repo", settings=settings) + ) + + kwargs = create.call_args.kwargs + assert kwargs.get("encrypted_ports", []) == expected + assert all(type(port) is int for port in kwargs.get("encrypted_ports", [])) + assert kwargs["env"].get(EXPECTED_TUNNEL_PORTS_ENV_VAR) == ( + ",".join(str(port) for port in expected) if expected else None + ) + assert handle.tunnel_urls == (urls or None) + if not expected: + sandbox.tunnels.assert_not_called() diff --git a/packages/modal-infra/tests/test_sandbox_resources.py b/packages/modal-infra/tests/test_sandbox_resources.py index cb8d7a1b90..2505511766 100644 --- a/packages/modal-infra/tests/test_sandbox_resources.py +++ b/packages/modal-infra/tests/test_sandbox_resources.py @@ -6,6 +6,7 @@ from src.sandbox.launch_policy import launch_kwargs, parse_launch from src.sandbox.manager import SandboxConfig, SandboxManager +from src.sandbox.tunnels import SandboxTunnels, TunnelUrls class TestResourceKwargs: @@ -47,11 +48,11 @@ class TestCreateSandboxResources: @pytest.mark.asyncio async def test_create_sandbox_passes_cpu_and_memory(self, monkeypatch): captured: dict = {} - monkeypatch.setattr("src.sandbox.manager.modal.Sandbox.create", _fake_create(captured)) + monkeypatch.setattr("src.sandbox.launch.modal.Sandbox.create", _fake_create(captured)) monkeypatch.setattr( - SandboxManager, - "_resolve_and_setup_tunnels", - AsyncMock(return_value=(None, None, None, None)), + SandboxTunnels, + "resolve", + AsyncMock(return_value=TunnelUrls(None, None, None, None)), ) manager = SandboxManager() @@ -74,13 +75,13 @@ class FakeImage: object_id = "img-1" monkeypatch.setattr( - "src.sandbox.manager.modal.Image.from_id", lambda *_a, **_kw: FakeImage() + "src.sandbox.launch.modal.Image.from_id", lambda *_a, **_kw: FakeImage() ) - monkeypatch.setattr("src.sandbox.manager.modal.Sandbox.create", _fake_create(captured)) + monkeypatch.setattr("src.sandbox.launch.modal.Sandbox.create", _fake_create(captured)) monkeypatch.setattr( - SandboxManager, - "_resolve_and_setup_tunnels", - AsyncMock(return_value=(None, None, None, None)), + SandboxTunnels, + "resolve", + AsyncMock(return_value=TunnelUrls(None, None, None, None)), ) manager = SandboxManager() diff --git a/packages/modal-infra/tests/test_ttyd.py b/packages/modal-infra/tests/test_ttyd.py index 015cb456e5..3ddd6ae978 100644 --- a/packages/modal-infra/tests/test_ttyd.py +++ b/packages/modal-infra/tests/test_ttyd.py @@ -4,74 +4,38 @@ import pytest -from sandbox_runtime.constants import NOVNC_PORT, TTYD_PORT +from sandbox_runtime.constants import ( + EXPECTED_TUNNEL_PORTS_ENV_VAR, + NOVNC_PORT, + TTYD_PORT, +) from src.sandbox.manager import ( CODE_SERVER_PORT, TTYD_PROXY_PORT, SandboxConfig, SandboxManager, ) +from src.sandbox.tunnels import SandboxTunnels, TunnelUrls -class TestCollectExposedPortsTerminal: - """_collect_exposed_ports with terminal_enabled flag.""" +@pytest.mark.parametrize("code_server", [False, True]) +@pytest.mark.parametrize("terminal", [False, True]) +def test_terminal_port_ownership(code_server, terminal): + tunnels = SandboxTunnels( + code_server_enabled=code_server, + settings={"terminalEnabled": terminal}, + ) + assert (TTYD_PROXY_PORT in tunnels.exposed_ports) is terminal + assert (CODE_SERVER_PORT in tunnels.exposed_ports) is code_server + assert TTYD_PORT not in tunnels.exposed_ports - def test_terminal_enabled_includes_proxy_port(self): - exposed, _extra = SandboxManager._collect_exposed_ports( - code_server_enabled=False, - vnc_enabled=False, - terminal_enabled=True, - settings=None, - code_server_port=CODE_SERVER_PORT, - novnc_port=NOVNC_PORT, - ttyd_proxy_port=TTYD_PROXY_PORT, - ) - assert TTYD_PROXY_PORT in exposed - # ttyd raw port should NOT be exposed (only the proxy port) - assert TTYD_PORT not in exposed - def test_terminal_disabled_excludes_proxy_port(self): - exposed, _extra = SandboxManager._collect_exposed_ports( - code_server_enabled=False, - vnc_enabled=False, - terminal_enabled=False, - settings=None, - code_server_port=CODE_SERVER_PORT, - novnc_port=NOVNC_PORT, - ttyd_proxy_port=TTYD_PROXY_PORT, - ) - assert TTYD_PROXY_PORT not in exposed - - def test_terminal_and_code_server_both_enabled(self): - exposed, _extra = SandboxManager._collect_exposed_ports( - code_server_enabled=True, - vnc_enabled=False, - terminal_enabled=True, - settings=None, - code_server_port=CODE_SERVER_PORT, - novnc_port=NOVNC_PORT, - ttyd_proxy_port=TTYD_PROXY_PORT, - ) - assert CODE_SERVER_PORT in exposed - assert TTYD_PROXY_PORT in exposed - - def test_terminal_port_deduped_from_tunnel_ports(self): - """If user explicitly lists TTYD_PROXY_PORT in tunnelPorts, it should not duplicate.""" - settings = {"tunnelPorts": [TTYD_PROXY_PORT, 3000]} - exposed, extra = SandboxManager._collect_exposed_ports( - code_server_enabled=False, - vnc_enabled=False, - terminal_enabled=True, - settings=settings, - code_server_port=CODE_SERVER_PORT, - novnc_port=NOVNC_PORT, - ttyd_proxy_port=TTYD_PROXY_PORT, - ) - assert exposed.count(TTYD_PROXY_PORT) == 1 - assert 3000 in exposed - # TTYD_PROXY_PORT should not be in extra (reserved) - assert TTYD_PROXY_PORT not in extra - assert 3000 in extra +def test_terminal_port_is_not_duplicated_by_extra_tunnel(): + tunnels = SandboxTunnels( + settings={"terminalEnabled": True, "tunnelPorts": [TTYD_PROXY_PORT, 3000]} + ) + assert tunnels.exposed_ports == [TTYD_PROXY_PORT, 3000] + assert tunnels.environment[EXPECTED_TUNNEL_PORTS_ENV_VAR] == "3000" class TestResolveTunnelsTerminal: @@ -85,17 +49,17 @@ async def test_returns_ttyd_url_when_terminal_enabled(self): sandbox = MagicMock() sandbox.tunnels.return_value = {TTYD_PROXY_PORT: tunnel} - cs_url, vnc_url, ttyd_url, extra = await SandboxManager._resolve_and_setup_tunnels( - sandbox, - "sb-123", + cs_url, vnc_url, ttyd_url, extra = await SandboxTunnels( code_server_enabled=False, vnc_enabled=False, - terminal_enabled=True, - extra_ports=[], - code_server_port=CODE_SERVER_PORT, - novnc_port=NOVNC_PORT, - ttyd_proxy_port=TTYD_PROXY_PORT, - ) + settings={ + "terminalEnabled": True, + "tunnelPorts": [], + "codeServerPort": CODE_SERVER_PORT, + "vncPort": NOVNC_PORT, + "terminalPort": TTYD_PROXY_PORT, + }, + ).resolve(sandbox, "sb-123") assert cs_url is None assert vnc_url is None assert ttyd_url == "https://ttyd.example.com" @@ -104,17 +68,17 @@ async def test_returns_ttyd_url_when_terminal_enabled(self): @pytest.mark.asyncio async def test_returns_none_when_terminal_disabled(self): sandbox = MagicMock() - cs_url, vnc_url, ttyd_url, extra = await SandboxManager._resolve_and_setup_tunnels( - sandbox, - "sb-123", + cs_url, vnc_url, ttyd_url, extra = await SandboxTunnels( code_server_enabled=False, vnc_enabled=False, - terminal_enabled=False, - extra_ports=[], - code_server_port=CODE_SERVER_PORT, - novnc_port=NOVNC_PORT, - ttyd_proxy_port=TTYD_PROXY_PORT, - ) + settings={ + "terminalEnabled": False, + "tunnelPorts": [], + "codeServerPort": CODE_SERVER_PORT, + "vncPort": NOVNC_PORT, + "terminalPort": TTYD_PROXY_PORT, + }, + ).resolve(sandbox, "sb-123") assert cs_url is None assert vnc_url is None assert ttyd_url is None @@ -133,17 +97,17 @@ async def test_both_code_server_and_terminal(self): TTYD_PROXY_PORT: ttyd_tunnel, } - cs_url, vnc_url, ttyd_url, extra = await SandboxManager._resolve_and_setup_tunnels( - sandbox, - "sb-123", + cs_url, vnc_url, ttyd_url, extra = await SandboxTunnels( code_server_enabled=True, vnc_enabled=False, - terminal_enabled=True, - extra_ports=[], - code_server_port=CODE_SERVER_PORT, - novnc_port=NOVNC_PORT, - ttyd_proxy_port=TTYD_PROXY_PORT, - ) + settings={ + "terminalEnabled": True, + "tunnelPorts": [], + "codeServerPort": CODE_SERVER_PORT, + "vncPort": NOVNC_PORT, + "terminalPort": TTYD_PROXY_PORT, + }, + ).resolve(sandbox, "sb-123") assert cs_url == "https://cs.example.com" assert vnc_url is None assert ttyd_url == "https://ttyd.example.com" @@ -169,12 +133,12 @@ class FakeSandbox: fake_create = MagicMock() fake_create.aio = fake_create_aio - monkeypatch.setattr("src.sandbox.manager.modal.Sandbox.create", fake_create) + monkeypatch.setattr("src.sandbox.launch.modal.Sandbox.create", fake_create) monkeypatch.setattr( - SandboxManager, - "_resolve_and_setup_tunnels", - AsyncMock(return_value=(None, None, "https://ttyd.example.com", None)), + SandboxTunnels, + "resolve", + AsyncMock(return_value=TunnelUrls(None, None, "https://ttyd.example.com", None)), ) manager = SandboxManager() @@ -209,10 +173,10 @@ class FakeSandbox: fake_create = MagicMock() fake_create.aio = fake_create_aio - monkeypatch.setattr("src.sandbox.manager.modal.Sandbox.create", fake_create) + monkeypatch.setattr("src.sandbox.launch.modal.Sandbox.create", fake_create) - tunnel_mock = AsyncMock(return_value=(None, None, None, None)) - monkeypatch.setattr(SandboxManager, "_resolve_and_setup_tunnels", tunnel_mock) + tunnel_mock = AsyncMock(return_value=TunnelUrls(None, None, None, None)) + monkeypatch.setattr(SandboxTunnels, "resolve", tunnel_mock) manager = SandboxManager() config = SandboxConfig( @@ -255,12 +219,14 @@ class FakeSandbox: fake_create = MagicMock() fake_create.aio = fake_create_aio - monkeypatch.setattr("src.sandbox.manager.modal.Image.from_id", fake_from_id) - monkeypatch.setattr("src.sandbox.manager.modal.Sandbox.create", fake_create) + monkeypatch.setattr("src.sandbox.launch.modal.Image.from_id", fake_from_id) + monkeypatch.setattr("src.sandbox.launch.modal.Sandbox.create", fake_create) monkeypatch.setattr( - SandboxManager, - "_resolve_and_setup_tunnels", - AsyncMock(return_value=(None, None, "https://ttyd-restored.example.com", None)), + SandboxTunnels, + "resolve", + AsyncMock( + return_value=TunnelUrls(None, None, "https://ttyd-restored.example.com", None) + ), ) manager = SandboxManager() @@ -305,10 +271,10 @@ class FakeSandbox: fake_create = MagicMock() fake_create.aio = fake_create_aio - monkeypatch.setattr("src.sandbox.manager.modal.Image.from_id", fake_from_id) - monkeypatch.setattr("src.sandbox.manager.modal.Sandbox.create", fake_create) - tunnel_mock = AsyncMock(return_value=(None, None, None, None)) - monkeypatch.setattr(SandboxManager, "_resolve_and_setup_tunnels", tunnel_mock) + monkeypatch.setattr("src.sandbox.launch.modal.Image.from_id", fake_from_id) + monkeypatch.setattr("src.sandbox.launch.modal.Sandbox.create", fake_create) + tunnel_mock = AsyncMock(return_value=TunnelUrls(None, None, None, None)) + monkeypatch.setattr(SandboxTunnels, "resolve", tunnel_mock) manager = SandboxManager() handle = await manager.restore_from_snapshot( diff --git a/packages/modal-infra/tests/test_tunnel_ports.py b/packages/modal-infra/tests/test_tunnel_ports.py index 68e51ab939..ead6c6b51a 100644 --- a/packages/modal-infra/tests/test_tunnel_ports.py +++ b/packages/modal-infra/tests/test_tunnel_ports.py @@ -14,6 +14,7 @@ TUNNEL_ENV_SANDBOX_ID_KEY, ) from src.sandbox.manager import CODE_SERVER_PORT, SandboxConfig, SandboxManager +from src.sandbox.tunnels import SandboxTunnels, TunnelUrls def _mock_sandbox_with_filesystem() -> tuple[MagicMock, AsyncMock]: @@ -27,7 +28,7 @@ def _mock_sandbox_with_filesystem() -> tuple[MagicMock, AsyncMock]: class TestResolveTunnels: - """SandboxManager._resolve_tunnels tests.""" + """SandboxTunnels._resolve_tunnels tests.""" @pytest.mark.asyncio async def test_resolves_all_ports(self): @@ -39,7 +40,7 @@ async def test_resolves_all_ports(self): sandbox = MagicMock() sandbox.tunnels.return_value = {3000: tunnel_3000, 3001: tunnel_3001} - result = await SandboxManager._resolve_tunnels(sandbox, "sb-1", [3000, 3001]) + result = await SandboxTunnels._resolve_tunnels(sandbox, "sb-1", [3000, 3001]) assert result == { 3000: "https://tunnel-3000.example.com", 3001: "https://tunnel-3001.example.com", @@ -53,9 +54,9 @@ async def test_returns_partial_on_missing_port(self): sandbox = MagicMock() sandbox.tunnels.return_value = {3000: tunnel_3000} - with patch("src.sandbox.manager.asyncio.sleep", new_callable=AsyncMock): - result = await SandboxManager._resolve_tunnels( - sandbox, "sb-1", [3000, 3001], retries=2, backoff=0.0 + with patch("src.sandbox.tunnels.asyncio.sleep", new_callable=AsyncMock): + result = await SandboxTunnels._resolve_tunnels( + sandbox, "sb-1", [3000, 3001], retries=2, backoff_seconds=0.0 ) assert result == {3000: "https://tunnel-3000.example.com"} @@ -64,9 +65,9 @@ async def test_returns_empty_on_exception_after_retries(self): sandbox = MagicMock() sandbox.tunnels.side_effect = Exception("tunnel unavailable") - with patch("src.sandbox.manager.asyncio.sleep", new_callable=AsyncMock): - result = await SandboxManager._resolve_tunnels( - sandbox, "sb-1", [3000], retries=3, backoff=0.0 + with patch("src.sandbox.tunnels.asyncio.sleep", new_callable=AsyncMock): + result = await SandboxTunnels._resolve_tunnels( + sandbox, "sb-1", [3000], retries=3, backoff_seconds=0.0 ) assert result == {} @@ -83,9 +84,9 @@ async def test_retries_on_partial_resolution(self): {3000: tunnel_3000, 3001: tunnel_3001}, ] - with patch("src.sandbox.manager.asyncio.sleep", new_callable=AsyncMock): - result = await SandboxManager._resolve_tunnels( - sandbox, "sb-1", [3000, 3001], retries=3, backoff=0.0 + with patch("src.sandbox.tunnels.asyncio.sleep", new_callable=AsyncMock): + result = await SandboxTunnels._resolve_tunnels( + sandbox, "sb-1", [3000, 3001], retries=3, backoff_seconds=0.0 ) assert result == { 3000: "https://tunnel-3000.example.com", @@ -100,17 +101,17 @@ class TestResolveAndSetupTunnels: @pytest.mark.asyncio async def test_returns_none_none_none_for_no_ports(self): sandbox = MagicMock() - cs_url, vnc_url, ttyd_url, extra = await SandboxManager._resolve_and_setup_tunnels( - sandbox, - "sb-1", - False, - False, - False, - [], - code_server_port=CODE_SERVER_PORT, - novnc_port=NOVNC_PORT, - ttyd_proxy_port=TTYD_PROXY_PORT, - ) + cs_url, vnc_url, ttyd_url, extra = await SandboxTunnels( + code_server_enabled=False, + vnc_enabled=False, + settings={ + "terminalEnabled": False, + "tunnelPorts": [], + "codeServerPort": CODE_SERVER_PORT, + "vncPort": NOVNC_PORT, + "terminalPort": TTYD_PROXY_PORT, + }, + ).resolve(sandbox, "sb-1") assert cs_url is None assert vnc_url is None assert ttyd_url is None @@ -122,22 +123,22 @@ async def test_resolves_extra_ports(self): sandbox, _write_text = _mock_sandbox_with_filesystem() with patch.object( - SandboxManager, + SandboxTunnels, "_resolve_tunnels", new_callable=AsyncMock, return_value=tunnel_urls, ): - cs_url, vnc_url, ttyd_url, extra = await SandboxManager._resolve_and_setup_tunnels( - sandbox, - "sb-1", - False, - False, - False, - [3000], - code_server_port=CODE_SERVER_PORT, - novnc_port=NOVNC_PORT, - ttyd_proxy_port=TTYD_PROXY_PORT, - ) + cs_url, vnc_url, ttyd_url, extra = await SandboxTunnels( + code_server_enabled=False, + vnc_enabled=False, + settings={ + "terminalEnabled": False, + "tunnelPorts": [3000], + "codeServerPort": CODE_SERVER_PORT, + "vncPort": NOVNC_PORT, + "terminalPort": TTYD_PROXY_PORT, + }, + ).resolve(sandbox, "sb-1") assert cs_url is None assert vnc_url is None @@ -154,22 +155,22 @@ async def test_splits_code_server_from_extra_ports(self): sandbox, _write_text = _mock_sandbox_with_filesystem() with patch.object( - SandboxManager, + SandboxTunnels, "_resolve_tunnels", new_callable=AsyncMock, return_value=resolved, ): - cs_url, vnc_url, ttyd_url, extra = await SandboxManager._resolve_and_setup_tunnels( - sandbox, - "sb-1", - True, - False, - False, - [3000], - code_server_port=CODE_SERVER_PORT, - novnc_port=NOVNC_PORT, - ttyd_proxy_port=TTYD_PROXY_PORT, - ) + cs_url, vnc_url, ttyd_url, extra = await SandboxTunnels( + code_server_enabled=True, + vnc_enabled=False, + settings={ + "terminalEnabled": False, + "tunnelPorts": [3000], + "codeServerPort": CODE_SERVER_PORT, + "vncPort": NOVNC_PORT, + "terminalPort": TTYD_PROXY_PORT, + }, + ).resolve(sandbox, "sb-1") assert cs_url == "https://cs.example.com" assert vnc_url is None @@ -183,22 +184,22 @@ async def test_keeps_code_server_port_tunnel_when_code_server_disabled(self): sandbox, _write_text = _mock_sandbox_with_filesystem() with patch.object( - SandboxManager, + SandboxTunnels, "_resolve_tunnels", new_callable=AsyncMock, return_value=resolved, ): - cs_url, vnc_url, ttyd_url, extra = await SandboxManager._resolve_and_setup_tunnels( - sandbox, - "sb-1", - False, - False, - False, - [CODE_SERVER_PORT], - code_server_port=CODE_SERVER_PORT, - novnc_port=NOVNC_PORT, - ttyd_proxy_port=TTYD_PROXY_PORT, - ) + cs_url, vnc_url, ttyd_url, extra = await SandboxTunnels( + code_server_enabled=False, + vnc_enabled=False, + settings={ + "terminalEnabled": False, + "tunnelPorts": [CODE_SERVER_PORT], + "codeServerPort": CODE_SERVER_PORT, + "vncPort": NOVNC_PORT, + "terminalPort": TTYD_PROXY_PORT, + }, + ).resolve(sandbox, "sb-1") assert cs_url is None assert vnc_url is None @@ -215,35 +216,35 @@ async def test_splits_custom_code_server_port_from_user_tunnel(self): sandbox, _write_text = _mock_sandbox_with_filesystem() with patch.object( - SandboxManager, + SandboxTunnels, "_resolve_tunnels", new_callable=AsyncMock, return_value=resolved, ): - cs_url, _vnc_url, _ttyd_url, extra = await SandboxManager._resolve_and_setup_tunnels( - sandbox, - "sb-1", - True, - False, - False, - [CODE_SERVER_PORT], - code_server_port=8081, - novnc_port=NOVNC_PORT, - ttyd_proxy_port=TTYD_PROXY_PORT, - ) + cs_url, _vnc_url, _ttyd_url, extra = await SandboxTunnels( + code_server_enabled=True, + vnc_enabled=False, + settings={ + "terminalEnabled": False, + "tunnelPorts": [CODE_SERVER_PORT], + "codeServerPort": 8081, + "vncPort": NOVNC_PORT, + "terminalPort": TTYD_PROXY_PORT, + }, + ).resolve(sandbox, "sb-1") assert cs_url == "https://cs.example.com" assert extra == {CODE_SERVER_PORT: "https://my-app.example.com"} class TestWriteTunnelEnvFile: - """SandboxManager._write_tunnel_env_file tests.""" + """SandboxTunnels._write_tunnel_env_file tests.""" @pytest.mark.asyncio async def test_writes_dotenv_format_to_expected_path(self): sandbox, write_text = _mock_sandbox_with_filesystem() - await SandboxManager._write_tunnel_env_file( + await SandboxTunnels._write_tunnel_env_file( sandbox, "sb-1", { @@ -267,8 +268,8 @@ async def test_write_failure_does_not_raise(self): sandbox, write_text = _mock_sandbox_with_filesystem() write_text.side_effect = Exception("write failed") - with patch("src.sandbox.manager.log") as mock_log: - await SandboxManager._write_tunnel_env_file( + with patch("src.sandbox.tunnels.log") as mock_log: + await SandboxTunnels._write_tunnel_env_file( sandbox, "sb-1", {3000: "https://tunnel-3000.example.com"} ) @@ -285,22 +286,22 @@ async def test_writes_file_when_extra_urls_present(self): tunnel_urls = {3000: "https://tunnel-3000.example.com"} with patch.object( - SandboxManager, + SandboxTunnels, "_resolve_tunnels", new_callable=AsyncMock, return_value=tunnel_urls, ): - await SandboxManager._resolve_and_setup_tunnels( - sandbox, - "sb-1", - False, - False, - False, - [3000], - code_server_port=CODE_SERVER_PORT, - novnc_port=NOVNC_PORT, - ttyd_proxy_port=TTYD_PROXY_PORT, - ) + await SandboxTunnels( + code_server_enabled=False, + vnc_enabled=False, + settings={ + "terminalEnabled": False, + "tunnelPorts": [3000], + "codeServerPort": CODE_SERVER_PORT, + "vncPort": NOVNC_PORT, + "terminalPort": TTYD_PROXY_PORT, + }, + ).resolve(sandbox, "sb-1") write_text.assert_awaited_once() written = write_text.call_args[0][0] @@ -312,22 +313,22 @@ async def test_does_not_write_file_when_no_extra_urls(self): sandbox, write_text = _mock_sandbox_with_filesystem() with patch.object( - SandboxManager, + SandboxTunnels, "_resolve_tunnels", new_callable=AsyncMock, return_value={}, ): - _cs, _vnc, _ttyd, extra = await SandboxManager._resolve_and_setup_tunnels( - sandbox, - "sb-1", - False, - False, - False, - [3000], - code_server_port=CODE_SERVER_PORT, - novnc_port=NOVNC_PORT, - ttyd_proxy_port=TTYD_PROXY_PORT, - ) + _cs, _vnc, _ttyd, extra = await SandboxTunnels( + code_server_enabled=False, + vnc_enabled=False, + settings={ + "terminalEnabled": False, + "tunnelPorts": [3000], + "codeServerPort": CODE_SERVER_PORT, + "vncPort": NOVNC_PORT, + "terminalPort": TTYD_PROXY_PORT, + }, + ).resolve(sandbox, "sb-1") assert extra is None write_text.assert_not_awaited() @@ -338,22 +339,22 @@ async def test_does_not_write_file_for_only_reserved_ports(self): sandbox, write_text = _mock_sandbox_with_filesystem() with patch.object( - SandboxManager, + SandboxTunnels, "_resolve_tunnels", new_callable=AsyncMock, return_value={CODE_SERVER_PORT: "https://cs.example.com"}, ): - await SandboxManager._resolve_and_setup_tunnels( - sandbox, - "sb-1", - True, - False, - False, - [], - code_server_port=CODE_SERVER_PORT, - novnc_port=NOVNC_PORT, - ttyd_proxy_port=TTYD_PROXY_PORT, - ) + await SandboxTunnels( + code_server_enabled=True, + vnc_enabled=False, + settings={ + "terminalEnabled": False, + "tunnelPorts": [], + "codeServerPort": CODE_SERVER_PORT, + "vncPort": NOVNC_PORT, + "terminalPort": TTYD_PROXY_PORT, + }, + ).resolve(sandbox, "sb-1") write_text.assert_not_awaited() @@ -364,24 +365,24 @@ async def test_write_failure_does_not_block_return(self): with ( patch.object( - SandboxManager, + SandboxTunnels, "_resolve_tunnels", new_callable=AsyncMock, return_value={3000: "https://tunnel-3000.example.com"}, ), - patch("src.sandbox.manager.log"), + patch("src.sandbox.tunnels.log"), ): - _cs, _vnc, _ttyd, extra = await SandboxManager._resolve_and_setup_tunnels( - sandbox, - "sb-1", - False, - False, - False, - [3000], - code_server_port=CODE_SERVER_PORT, - novnc_port=NOVNC_PORT, - ttyd_proxy_port=TTYD_PROXY_PORT, - ) + _cs, _vnc, _ttyd, extra = await SandboxTunnels( + code_server_enabled=False, + vnc_enabled=False, + settings={ + "terminalEnabled": False, + "tunnelPorts": [3000], + "codeServerPort": CODE_SERVER_PORT, + "vncPort": NOVNC_PORT, + "terminalPort": TTYD_PROXY_PORT, + }, + ).resolve(sandbox, "sb-1") assert extra == {3000: "https://tunnel-3000.example.com"} @@ -403,11 +404,11 @@ class FakeSandbox: return FakeSandbox() fake_create_aio.aio = fake_create_aio - monkeypatch.setattr("src.sandbox.manager.modal.Sandbox.create", fake_create_aio) + monkeypatch.setattr("src.sandbox.launch.modal.Sandbox.create", fake_create_aio) monkeypatch.setattr( - SandboxManager, - "_resolve_and_setup_tunnels", - AsyncMock(return_value=(None, None, None, None)), + SandboxTunnels, + "resolve", + AsyncMock(return_value=TunnelUrls(None, None, None, None)), ) manager = SandboxManager() @@ -435,11 +436,11 @@ class FakeSandbox: return FakeSandbox() fake_create_aio.aio = fake_create_aio - monkeypatch.setattr("src.sandbox.manager.modal.Sandbox.create", fake_create_aio) + monkeypatch.setattr("src.sandbox.launch.modal.Sandbox.create", fake_create_aio) monkeypatch.setattr( - SandboxManager, - "_resolve_and_setup_tunnels", - AsyncMock(return_value=(None, None, None, None)), + SandboxTunnels, + "resolve", + AsyncMock(return_value=TunnelUrls(None, None, None, None)), ) manager = SandboxManager() @@ -467,13 +468,13 @@ class FakeSandbox: fake_create_aio.aio = fake_create_aio monkeypatch.setattr( - "src.sandbox.manager.modal.Image.from_id", lambda *_a, **_kw: FakeImage() + "src.sandbox.launch.modal.Image.from_id", lambda *_a, **_kw: FakeImage() ) - monkeypatch.setattr("src.sandbox.manager.modal.Sandbox.create", fake_create_aio) + monkeypatch.setattr("src.sandbox.launch.modal.Sandbox.create", fake_create_aio) monkeypatch.setattr( - SandboxManager, - "_resolve_and_setup_tunnels", - AsyncMock(return_value=(None, None, None, None)), + SandboxTunnels, + "resolve", + AsyncMock(return_value=TunnelUrls(None, None, None, None)), ) manager = SandboxManager() @@ -486,128 +487,61 @@ class FakeSandbox: assert captured["env"][EXPECTED_TUNNEL_PORTS_ENV_VAR] == "3000" -class TestCollectExposedPorts: - """SandboxManager._collect_exposed_ports tests.""" - - def test_no_ports_when_no_settings(self): - exposed, tunnel = SandboxManager._collect_exposed_ports( - False, False, False, None, CODE_SERVER_PORT, NOVNC_PORT, TTYD_PROXY_PORT - ) - assert exposed == [] - assert tunnel == [] - - def test_code_server_only(self): - exposed, tunnel = SandboxManager._collect_exposed_ports( - True, False, False, None, CODE_SERVER_PORT, NOVNC_PORT, TTYD_PROXY_PORT - ) - assert exposed == [CODE_SERVER_PORT] - assert tunnel == [] - - def test_tunnel_ports_only(self): - exposed, tunnel = SandboxManager._collect_exposed_ports( +@pytest.mark.parametrize( + "code_server, settings, exposed, expected_extras", + [ + (False, None, [], None), + (True, None, [CODE_SERVER_PORT], None), + (False, {"tunnelPorts": [3000, 5173]}, [3000, 5173], "3000,5173"), + (True, {"tunnelPorts": [3000]}, [CODE_SERVER_PORT, 3000], "3000"), + (False, {"terminalEnabled": True}, [TTYD_PROXY_PORT], None), + ( False, - False, - False, - {"tunnelPorts": [3000, 5173]}, - CODE_SERVER_PORT, - NOVNC_PORT, - TTYD_PROXY_PORT, - ) - assert exposed == [3000, 5173] - assert tunnel == [3000, 5173] - - def test_combined_code_server_and_tunnels(self): - exposed, tunnel = SandboxManager._collect_exposed_ports( + {"terminalEnabled": True, "tunnelPorts": [TTYD_PROXY_PORT, 3000]}, + [TTYD_PROXY_PORT, 3000], + "3000", + ), + (True, {"tunnelPorts": [CODE_SERVER_PORT, 3000]}, [CODE_SERVER_PORT, 3000], "3000"), + ( True, + {"codeServerPort": 8081, "tunnelPorts": [CODE_SERVER_PORT]}, + [8081, CODE_SERVER_PORT], + str(CODE_SERVER_PORT), + ), + ( False, - False, - {"tunnelPorts": [3000]}, - CODE_SERVER_PORT, - NOVNC_PORT, - TTYD_PROXY_PORT, - ) - assert exposed == [CODE_SERVER_PORT, 3000] - assert tunnel == [3000] - - def test_terminal_only(self): - exposed, tunnel = SandboxManager._collect_exposed_ports( - False, False, True, None, CODE_SERVER_PORT, NOVNC_PORT, TTYD_PROXY_PORT - ) - assert exposed == [TTYD_PROXY_PORT] - assert tunnel == [] - - def test_deduplicates_ttyd_port_from_tunnels(self): - exposed, tunnel = SandboxManager._collect_exposed_ports( - False, - False, - True, - {"tunnelPorts": [TTYD_PROXY_PORT, 3000]}, - CODE_SERVER_PORT, - NOVNC_PORT, - TTYD_PROXY_PORT, - ) - assert exposed == [TTYD_PROXY_PORT, 3000] - assert tunnel == [3000] - - def test_deduplicates_code_server_port_from_tunnels(self): - exposed, tunnel = SandboxManager._collect_exposed_ports( - True, - False, - False, - {"tunnelPorts": [CODE_SERVER_PORT, 3000]}, - CODE_SERVER_PORT, - NOVNC_PORT, - TTYD_PROXY_PORT, - ) - assert exposed == [CODE_SERVER_PORT, 3000] - assert tunnel == [3000] - - def test_custom_code_server_port_frees_default_for_tunnel(self): - # code-server moved to 8081 → the default 8080 is free as a user tunnel. - exposed, tunnel = SandboxManager._collect_exposed_ports( - True, - False, - False, - {"tunnelPorts": [CODE_SERVER_PORT]}, - 8081, - NOVNC_PORT, - TTYD_PROXY_PORT, - ) - assert exposed == [8081, CODE_SERVER_PORT] - assert tunnel == [CODE_SERVER_PORT] - - def test_custom_terminal_port_frees_default_for_tunnel(self): - exposed, tunnel = SandboxManager._collect_exposed_ports( - False, - False, - True, - {"tunnelPorts": [TTYD_PROXY_PORT, 3000]}, - CODE_SERVER_PORT, - NOVNC_PORT, - 7000, - ) - assert exposed == [7000, TTYD_PROXY_PORT, 3000] - assert tunnel == [TTYD_PROXY_PORT, 3000] + {"terminalEnabled": True, "terminalPort": 7000, "tunnelPorts": [TTYD_PROXY_PORT, 3000]}, + [7000, TTYD_PROXY_PORT, 3000], + f"{TTYD_PROXY_PORT},3000", + ), + ], +) +def test_exposed_ports_and_runtime_expectations_agree( + code_server, settings, exposed, expected_extras +): + tunnels = SandboxTunnels(code_server_enabled=code_server, settings=settings) + assert tunnels.exposed_ports == exposed + assert tunnels.environment.get(EXPECTED_TUNNEL_PORTS_ENV_VAR) == expected_extras class TestValidatePorts: - """SandboxManager._validate_ports tests.""" + """SandboxTunnels._validate_ports tests.""" def test_accepts_valid_ports(self): - assert SandboxManager._validate_ports([80, 3000, 65535]) == [80, 3000, 65535] + assert SandboxTunnels._validate_ports([80, 3000, 65535]) == [80, 3000, 65535] def test_rejects_out_of_range(self): - assert SandboxManager._validate_ports([0, -1, 65536, 3000]) == [3000] + assert SandboxTunnels._validate_ports([0, -1, 65536, 3000]) == [3000] def test_rejects_non_integers(self): - assert SandboxManager._validate_ports(["3000", 3.5, None, 8080]) == [8080] + assert SandboxTunnels._validate_ports(["3000", 3.5, None, 8080]) == [8080] def test_caps_at_ten(self): ports = list(range(1, 20)) - assert len(SandboxManager._validate_ports(ports)) == 10 + assert len(SandboxTunnels._validate_ports(ports)) == 10 def test_empty_list(self): - assert SandboxManager._validate_ports([]) == [] + assert SandboxTunnels._validate_ports([]) == [] def _patch_sandbox_create(monkeypatch, captured: dict) -> None: @@ -623,40 +557,40 @@ class FakeSandbox: return FakeSandbox() fake_create_aio.aio = fake_create_aio - monkeypatch.setattr("src.sandbox.manager.modal.Sandbox.create", fake_create_aio) + monkeypatch.setattr("src.sandbox.launch.modal.Sandbox.create", fake_create_aio) monkeypatch.setattr( - SandboxManager, - "_resolve_and_setup_tunnels", - AsyncMock(return_value=(None, None, None, None)), + SandboxTunnels, + "resolve", + AsyncMock(return_value=TunnelUrls(None, None, None, None)), ) class TestResolveServicePorts: - """SandboxManager._resolve_service_ports tests.""" + """SandboxTunnels._resolve_service_ports tests.""" def test_defaults_when_unset(self): - assert SandboxManager._resolve_service_ports(None) == ( + assert SandboxTunnels._resolve_service_ports(None) == ( CODE_SERVER_PORT, NOVNC_PORT, TTYD_PROXY_PORT, ) - assert SandboxManager._resolve_service_ports({}) == ( + assert SandboxTunnels._resolve_service_ports({}) == ( CODE_SERVER_PORT, NOVNC_PORT, TTYD_PROXY_PORT, ) def test_uses_configured_ports(self): - assert SandboxManager._resolve_service_ports( + assert SandboxTunnels._resolve_service_ports( {"codeServerPort": 9000, "vncPort": 9001, "terminalPort": 9002} ) == (9000, 9001, 9002) def test_falls_back_on_invalid(self): - assert SandboxManager._resolve_service_ports( + assert SandboxTunnels._resolve_service_ports( {"codeServerPort": 0, "vncPort": -1, "terminalPort": 99999} ) == (CODE_SERVER_PORT, NOVNC_PORT, TTYD_PROXY_PORT) # strings and bools are not valid in-range ints - assert SandboxManager._resolve_service_ports( + assert SandboxTunnels._resolve_service_ports( {"codeServerPort": "8081", "vncPort": False, "terminalPort": True} ) == (CODE_SERVER_PORT, NOVNC_PORT, TTYD_PROXY_PORT) diff --git a/packages/modal-infra/tests/test_vnc.py b/packages/modal-infra/tests/test_vnc.py index cca7044ac6..18e2412789 100644 --- a/packages/modal-infra/tests/test_vnc.py +++ b/packages/modal-infra/tests/test_vnc.py @@ -5,6 +5,7 @@ import pytest from sandbox_runtime.constants import ( + EXPECTED_TUNNEL_PORTS_ENV_VAR, NOVNC_PORT, NOVNC_PORT_ENV_VAR, VNC_PASSWORD_ENV_VAR, @@ -12,6 +13,7 @@ VNC_PORT, ) from src.sandbox.manager import CODE_SERVER_PORT, TTYD_PROXY_PORT, SandboxConfig, SandboxManager +from src.sandbox.tunnels import SandboxTunnels, TunnelUrls def _patch_sandbox_create(monkeypatch, captured: dict) -> None: @@ -27,7 +29,7 @@ class FakeSandbox: fake_create = MagicMock() fake_create.aio = fake_create_aio - monkeypatch.setattr("src.sandbox.manager.modal.Sandbox.create", fake_create) + monkeypatch.setattr("src.sandbox.launch.modal.Sandbox.create", fake_create) class TestCreateSandboxVnc: @@ -36,9 +38,9 @@ async def test_returns_url_and_password_and_exposes_only_novnc(self, monkeypatch captured = {} _patch_sandbox_create(monkeypatch, captured) monkeypatch.setattr( - SandboxManager, - "_resolve_and_setup_tunnels", - AsyncMock(return_value=(None, "https://vnc.example.com", None, None)), + SandboxTunnels, + "resolve", + AsyncMock(return_value=TunnelUrls(None, "https://vnc.example.com", None, None)), ) handle = await SandboxManager().create_sandbox( @@ -63,9 +65,9 @@ async def test_disabled_vnc_has_no_credentials_or_port(self, monkeypatch): captured = {} _patch_sandbox_create(monkeypatch, captured) monkeypatch.setattr( - SandboxManager, - "_resolve_and_setup_tunnels", - AsyncMock(return_value=(None, None, None, None)), + SandboxTunnels, + "resolve", + AsyncMock(return_value=TunnelUrls(None, None, None, None)), ) handle = await SandboxManager().create_sandbox( @@ -84,11 +86,13 @@ class TestRestoreSandboxVnc: async def test_generates_credentials_and_returns_them_with_url(self, monkeypatch): captured = {} _patch_sandbox_create(monkeypatch, captured) - monkeypatch.setattr("src.sandbox.manager.modal.Image.from_id", lambda *_args: MagicMock()) + monkeypatch.setattr("src.sandbox.launch.modal.Image.from_id", lambda *_args: MagicMock()) monkeypatch.setattr( - SandboxManager, - "_resolve_and_setup_tunnels", - AsyncMock(return_value=(None, "https://restored-vnc.example.com", None, None)), + SandboxTunnels, + "resolve", + AsyncMock( + return_value=TunnelUrls(None, "https://restored-vnc.example.com", None, None) + ), ) handle = await SandboxManager().restore_from_snapshot( @@ -108,37 +112,28 @@ async def test_generates_credentials_and_returns_them_with_url(self, monkeypatch async def test_resolves_custom_novnc_tunnel(): sandbox = MagicMock() with patch.object( - SandboxManager, + SandboxTunnels, "_resolve_tunnels", new_callable=AsyncMock, return_value={6081: "https://vnc.example.com"}, ) as resolve_tunnels: - result = await SandboxManager._resolve_and_setup_tunnels( - sandbox, - "sandbox-vnc", - False, - True, - False, - [], - code_server_port=CODE_SERVER_PORT, - novnc_port=6081, - ttyd_proxy_port=TTYD_PROXY_PORT, - ) + result = await SandboxTunnels( + code_server_enabled=False, + vnc_enabled=True, + settings={ + "terminalEnabled": False, + "tunnelPorts": [], + "codeServerPort": CODE_SERVER_PORT, + "vncPort": 6081, + "terminalPort": TTYD_PROXY_PORT, + }, + ).resolve(sandbox, "sandbox-vnc") resolve_tunnels.assert_awaited_once_with(sandbox, "sandbox-vnc", [6081]) assert result == (None, "https://vnc.example.com", None, None) def test_raw_vnc_port_is_never_exposed_as_an_extra_tunnel(): - exposed, extras = SandboxManager._collect_exposed_ports( - False, - False, - False, - {"tunnelPorts": [VNC_PORT, 3000]}, - CODE_SERVER_PORT, - NOVNC_PORT, - TTYD_PROXY_PORT, - ) - - assert exposed == [3000] - assert extras == [3000] + tunnels = SandboxTunnels(settings={"tunnelPorts": [VNC_PORT, 3000]}) + assert tunnels.exposed_ports == [3000] + assert tunnels.environment[EXPECTED_TUNNEL_PORTS_ENV_VAR] == "3000"