From 63546b2c8ecaf1a499e096361446adb38645e65c Mon Sep 17 00:00:00 2001 From: Cole Murray Date: Tue, 22 Sep 2026 00:58:49 -0700 Subject: [PATCH 1/6] refactor: split Modal sandbox manager responsibilities --- docs/plans/modal-sandbox-manager-refactor.md | 94 ++++ packages/modal-infra/src/sandbox/launch.py | 225 ++++++++ packages/modal-infra/src/sandbox/manager.py | 524 ++---------------- packages/modal-infra/src/sandbox/models.py | 53 ++ packages/modal-infra/src/sandbox/tunnels.py | 213 +++++++ .../tests/test_agent_slack_notify_env.py | 11 +- .../modal-infra/tests/test_code_server.py | 65 +-- .../modal-infra/tests/test_llm_secrets.py | 4 +- .../tests/test_sandbox_env_vars.py | 59 +- .../modal-infra/tests/test_sandbox_launch.py | 184 ++++-- .../tests/test_sandbox_resources.py | 22 +- packages/modal-infra/tests/test_ttyd.py | 170 +++--- .../modal-infra/tests/test_tunnel_ports.py | 439 +++++++-------- packages/modal-infra/tests/test_vnc.py | 69 ++- 14 files changed, 1135 insertions(+), 997 deletions(-) create mode 100644 docs/plans/modal-sandbox-manager-refactor.md create mode 100644 packages/modal-infra/src/sandbox/launch.py create mode 100644 packages/modal-infra/src/sandbox/models.py create mode 100644 packages/modal-infra/src/sandbox/tunnels.py diff --git a/docs/plans/modal-sandbox-manager-refactor.md b/docs/plans/modal-sandbox-manager-refactor.md new file mode 100644 index 0000000000..b16102a1c3 --- /dev/null +++ b/docs/plans/modal-sandbox-manager-refactor.md @@ -0,0 +1,94 @@ +# Modal sandbox manager refactor + +## Assessment + +At base commit `232bb74c5`, `packages/modal-infra/src/sandbox/manager.py` is 722 lines. Its largest +responsibilities are launch translation and tunnel setup, not lifecycle operations. + +| Responsibility | Before | Owner after refactoring | +| ------------------------------------------------------------------------------- | ----------------------------------------------- | -------------------------------------------------- | +| Configuration and returned sandbox records | `SandboxConfig`, `SandboxHandle` in manager | `models.py`; existing manager imports remain valid | +| Create/restore normalization, repository validation, operation logging | `create_sandbox`, `restore_from_snapshot` | `SandboxManager` | +| Base/repository/snapshot image resolution and missing-image classification | `_launch_sandbox` | `SandboxLauncher` | +| Environment precedence, reserved keys, session serialization, VCS compatibility | `_launch_sandbox` | `SandboxLauncher` | +| Credentials, resource translation, Modal creation, handle assembly | Password/resource helpers and `_launch_sandbox` | `SandboxLauncher` | +| Service/extra-port ownership, retries, URL routing, tunnel-file publication | Six networking helpers plus launch assembly | `SandboxTunnels` | +| Bounded filesystem snapshot capture | `take_snapshot` | `SandboxManager` | +| Provider lookup and confirmed termination | `get_sandbox_by_id`, `stop_sandbox` | `SandboxManager` | + +### Structural problems + +- **Divergent change:** provider image, environment, networking, and lifecycle changes all require + editing the same class. Extract the substantial launch and networking decisions. +- **Duplicated knowledge:** exposed service ports are reconstructed for URL resolution. Compute + ownership once and use it for exposure, runtime environment, and returned URLs. +- **Leaky interface:** tunnel setup requires nine arguments and returns an anonymous tuple. Bind the + launch's port configuration to one object and return named URL fields. +- **Dependency direction:** moving collaborators without moving shared records would make them + import their manager. Put configuration and handles in a dependency-leaf module. + +## Implementation plan + +1. Establish the existing Modal test-suite baseline. +2. Move configuration/handle records to `models.py`, explicitly preserving public manager exports. +3. Extract `SandboxTunnels`. Its constructor resolves port ownership; `environment` and + `exposed_ports` describe the launch; `resolve` handles best-effort URL resolution/publication. +4. Extract `SandboxLauncher`, retaining one shared launch path and the existing typed image-source + variants. Keep image failures, credential generation, and environment precedence unchanged. +5. Retain create/restore normalization, snapshots, lookup, and termination in `SandboxManager`. +6. Migrate existing tests to the owning modules and exercise the complete manager/launcher/tunnel + path with only provider I/O mocked. +7. Run the complete Modal suite, Ruff lint/format, a base-versus-head MyPy comparison, and PR CI. + +The dependency direction is `manager -> launch -> tunnels`; manager and launch also use `models`. +Collaborators never import the manager. No generic provider interface, registry, +dependency-injection framework, or new lifecycle authority is introduced. + +## Compatibility and verification + +- Preserve all five public manager method signatures and configuration/handle fields. +- Preserve fresh/repository/snapshot environment rules, unknown session-config fields, legacy + restore credentials, generated passwords, resource settings, and image-error classification. +- Preserve partial tunnel results, retry delays, non-fatal file-write failures, disabled-service + port ownership, and raw-VNC exclusion from extra tunnels. +- Preserve snapshot deadline rounding/capping and wait-for-termination behavior. +- Keep existing logger names and event identifiers. +- Verify no fallback image or automatic spawn retry is introduced by error handling. +- This is a provider-local refactor; it does not change the control-plane lifecycle policy, runtime + protocol, deployment configuration, or image-build sandbox service. +- Mocked provider tests establish translation and orchestration behavior, not live Modal behavior. + No production deployment or billable provider canary is part of this change. + +## Simplification Analysis + +### Core Purpose + +Separate launch translation and networking from existing-sandbox lifecycle operations. + +### Unnecessary Complexity Found + +- The old tunnel helpers repeated service-port selection and passed that knowledge between methods. + `SandboxTunnels` now owns the selection and its consumers. + +### Code to Remove + +- Remove the manager's networking helpers and embedded launch implementation by extracting their + actual responsibilities; do not retain private forwarding wrappers. +- Remove duplicated port reconstruction rather than adding another configuration layer. + +### Simplification Recommendations + +1. Keep the small snapshot, lookup, and termination operations in the manager. +2. Keep the existing common launch path and typed image variants. +3. Use concrete collaborators and named results, with no speculative extension points. + +### YAGNI Violations + +None introduced. Small lifecycle operations do not warrant additional service classes. + +### Final Assessment + +`manager.py` is 264 lines after refactoring: 458 fewer lines (63% smaller). The two collaborators +are 225 and 213 lines; shared records occupy 53 lines. Across these four production files, the total +grows by 33 lines for module boundaries, explicit exports, and the named tunnel result. Complexity +is low; proceed with this split. diff --git a/packages/modal-infra/src/sandbox/launch.py b/packages/modal-infra/src/sandbox/launch.py new file mode 100644 index 0000000000..3b354158c1 --- /dev/null +++ b/packages/modal-infra/src/sandbox/launch.py @@ -0,0 +1,225 @@ +"""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 ( + NOVNC_PORT_ENV_VAR, + SANDBOX_TIMEOUT_ENV_VAR, + VNC_PASSWORD_ENV_VAR, + VNC_PASSWORD_MAX_BYTES, +) +from sandbox_runtime.types import SandboxStatus + +from ..app import app, llm_secrets +from ..images.base import base_image +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, +} + + +class RepositoryImageUnavailableError(RuntimeError): + """The selected repository image no longer exists in Modal.""" + + +def _resource_kwargs(settings: dict[str, Any] | None) -> dict[str, Any]: + """Map sandbox settings to Modal resource kwargs. + + `cpuCores` -> Modal `cpu` (cores, fractional allowed), `memoryMib` -> Modal + `memory` (MiB). The control plane owns normalization; this only maps + already-normalized settings into provider-specific argument names. + """ + if not settings: + return {} + + kwargs: dict[str, Any] = {} + + cpu_cores = settings.get("cpuCores") + if cpu_cores is not None: + kwargs["cpu"] = float(cpu_cores) + + memory_mib = settings.get("memoryMib") + if memory_mib is not None: + kwargs["memory"] = memory_mib + + return kwargs + + +@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 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)}" + + 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 "", + } + ) + + clone_token: str | None = None + include_github_cli_aliases = False + snapshot_id: str | None = None + if isinstance(spec.source, BaseImageSource): + image = 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) + + create_kwargs: dict[str, Any] = { + "image": image, + "app": app, + "secrets": [llm_secrets], + "timeout": config.timeout_seconds, + "workdir": "/workspace", + "env": env_vars, + **_resource_kwargs(config.settings), + } + if tunnels.exposed_ports: + create_kwargs["encrypted_ports"] = tunnels.exposed_ports + + try: + sandbox = await modal.Sandbox.create.aio( + "python", + "-m", + "sandbox_runtime.entrypoint", + **create_kwargs, + ) + except modal.exception.NotFoundError as e: + if isinstance(spec.source, RepositoryImageSource): + raise RepositoryImageUnavailableError("repository image is unavailable") from e + raise + 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, + ) diff --git a/packages/modal-infra/src/sandbox/manager.py b/packages/modal-infra/src/sandbox/manager.py index 45aaacfd05..b197975a3c 100644 --- a/packages/modal-infra/src/sandbox/manager.py +++ b/packages/modal-infra/src/sandbox/manager.py @@ -1,65 +1,39 @@ -""" -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 -from sandbox_runtime.constants import ( - CODE_SERVER_PORT, - CODE_SERVER_PORT_ENV_VAR, - DEFAULT_SANDBOX_TIMEOUT_SECONDS, - EXPECTED_TUNNEL_PORTS_ENV_VAR, - NOVNC_PORT, - NOVNC_PORT_ENV_VAR, - SANDBOX_TIMEOUT_ENV_VAR, - 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, -) +from sandbox_runtime.constants import DEFAULT_SANDBOX_TIMEOUT_SECONDS from sandbox_runtime.log_config import get_logger from sandbox_runtime.types import SandboxStatus, SessionConfig -from ..app import app, llm_secrets -from ..images.base import base_image -from .vcs_env import inject_vcs_env_vars +from .launch import ( + BaseImageSource, + RepositoryImageSource, + RepositoryImageUnavailableError, + SandboxImageSource, + SandboxLauncher, + SandboxLaunchSpec, + SnapshotImageSource, +) +from .models import DEFAULT_VNC_ENABLED, SandboxConfig, SandboxHandle + +# Preserve the existing public imports after moving their implementations. +__all__ = [ + "DEFAULT_SANDBOX_TIMEOUT_SECONDS", + "DEFAULT_VNC_ENABLED", + "SNAPSHOT_FILESYSTEM_TIMEOUT_SECONDS", + "RepositoryImageUnavailableError", + "SandboxConfig", + "SandboxHandle", + "SandboxManager", +] log = get_logger("manager") SNAPSHOT_FILESYSTEM_TIMEOUT_SECONDS = 300 -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, -} - - -class RepositoryImageUnavailableError(RuntimeError): - """The selected repository image no longer exists in Modal.""" def _has_repository(repo_owner: str | None, repo_name: str | None) -> bool: @@ -70,445 +44,13 @@ def _has_repository(repo_owner: str | None, repo_name: str | None) -> bool: return has_owner -def _resource_kwargs(settings: dict[str, Any] | None) -> dict: - """Map sandbox settings to Modal resource kwargs. - - `cpuCores` -> Modal `cpu` (cores, fractional allowed), `memoryMib` -> Modal - `memory` (MiB). The control plane owns normalization; this only maps - already-normalized settings into provider-specific argument names. - """ - if not settings: - return {} - - kwargs: dict = {} - - cpu_cores = settings.get("cpuCores") - if cpu_cores is not None: - kwargs["cpu"] = float(cpu_cores) - - memory_mib = settings.get("memoryMib") - if memory_mib is not None: - kwargs["memory"] = memory_mib - - return kwargs - - -@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 - ) - - -@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 - - -@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)}" - - 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 "", - } - ) - - clone_token: str | None = None - include_github_cli_aliases = False - snapshot_id: str | None = None - if isinstance(spec.source, _BaseImageSource): - image = 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 - - 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) - - create_kwargs: dict[str, Any] = { - "image": image, - "app": app, - "secrets": [llm_secrets], - "timeout": config.timeout_seconds, - "workdir": "/workspace", - "env": env_vars, - **_resource_kwargs(config.settings), - } - if exposed_ports: - create_kwargs["encrypted_ports"] = exposed_ports - - try: - sandbox = await modal.Sandbox.create.aio( - "python", - "-m", - "sandbox_runtime.entrypoint", - **create_kwargs, - ) - except modal.exception.NotFoundError as e: - if isinstance(spec.source, _RepositoryImageSource): - raise RepositoryImageUnavailableError("repository image is unavailable") from e - raise - 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, - ) - async def create_sandbox( self, config: SandboxConfig, @@ -518,7 +60,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 @@ -530,14 +72,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( @@ -684,8 +226,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, @@ -700,7 +242,7 @@ async def restore_from_snapshot( agent_slack_notify_enabled=agent_slack_notify_enabled, settings=settings, ), - 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..0ebfbd8cc1 --- /dev/null +++ b/packages/modal-infra/src/sandbox/models.py @@ -0,0 +1,53 @@ +"""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 + +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 + ) + + +@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 diff --git a/packages/modal-infra/src/sandbox/tunnels.py b/packages/modal-infra/src/sandbox/tunnels.py new file mode 100644 index 0000000000..fb075864a9 --- /dev/null +++ b/packages/modal-infra/src/sandbox/tunnels.py @@ -0,0 +1,213 @@ +"""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 + + +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 = 3, + backoff_seconds: 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_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 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..9ee0f62fc8 100644 --- a/packages/modal-infra/tests/test_code_server.py +++ b/packages/modal-infra/tests/test_code_server.py @@ -4,25 +4,28 @@ import pytest -from src.sandbox.manager import CODE_SERVER_PORT, SandboxConfig, SandboxManager +from sandbox_runtime.constants import CODE_SERVER_PORT +from src.sandbox.launch import SandboxLauncher +from src.sandbox.manager import 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 +35,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 +43,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 +55,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 +72,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 +99,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 +143,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 +190,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 +241,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 20d1e373ad..4befa6fb49 100644 --- a/packages/modal-infra/tests/test_llm_secrets.py +++ b/packages/modal-infra/tests/test_llm_secrets.py @@ -27,7 +27,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 @@ -41,7 +41,7 @@ async def test_restore_attaches_the_deployment_wide_secret(captured_launch, monk 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_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 e2f7e2d06b..be83dac1a5 100644 --- a/packages/modal-infra/tests/test_sandbox_launch.py +++ b/packages/modal-infra/tests/test_sandbox_launch.py @@ -5,15 +5,19 @@ from unittest.mock import AsyncMock, Mock import pytest +from modal.exception import NotFoundError from sandbox_runtime.constants import ( CODE_SERVER_PORT_ENV_VAR, 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.manager import ( RepositoryImageUnavailableError, SandboxConfig, @@ -25,7 +29,18 @@ 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 @@ -42,27 +57,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, + SandboxLauncher, "_generate_code_server_password", staticmethod(lambda: "code-password") ) - monkeypatch.setattr( - SandboxManager, "_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 = { @@ -176,23 +178,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( @@ -203,18 +199,128 @@ 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): - 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) + ) - monkeypatch.setattr("src.sandbox.manager.modal.Image.from_id", lambda _image_id: object()) - create = SimpleNamespace(aio=AsyncMock(side_effect=NotFoundError("image not found"))) - monkeypatch.setattr("src.sandbox.manager.modal.Sandbox.create", create) + 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") + ) - with pytest.raises(RepositoryImageUnavailableError): - await SandboxManager().create_sandbox( + 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="repo-image-missing", + repo_image_id="image-1" if image_source == "repository" else None, + **common, ) ) + + 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, + ) diff --git a/packages/modal-infra/tests/test_sandbox_resources.py b/packages/modal-infra/tests/test_sandbox_resources.py index cc2aa8b105..8636c70d23 100644 --- a/packages/modal-infra/tests/test_sandbox_resources.py +++ b/packages/modal-infra/tests/test_sandbox_resources.py @@ -4,7 +4,9 @@ import pytest -from src.sandbox.manager import SandboxConfig, SandboxManager, _resource_kwargs +from src.sandbox.launch import _resource_kwargs +from src.sandbox.manager import SandboxConfig, SandboxManager +from src.sandbox.tunnels import SandboxTunnels, TunnelUrls class TestResourceKwargs: @@ -46,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() @@ -73,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..db8419fb09 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 src.sandbox.manager import ( +from sandbox_runtime.constants import ( CODE_SERVER_PORT, + EXPECTED_TUNNEL_PORTS_ENV_VAR, + NOVNC_PORT, + TTYD_PORT, TTYD_PROXY_PORT, +) +from src.sandbox.manager import ( 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..087d2437d4 100644 --- a/packages/modal-infra/tests/test_tunnel_ports.py +++ b/packages/modal-infra/tests/test_tunnel_ports.py @@ -5,6 +5,7 @@ import pytest from sandbox_runtime.constants import ( + CODE_SERVER_PORT, CODE_SERVER_PORT_ENV_VAR, EXPECTED_TUNNEL_PORTS_ENV_VAR, NOVNC_PORT, @@ -13,7 +14,8 @@ TUNNEL_ENV_FILE_PATH, TUNNEL_ENV_SANDBOX_ID_KEY, ) -from src.sandbox.manager import CODE_SERVER_PORT, SandboxConfig, SandboxManager +from src.sandbox.manager import SandboxConfig, SandboxManager +from src.sandbox.tunnels import SandboxTunnels, TunnelUrls def _mock_sandbox_with_filesystem() -> tuple[MagicMock, AsyncMock]: @@ -27,7 +29,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 +41,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 +55,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 +66,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 +85,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 +102,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 +124,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 +156,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 +185,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 +217,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 +269,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 +287,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 +314,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 +340,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 +366,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 +405,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 +437,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 +469,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 +488,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 +558,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..86a01856eb 100644 --- a/packages/modal-infra/tests/test_vnc.py +++ b/packages/modal-infra/tests/test_vnc.py @@ -5,13 +5,17 @@ import pytest from sandbox_runtime.constants import ( + CODE_SERVER_PORT, + EXPECTED_TUNNEL_PORTS_ENV_VAR, NOVNC_PORT, NOVNC_PORT_ENV_VAR, + TTYD_PROXY_PORT, VNC_PASSWORD_ENV_VAR, VNC_PASSWORD_MAX_BYTES, VNC_PORT, ) -from src.sandbox.manager import CODE_SERVER_PORT, TTYD_PROXY_PORT, SandboxConfig, SandboxManager +from src.sandbox.manager import SandboxConfig, SandboxManager +from src.sandbox.tunnels import SandboxTunnels, TunnelUrls def _patch_sandbox_create(monkeypatch, captured: dict) -> None: @@ -27,7 +31,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 +40,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 +67,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 +88,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 +114,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" From 0dc066e3a2e5f3e5b28b10818ed646b9764f0139 Mon Sep 17 00:00:00 2001 From: Cole Murray Date: Tue, 22 Sep 2026 01:05:21 -0700 Subject: [PATCH 2/6] docs: keep sandbox manager refactor plan local --- docs/plans/modal-sandbox-manager-refactor.md | 94 -------------------- 1 file changed, 94 deletions(-) delete mode 100644 docs/plans/modal-sandbox-manager-refactor.md diff --git a/docs/plans/modal-sandbox-manager-refactor.md b/docs/plans/modal-sandbox-manager-refactor.md deleted file mode 100644 index b16102a1c3..0000000000 --- a/docs/plans/modal-sandbox-manager-refactor.md +++ /dev/null @@ -1,94 +0,0 @@ -# Modal sandbox manager refactor - -## Assessment - -At base commit `232bb74c5`, `packages/modal-infra/src/sandbox/manager.py` is 722 lines. Its largest -responsibilities are launch translation and tunnel setup, not lifecycle operations. - -| Responsibility | Before | Owner after refactoring | -| ------------------------------------------------------------------------------- | ----------------------------------------------- | -------------------------------------------------- | -| Configuration and returned sandbox records | `SandboxConfig`, `SandboxHandle` in manager | `models.py`; existing manager imports remain valid | -| Create/restore normalization, repository validation, operation logging | `create_sandbox`, `restore_from_snapshot` | `SandboxManager` | -| Base/repository/snapshot image resolution and missing-image classification | `_launch_sandbox` | `SandboxLauncher` | -| Environment precedence, reserved keys, session serialization, VCS compatibility | `_launch_sandbox` | `SandboxLauncher` | -| Credentials, resource translation, Modal creation, handle assembly | Password/resource helpers and `_launch_sandbox` | `SandboxLauncher` | -| Service/extra-port ownership, retries, URL routing, tunnel-file publication | Six networking helpers plus launch assembly | `SandboxTunnels` | -| Bounded filesystem snapshot capture | `take_snapshot` | `SandboxManager` | -| Provider lookup and confirmed termination | `get_sandbox_by_id`, `stop_sandbox` | `SandboxManager` | - -### Structural problems - -- **Divergent change:** provider image, environment, networking, and lifecycle changes all require - editing the same class. Extract the substantial launch and networking decisions. -- **Duplicated knowledge:** exposed service ports are reconstructed for URL resolution. Compute - ownership once and use it for exposure, runtime environment, and returned URLs. -- **Leaky interface:** tunnel setup requires nine arguments and returns an anonymous tuple. Bind the - launch's port configuration to one object and return named URL fields. -- **Dependency direction:** moving collaborators without moving shared records would make them - import their manager. Put configuration and handles in a dependency-leaf module. - -## Implementation plan - -1. Establish the existing Modal test-suite baseline. -2. Move configuration/handle records to `models.py`, explicitly preserving public manager exports. -3. Extract `SandboxTunnels`. Its constructor resolves port ownership; `environment` and - `exposed_ports` describe the launch; `resolve` handles best-effort URL resolution/publication. -4. Extract `SandboxLauncher`, retaining one shared launch path and the existing typed image-source - variants. Keep image failures, credential generation, and environment precedence unchanged. -5. Retain create/restore normalization, snapshots, lookup, and termination in `SandboxManager`. -6. Migrate existing tests to the owning modules and exercise the complete manager/launcher/tunnel - path with only provider I/O mocked. -7. Run the complete Modal suite, Ruff lint/format, a base-versus-head MyPy comparison, and PR CI. - -The dependency direction is `manager -> launch -> tunnels`; manager and launch also use `models`. -Collaborators never import the manager. No generic provider interface, registry, -dependency-injection framework, or new lifecycle authority is introduced. - -## Compatibility and verification - -- Preserve all five public manager method signatures and configuration/handle fields. -- Preserve fresh/repository/snapshot environment rules, unknown session-config fields, legacy - restore credentials, generated passwords, resource settings, and image-error classification. -- Preserve partial tunnel results, retry delays, non-fatal file-write failures, disabled-service - port ownership, and raw-VNC exclusion from extra tunnels. -- Preserve snapshot deadline rounding/capping and wait-for-termination behavior. -- Keep existing logger names and event identifiers. -- Verify no fallback image or automatic spawn retry is introduced by error handling. -- This is a provider-local refactor; it does not change the control-plane lifecycle policy, runtime - protocol, deployment configuration, or image-build sandbox service. -- Mocked provider tests establish translation and orchestration behavior, not live Modal behavior. - No production deployment or billable provider canary is part of this change. - -## Simplification Analysis - -### Core Purpose - -Separate launch translation and networking from existing-sandbox lifecycle operations. - -### Unnecessary Complexity Found - -- The old tunnel helpers repeated service-port selection and passed that knowledge between methods. - `SandboxTunnels` now owns the selection and its consumers. - -### Code to Remove - -- Remove the manager's networking helpers and embedded launch implementation by extracting their - actual responsibilities; do not retain private forwarding wrappers. -- Remove duplicated port reconstruction rather than adding another configuration layer. - -### Simplification Recommendations - -1. Keep the small snapshot, lookup, and termination operations in the manager. -2. Keep the existing common launch path and typed image variants. -3. Use concrete collaborators and named results, with no speculative extension points. - -### YAGNI Violations - -None introduced. Small lifecycle operations do not warrant additional service classes. - -### Final Assessment - -`manager.py` is 264 lines after refactoring: 458 fewer lines (63% smaller). The two collaborators -are 225 and 213 lines; shared records occupy 53 lines. Across these four production files, the total -grows by 33 lines for module boundaries, explicit exports, and the named tunnel result. Complexity -is low; proceed with this split. From 8031db18f121ecea0dd3a4341ffdfa1deacaef06 Mon Sep 17 00:00:00 2001 From: Cole Murray Date: Tue, 22 Sep 2026 09:44:28 -0700 Subject: [PATCH 3/6] fix: restore sandbox manager constant exports --- packages/modal-infra/src/sandbox/manager.py | 32 ++++++++++++++++++++- 1 file changed, 31 insertions(+), 1 deletion(-) diff --git a/packages/modal-infra/src/sandbox/manager.py b/packages/modal-infra/src/sandbox/manager.py index b197975a3c..c63776b3fa 100644 --- a/packages/modal-infra/src/sandbox/manager.py +++ b/packages/modal-infra/src/sandbox/manager.py @@ -5,7 +5,22 @@ import modal -from sandbox_runtime.constants import DEFAULT_SANDBOX_TIMEOUT_SECONDS +from sandbox_runtime.constants import ( + CODE_SERVER_PORT, + CODE_SERVER_PORT_ENV_VAR, + DEFAULT_SANDBOX_TIMEOUT_SECONDS, + EXPECTED_TUNNEL_PORTS_ENV_VAR, + NOVNC_PORT, + NOVNC_PORT_ENV_VAR, + SANDBOX_TIMEOUT_ENV_VAR, + 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, +) from sandbox_runtime.log_config import get_logger from sandbox_runtime.types import SandboxStatus, SessionConfig @@ -19,12 +34,27 @@ SnapshotImageSource, ) from .models import DEFAULT_VNC_ENABLED, SandboxConfig, SandboxHandle +from .tunnels import MAX_TUNNEL_PORTS # Preserve the existing public imports after moving their implementations. __all__ = [ + "CODE_SERVER_PORT", + "CODE_SERVER_PORT_ENV_VAR", "DEFAULT_SANDBOX_TIMEOUT_SECONDS", "DEFAULT_VNC_ENABLED", + "EXPECTED_TUNNEL_PORTS_ENV_VAR", + "MAX_TUNNEL_PORTS", + "NOVNC_PORT", + "NOVNC_PORT_ENV_VAR", + "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", From ddcb99a745fd7246e45e4b0b291349b0bf94970e Mon Sep 17 00:00:00 2001 From: Cole Murray Date: Tue, 22 Sep 2026 09:46:02 -0700 Subject: [PATCH 4/6] test: cover legacy sandbox manager imports --- .../modal-infra/tests/test_code_server.py | 3 +- .../modal-infra/tests/test_manager_exports.py | 36 +++++++++++++++++++ packages/modal-infra/tests/test_ttyd.py | 4 +-- .../modal-infra/tests/test_tunnel_ports.py | 3 +- packages/modal-infra/tests/test_vnc.py | 4 +-- 5 files changed, 41 insertions(+), 9 deletions(-) create mode 100644 packages/modal-infra/tests/test_manager_exports.py diff --git a/packages/modal-infra/tests/test_code_server.py b/packages/modal-infra/tests/test_code_server.py index 9ee0f62fc8..302d8c7675 100644 --- a/packages/modal-infra/tests/test_code_server.py +++ b/packages/modal-infra/tests/test_code_server.py @@ -4,9 +4,8 @@ import pytest -from sandbox_runtime.constants import CODE_SERVER_PORT from src.sandbox.launch import SandboxLauncher -from src.sandbox.manager import SandboxConfig, SandboxManager +from src.sandbox.manager import CODE_SERVER_PORT, SandboxConfig, SandboxManager from src.sandbox.tunnels import SandboxTunnels, TunnelUrls 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..c244131a1d --- /dev/null +++ b/packages/modal-infra/tests/test_manager_exports.py @@ -0,0 +1,36 @@ +"""Compatibility coverage for the manager's pre-refactor constant imports.""" + +import pytest + +from sandbox_runtime import constants +from src.sandbox import manager +from src.sandbox.models import DEFAULT_VNC_ENABLED +from src.sandbox.tunnels import MAX_TUNNEL_PORTS + + +@pytest.mark.parametrize( + "name, expected", + [ + ("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), + ("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), + ("SANDBOX_TIMEOUT_ENV_VAR", constants.SANDBOX_TIMEOUT_ENV_VAR), + ("SNAPSHOT_FILESYSTEM_TIMEOUT_SECONDS", 300), + ("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_ttyd.py b/packages/modal-infra/tests/test_ttyd.py index db8419fb09..3ddd6ae978 100644 --- a/packages/modal-infra/tests/test_ttyd.py +++ b/packages/modal-infra/tests/test_ttyd.py @@ -5,13 +5,13 @@ import pytest from sandbox_runtime.constants import ( - CODE_SERVER_PORT, EXPECTED_TUNNEL_PORTS_ENV_VAR, NOVNC_PORT, TTYD_PORT, - TTYD_PROXY_PORT, ) from src.sandbox.manager import ( + CODE_SERVER_PORT, + TTYD_PROXY_PORT, SandboxConfig, SandboxManager, ) diff --git a/packages/modal-infra/tests/test_tunnel_ports.py b/packages/modal-infra/tests/test_tunnel_ports.py index 087d2437d4..ead6c6b51a 100644 --- a/packages/modal-infra/tests/test_tunnel_ports.py +++ b/packages/modal-infra/tests/test_tunnel_ports.py @@ -5,7 +5,6 @@ import pytest from sandbox_runtime.constants import ( - CODE_SERVER_PORT, CODE_SERVER_PORT_ENV_VAR, EXPECTED_TUNNEL_PORTS_ENV_VAR, NOVNC_PORT, @@ -14,7 +13,7 @@ TUNNEL_ENV_FILE_PATH, TUNNEL_ENV_SANDBOX_ID_KEY, ) -from src.sandbox.manager import SandboxConfig, SandboxManager +from src.sandbox.manager import CODE_SERVER_PORT, SandboxConfig, SandboxManager from src.sandbox.tunnels import SandboxTunnels, TunnelUrls diff --git a/packages/modal-infra/tests/test_vnc.py b/packages/modal-infra/tests/test_vnc.py index 86a01856eb..18e2412789 100644 --- a/packages/modal-infra/tests/test_vnc.py +++ b/packages/modal-infra/tests/test_vnc.py @@ -5,16 +5,14 @@ import pytest from sandbox_runtime.constants import ( - CODE_SERVER_PORT, EXPECTED_TUNNEL_PORTS_ENV_VAR, NOVNC_PORT, NOVNC_PORT_ENV_VAR, - TTYD_PROXY_PORT, VNC_PASSWORD_ENV_VAR, VNC_PASSWORD_MAX_BYTES, VNC_PORT, ) -from src.sandbox.manager import SandboxConfig, SandboxManager +from src.sandbox.manager import CODE_SERVER_PORT, TTYD_PROXY_PORT, SandboxConfig, SandboxManager from src.sandbox.tunnels import SandboxTunnels, TunnelUrls From 958bbc6b8ea3225fd8e71be346fa453fe2c7a3cb Mon Sep 17 00:00:00 2001 From: Cole Murray Date: Tue, 22 Sep 2026 09:47:48 -0700 Subject: [PATCH 5/6] fix: reject boolean sandbox tunnel ports --- packages/modal-infra/src/sandbox/tunnels.py | 2 +- .../modal-infra/tests/test_sandbox_launch.py | 47 +++++++++++++++++++ 2 files changed, 48 insertions(+), 1 deletion(-) diff --git a/packages/modal-infra/src/sandbox/tunnels.py b/packages/modal-infra/src/sandbox/tunnels.py index fb075864a9..3aa16b3018 100644 --- a/packages/modal-infra/src/sandbox/tunnels.py +++ b/packages/modal-infra/src/sandbox/tunnels.py @@ -152,7 +152,7 @@ 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 1 <= p <= 65535: + if isinstance(p, int) and not isinstance(p, bool) and 1 <= p <= 65535: ports.append(p) if len(ports) >= MAX_TUNNEL_PORTS: break diff --git a/packages/modal-infra/tests/test_sandbox_launch.py b/packages/modal-infra/tests/test_sandbox_launch.py index be83dac1a5..830b9bc367 100644 --- a/packages/modal-infra/tests/test_sandbox_launch.py +++ b/packages/modal-infra/tests/test_sandbox_launch.py @@ -324,3 +324,50 @@ async def test_launch_returns_handle_despite_tunnel_failures(monkeypatch, image_ f"{TUNNEL_ENV_SANDBOX_ID_KEY}=sandbox-partial\nTUNNEL_3000=https://app.example\n", TUNNEL_ENV_FILE_PATH, ) + + +@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() From 8788738a3b7bdf400e1e9e59bd14346348e89b8d Mon Sep 17 00:00:00 2001 From: Cole Murray Date: Tue, 22 Sep 2026 09:48:38 -0700 Subject: [PATCH 6/6] refactor: name tunnel resolution retry defaults --- packages/modal-infra/src/sandbox/tunnels.py | 6 ++++-- 1 file changed, 4 insertions(+), 2 deletions(-) diff --git a/packages/modal-infra/src/sandbox/tunnels.py b/packages/modal-infra/src/sandbox/tunnels.py index 3aa16b3018..17210838aa 100644 --- a/packages/modal-infra/src/sandbox/tunnels.py +++ b/packages/modal-infra/src/sandbox/tunnels.py @@ -22,6 +22,8 @@ # 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): @@ -114,8 +116,8 @@ async def _resolve_tunnels( sandbox: modal.Sandbox, sandbox_id: str, ports: list[int], - retries: int = 3, - backoff_seconds: float = 1.0, + 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] = {}