diff --git a/alembic/versions/0002_host_sizing.py b/alembic/versions/0002_host_sizing.py new file mode 100644 index 0000000..b85e55b --- /dev/null +++ b/alembic/versions/0002_host_sizing.py @@ -0,0 +1,25 @@ +"""per-request host sizing + +Revision ID: 0002_host_sizing +Revises: 0001_initial +""" + +from collections.abc import Sequence + +from alembic import op +import sqlalchemy as sa + +revision: str = '0002_host_sizing' +down_revision: str | None = '0001_initial' +branch_labels: str | Sequence[str] | None = None +depends_on: str | Sequence[str] | None = None + + +def upgrade() -> None: + op.add_column('hosts', sa.Column('instance_type', sa.Text(), nullable=True)) + op.add_column('hosts', sa.Column('disk_gb', sa.Integer(), nullable=True)) + + +def downgrade() -> None: + op.drop_column('hosts', 'disk_gb') + op.drop_column('hosts', 'instance_type') diff --git a/docs/architecture.md b/docs/architecture.md index 646caf9..cc7c19a 100644 --- a/docs/architecture.md +++ b/docs/architecture.md @@ -47,6 +47,9 @@ its exception types. `providers.base.VMProvider` is the whole interface: - `name` / `diagnose_hint` class vars +- `supports_instance_type` / `supports_disk_gb` class vars — which + per-request sizing fields `create_vm` honors; the host service + refuses unsupported sizing with a 400 before any row or VM exists - `default_image` / `bootstrap_ssh_timeout_seconds` properties - `create_vm(...) -> VMCreateResult` - `delete_vm(name)` @@ -106,7 +109,10 @@ Two maintenance commands run as cron jobs from the same image: `hosts.janitor` reaps expired and orphaned hosts, `hosts.pool` keeps a warm pool of pre-provisioned hosts per provider (`POOL_SIZES`, with `POOL_SIZE` as the default provider's target) to hide provider cold -starts. +starts. Pool members are warmed with the provider's default image and +size, so a request that customizes its host — `image`, `env`, +`instance_type`, or `disk_gb` — always provisions fresh instead of +claiming a warm host. ## Diagnostics diff --git a/docs/deploy.md b/docs/deploy.md index 402f231..ed4412c 100644 --- a/docs/deploy.md +++ b/docs/deploy.md @@ -175,8 +175,8 @@ AWS provider: | --- | --- | --- | | `AWS_REGION` | — (required) | Region for the EC2 client and launches. | | `AWS_DEFAULT_IMAGE` | — (required) | AMI id or SSM parameter path used when the caller omits `image`. | -| `AWS_INSTANCE_TYPE` | `t3.medium` | EC2 instance type. | -| `AWS_ROOT_GB` | `100` | Root EBS volume size (gp3, encrypted). | +| `AWS_INSTANCE_TYPE` | `t3.medium` | EC2 instance type when the caller omits `instance_type`. | +| `AWS_ROOT_GB` | `100` | Root EBS volume size (gp3, encrypted) when the caller omits `disk_gb`. | | `AWS_SUBNET_ID` | — | Optional subnet; default VPC's otherwise. | | `AWS_SECURITY_GROUP_ID` | — | Pre-existing SG; unset → drukbox manages `drukbox-managed`. | | `AWS_SSH_CIDRS` | — | SSH ingress CIDRs. Authoritative when set; unset → detected egress `/32`, falling back to `0.0.0.0/0`. | @@ -191,7 +191,7 @@ Hetzner provider: | `HETZNER_API_TOKEN` | — (required) | Bearer token for the Hetzner Cloud API. | | `HETZNER_LOCATION` | — (required) | Location for launches, e.g. `nbg1`, `fsn1`, `hel1`, `ash`. | | `HETZNER_DEFAULT_IMAGE` | `ubuntu-24.04` | Image name/id used when the caller omits `image`. | -| `HETZNER_SERVER_TYPE` | `cx23` | Server type, e.g. `cx23`, `cx33`. Hetzner retires older generations (e.g. `cx22`); a deprecated type fails provisioning with a 422. | +| `HETZNER_SERVER_TYPE` | `cx23` | Server type when the caller omits `instance_type`, e.g. `cx23`, `cx33`. Hetzner retires older generations (e.g. `cx22`); a deprecated type fails provisioning with a 422. | | `HETZNER_API_TIMEOUT` | `30.0` | Timeout for Hetzner API calls. | | `HETZNER_BOOTSTRAP_SSH_TIMEOUT_SECONDS` | `120.0` | ssh-keyscan retry budget for a fresh server. | | `HETZNER_SSH_USERNAME` | `root` | In-VM user callers SSH as. | diff --git a/src/hosts/api.py b/src/hosts/api.py index a2547ae..8643520 100644 --- a/src/hosts/api.py +++ b/src/hosts/api.py @@ -12,7 +12,7 @@ from hosts.schemas import HostCreate, HostOut, HostRenew from hosts.service import HostService from networking.tailscale import NetworkError -from providers.exceptions import ProviderError, UnknownProviderError +from providers.exceptions import ProviderError, UnknownProviderError, UnsupportedSizingError log = logging.getLogger(__name__) @@ -57,8 +57,10 @@ async def create_host( expires_at=expires_at, idempotency_key=idempotency_key, provider=host_create.provider, + instance_type=host_create.instance_type, + disk_gb=host_create.disk_gb, ) - except UnknownProviderError as exc: + except (UnknownProviderError, UnsupportedSizingError) as exc: raise HTTPException(status_code=400, detail=str(exc)) from exc except SQLAlchemyError as exc: log.exception("unexpected database error during host provisioning") diff --git a/src/hosts/models.py b/src/hosts/models.py index e3065e5..5bf6329 100644 --- a/src/hosts/models.py +++ b/src/hosts/models.py @@ -75,6 +75,10 @@ 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. + 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 diff --git a/src/hosts/schemas.py b/src/hosts/schemas.py index 199d93a..8dd4efd 100644 --- a/src/hosts/schemas.py +++ b/src/hosts/schemas.py @@ -2,7 +2,7 @@ import uuid from datetime import UTC, datetime -from pydantic import BaseModel, ConfigDict, Field, field_validator +from pydantic import BaseModel, ConfigDict, Field, ValidationInfo, field_validator RESERVED_HOST_ENV_KEYS = frozenset({"TAILSCALE_AUTHKEY"}) _ENV_KEY_RE = re.compile(r"[A-Za-z_][A-Za-z0-9_]*") @@ -26,13 +26,25 @@ class HostCreate(BaseModel): default=None, description="VM provider to provision on. Omit to use the service default.", ) + instance_type: str | None = Field( + default=None, + description=( + "Provider-native instance size, e.g. AWS `t3.xlarge` or Hetzner " + "`cx33`. Omit to use the provider's configured default." + ), + ) + disk_gb: int | None = Field( + default=None, + ge=1, + description="Root disk size in GB. Omit to use the provider's configured default.", + ) - @field_validator("image") + @field_validator("image", "instance_type") @classmethod - def reject_blank_image(cls, image: str | None) -> str | None: - if image is not None and not image.strip(): - raise ValueError("image must not be blank") - return image + def reject_blank(cls, value: str | None, info: ValidationInfo) -> str | None: + if value is not None and not value.strip(): + raise ValueError(f"{info.field_name} must not be blank") + return value @field_validator("env") @classmethod @@ -69,6 +81,8 @@ class HostOut(BaseModel): status: str provider: str image: str + instance_type: str | None + disk_gb: int | None external_ssh_host: str external_ssh_port: int ssh_username: str diff --git a/src/hosts/service.py b/src/hosts/service.py index b6c9837..452c2fe 100644 --- a/src/hosts/service.py +++ b/src/hosts/service.py @@ -26,6 +26,7 @@ ProviderNotFoundError, ProviderTransportError, UnknownProviderError, + UnsupportedSizingError, ) from providers.registry import get_provider_names, get_vm_provider @@ -93,6 +94,8 @@ async def get_or_create_host( expires_at: datetime | None | EllipsisType = ..., idempotency_key: str | None = None, provider: str | None = None, + 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 @@ -113,15 +116,22 @@ async def get_or_create_host( 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 - # doesn't customize the host: default image and no env. + # doesn't customize the host: default image, no env, and no per-request + # sizing — pool members are warmed at the provider's default size. requested_provider = provider or self.settings.default_host_provider - if not env and image is None and self.settings.get_pool_targets().get(requested_provider): + customized = env or image is not None or instance_type or disk_gb + if not customized and self.settings.get_pool_targets().get(requested_provider): host = await self._try_claim_pool_host( provider=requested_provider, expires_at=expires_at ) if host is None: host = await self.create_host( - env=env, image=image, expires_at=expires_at, provider=provider + env=env, + image=image, + expires_at=expires_at, + provider=provider, + instance_type=instance_type, + disk_gb=disk_gb, ) if idempotency_key and not await self._record_idempotency_key(idempotency_key, host): @@ -190,12 +200,22 @@ async def create_host( image: str | None, expires_at: datetime | None | EllipsisType = ..., provider: str | None = None, + instance_type: str | None = None, + 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( + f"provider {vm.name!r} does not support a per-request instance_type" + ) + if disk_gb and not vm.supports_disk_gb: + raise UnsupportedSizingError( + f"provider {vm.name!r} does not support a per-request disk_gb" + ) uid = uuid7() name = Host.build_name(uid) now = utc_now() @@ -217,6 +237,8 @@ async def create_host( name=name, provider=vm.name, image=host_image, + instance_type=instance_type, + disk_gb=disk_gb, status=HostStatus.PROVISIONING.value, created_at=now, updated_at=now, @@ -445,6 +467,8 @@ async def provision(self, host_id: str) -> None: image=host.image, env=environment, setup_script=setup_script, + instance_type=host.instance_type, + disk_gb=host.disk_gb, ) except (ProviderCommandError, ProviderTransportError) as exc: await self.mark_failed(host, exc) diff --git a/src/hosts/tests/conftest.py b/src/hosts/tests/conftest.py index c6568c2..e982ff0 100644 --- a/src/hosts/tests/conftest.py +++ b/src/hosts/tests/conftest.py @@ -13,6 +13,8 @@ class StubVMProvider(VMProvider): name = "stub" diagnose_hint = "check_stub" + supports_instance_type = True + supports_disk_gb = True def __init__(self) -> None: self.deleted: list[str] = [] @@ -36,6 +38,8 @@ async def create_vm( image: str, env: dict[str, str] | None = None, setup_script: str | None = None, + instance_type: str | None = None, + disk_gb: int | None = None, ) -> VMCreateResult: return VMCreateResult(provider_id=name, name=name, ssh_port=22, ssh_username="stub") diff --git a/src/hosts/tests/test_api.py b/src/hosts/tests/test_api.py index 143b4a5..21d84c5 100644 --- a/src/hosts/tests/test_api.py +++ b/src/hosts/tests/test_api.py @@ -2,6 +2,7 @@ from datetime import datetime, timedelta from unittest.mock import AsyncMock +from sqlalchemy import func, select from uuid6 import uuid7 from core.database import async_session_factory @@ -541,6 +542,74 @@ async def test_create_host_idempotency_key_expired_creates_new_host(client, monk assert response.json()["id"] != str(seed_host.id) +async def test_create_host_stores_and_returns_sizing(client, monkeypatch, stub_provider): + """A sized request lands on the row and is reflected by both POST and GET.""" + monkeypatch.setattr("hosts.service.HostService.provision", AsyncMock()) + + response = await client.post( + "/hosts", + headers={"Authorization": "Bearer service-token"}, + json={"provider": "stub", "instance_type": "stub-large", "disk_gb": 200}, + ) + + assert response.status_code == 201 + payload = response.json() + assert payload["instance_type"] == "stub-large" + assert payload["disk_gb"] == 200 + + fetched = await client.get( + f"/hosts/{payload['id']}", + headers={"Authorization": "Bearer service-token"}, + ) + + assert fetched.status_code == 200 + assert fetched.json()["instance_type"] == "stub-large" + assert fetched.json()["disk_gb"] == 200 + + +async def test_create_host_without_sizing_reflects_null(client, monkeypatch): + """Omitted sizing means the provider's configured default, surfaced as null.""" + monkeypatch.setattr("hosts.service.HostService.provision", AsyncMock()) + + response = await client.post("/hosts", headers={"Authorization": "Bearer service-token"}) + + assert response.status_code == 201 + payload = response.json() + assert payload["instance_type"] is None + assert payload["disk_gb"] is None + + +async def test_create_host_rejects_sizing_the_provider_does_not_support(client): + """exe exposes no per-request sizing; the request 400s before any row exists.""" + for field, value in (("instance_type", "t3.xlarge"), ("disk_gb", 200)): + response = await client.post( + "/hosts", + headers={"Authorization": "Bearer service-token"}, + json={field: value}, + ) + + assert response.status_code == 400 + detail = response.json()["detail"] + assert "'exe'" in detail + assert field in detail + + async with async_session_factory() as session: + assert await session.scalar(select(func.count()).select_from(Host)) == 0 + + +async def test_create_host_rejects_blank_instance_type(client): + response = await client.post( + "/hosts", + headers={"Authorization": "Bearer service-token"}, + json={"instance_type": " "}, + ) + + assert response.status_code == 422 + detail = response.json()["detail"] + assert detail[0]["loc"] == ["body", "instance_type"] + assert "instance_type must not be blank" in detail[0]["msg"] + + async def test_create_host_rejects_blank_image(client): response = await client.post( "/hosts", diff --git a/src/hosts/tests/test_models.py b/src/hosts/tests/test_models.py index dfcea02..02d8d6e 100644 --- a/src/hosts/tests/test_models.py +++ b/src/hosts/tests/test_models.py @@ -102,6 +102,10 @@ async def test_provision_walks_host_to_active(monkeypatch): create_vm_kwargs = mocks["create_vm"].await_args.kwargs assert create_vm_kwargs["name"] == "lb-sandbox-test" assert create_vm_kwargs["image"] == "ghcr.io/drukbox/custom-sandbox:provision-test" + # No per-request sizing on the row → the provider falls back to its + # configured default size. + assert create_vm_kwargs["instance_type"] is None + assert create_vm_kwargs["disk_gb"] is None assert create_vm_kwargs["env"]["TAILSCALE_AUTHKEY"] == "tskey-secret" assert create_vm_kwargs["env"]["TAILSCALE_ADVERTISE_TAGS"] == "tag:sandbox" assert create_vm_kwargs["env"]["SANDBOX_GATEWAY_URL"] == "wss://gateway.example.ts.net/daemon" @@ -119,6 +123,26 @@ async def test_provision_walks_host_to_active(monkeypatch): assert wait_kwargs["host_name"] == "exe-runtime-1" +async def test_provision_forwards_sizing_from_the_row_to_create_vm(monkeypatch): + # provision() reads sizing off the host row, not from request state, so a + # janitor-retried or resumed provision still launches the requested size. + host = await create_host_record( + name="lb-sized", + status="provisioning", + instance_type="t3.xlarge", + disk_gb=250, + ) + mocks = await _patch_provision_happy_path(monkeypatch) + + async with async_session_factory() as session: + await HostService(session).provision(str(host.id)) + + assert mocks["create_vm"].await_args is not None + create_vm_kwargs = mocks["create_vm"].await_args.kwargs + assert create_vm_kwargs["instance_type"] == "t3.xlarge" + assert create_vm_kwargs["disk_gb"] == 250 + + async def test_provision_threads_ssh_username_from_vm_result_onto_host(monkeypatch): # The provider knows which in-VM user the image runs as; provision() # must propagate it so the HostOut response (and subsequent GETs) @@ -488,6 +512,8 @@ async def create_host_record( status: str, env: dict[str, str] | None = None, image: str | None = None, + instance_type: str | None = None, + disk_gb: int | None = None, tailscale_device_id: str | None = None, ) -> Host: now = utc_now() @@ -496,6 +522,8 @@ async def create_host_record( name=name, status=status, image=image or ExeSettings().default_image, # pyright: ignore[reportCallIssue] + instance_type=instance_type, + disk_gb=disk_gb, internal_ssh_host=f"{name}.example.ts.net", external_ssh_host="", env=env or {}, diff --git a/src/hosts/tests/test_pool.py b/src/hosts/tests/test_pool.py index 7463471..7d7adb8 100644 --- a/src/hosts/tests/test_pool.py +++ b/src/hosts/tests/test_pool.py @@ -208,6 +208,32 @@ async def test_create_host_never_claims_another_providers_pool_member( mocked_provision.assert_awaited_once() +async def test_create_host_skips_pool_when_sized(multi_pool_settings, monkeypatch, stub_provider): + # Pool members are warmed at the provider's default size, so a sized + # request must provision fresh even when its provider has a warm host. + pool_host = await _seed_pool_host(name="lb-pool-stub", provider="stub") + mocked_provision = AsyncMock() + monkeypatch.setattr("hosts.service.HostService.provision", mocked_provision) + + for instance_type, disk_gb in (("stub-large", None), (None, 200)): + async with async_session_factory() as session: + service = HostService(session, settings=multi_pool_settings) + result = await service.get_or_create_host( + env={}, + image=None, + provider="stub", + instance_type=instance_type, + disk_gb=disk_gb, + ) + + assert result.id != pool_host.id + assert result.claimed_at is None + assert result.instance_type == instance_type + assert result.disk_gb == disk_gb + + assert mocked_provision.await_count == 2 + + async def test_create_host_skips_pool_when_image_override(pooled_settings, monkeypatch): await _seed_pool_host(name="lb-pool-1") mocked_provision = AsyncMock() diff --git a/src/providers/aws/provider.py b/src/providers/aws/provider.py index e9fa301..fda5115 100644 --- a/src/providers/aws/provider.py +++ b/src/providers/aws/provider.py @@ -23,6 +23,8 @@ class AWSProvider(VMProvider): name: ClassVar[str] = "aws" diagnose_hint: ClassVar[str] = "check_aws_credentials_and_region" + supports_instance_type = True + supports_disk_gb = True def __init__( self, @@ -64,6 +66,8 @@ async def create_vm( 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 image.startswith("ami-"): ami_id = image @@ -108,9 +112,9 @@ async def create_vm( instance_id = await self.api.run_instance( client_token=name, ami_id=ami_id, - instance_type=self.settings.instance_type, + instance_type=instance_type or self.settings.instance_type, key_name=key_name, - root_gb=self.settings.root_gb, + root_gb=disk_gb or self.settings.root_gb, tags=tags, user_data=user_data, associate_public_ip=associate_public_ip, diff --git a/src/providers/aws/tests/test_provider.py b/src/providers/aws/tests/test_provider.py index 3ad694d..c823a80 100644 --- a/src/providers/aws/tests/test_provider.py +++ b/src/providers/aws/tests/test_provider.py @@ -52,6 +52,7 @@ async def test_create_vm_with_tailscale_on_skips_keypair_sg_and_public_ip(): assert kwargs["security_group_id"] is None assert kwargs["tags"]["managed-by"] == "drukbox" assert kwargs["client_token"] == "sb-test" + assert kwargs["instance_type"] == "t3.medium" assert kwargs["root_gb"] == 100 assert result.private_key is None assert result.ssh_host == "" @@ -253,6 +254,25 @@ async def test_create_vm_passes_custom_root_gb_to_run_instance(): assert api.run_instance.await_args.kwargs["root_gb"] == 250 +@pytest.mark.asyncio +async def test_create_vm_honors_per_request_sizing_over_settings(): + api = _api_mock() + provider = AWSProvider(api, _settings(), tailscale_enabled=True) + + await provider.create_vm( + name="sb-test", + image="ami-deadbeef", + env={}, + setup_script="echo hi", + instance_type="t3.xlarge", + disk_gb=250, + ) + + kwargs = api.run_instance.await_args.kwargs + assert kwargs["instance_type"] == "t3.xlarge" + assert kwargs["root_gb"] == 250 + + @pytest.mark.asyncio async def test_create_vm_populates_ssh_username_from_settings(): api = _api_mock() diff --git a/src/providers/base.py b/src/providers/base.py index f99d2cc..491d84c 100644 --- a/src/providers/base.py +++ b/src/providers/base.py @@ -18,6 +18,11 @@ class VMProvider(abc.ABC): # Remediation slug attached to a failed /doctor probe. Owned here because # the provider is what knows how its own dependency gets fixed. diagnose_hint: ClassVar[str] + # Which per-request sizing fields create_vm honors. HostService rejects a + # sized request up front — before any host row or VM exists — when the + # target provider leaves these False. + supports_instance_type: ClassVar[bool] = False + supports_disk_gb: ClassVar[bool] = False @classmethod @abc.abstractmethod @@ -45,6 +50,8 @@ async def create_vm( image: str, env: dict[str, str] | None = None, setup_script: str | None = None, + instance_type: str | None = None, + disk_gb: int | None = None, ) -> VMCreateResult: ... @abc.abstractmethod diff --git a/src/providers/docker/provider.py b/src/providers/docker/provider.py index 7cb94e6..92bf100 100644 --- a/src/providers/docker/provider.py +++ b/src/providers/docker/provider.py @@ -58,6 +58,8 @@ async def create_vm( image: str, env: dict[str, str] | None = None, setup_script: str | None = None, + instance_type: str | None = None, + disk_gb: int | None = None, ) -> VMCreateResult: # A setup script only ever arrives when Tailscale is enabled, and a # local container has no path onto the tailnet. Fail loud rather than diff --git a/src/providers/exceptions.py b/src/providers/exceptions.py index 8074404..9e91f57 100644 --- a/src/providers/exceptions.py +++ b/src/providers/exceptions.py @@ -36,3 +36,7 @@ class UnknownProviderError(ProviderError): class CapabilityUnsupportedError(ProviderError): pass + + +class UnsupportedSizingError(ProviderError): + pass diff --git a/src/providers/exe/provider.py b/src/providers/exe/provider.py index 83d2f56..38d753c 100644 --- a/src/providers/exe/provider.py +++ b/src/providers/exe/provider.py @@ -57,6 +57,8 @@ async def create_vm( image: str, env: dict[str, str] | None = None, setup_script: str | None = None, + instance_type: str | None = None, + disk_gb: int | None = None, ) -> VMCreateResult: payload = await self.api.create_vm( name=name, diff --git a/src/providers/hetzner/api.py b/src/providers/hetzner/api.py index e74aca8..9fe8ca4 100644 --- a/src/providers/hetzner/api.py +++ b/src/providers/hetzner/api.py @@ -72,10 +72,11 @@ async def create_server( ssh_key_name: str, user_data: str, labels: dict[str, str], + server_type: str | None = None, ) -> str: body: dict[str, Any] = { "name": name, - "server_type": self.server_type, + "server_type": server_type or self.server_type, "image": image, "location": self.location, "ssh_keys": [ssh_key_name], diff --git a/src/providers/hetzner/provider.py b/src/providers/hetzner/provider.py index 471e3de..cd92336 100644 --- a/src/providers/hetzner/provider.py +++ b/src/providers/hetzner/provider.py @@ -14,6 +14,9 @@ class HetznerProvider(VMProvider): name: ClassVar[str] = "hetzner" diagnose_hint: ClassVar[str] = "check_hetzner_api_token_and_location" + # instance_type maps onto Hetzner's server type; root disk size is fixed + # by the server type, so disk_gb stays unsupported. + supports_instance_type = True def __init__( self, @@ -51,6 +54,8 @@ async def create_vm( image: str, env: dict[str, str] | None = None, setup_script: str | None = None, + instance_type: str | None = None, + disk_gb: int | None = None, ) -> VMCreateResult: # Always mint a per-VM key, in both networking modes: Hetzner servers # have a public IP and no firewall by default, and attaching a key both @@ -70,6 +75,7 @@ async def create_vm( server_id = await self.api.create_server( name=name, image=image, + server_type=instance_type, ssh_key_name=key_name, user_data=user_data, labels=labels, diff --git a/src/providers/hetzner/tests/test_api.py b/src/providers/hetzner/tests/test_api.py index 03e3d38..f9b0d8b 100644 --- a/src/providers/hetzner/tests/test_api.py +++ b/src/providers/hetzner/tests/test_api.py @@ -66,6 +66,25 @@ async def test_create_server_posts_body_and_returns_id(respx_mock): assert b'"image":"ubuntu-24.04"' in body +@pytest.mark.asyncio +@respx.mock(base_url=BASE_URL) +async def test_create_server_prefers_explicit_server_type(respx_mock): + route = respx_mock.post("/servers").mock( + return_value=httpx.Response(201, json={"server": {"id": 1}}), + ) + + await _api().create_server( + name="sb", + image="ubuntu-24.04", + ssh_key_name="k", + user_data="", + labels={}, + server_type="cx33", + ) + + assert b'"server_type":"cx33"' in route.calls.last.request.read() + + @pytest.mark.asyncio @respx.mock(base_url=BASE_URL) async def test_create_server_omits_user_data_when_empty(respx_mock): diff --git a/src/providers/hetzner/tests/test_provider.py b/src/providers/hetzner/tests/test_provider.py index 887bc02..808d5b9 100644 --- a/src/providers/hetzner/tests/test_provider.py +++ b/src/providers/hetzner/tests/test_provider.py @@ -58,6 +58,22 @@ async def test_create_vm_mints_key_and_returns_public_ip_and_private_key(): assert "-----BEGIN OPENSSH PRIVATE KEY-----" in result.private_key +@pytest.mark.asyncio +async def test_create_vm_passes_instance_type_as_server_type(): + api = _api_mock() + provider = HetznerProvider(api, _settings()) + + await provider.create_vm( + name="sb-test", + image="ubuntu-24.04", + env={}, + setup_script="echo hi", + instance_type="cx33", + ) + + assert api.create_server.await_args.kwargs["server_type"] == "cx33" + + @pytest.mark.asyncio async def test_create_vm_deletes_key_when_create_server_fails(): api = _api_mock() diff --git a/src/providers/tests/test_registry.py b/src/providers/tests/test_registry.py index 29699ed..0c3be10 100644 --- a/src/providers/tests/test_registry.py +++ b/src/providers/tests/test_registry.py @@ -33,6 +33,8 @@ async def create_vm( image: str, env: dict[str, str] | None = None, setup_script: str | None = None, + instance_type: str | None = None, + disk_gb: int | None = None, ) -> VMCreateResult: return VMCreateResult(provider_id=name, name=name, ssh_port=22, ssh_username="stub")