diff --git a/.github/workflows/vercel.yml b/.github/workflows/vercel.yml new file mode 100644 index 0000000..56d4a6a --- /dev/null +++ b/.github/workflows/vercel.yml @@ -0,0 +1,18 @@ +name: Validate Vercel Sandbox Image + +on: + pull_request: + paths: + - images/vercel/** + - .github/workflows/vercel.yml + +permissions: + contents: read + +jobs: + image: + runs-on: ubuntu-latest + steps: + - uses: actions/checkout@v6 + - run: docker build -t drukbox-vercel:check images/vercel + - run: docker run --rm drukbox-vercel:check tailscale version diff --git a/alembic/versions/0008_host_lease_deadline.py b/alembic/versions/0008_host_lease_deadline.py new file mode 100644 index 0000000..f3fb27b --- /dev/null +++ b/alembic/versions/0008_host_lease_deadline.py @@ -0,0 +1,23 @@ +"""Store the provider lifetime limit for each host. + +Revision ID: 0008_host_lease_deadline +Revises: 0007_host_service_account +""" + +from collections.abc import Sequence + +import sqlalchemy as sa +from alembic import op + +revision: str = "0008_host_lease_deadline" +down_revision: str | None = "0007_host_service_account" +branch_labels: str | Sequence[str] | None = None +depends_on: str | Sequence[str] | None = None + + +def upgrade() -> None: + op.add_column("hosts", sa.Column("lease_deadline", sa.DateTime(timezone=True), nullable=True)) + + +def downgrade() -> None: + op.drop_column("hosts", "lease_deadline") diff --git a/docs/add-a-provider.md b/docs/add-a-provider.md index 69f10f7..8a30f0f 100644 --- a/docs/add-a-provider.md +++ b/docs/add-a-provider.md @@ -95,6 +95,13 @@ provider with a secret store of its own, such as docker-sbx, implements Do not add provider-specific fields to the host schema. Add a capability instead. +Set `max_lifetime` to a `timedelta` when the provider has a fixed VM +lifetime. Drukbox stores `lease_deadline` before provisioning, caps default +leases and pool ages, and rejects explicit leases beyond that limit. +The value must be conservative: the provider must grant at least that +lifetime after creation begins. Leave it `None` for providers with no +fixed lifetime. + ## 7. Tests - Unit-test the provider with a mocked api object diff --git a/docs/api.md b/docs/api.md index d282604..fba60ac 100644 --- a/docs/api.md +++ b/docs/api.md @@ -40,6 +40,20 @@ or claimed the host, `admin` for an admin key, or `null` for an unclaimed warm host. Callers cannot set it. An `Idempotency-Key` belongs to the service account that first used it. Another one reusing it gets `409`. +## Host leases + +`POST /hosts` without `expires_at` uses the default lease. An explicit +`null` requests a permanent host. `POST /hosts/{id}/renew` with an empty +body renews from now; supply `expires_at` to request a specific expiry. + +A host response includes `lease_deadline`. A date means that the provider +will stop the VM after a fixed lifetime. Default leases and renewals are +capped at that date. A permanent lease or an explicit expiry beyond the +limit returns `400` with `HOST_LEASE`. The limit is stored per host and +does not change when provider settings change. `null` means there is no +fixed provider lifetime. Renewal does not restart a VM or extend its +provider lifetime. + ## Refresh a host secret `POST /hosts/{host_id}/secrets/{service}/refresh` makes the exchange drop diff --git a/docs/deploy.md b/docs/deploy.md index e537ca7..0e5dde6 100644 --- a/docs/deploy.md +++ b/docs/deploy.md @@ -100,6 +100,7 @@ account token returns `503`. See [API](api.md#service-accounts). | `aws` | EC2 instances | Remote | | `hetzner` | Hetzner Cloud VMs | Remote | | `exoscale` | Exoscale VMs | Remote | +| `vercel` | [Vercel sandboxes](vercel.md) | Remote, Tailscale required | | `docker` | Containers ([Local sandboxes with Docker](#local-sandboxes-with-docker)) | Local, no external account | | `docker-sbx` | microVMs ([Local microVMs with Docker Sandboxes](#local-microvms-with-docker-sandboxes)) | Local | @@ -583,7 +584,7 @@ Core, optional: | `TEMPLATE_BUILD_TIMEOUT` | `3600` | Max age in seconds of an unfinished template build before the janitor marks it failed. | | `TEMPLATE_FAILED_RETENTION` | `86400` | Seconds that failed template records and diagnostics remain before the janitor deletes them. | | `TEMPLATE_UNUSED_TTL` | `1209600` | Seconds that an available template remains after its last use, or creation when never used. | -| `LEASE_DEFAULT_TTL` | `86400` | Lease TTL in seconds for hosts created without an explicit `expires_at`, and the extension applied by an empty `POST /hosts/{id}/renew`. An explicit `expires_at: null` at create time opts out of expiry entirely. | +| `LEASE_DEFAULT_TTL` | `86400` | Lease TTL in seconds for hosts created without an explicit `expires_at`, and the extension applied by an empty `POST /hosts/{id}/renew`. An explicit `expires_at: null` at create time opts out where the provider permits it. Defaults are capped by the host's `lease_deadline`. | | `IDEMPOTENCY_KEY_TTL_HOURS` | `24` | Retention period for successful `Idempotency-Key` mappings. | | `POOL_SIZES` | `{}` | Warm hosts to keep ready per provider, as JSON (e.g. `{"exe": 2, "hetzner": 1}`). Overrides `POOL_SIZE` for the providers it names. | | `POOL_SIZE` | `0` | Warm hosts to keep ready for the default provider. `0` disables its pool. | diff --git a/docs/vercel.md b/docs/vercel.md new file mode 100644 index 0000000..a42a755 --- /dev/null +++ b/docs/vercel.md @@ -0,0 +1,109 @@ +# Vercel sandboxes + +Send `{"provider": "vercel"}` to `POST /hosts`. Drukbox creates a named +Vercel sandbox, starts Tailscale, and returns its `internal_ssh_host`. +The API and callers connect with Tailscale SSH as root. There is no +public SSH endpoint or gateway. + +## Configure + +Build the supplied `images/vercel/Dockerfile` as a `linux/amd64` image +and publish it to the Vercel Container Registry (VCR) of your project. +For example, from `images/vercel` with the project linked in the Vercel CLI: + +```bash +vercel vcr login docker +vercel vcr build docker . drukbox-sandbox:v1 --push +``` + +Wait for the image to show `Ready` in VCR. Then configure Drukbox: + +```dotenv +VERCEL_TOKEN=YOUR-VERCEL-ACCESS-TOKEN +VERCEL_TEAM_ID=team_YOUR_TEAM +VERCEL_PROJECT_ID=prj_YOUR_PROJECT +VERCEL_DEFAULT_IMAGE=drukbox-sandbox:v1 +TAILSCALE_ENABLED=true +``` + +Use an access token with sandbox access to this team and project. These +settings also accept files in `/run/secrets`, including `VERCEL_TOKEN`. +Configure the Tailscale OAuth credentials and tag policy as described in +[Networking](networking.md). Permit Tailscale SSH as root from the API +and callers to the sandbox tag. A separate Vercel project per Drukbox +deployment keeps ownership clear. + +The image contains bash, Tailscale, jq, sudo, git, gh, and CA tools. +Bootstrap uses userspace Tailscale when systemd is absent. It needs no +TUN device. Custom images must contain these tools; a Vercel managed +image alone does not satisfy this contract. Vercel ignores image +`ENTRYPOINT` and `CMD`; Drukbox starts bootstrap through the command API. + +Userspace Tailscale provides incoming SSH but does not add kernel routes +for application traffic. The secrets proxy address must be reachable +through the sandbox's normal outbound network. A tailnet-only proxy +address is not sufficient. + +| Setting | Default | Purpose | +| --- | --- | --- | +| `VERCEL_DEFAULT_IMAGE` | Required | VCR image reference; use a digest for a fixed image | +| `VERCEL_VCPUS` | `2` | Default virtual CPU count | +| `VERCEL_SESSION_TIMEOUT_SECONDS` | `2700` | Fixed session duration, including bootstrap | +| `VERCEL_API_TIMEOUT` | `150` | HTTP request timeout, in seconds | +| `VERCEL_BOOTSTRAP_SSH_TIMEOUT_SECONDS` | `120` | SSH readiness timeout | + +An `image` in `POST /hosts` selects another prepared VCR image. An +`instance_type`, such as `"4"`, sets the vCPU count. Drukbox accepts 1–32; +Vercel enforces the account's plan limit. Per-host disk sizing and Drukbox +template builds are not supported. Bootstrap has a 120-second command +limit and a 130-second client deadline. + +## Leases and cleanup + +Vercel limits each uninterrupted session to 45 minutes on Hobby and +24 hours on Pro and Enterprise. Set `VERCEL_SESSION_TIMEOUT_SECONDS` +within the plan limit. Drukbox accepts 600–86400 seconds and defaults to +the Hobby limit. Each sandbox has persistence disabled. Drukbox does not +snapshot, resume, or replace a stopped sandbox. + +Run `uv run alembic upgrade head` before using this version. The migration +adds a nullable `lease_deadline` to hosts. At creation, Drukbox stores the +creation time plus the configured session duration, less 60 seconds. +This conservative limit includes provisioning time and is fixed for that +host. A change to provider settings cannot extend an existing host. + +Omitted leases and empty renewals are capped at `lease_deadline`. An +explicit permanent lease or a date beyond the limit returns `400` with +`HOST_LEASE`. Pool hosts use the same limit; a claim cannot extend it. +Renewal does not call Vercel or increase the session duration. Copy work +out before expiry. + +Deletion removes the named sandbox and its orphan snapshots. A failed +bootstrap or an uncertain create response triggers a deletion attempt. +The generated name makes cleanup possible even if a create response is +lost. If cleanup fails, Drukbox retains the host row for the janitor. +Keep the janitor running to remove expired provider records and Tailscale +devices after a session stops. + +`GET /doctor` checks read access to the project's sandbox list. It does +not create a sandbox, check the image, or test Tailscale SSH. + +## Verify + +```bash +uv run pytest src/providers/vercel/tests src/hosts/tests/test_lease_deadline.py +uv run ruff check +uv run ruff format --check +uv run pyright +``` + +The provider tests mock Vercel HTTP responses. Before production use, +verify a real create, Tailscale SSH connection, renewal, expiry, and +janitor deletion in the target project. + +## References + +- [Vercel sandbox images](https://vercel.com/docs/sandbox/concepts/images) +- [Vercel session duration and persistence](https://vercel.com/kb/guide/vercel-sandbox-duration-and-persistence) +- [Vercel SDK API client](https://github.com/vercel/sandbox/blob/main/packages/vercel-sandbox/src/api-client/api-client.ts) +- [Tailscale userspace networking](https://tailscale.com/kb/1112/userspace-networking) diff --git a/images/vercel/Dockerfile b/images/vercel/Dockerfile new file mode 100644 index 0000000..0b94880 --- /dev/null +++ b/images/vercel/Dockerfile @@ -0,0 +1,8 @@ +FROM tailscale/tailscale:v1.102.5 AS tailscale +FROM ubuntu:24.04 + +RUN apt-get update \ + && apt-get install -y --no-install-recommends bash ca-certificates sudo jq git gh openssh-server \ + && rm -rf /var/lib/apt/lists/* + +COPY --from=tailscale /usr/local/bin/tailscale /usr/local/bin/tailscaled /usr/local/bin/ diff --git a/src/hosts/exceptions.py b/src/hosts/exceptions.py index c2c558c..ea4422a 100644 --- a/src/hosts/exceptions.py +++ b/src/hosts/exceptions.py @@ -19,3 +19,8 @@ class HostTeardownError(AppException): class ProvisioningFailedError(AppException): status_code = 502 error_code = "PROVISIONING_FAILED" + + +class HostLeaseError(AppException): + status_code = 400 + error_code = "HOST_LEASE" diff --git a/src/hosts/models.py b/src/hosts/models.py index bf7b4d8..66863c7 100644 --- a/src/hosts/models.py +++ b/src/hosts/models.py @@ -1,6 +1,7 @@ import uuid from datetime import UTC, datetime from enum import StrEnum +from types import EllipsisType from sqlalchemy import JSON, DateTime, ForeignKey, String, Text, TypeDecorator, Uuid from sqlalchemy.dialects.postgresql import JSONB @@ -9,10 +10,9 @@ from uuid6 import uuid7 from core.database import Base +from hosts.exceptions import HostLeaseError -# Use JSONB on Postgres (indexable, binary storage); fall back to JSON -# (TEXT-backed) on SQLite and other dialects so the OSS quickstart works -# without Postgres. +# JSONB is available on Postgres; SQLite keeps the local development database usable. _JSONType = JSON().with_variant(JSONB(), "postgresql") @@ -31,8 +31,7 @@ class _UTCDateTime(TypeDecorator[datetime]): def process_bind_param(self, value: datetime | None, dialect: object) -> datetime | None: if value and value.tzinfo: - # SQLite drops the offset, so store the equivalent UTC instant — - # otherwise 00:00-08:00 reads back as 00:00Z, not 08:00Z. + # SQLite drops timezone offsets, so normalize before storage. return value.astimezone(UTC) return value @@ -45,11 +44,7 @@ def process_result_value(self, value: datetime | None, dialect: object) -> datet UTCDateTime = _UTCDateTime() _HOST_NAME_PREFIX = "sb-" -# 48 bits of UUIDv7 entropy for a short readable name. UUIDv7's leading 48 -# bits are the millisecond timestamp — concurrent creates in the same ms -# share an identical leading-hex prefix and collided on the unique index. -# Slice from the trailing random segment (rand_b, bits 64..125) so names -# are derived from actual entropy, not a clock reading. +# Use UUIDv7's random suffix because its timestamp prefix is shared by concurrent creates. _HOST_NAME_UID_CHARS = 12 @@ -64,10 +59,7 @@ class HostStatus(StrEnum): class Host(Base): __tablename__ = "hosts" - # Allow non-Mapped[] annotations on this class (we use it for - # `private_key`, a transient per-instance attribute that must never - # be persisted). Without this flag SQLAlchemy 2.0's annotated - # declarative mapper rejects plain annotations. + # The private key exists only on the create response and must never be stored. __allow_unmapped__ = True id: Mapped[uuid.UUID] = mapped_column(Uuid(as_uuid=True), primary_key=True, default=uuid7) @@ -78,20 +70,14 @@ class Host(Base): status: Mapped[str] = mapped_column(String(32), default=HostStatus.PROVISIONING.value) provider: Mapped[str] = mapped_column(String(20), default="exe") image: Mapped[str] = mapped_column(Text) - # Per-request sizing, provider-native values (EC2 instance type, Hetzner - # server type). NULL means the provider's configured default size. + # Null sizing selects the provider default. instance_type: Mapped[str | None] = mapped_column(Text, nullable=True, default=None) disk_gb: Mapped[int | None] = mapped_column(nullable=True, default=None) - # Reachable SSH addresses. Both populated when Tailscale is enabled - # (internal = MagicDNS name, external = provider-given address); only - # external_ssh_host is populated when Tailscale is disabled. The - # internal path is always reached on port 22 by Tailscale convention, - # so no internal_ssh_port column. + # The internal Tailscale SSH port is always 22. external_ssh_host: Mapped[str] = mapped_column(Text, default="") external_ssh_port: Mapped[int] = mapped_column(default=22) ssh_username: Mapped[str] = mapped_column(Text, default="") - # The public half of the per-host keypair. The gateway authenticates - # callers against it. The private half is returned once and never stored. + # Gateway callers authenticate against this public key. public_key: Mapped[str] = mapped_column(Text, default="") internal_ssh_host: Mapped[str | None] = mapped_column(Text, nullable=True, default=None) known_hosts: Mapped[str] = mapped_column(Text, default="") @@ -108,25 +94,32 @@ class Host(Base): nullable=True, default=None, ) + lease_deadline: Mapped[datetime | None] = mapped_column(UTCDateTime, nullable=True) claimed_at: Mapped[datetime | None] = mapped_column( UTCDateTime, nullable=True, default=None, ) - # True only for hosts the pool maintainer warmed. Demand-provisioned hosts - # are False so pool claim/count/shed never hand out or delete a caller-owned - # sandbox — both kinds start with claimed_at NULL, so claimed_at alone can't - # tell them apart. + # An unclaimed caller-owned host must never be counted or removed as pool capacity. pool_member: Mapped[bool] = mapped_column(default=False) last_error: Mapped[str] = mapped_column(Text, default="") - # Non-persisted, transient per-instance attribute. provision() assigns - # the freshly-minted private key here so HostOut returns it exactly - # once at create time; a subsequent GET reads a row from disk where - # this attribute falls back to None. `__allow_unmapped__` above lets - # SQLAlchemy treat the plain annotation as a class attribute instead - # of a missing column. private_key: str | None = None + def lease_expiry( + self, requested: datetime | None | EllipsisType, *, default: datetime + ) -> datetime | None: + expiry = default if requested is ... else requested + if self.lease_deadline: + if self.lease_deadline <= datetime.now(UTC): + raise HostLeaseError("The provider lifetime has ended") + if requested is ...: + return min(default, self.lease_deadline) + if not expiry or expiry > self.lease_deadline: + raise HostLeaseError( + f"expires_at must be at or before {self.lease_deadline.isoformat()}" + ) + return expiry + def __str__(self) -> str: return f"{self.provider}:{self.name}" diff --git a/src/hosts/schemas.py b/src/hosts/schemas.py index 3f80ca8..6529ded 100644 --- a/src/hosts/schemas.py +++ b/src/hosts/schemas.py @@ -143,3 +143,6 @@ class HostOut(BaseModel): updated_at: datetime activated_at: datetime | None expires_at: datetime | None + lease_deadline: datetime | None = Field( + description="Latest allowed lease expiry; null when the provider has no fixed lifetime." + ) diff --git a/src/hosts/scripts/sandbox_bootstrap.sh b/src/hosts/scripts/sandbox_bootstrap.sh index b7aa0bb..05e3fbd 100644 --- a/src/hosts/scripts/sandbox_bootstrap.sh +++ b/src/hosts/scripts/sandbox_bootstrap.sh @@ -1,19 +1,4 @@ #!/usr/bin/env bash -# Sandbox host first-boot bootstrap. Delivered to the VM at create time via the -# VM provider's setup-script mechanism (exe.dev's --setup-script today). -# -# The VM brings itself onto the tailnet and exits. Drukbox observes the new -# device by polling Tailscale's API from outside; no callback into drukbox -# is made from this script, and drukbox is never reachable from the box. -# -# Required env: -# TAILSCALE_AUTHKEY -# TAILSCALE_HOSTNAME -# -# Optional env: -# TAILSCALE_ADVERTISE_TAGS (default: tag:sandbox) -# TAILSCALE_LOGIN_SERVER (default: unset) - set -euo pipefail state_dir=/var/lib/sandbox @@ -35,9 +20,6 @@ run_privileged() { run_privileged install -d -m 755 -o "$(id -u)" -g "$(id -g)" "$state_dir" -# Defensive against re-runs: --setup-script is documented as run-once, and the -# legacy in-image systemd unit (if present on transitional images) also gates on -# this flag. Either path converges on the same end state. if [[ -f "$done_path" ]]; then exit 0 fi @@ -54,7 +36,22 @@ require_var TAILSCALE_HOSTNAME advertise_tags="${TAILSCALE_ADVERTISE_TAGS:-tag:sandbox}" -run_privileged systemctl enable --now tailscaled.service +if [[ -d /run/systemd/system ]]; then + run_privileged systemctl enable --now tailscaled.service +else + run_privileged install -d -m 755 /var/run/tailscale + run_privileged sh -c 'nohup tailscaled --tun=userspace-networking --state=mem: /var/log/tailscaled.log 2>&1 &' + for attempt in {1..60}; do + if [[ -S /var/run/tailscale/tailscaled.sock ]]; then + break + fi + sleep 0.5 + done + if [[ ! -S /var/run/tailscale/tailscaled.sock ]]; then + echo "Tailscale did not create its control socket." >&2 + exit 1 + fi +fi tailscale_running() { tailscale status --json 2>/dev/null | jq -e '.BackendState == "Running"' >/dev/null diff --git a/src/hosts/service.py b/src/hosts/service.py index 0ef6a16..edd1fd6 100644 --- a/src/hosts/service.py +++ b/src/hosts/service.py @@ -20,7 +20,12 @@ from host_secrets import catalog from host_secrets.exceptions import SecretsProxyNotConfiguredError, SecretStaticError from host_secrets.placeholder import Placeholder -from hosts.exceptions import HostStateError, IdempotencyKeyConflictError, ProvisioningFailedError +from hosts.exceptions import ( + HostLeaseError, + HostStateError, + IdempotencyKeyConflictError, + ProvisioningFailedError, +) from hosts.models import Host, HostStatus, IdempotencyKey from networking.tailscale import ( DeviceDiscoveryTimeoutError, @@ -84,11 +89,6 @@ def __init__( ) -> None: self.session = session self.settings = settings or get_settings() - # Construct Tailscale only when explicitly enabled. Settings' - # model_validator guarantees the credentials are present whenever - # tailscale_enabled is true. Tests can inject a mock Tailscale via - # the kwarg regardless of the flag — useful for exercising the - # tailnet path without real credentials. if tailscale: self.tailscale: Tailscale | None = tailscale elif self.settings.tailscale_enabled: @@ -113,11 +113,7 @@ async def get_or_create_host( instance_type: str | None = None, disk_gb: int | None = None, ) -> Host: - # ``...`` (omitted) means "default lease"; an explicit None is the - # caller's deliberate opt-in to a permanent, never-reaped host. The - # sentinel travels to the point where a lease is actually stamped - # (pool claim, or the post-provision rewrite) so the in-flight row - # keeps the short provisioning safety TTL. + # Omission and explicit null must stay distinct until provisioning completes. if provider: registered = get_provider_names() if provider not in registered: @@ -130,10 +126,7 @@ async def get_or_create_host( return existing host: Host | None = None - # Warm hosts are provider-specific, so the claim is scoped to the - # requested provider's pool. A request is pool-eligible only when it - # does not customize the host: no image, template, env, secrets, or - # per-request sizing — pool members are warmed at the provider's defaults. + # Only requests for the provider defaults can claim a warm host. requested_provider = provider or self.settings.default_host_provider customized = env or secrets or image or template or instance_type or disk_gb if not customized and self.settings.get_pool_targets().get(requested_provider): @@ -174,36 +167,33 @@ async def _try_claim_pool_host( provider: str, expires_at: datetime | None | EllipsisType, ) -> Host | None: - # Pick a candidate, then atomically claim it with UPDATE ... WHERE - # claimed_at IS NULL ... RETURNING. The WHERE predicate is the actual - # race guard — concurrent claimants resolve to a single winner per - # row regardless of dialect (PG: MVCC + WHERE filter; SQLite: write - # lock + WHERE filter). Losers return None and the caller falls - # through to fresh provisioning. + # The conditional UPDATE gives concurrent claimants one winner. now = utc_now() - candidate_id = ( - await self.session.execute( - select(Host.id) - .where(Host.provider == provider) - .where(Host.pool_member.is_(True)) - .where(Host.claimed_at.is_(None)) - .where(Host.status == HostStatus.ACTIVE.value) - .where(or_(Host.expires_at.is_(None), Host.expires_at > now)) - .order_by(Host.created_at.asc()) - .limit(1) - ) - ).scalar_one_or_none() - if not candidate_id: + candidates = ( + select(Host) + .where(Host.provider == provider) + .where(Host.pool_member.is_(True)) + .where(Host.claimed_at.is_(None)) + .where(Host.status == HostStatus.ACTIVE.value) + .where(or_(Host.expires_at.is_(None), Host.expires_at > now)) + .where(or_(Host.lease_deadline.is_(None), Host.lease_deadline > now)) + .order_by(Host.created_at.asc()) + .limit(1) + ) + if expires_at is not ...: + if expires_at: + candidates = candidates.where( + or_(Host.lease_deadline.is_(None), Host.lease_deadline >= expires_at) + ) + else: + candidates = candidates.where(Host.lease_deadline.is_(None)) + candidate = (await self.session.execute(candidates)).scalar_one_or_none() + if not candidate: return - - if expires_at is ...: - expires_at = self._default_lease_expires_at() - # The claim replaces the warm-pool max-age TTL with the caller's lease: - # a concrete window (explicit or the default), or None for a caller - # who deliberately opted into a permanent host. + expires_at = candidate.lease_expiry(expires_at, default=self._default_lease_expires_at()) result = await self.session.execute( update(Host) - .where(Host.id == candidate_id) + .where(Host.id == candidate.id) .where(Host.claimed_at.is_(None)) .values( service_account=service_account, @@ -215,11 +205,9 @@ async def _try_claim_pool_host( ) host = result.scalar_one_or_none() await self.session.commit() - if not host: - # Lost the race to another claimant; let the caller fall through. - return - logger.info("pool: claimed host_id=%s name=%s", host.id, host.name) - return host + if host: + logger.info("pool: claimed host_id=%s name=%s", host.id, host.name) + return host async def create_host( self, @@ -235,9 +223,6 @@ async def create_host( disk_gb: int | None = None, pool_member: bool = False, ) -> Host: - # Always provisions a brand-new VM; the pool maintainer calls this - # directly (with pool_member=True) so it never recursively claims its - # own pool members. vm = get_vm_provider(provider) if instance_type and not vm.supports_instance_type: raise UnsupportedSizingError( @@ -259,15 +244,11 @@ async def create_host( name = Host.build_name(uid) now = utc_now() host_image = image or vm.default_image - # Safety TTL covers the strand window: if the client disconnects - # mid-provision, this is what makes the janitor reap the row + VM. - # Replaced with the caller's value after provisioning succeeds. A - # default-lease create keeps just the safety TTL in flight — the - # lease is stamped only once the host is usable. + # The janitor must be able to reap a VM after a client disconnects during provisioning. safety_expires_at = now + timedelta(seconds=self.settings.provisioning_grace_seconds) initial_expires_at = ( max(expires_at, safety_expires_at) - if isinstance(expires_at, datetime) + if expires_at is not ... and expires_at else safety_expires_at ) host = Host( @@ -285,7 +266,14 @@ async def create_host( updated_at=now, expires_at=initial_expires_at, pool_member=pool_member, + lease_deadline=now + vm.max_lifetime if vm.max_lifetime else None, ) + if pool_member: + expires_at = host.lease_expiry( + ..., default=now + timedelta(hours=self.settings.pool_host_max_age_hours) + ) + else: + host.lease_expiry(expires_at, default=self._default_lease_expires_at()) self.session.add(host) await self.session.commit() await self.session.refresh(host) @@ -296,14 +284,12 @@ async def create_host( if host.status == HostStatus.ERROR.value: raise ProvisioningFailedError(host.last_error or "provisioning failed") - # Provisioning won: replace the safety TTL with the caller's intent - # in a dedicated session so we don't extend ``self.session``'s - # transaction (which can perturb advisory-lock-bearing callers like - # the pool maintainer). Guarded on the in-flight value: a renewal - # that landed while the host was bootstrapping is newer intent and - # must not be clobbered. - if expires_at is ...: - expires_at = self._default_lease_expires_at() + # Use a separate session to preserve pool advisory locks. A concurrent renewal wins. + try: + expires_at = host.lease_expiry(expires_at, default=self._default_lease_expires_at()) + except HostLeaseError as exc: + await self.mark_failed(host, exc) + raise ProvisioningFailedError(str(exc)) from exc async with async_session_factory() as ttl_session: await ttl_session.execute( update(Host) @@ -352,9 +338,7 @@ async def _lookup_idempotency_key(self, key: str, service_account: str | None) - f"idempotency key {key} belongs to another service account" ) return host - # Stale: expired, or the host vanished without the FK cascade firing. - # GC in a dedicated session so we don't autoflush the caller's pending - # state on `self.session`. + # A separate session avoids flushing pending host changes during key cleanup. async with async_session_factory() as gc_session: await gc_session.execute(delete(IdempotencyKey).where(IdempotencyKey.key == key)) await gc_session.commit() @@ -383,9 +367,6 @@ async def _record_idempotency_key(self, key: str, host: Host) -> bool: return True async def _release_idempotency_loser(self, host: Host) -> None: - # Two shapes of loser: claimed pool host → return to pool with a - # fresh max-age TTL; freshly-created host → mark expired so the - # janitor reaps it (delete_host refuses PROVISIONING). async with async_session_factory() as fix_session: fresh = await fix_session.get(Host, host.id) if not fresh: @@ -394,7 +375,10 @@ async def _release_idempotency_loser(self, host: Host) -> None: if fresh.claimed_at: fresh.claimed_at = None fresh.service_account = None - fresh.expires_at = now + timedelta(hours=self.settings.pool_host_max_age_hours) + pool_expiry = now + timedelta(hours=self.settings.pool_host_max_age_hours) + fresh.expires_at = ( + min(pool_expiry, fresh.lease_deadline) if fresh.lease_deadline else pool_expiry + ) fresh.updated_at = now logger.info( "idempotency: returned pool host_id=%s to pool after lost race", @@ -434,7 +418,9 @@ async def renew_host(self, host_id: uuid.UUID, *, expires_at: datetime | None = if host.status not in RENEWABLE_STATUSES: raise HostStateError(f"cannot renew a host in status {host.status}") - host.expires_at = expires_at or self._default_lease_expires_at() + host.expires_at = host.lease_expiry( + expires_at or ..., default=self._default_lease_expires_at() + ) host.updated_at = utc_now() await self.session.commit() await self.session.refresh(host) @@ -469,42 +455,30 @@ async def delete_host( raise ResourceNotFoundError("host not found") if pool_shed and host.claimed_at: - # A caller claimed this host between the maintainer selecting it as - # excess and this locked read — leave it for its owner, don't reap it. + # A claim can occur after pool maintenance selects an excess host. return False if expired_only and (not host.expires_at or host.expires_at > utc_now()): - # The owner renewed this host between the janitor selecting it as - # expired and this locked read — the lease is live again, spare it. + # A renewal can occur after the janitor selects an expired host. return False if not force and host.status in DELETE_BLOCKED_STATUSES: raise HostStateError("host is still provisioning") if force or host.status in VM_BACKED_STATUSES: - # force is the janitor reaping an abandoned provision: attempt - # teardown even from an early state, since a row stranded in - # CREATING_VM may already have a VM (delete_vm no-ops if it doesn't). + # An abandoned create can own a VM before its state reaches BOOTSTRAPPING. if host.tailscale_device_id and self.tailscale: - # Clear and commit the device_id before deleting the VM: - # a later delete_vm transport error must not retry the - # already-completed release. Hosts provisioned under - # Tailscale but reaped after the operator turned it off - # fall through and let the auth-key TTL expire the device. + # Commit device release so a failed VM deletion does not repeat it. await self.tailscale.release_device(host.tailscale_device_id) host.tailscale_device_id = None host.updated_at = utc_now() await self.session.commit() - # Secrets go before the VM. A provider failure keeps the row for a retry. + # Keep the row if secret deletion fails, so cleanup can be retried. vm = get_vm_provider(host.provider) await vm.secrets.delete_secrets(vm=host.name) try: await vm.delete_vm(host.name) except ProviderNotFoundError: - # VM already absent at the provider — exe.dev may have evicted - # it, or a previous delete partially succeeded. Treat as done - # so we can clean up the DB row, but log so unexpected - # evictions are visible. logger.warning( "host VM already absent at provider during teardown: " "host_id=%s name=%s provider=%s", @@ -534,8 +508,6 @@ async def provision(self, host_id: str) -> None: environment = dict(host.env) setup_script: str | None = None if tailscale: - # The bootstrap script hard-requires TAILSCALE_AUTHKEY; only - # deliver it (and mint a key) when Tailscale is in play. try: join_credentials = await tailscale.issue_join_credentials(host_name=host.name) except NetworkError as exc: @@ -557,9 +529,6 @@ async def provision(self, host_id: str) -> None: gateway = GatewaySettings() if vm.gateway_process_class and not gateway.ssh_host: - # A gateway-provider host is reachable only through the gateway; - # provisioning one without an address would hand out dead - # coordinates. await self.mark_failed( host, ProvisioningFailedError( @@ -586,15 +555,9 @@ async def provision(self, host_id: str) -> None: host.ssh_username = vm_result.ssh_username host.public_key = vm_result.public_key or "" if vm.gateway_process_class: - # The gateway is the SSH path for hosts of a gateway provider. - # The username names the host; the per-host key is the credential. host.external_ssh_host = gateway.ssh_host host.external_ssh_port = gateway.ssh_port host.ssh_username = host.name - # Stamp the per-VM key onto this instance so the POST response - # carries it. There's no column behind `private_key`, so a later - # GET that loads a fresh row sees the class default (None) and - # never echoes the key back. host.private_key = vm_result.private_key if tailscale: host.internal_ssh_host = tailscale.build_ssh_host(host.name) @@ -660,24 +623,14 @@ async def mark_failed(self, host: Host, exc: Exception) -> None: host.status, ) host.status = HostStatus.ERROR.value - # Client-safe summary, not the raw traceback: last_error is echoed back - # to callers, while the full traceback stays in the log above. host.last_error = f"{type(exc).__name__}: {exc}" now = utc_now() - # An errored host is dead weight (its VM may be half-created). Expire it - # now so the janitor is the single owner of teardown; the POST caller - # already got last_error in the 502. host.expires_at = now host.updated_at = now await self.session.commit() async def scan_known_hosts(self, host: Host) -> bytes: - # Scan every reachable address. tailscaled-SSH (internal) and the - # provider's edge sshd (external) present different host keys, so - # callers picking either path need both entries to verify. Each address - # carries its own port: the internal path is always 22 by Tailscale - # convention, while the external sshd may be remapped (e.g. a published - # container port), so they're scanned separately. + # Tailscale SSH and public SSH can present different host keys. targets: list[tuple[str, int]] = [] if host.internal_ssh_host: targets.append((host.internal_ssh_host, 22)) @@ -691,8 +644,7 @@ async def scan_known_hosts(self, host: Host) -> bytes: collected = b"".join(stdout for stdout, _ in scans) if all(ssh_host.encode() in collected for ssh_host, _ in targets): return collected - # Tailscaled-SSH lags device discovery; keyscan can connect but - # read nothing during the gap. Retry within the budget. + # Tailscale SSH can become ready after device discovery. last_detail = "; ".join(error for _, error in scans if error) or "empty output" if time.monotonic() >= deadline: raise RuntimeError(f"ssh-keyscan never returned host keys: {last_detail}") @@ -710,9 +662,6 @@ async def _keyscan(ssh_host: str, ssh_port: int) -> tuple[bytes, str]: stderr=asyncio.subprocess.PIPE, ) except OSError as error: - # ssh-keyscan missing from the image (or otherwise unspawnable) raises - # here; translate to the RuntimeError provision() routes through - # mark_failed, so it surfaces as a provisioning failure, not a 500. raise RuntimeError(f"could not run ssh-keyscan: {error}") from error stdout, stderr = await process.communicate() return stdout, stderr.decode().strip() diff --git a/src/hosts/tests/test_lease_deadline.py b/src/hosts/tests/test_lease_deadline.py new file mode 100644 index 0000000..00ba0dd --- /dev/null +++ b/src/hosts/tests/test_lease_deadline.py @@ -0,0 +1,149 @@ +from datetime import UTC, datetime, timedelta +from unittest.mock import AsyncMock +from uuid import UUID + +import pytest +from httpx import AsyncClient +from sqlalchemy import func, select + +from core.database import async_session_factory +from hosts.exceptions import HostLeaseError +from hosts.models import Host, HostStatus +from hosts.service import HostService + +AUTH = {"Authorization": "Bearer service-token"} + + +@pytest.fixture +def limited_provider(stub_provider, monkeypatch): + stub_provider.max_lifetime = timedelta(minutes=44) + monkeypatch.setattr(HostService, "provision", AsyncMock()) + return stub_provider + + +async def test_create_caps_default_lease_and_stores_deadline(client: AsyncClient, limited_provider): + response = await client.post("/hosts", headers=AUTH, json={"provider": "stub"}) + assert response.status_code == 201 + body = response.json() + assert body["expires_at"] == body["lease_deadline"] + assert datetime.fromisoformat(body["lease_deadline"]) == datetime.fromisoformat( + body["created_at"] + ) + timedelta(minutes=44) + async with async_session_factory() as session: + host = await session.get(Host, UUID(body["id"])) + assert host + assert host.expires_at == host.lease_deadline + + +@pytest.mark.parametrize("expiry", [None, (datetime.now(UTC) + timedelta(days=1)).isoformat()]) +async def test_unavailable_lease_is_rejected_before_a_host_row( + client: AsyncClient, limited_provider, expiry +): + response = await client.post( + "/hosts", headers=AUTH, json={"provider": "stub", "expires_at": expiry} + ) + assert response.status_code == 400 + assert response.json()["error_code"] == "HOST_LEASE" + async with async_session_factory() as session: + assert await session.scalar(select(func.count()).select_from(Host)) == 0 + + +async def test_explicit_short_lease_is_preserved(client: AsyncClient, limited_provider): + expiry = datetime.now(UTC) + timedelta(minutes=5) + response = await client.post( + "/hosts", headers=AUTH, json={"provider": "stub", "expires_at": expiry.isoformat()} + ) + assert response.status_code == 201 + assert datetime.fromisoformat(response.json()["expires_at"]) == expiry + + +async def test_renew_uses_stored_deadline_after_provider_setting_changes( + client: AsyncClient, limited_provider +): + response = await client.post("/hosts", headers=AUTH, json={"provider": "stub"}) + body = response.json() + async with async_session_factory() as session: + host = await session.get(Host, UUID(body["id"])) + assert host + host.status = HostStatus.ACTIVE.value + await session.commit() + limited_provider.max_lifetime = timedelta(days=1) + renewed = await client.post(f"/hosts/{body['id']}/renew", headers=AUTH, json={}) + assert renewed.status_code == 200 + assert renewed.json()["expires_at"] == body["lease_deadline"] + rejected = await client.post( + f"/hosts/{body['id']}/renew", + headers=AUTH, + json={"expires_at": (datetime.now(UTC) + timedelta(hours=1)).isoformat()}, + ) + assert rejected.status_code == 400 + assert rejected.json()["error_code"] == "HOST_LEASE" + + +async def test_warm_pool_expiry_and_claim_stay_within_lifetime(limited_provider): + async with async_session_factory() as session: + service = HostService(session) + warm = await service.create_host( + env={}, + image=None, + provider="stub", + pool_member=True, + expires_at=datetime.now(UTC) + timedelta(hours=4), + ) + assert warm.expires_at == warm.lease_deadline + warm.status = HostStatus.ACTIVE.value + await session.commit() + claimed = await service._try_claim_pool_host( + service_account="admin", provider="stub", expires_at=... + ) + assert claimed + assert claimed.id == warm.id + assert claimed.expires_at == warm.lease_deadline + await service._release_idempotency_loser(claimed) + await session.refresh(warm) + assert not warm.claimed_at + assert warm.expires_at == warm.lease_deadline + + +@pytest.mark.parametrize("expiry", [None, datetime.now(UTC) + timedelta(hours=1)]) +async def test_pool_does_not_claim_a_host_that_cannot_meet_requested_lease( + limited_provider, expiry +): + async with async_session_factory() as session: + service = HostService(session) + warm = await service.create_host(env={}, image=None, provider="stub", pool_member=True) + warm.status = HostStatus.ACTIVE.value + await session.commit() + claimed = await service._try_claim_pool_host( + service_account="admin", provider="stub", expires_at=expiry + ) + assert not claimed + await session.refresh(warm) + assert not warm.claimed_at + + +def test_expired_lifetime_cannot_be_renewed(): + now = datetime.now(UTC) + host = Host(lease_deadline=now - timedelta(seconds=1)) + with pytest.raises(HostLeaseError, match="lifetime has ended"): + host.lease_expiry(..., default=now + timedelta(hours=1)) + + +async def test_provisioning_past_the_lifetime_marks_host_failed( + client: AsyncClient, limited_provider, monkeypatch +): + async def provision(service: HostService, host_id: str) -> None: + host = await service.get_host(UUID(host_id)) + assert host + host.lease_deadline = datetime.now(UTC) - timedelta(seconds=1) + host.status = HostStatus.ACTIVE.value + await service.session.commit() + + monkeypatch.setattr(HostService, "provision", provision) + response = await client.post("/hosts", headers=AUTH, json={"provider": "stub"}) + assert response.status_code == 502 + async with async_session_factory() as session: + host = (await session.execute(select(Host))).scalar_one() + assert host.status == HostStatus.ERROR.value + assert host.expires_at + assert host.expires_at <= datetime.now(UTC) diff --git a/src/hosts/tests/test_userspace_bootstrap.py b/src/hosts/tests/test_userspace_bootstrap.py new file mode 100644 index 0000000..cc048c9 --- /dev/null +++ b/src/hosts/tests/test_userspace_bootstrap.py @@ -0,0 +1,71 @@ +import asyncio +import os +import sys +from collections.abc import Iterator +from pathlib import Path +from tempfile import TemporaryDirectory + +import pytest + +from hosts.service import _SANDBOX_BOOTSTRAP_SCRIPT + + +@pytest.fixture +def tmp_path() -> Iterator[Path]: + with TemporaryDirectory(prefix="dbx-", dir="/tmp") as directory: + yield Path(directory) + + +@pytest.mark.parametrize("systemd", [False, True]) +async def test_bootstrap_starts_tailscale_and_enables_ssh(tmp_path: Path, systemd: bool): + commands = tmp_path / "bin" + commands.mkdir() + log = tmp_path / "commands.log" + socket_path = tmp_path / "var/run/tailscale/tailscaled.sock" + scripts = { + "sudo": '#!/bin/sh\nshift\nexec "$@"\n', + "systemctl": f'#!/bin/sh\necho "systemctl $*" >> {log}\n', + "tailscale": f'#!/bin/sh\necho "tailscale $*" >> {log}\n', + "jq": "#!/bin/sh\ncat >/dev/null\nexit 1\n", + "tailscaled": ( + f"#!{sys.executable}\nimport socket, sys\n" + f"with open({str(log)!r}, 'a') as log:\n" + " log.write('tailscaled ' + ' '.join(sys.argv[1:]) + '\\n')\n" + f"connection = socket.socket(socket.AF_UNIX)\nconnection.bind({str(socket_path)!r})\n" + ), + } + for name, contents in scripts.items(): + path = commands / name + path.write_text(contents) + path.chmod(0o755) + (tmp_path / "var/log").mkdir(parents=True) + if systemd: + (tmp_path / "run/systemd/system").mkdir(parents=True) + script = _SANDBOX_BOOTSTRAP_SCRIPT.replace("/var/", f"{tmp_path}/var/").replace( + "[[ -d /run/systemd/system ]]", f"[[ -d {tmp_path}/run/systemd/system ]]" + ) + process = await asyncio.create_subprocess_exec( + "bash", + "-c", + script, + env={ + **os.environ, + "PATH": f"{commands}:{os.environ['PATH']}", + "TAILSCALE_AUTHKEY": "test-key", + "TAILSCALE_HOSTNAME": "sb-test", + }, + stdout=asyncio.subprocess.PIPE, + stderr=asyncio.subprocess.PIPE, + ) + stdout, stderr = await asyncio.wait_for(process.communicate(), timeout=10) + assert process.returncode == 0, (stdout, stderr) + operations = log.read_text() + assert "--ssh" in operations + assert "--hostname=sb-test" in operations + assert (tmp_path / "var/lib/sandbox/bootstrap.done").exists() + if systemd: + assert "systemctl enable --now tailscaled.service" in operations + assert "--tun=userspace-networking" not in operations + else: + assert "--tun=userspace-networking --state=mem:" in operations + assert "systemctl" not in operations diff --git a/src/providers/__init__.py b/src/providers/__init__.py index 4091128..9410318 100644 --- a/src/providers/__init__.py +++ b/src/providers/__init__.py @@ -3,4 +3,5 @@ import providers.docker_sbx import providers.exe import providers.exoscale -import providers.hetzner # noqa: F401 +import providers.hetzner +import providers.vercel # noqa: F401 diff --git a/src/providers/base.py b/src/providers/base.py index a6fc2d0..9421d9e 100644 --- a/src/providers/base.py +++ b/src/providers/base.py @@ -1,5 +1,6 @@ import abc from dataclasses import dataclass +from datetime import timedelta from typing import ClassVar, NamedTuple, Self import asyncssh @@ -79,6 +80,7 @@ class VMProvider(abc.ABC): supports_disk_gb: ClassVar[bool] = False # A local provider's hosts keep the external path only. supports_tailnet: ClassVar[bool] = True + max_lifetime: timedelta | None = None @classmethod @abc.abstractmethod diff --git a/src/providers/vercel/__init__.py b/src/providers/vercel/__init__.py new file mode 100644 index 0000000..e1b1f9b --- /dev/null +++ b/src/providers/vercel/__init__.py @@ -0,0 +1,4 @@ +from providers.registry import register_vm_provider +from providers.vercel.provider import VercelProvider + +register_vm_provider(VercelProvider) diff --git a/src/providers/vercel/api.py b/src/providers/vercel/api.py new file mode 100644 index 0000000..8851c96 --- /dev/null +++ b/src/providers/vercel/api.py @@ -0,0 +1,126 @@ +import asyncio +from typing import Any +from urllib.parse import quote + +import httpx + +from providers.exceptions import ( + ProviderAuthError, + ProviderCommandError, + ProviderNotFoundError, + ProviderTransportError, +) + +from .settings import VercelSettings + + +class VercelAPI: + def __init__(self, settings: VercelSettings) -> None: + self.settings = settings + self.client = httpx.AsyncClient( + base_url="https://vercel.com/api/", + headers={"Authorization": f"Bearer {settings.token}"}, + timeout=httpx.Timeout(settings.api_timeout, connect=5), + ) + + async def create_sandbox(self, name: str, *, image: str, vcpus: int, label: str) -> str: + data = await self.request( + "POST", + "v3/sandboxes", + json={ + "name": name, + "projectId": self.settings.project_id, + "image": image, + "resources": {"vcpus": vcpus}, + "timeout": self.settings.session_timeout_seconds * 1000, + "persistent": False, + "tags": {"managed-by": label}, + }, + ) + try: + session = data["session"] + session_id = session["id"] + if ( + not isinstance(session_id, str) + or not session_id + or data["sandbox"]["name"] != name + or data["sandbox"]["persistent"] is not False + or session["timeout"] < self.settings.session_timeout_seconds * 1000 + ): + raise ProviderTransportError("Unexpected Vercel sandbox configuration") + except (KeyError, TypeError) as exc: + raise ProviderTransportError("Invalid Vercel sandbox response") from exc + return session_id + + async def bootstrap(self, session_id: str, script: str) -> None: + path = f"v2/sandboxes/sessions/{quote(session_id, safe='')}/cmd" + try: + async with asyncio.timeout(130): + data = await self.request( + "POST", + path, + json={ + "command": "bash", + "args": ["-e", "-c", script], + "sudo": True, + "timeout": 120000, + }, + ) + command_id = data["command"]["id"] + if not isinstance(command_id, str) or not command_id: + raise ProviderTransportError("Invalid Vercel command ID") + result = await self.request( + "GET", f"{path}/{quote(command_id, safe='')}", params={"wait": "true"} + ) + if result["command"]["exitCode"] != 0: + raise ProviderCommandError("Vercel sandbox bootstrap failed") + except (KeyError, TypeError) as exc: + raise ProviderTransportError("Invalid Vercel command response") from exc + except TimeoutError as exc: + raise ProviderTransportError("Vercel sandbox bootstrap timed out") from exc + + async def delete_sandbox(self, name: str) -> None: + await self.request( + "DELETE", + f"v2/sandboxes/{quote(name, safe='')}", + params={"projectId": self.settings.project_id, "deleteOrphanSnapshots": "true"}, + ) + + async def diagnose(self) -> str: + await self.request( + "GET", "v2/sandboxes", params={"project": self.settings.project_id, "limit": "1"} + ) + return "Vercel sandbox API authentication succeeded" + + async def aclose(self) -> None: + await self.client.aclose() + + async def request( + self, + method: str, + path: str, + *, + json: dict[str, Any] | None = None, + params: dict[str, str] | None = None, + ) -> dict[str, Any]: + try: + response = await self.client.request( + method, path, json=json, params={"teamId": self.settings.team_id, **(params or {})} + ) + except httpx.RequestError as exc: + raise ProviderTransportError("Vercel API transport failed") from exc + if response.status_code in {401, 403}: + raise ProviderAuthError("Vercel API authentication failed") + if response.status_code == 404: + raise ProviderNotFoundError("Vercel sandbox resource not found") + if response.status_code >= 500 or response.status_code == 429: + raise ProviderTransportError(f"Vercel API returned HTTP {response.status_code}") + if not 200 <= response.status_code < 300: + raise ProviderCommandError(f"Vercel API rejected request: HTTP {response.status_code}") + try: + data = response.json() + except ValueError as exc: + raise ProviderTransportError("Invalid Vercel API response") from exc + if not isinstance(data, dict): + raise ProviderTransportError("Invalid Vercel API response") + return data diff --git a/src/providers/vercel/provider.py b/src/providers/vercel/provider.py new file mode 100644 index 0000000..567a7b0 --- /dev/null +++ b/src/providers/vercel/provider.py @@ -0,0 +1,95 @@ +import contextlib +from datetime import timedelta +from typing import ClassVar, Self + +from core.settings import get_settings +from providers import environment +from providers.base import VMCreateResult, VMProvider +from providers.exceptions import ( + ProviderAuthError, + ProviderCommandError, + ProviderError, + ProviderNotFoundError, + ProviderTransportError, +) + +from .api import VercelAPI +from .settings import VercelSettings + + +class VercelProvider(VMProvider): + name: ClassVar[str] = "vercel" + diagnose_hint: ClassVar[str] = "check_vercel_and_tailscale_settings" + supports_instance_type = True + + def __init__(self, api: VercelAPI, settings: VercelSettings, *, service_label: str) -> None: + self.api = api + self.settings = settings + self.service_label = service_label + self.max_lifetime = timedelta(seconds=settings.session_timeout_seconds - 60) + + @classmethod + def from_settings(cls) -> Self: + settings = VercelSettings() # pyright: ignore[reportCallIssue] + return cls(VercelAPI(settings), settings, service_label=get_settings().service_label) + + @property + def default_image(self) -> str: + return self.settings.default_image + + @property + def bootstrap_ssh_timeout_seconds(self) -> float: + return self.settings.bootstrap_ssh_timeout_seconds + + async def create_vm( + self, + *, + name: str, + image: str, + env: dict[str, str] | None = None, + setup_script: str | None = None, + instance_type: str | None = None, + disk_gb: int | None = None, + ) -> VMCreateResult: + if not setup_script: + raise ProviderCommandError("Vercel requires TAILSCALE_ENABLED=true") + try: + vcpus = int(instance_type) if instance_type else self.settings.vcpus + script = environment.get_cloud_init(setup_script, env) + except ValueError as exc: + raise ProviderCommandError("Invalid Vercel sizing or environment") from exc + if not 1 <= vcpus <= 32: + raise ProviderCommandError("Vercel instance_type must be a vCPU count from 1 to 32") + try: + session_id = await self.api.create_sandbox( + name, + image=image, + vcpus=vcpus, + label=self.service_label, + ) + except ProviderAuthError as exc: + raise ProviderCommandError("Vercel API authentication failed") from exc + except ProviderNotFoundError as exc: + raise ProviderCommandError("Vercel project or image was not found") from exc + except ProviderTransportError: + with contextlib.suppress(ProviderError): + await self.api.delete_sandbox(name) + raise + try: + await self.api.bootstrap(session_id, script) + except ProviderError as exc: + with contextlib.suppress(ProviderError): + await self.api.delete_sandbox(name) + raise ProviderTransportError("Vercel sandbox bootstrap failed") from exc + return VMCreateResult(provider_id=session_id, name=name, ssh_username="root") + + async def delete_vm(self, name: str) -> None: + await self.api.delete_sandbox(name) + + async def diagnose(self) -> str: + if not get_settings().tailscale_enabled: + raise ProviderCommandError("Vercel requires TAILSCALE_ENABLED=true") + return await self.api.diagnose() + + async def aclose(self) -> None: + await self.api.aclose() diff --git a/src/providers/vercel/settings.py b/src/providers/vercel/settings.py new file mode 100644 index 0000000..4bf28d9 --- /dev/null +++ b/src/providers/vercel/settings.py @@ -0,0 +1,23 @@ +from pydantic import Field +from pydantic_settings import BaseSettings, SettingsConfigDict + +from core.settings import get_secrets_dir + + +class VercelSettings(BaseSettings): + model_config = SettingsConfigDict( + env_file=".env", + env_prefix="VERCEL_", + extra="ignore", + hide_input_in_errors=True, + secrets_dir=get_secrets_dir(), + ) + + token: str + team_id: str + project_id: str + default_image: str + vcpus: int = Field(default=2, ge=1, le=32) + session_timeout_seconds: int = Field(default=2700, ge=600, le=86400) + api_timeout: float = Field(default=150, gt=0) + bootstrap_ssh_timeout_seconds: float = Field(default=120, gt=0) diff --git a/src/providers/vercel/tests/test_provider.py b/src/providers/vercel/tests/test_provider.py new file mode 100644 index 0000000..f4dd0cc --- /dev/null +++ b/src/providers/vercel/tests/test_provider.py @@ -0,0 +1,163 @@ +import json +from collections.abc import AsyncIterator +from datetime import timedelta + +import httpx +import pytest +import respx + +from providers.exceptions import ProviderCommandError, ProviderNotFoundError, ProviderTransportError +from providers.registry import get_provider_names +from providers.vercel.api import VercelAPI +from providers.vercel.provider import VercelProvider +from providers.vercel.settings import VercelSettings + +BASE = "https://vercel.com/api/" +SANDBOX = { + "sandbox": {"name": "sb-test", "persistent": False}, + "session": {"id": "session-1", "timeout": 2700000}, +} + + +@pytest.fixture +async def provider() -> AsyncIterator[VercelProvider]: + settings = VercelSettings( + token="test-token", team_id="team-1", project_id="project-1", default_image="sandbox:v1" + ) + provider = VercelProvider(VercelAPI(settings), settings, service_label="drukbox") + yield provider + await provider.aclose() + + +@respx.mock +async def test_create_bootstraps_named_image_and_returns_tailnet_host(provider: VercelProvider): + creation = respx.post(BASE + "v3/sandboxes").respond(201, json=SANDBOX) + command = respx.post(BASE + "v2/sandboxes/sessions/session-1/cmd").respond( + 200, json={"command": {"id": "command-1"}} + ) + completion = respx.get(BASE + "v2/sandboxes/sessions/session-1/cmd/command-1").respond( + 200, json={"command": {"exitCode": 0}} + ) + result = await provider.create_vm( + name="sb-test", + image="custom:v2", + setup_script="tailscale up", + env={"FOO": "hello world"}, + instance_type="4", + ) + request = creation.calls[0].request + assert request.headers["Authorization"] == "Bearer test-token" + assert request.url.params["teamId"] == "team-1" + assert json.loads(request.content) == { + "projectId": "project-1", + "name": "sb-test", + "image": "custom:v2", + "timeout": 2700000, + "resources": {"vcpus": 4}, + "persistent": False, + "tags": {"managed-by": "drukbox"}, + } + body = json.loads(command.calls[0].request.content) + assert body["sudo"] is True + assert body["timeout"] == 120000 + assert body["command"] == "bash" + assert "export FOO='hello world'" in body["args"][-1] + assert "tailscale up" in body["args"][-1] + assert completion.calls[0].request.url.params["wait"] == "true" + assert result.name == "sb-test" + assert result.provider_id == "session-1" + assert result.ssh_username == "root" + assert not result.ssh_host + assert provider.max_lifetime == timedelta(seconds=2640) + + +@respx.mock +async def test_requires_tailscale(provider: VercelProvider): + with pytest.raises(ProviderCommandError, match="TAILSCALE_ENABLED"): + await provider.create_vm(name="sb-test", image="sandbox:v1") + assert not respx.calls + + +@respx.mock +@pytest.mark.parametrize("size", ["small", "0", "33", "2.5"]) +async def test_invalid_size_makes_no_request(provider: VercelProvider, size: str): + with pytest.raises(ProviderCommandError): + await provider.create_vm( + name="sb-test", image="sandbox:v1", setup_script="true", instance_type=size + ) + assert not respx.calls + + +@respx.mock +@pytest.mark.parametrize("status", [400, 401, 403, 404, 409]) +async def test_rejected_create_does_not_delete_existing_sandbox( + provider: VercelProvider, status: int +): + respx.post(BASE + "v3/sandboxes").respond(status, text="sensitive details") + with pytest.raises(ProviderCommandError) as error: + await provider.create_vm(name="sb-test", image="sandbox:v1", setup_script="true") + assert "sensitive" not in str(error.value) + assert len(respx.calls) == 1 + + +@respx.mock +async def test_lost_create_response_deletes_by_stable_name(provider: VercelProvider): + respx.post(BASE + "v3/sandboxes").mock(side_effect=httpx.ReadTimeout("lost")) + deletion = respx.delete(BASE + "v2/sandboxes/sb-test").respond(200, json={"sandbox": {}}) + with pytest.raises(ProviderTransportError): + await provider.create_vm(name="sb-test", image="sandbox:v1", setup_script="true") + assert deletion.called + assert deletion.calls[0].request.url.params["projectId"] == "project-1" + assert deletion.calls[0].request.url.params["deleteOrphanSnapshots"] == "true" + + +@respx.mock +@pytest.mark.parametrize( + "response", + [ + {}, + {**SANDBOX, "session": {"id": "session-1", "timeout": 300000}}, + {**SANDBOX, "sandbox": {"name": "wrong", "persistent": False}}, + ], +) +async def test_invalid_create_contract_is_cleaned_up(provider: VercelProvider, response: dict): + respx.post(BASE + "v3/sandboxes").respond(201, json=response) + deletion = respx.delete(BASE + "v2/sandboxes/sb-test").respond(200, json={"sandbox": {}}) + with pytest.raises(ProviderTransportError): + await provider.create_vm(name="sb-test", image="sandbox:v1", setup_script="true") + assert deletion.called + + +@respx.mock +@pytest.mark.parametrize("exit_code", [1, 137]) +async def test_failed_bootstrap_deletes_sandbox(provider: VercelProvider, exit_code: int): + respx.post(BASE + "v3/sandboxes").respond(201, json=SANDBOX) + respx.post(BASE + "v2/sandboxes/sessions/session-1/cmd").respond( + 200, json={"command": {"id": "command-1"}} + ) + respx.get(BASE + "v2/sandboxes/sessions/session-1/cmd/command-1").respond( + 200, json={"command": {"exitCode": exit_code}} + ) + deletion = respx.delete(BASE + "v2/sandboxes/sb-test").respond(200, json={"sandbox": {}}) + with pytest.raises(ProviderTransportError, match="bootstrap"): + await provider.create_vm(name="sb-test", image="sandbox:v1", setup_script="true") + assert deletion.called + + +@respx.mock +async def test_missing_sandbox_translates_not_found(provider: VercelProvider): + respx.delete(BASE + "v2/sandboxes/sb-test").respond(404) + with pytest.raises(ProviderNotFoundError): + await provider.delete_vm("sb-test") + + +@respx.mock +async def test_diagnose_lists_only_one_sandbox(provider: VercelProvider): + listing = respx.get(BASE + "v2/sandboxes").respond(200, json={"sandboxes": []}) + assert "authentication succeeded" in await provider.diagnose() + assert listing.calls[0].request.url.params["project"] == "project-1" + assert listing.calls[0].request.url.params["limit"] == "1" + + +def test_provider_is_registered(): + assert "vercel" in get_provider_names()