Skip to content
Draft
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension


Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
7 changes: 5 additions & 2 deletions CHANGELOG.md
Original file line number Diff line number Diff line change
Expand Up @@ -1269,8 +1269,11 @@ and this project adheres to [Semantic Versioning](https://semver.org/spec/v2.0.0

- `integrations/inspect-ai` (`inspect-capsem-sandbox`) provides a standalone
[Inspect AI](https://inspect.aisi.org.uk/) `SandboxEnvironment` registered
under `capsem`, supporting direct VM execution
backed by the Capsem Python gateway SDK (`capsem>=0.7.0`).
under `capsem`, supporting direct VM execution (`execution_mode="vm"`) and
rootless OCI workload container execution (`execution_mode="container"`)
with single-service Docker Compose support, `SAMPLE_METADATA_*`
interpolation, and operator-owned host environment and bind-mount
allowlists, backed by the Capsem Python gateway SDK (`capsem>=0.7.0`).

- The profile catalog names its own defaults, one per runtime: the binary
compiles them from `config/profile-catalog.toml` (a runtime's default
Expand Down
107 changes: 107 additions & 0 deletions integrations/inspect-ai/inspect_capsem/_compose.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,107 @@
"""Inspect-facing `CapsemSandboxConfig` coercion and Compose file resolution."""

from __future__ import annotations

from collections.abc import Mapping
from pathlib import Path
from typing import Any

from inspect_ai.util import ComposeConfig, SandboxEnvironmentConfigType

from .config import CapsemSandboxConfig


def _is_dockerfile_string(s: str) -> bool:
name = Path(s).name.lower()
return (
name in ("dockerfile", "containerfile")
or name.startswith(("dockerfile.", "containerfile."))
or name.endswith((".dockerfile", ".containerfile"))
)


def _validate_direct_volumes(cfg: CapsemSandboxConfig) -> CapsemSandboxConfig:
if cfg.volumes:
from .containers import normalize_volumes

normalize_volumes(cfg.volumes, None, cfg.allowed_host_paths)
return cfg


def resolve_compose_file(
cfg: CapsemSandboxConfig, *, sample_metadata: Mapping[str, Any] | None = None
) -> CapsemSandboxConfig:
"""Resolve `cfg.compose_file` and return an updated `CapsemSandboxConfig`."""
if not cfg.compose_file or not cfg.compose_file.strip():
return _validate_direct_volumes(cfg)
from .containers import extract_capsem_compose_fields, parse_host_compose_yaml_file

compose_path = Path(cfg.compose_file)
if not compose_path.is_file():
raise FileNotFoundError(f"Compose file not found: {compose_path}")

def _extract(meta: Mapping[str, Any] | None) -> dict[str, Any]:
parsed = parse_host_compose_yaml_file(
compose_path, allowed_host_env=cfg.allowed_host_env, sample_metadata=meta
)
return extract_capsem_compose_fields(
parsed,
base_dir=compose_path.parent,
allowed_host_env=cfg.allowed_host_env,
allowed_host_paths=cfg.allowed_host_paths,
sample_metadata=meta,
)

overrides = _extract(sample_metadata)
base_overrides = _extract(None) if sample_metadata else {}
explicit = {
k
for k in cfg.model_fields_set
if not sample_metadata or getattr(cfg, k, None) != base_overrides.get(k)
}
data = {
**cfg.model_dump(exclude_unset=True),
**{k: v for k, v in overrides.items() if k not in explicit},
}
return CapsemSandboxConfig(**data)


def coerce_config(
config: SandboxEnvironmentConfigType | Mapping[str, Any] | None,
*,
resolve_compose: bool = True,
sample_metadata: Mapping[str, Any] | None = None,
) -> CapsemSandboxConfig:
"""Coerce an Inspect sandbox config argument into a `CapsemSandboxConfig`."""
if config is None:
return CapsemSandboxConfig()
if isinstance(config, (CapsemSandboxConfig, dict)):
cfg = config if isinstance(config, CapsemSandboxConfig) else CapsemSandboxConfig(**config)
return (
resolve_compose_file(cfg, sample_metadata=sample_metadata)
if (resolve_compose and cfg.compose_file)
else _validate_direct_volumes(cfg)
)
if isinstance(config, str):
s = config.strip()
if s.lower().endswith((".yaml", ".yml")):
cfg = CapsemSandboxConfig(execution_mode="container", compose_file=s)
return (
resolve_compose_file(cfg, sample_metadata=sample_metadata)
if resolve_compose
else cfg
)
if _is_dockerfile_string(s):
raise ValueError(
f"Dockerfile / Containerfile builds ({s!r}) are not supported in Capsem "
"OCI-workload mode; specify a pre-built OCI image reference instead."
)
return CapsemSandboxConfig(execution_mode="container", image=s)
if isinstance(config, ComposeConfig):
from .containers import extract_capsem_compose_fields

overrides = extract_capsem_compose_fields(
config.model_dump(exclude_none=True), base_dir=None, sample_metadata=sample_metadata
)
return CapsemSandboxConfig(**overrides)
raise TypeError(f"Unsupported Capsem sandbox config type: {type(config)!r}")
41 changes: 36 additions & 5 deletions integrations/inspect-ai/inspect_capsem/_controller.py
Original file line number Diff line number Diff line change
Expand Up @@ -6,7 +6,7 @@
import contextlib
import logging
import os
from collections.abc import Mapping
from collections.abc import Mapping, Sequence
from dataclasses import dataclass
from typing import Any, Protocol, runtime_checkable

Expand All @@ -16,14 +16,15 @@
ExecTimeoutError,
HttpError,
Hypervisor,
Registry,
models,
)
from capsem.execution import (
EXEC_TIMEOUT_CEILING_SECS,
decode_exec_output,
)

from inspect_capsem._transfer import _staged_download, _staged_upload
from inspect_capsem._transfer import _OCI_STAGE_DIR, _staged_download, _staged_upload

logger = logging.getLogger(__name__)

Expand Down Expand Up @@ -52,8 +53,11 @@ async def start_vm(
*,
cpu_count: int,
ram_gb: int,
image: str | None = None,
command: Sequence[str] | None = None,
env: dict[str, str] | None = None,
labels: Mapping[str, str] | None = None,
registry_ca_pem: str | None = None,
) -> str: ...
async def stop_vm(self, vm_id: str) -> None: ...
async def list_vms(self) -> list[models.SandboxInfo]: ...
Expand All @@ -75,6 +79,13 @@ def is_root_user_spec(user: str | None) -> bool:
return u.strip().lower() in ("", "root", "0") and g.strip().lower() in ("", "root", "0")


def _normalize_image_ref(image: str | None) -> str | None:
if not image or not image.strip():
return None
ref = image.strip()
return ref if "://" in ref else f"docker://{ref}"


def _managed_vm_prefix_slug() -> str:
raw = os.environ.get("CAPSEM_VM_PREFIX", "").strip()
if not raw:
Expand Down Expand Up @@ -123,11 +134,14 @@ def __init__(
*,
url: str | None = None,
token: str | None = None,
registry_ca_pem: str | None = None,
) -> None:
if hypervisor is None:
hypervisor = Hypervisor.connect(url, token, timeout=_SDK_CALL_TIMEOUT_SECS)
self._hypervisor: Hypervisor = hypervisor
self._registry_ca_pem = registry_ca_pem
self._sessions: dict[str, VM] = {}
self._oci_vms: set[str] = set()

async def close(self) -> None:
with contextlib.suppress(Exception):
Expand All @@ -143,14 +157,24 @@ async def start_vm(
*,
cpu_count: int,
ram_gb: int,
image: str | None = None,
command: Sequence[str] | None = None,
env: dict[str, str] | None = None,
labels: Mapping[str, str] | None = None,
registry_ca_pem: str | None = None,
) -> str:
norm_image = _normalize_image_ref(image)
kwargs: dict[str, Any] = {
"cpus": cpu_count,
"memory": ram_gb,
"labels": _managed_vm_labels(labels),
}
if norm_image:
kwargs["image"] = norm_image
if eff_ca := (registry_ca_pem or self._registry_ca_pem):
kwargs["registry"] = Registry(ca_pem=eff_ca)
if command is not None:
kwargs["command"] = list(command)
if env:
kwargs["env"] = dict(env)
try:
Expand All @@ -160,6 +184,8 @@ async def start_vm(
raise
vm_id = str(session.id)
self._sessions[vm_id] = session
if norm_image:
self._oci_vms.add(vm_id)
return vm_id

async def stop_vm(self, vm_id: str) -> None:
Expand All @@ -185,8 +211,9 @@ def _session_for(self, vm_id: str) -> VM:
async def exec_in_vm(self, vm_id: str, command: str, *, timeout: int = 120) -> CommandResult:
session = self._session_for(vm_id)
timeout = min(timeout, EXEC_TIMEOUT_CEILING_SECS)
target = models.ExecTarget.WORKLOAD if vm_id in self._oci_vms else models.ExecTarget.VM
try:
res = await session.exec(command, timeout_secs=timeout, target=models.ExecTarget.VM)
res = await session.exec(command, timeout_secs=timeout, target=target)
except ExecTimeoutError as exc:
raise TimeoutError(f"Capsem exec timed out after {timeout}s: {exc}") from exc
return CommandResult(
Expand All @@ -198,10 +225,14 @@ async def exec_in_vm(self, vm_id: str, command: str, *, timeout: int = 120) -> C

async def upload_to_vm(self, vm_id: str, guest_path: str, data: bytes) -> None:
files = self._session_for(vm_id).files
await _staged_upload(self, files, vm_id, guest_path, data)
stage_dir = _OCI_STAGE_DIR if vm_id in self._oci_vms else None
await _staged_upload(self, files, vm_id, guest_path, data, stage_dir=stage_dir)

async def download_from_vm(
self, vm_id: str, guest_path: str, *, max_bytes: int | None = None
) -> bytes:
files = self._session_for(vm_id).files
return await _staged_download(self, files, vm_id, guest_path, max_bytes=max_bytes)
stage_dir = _OCI_STAGE_DIR if vm_id in self._oci_vms else None
return await _staged_download(
self, files, vm_id, guest_path, max_bytes=max_bytes, stage_dir=stage_dir
)
Loading
Loading