diff --git a/mkdocs/docs/concepts/backends.md b/mkdocs/docs/concepts/backends.md index fd23b1ad5f..58cff86371 100644 --- a/mkdocs/docs/concepts/backends.md +++ b/mkdocs/docs/concepts/backends.md @@ -976,24 +976,29 @@ projects: ??? info "Required permissions" - The API key must have the following roles assigned: + The API key requires the `owner` role for your user and the `operator` and `user` roles for the team specified in `team_handle`. - * **Owner role for the user** - Required for creating and managing SSH keys - * **Operator role for the team** - Required for managing virtual machines within the team +??? info "Bare metal" + Bare metal servers (`bm-mi300x-8`) are experimental and disabled by default. To enable them, set `bare_metal: true` in the backend settings: -??? info "Pricing" - `dstack` shows the hourly price for Hot Aisle instances. Some instances also require an upfront payment for a minimum reservation period, which is usually a few hours. You will be charged for the full minimum period even if you stop the instance early. - - See the Hot Aisle API for the minimum reservation period for each instance type: - -
+
- ```shell - $ curl -H "Authorization: Token $API_KEY" https://admin.hotaisle.app/api/teams/$TEAM_HANDLE/virtual_machines/available/ | jq ".[] | {gpus: .Specs.gpus, MinimumReservationMinutes}" + ```yaml + projects: + - name: main + backends: + - type: hotaisle + team_handle: hotaisle-team-handle + creds: + type: api_key + api_key: 9c27a4bb7a8e472fae12ab34.3f2e3c1db75b9a0187fd2196c6b3e56d2b912e1c439ba08d89e7b6fcd4ef1d3f + bare_metal: true ```
+ Bare metal servers are prepaid for 8 hours. To use the full period, use a fleet with a fixed number of [`nodes`](../concepts/fleets.md#nodes), or at least set [`idle_duration`](../reference/dstack.yml/fleet.md#idle_duration) to cover it. Within the period, `dstack` doesn't terminate them but stops tracking them — terminate them manually in Hot Aisle. After the period, `dstack` terminates them automatically. + ### JarvisLabs Log into your [JarvisLabs](https://cloud.jarvislabs.ai/) account and create an API key. diff --git a/src/dstack/_internal/core/backends/base/offers.py b/src/dstack/_internal/core/backends/base/offers.py index 33d745c1ba..43d6cb484d 100644 --- a/src/dstack/_internal/core/backends/base/offers.py +++ b/src/dstack/_internal/core/backends/base/offers.py @@ -30,6 +30,7 @@ "gcp-dws-calendar-mode", "runpod-cpu", "runpod-cluster", + "hotaisle-bm", ] diff --git a/src/dstack/_internal/core/backends/hotaisle/api_client.py b/src/dstack/_internal/core/backends/hotaisle/api_client.py index a3cc355fcd..003953d80b 100644 --- a/src/dstack/_internal/core/backends/hotaisle/api_client.py +++ b/src/dstack/_internal/core/backends/hotaisle/api_client.py @@ -3,6 +3,7 @@ import requests from dstack._internal.core.backends.base.configurator import raise_invalid_credentials_error +from dstack._internal.core.errors import BackendError, NoCapacityError from dstack._internal.utils.logging import get_logger API_URL = "https://admin.hotaisle.app/api" @@ -88,6 +89,38 @@ def terminate_virtual_machine(self, vm_name: str) -> None: return response.raise_for_status() + def reserve_bare_metal_server(self, specs: Dict[str, Any], description: str) -> Dict[str, Any]: + url = f"{API_URL}/teams/{self.team_handle}/bare_metal/" + payload = {"specs": specs, "description": description} + response = self._make_request("POST", url, json=payload) + # 403: the team's bare metal server limit is reached or the API key lacks permissions. + # 404: no available server matches the specs, e.g. another team reserved it. + if response.status_code in [403, 404]: + raise NoCapacityError(response.text) + response.raise_for_status() + return response.json() + + def get_bare_metal_server(self, server_id: str) -> Dict[str, Any]: + url = f"{API_URL}/teams/{self.team_handle}/bare_metal/{server_id}/" + response = self._make_request("GET", url) + response.raise_for_status() + return response.json() + + def release_bare_metal_server(self, server_id: str, force: bool = True) -> None: + url = f"{API_URL}/teams/{self.team_handle}/bare_metal/{server_id}/" + # force releases even if min reservation time not met + params = {"force": "true"} if force else None + response = self._make_request("DELETE", url, params=params) + if response.status_code == 404: + logger.debug("Hot Aisle bare metal server %s not found", server_id) + return + if response.status_code == 400 and not force: + raise BackendError( + f"Hot Aisle refused to release bare metal server {server_id}" + f" before its minimum reservation period ends: {response.text}" + ) + response.raise_for_status() + def _make_request( self, method: str, diff --git a/src/dstack/_internal/core/backends/hotaisle/compute.py b/src/dstack/_internal/core/backends/hotaisle/compute.py index 2fbbb37da4..5e1d798396 100644 --- a/src/dstack/_internal/core/backends/hotaisle/compute.py +++ b/src/dstack/_internal/core/backends/hotaisle/compute.py @@ -1,7 +1,6 @@ import shlex import subprocess import tempfile -from threading import Thread from typing import Any, List, Optional import gpuhunt @@ -18,6 +17,7 @@ from dstack._internal.core.backends.base.offers import get_catalog_offers from dstack._internal.core.backends.hotaisle.api_client import HotAisleAPIClient from dstack._internal.core.backends.hotaisle.models import HotAisleConfig +from dstack._internal.core.errors import ProvisioningError from dstack._internal.core.models.backends.base import BackendType from dstack._internal.core.models.common import ( CoreModel, @@ -32,12 +32,16 @@ ) from dstack._internal.core.models.placement import PlacementGroup from dstack._internal.core.models.runs import JobProvisioningData +from dstack._internal.settings import FeatureFlags +from dstack._internal.utils.common import get_or_error from dstack._internal.utils.logging import get_logger logger = get_logger(__name__) SUPPORTED_GPUS = ["MI300X"] +SSH_CONNECT_TIMEOUT_SECONDS = 10 +SSH_LAUNCH_TIMEOUT_SECONDS = 60 class HotAisleCompute( @@ -63,7 +67,9 @@ def get_all_offers_with_availability( backend=BackendType.HOTAISLE, locations=self.config.regions or None, catalog=self.catalog, - extra_filter=_supported_instances, + extra_filter=lambda o: ( + _supported_instances(o) and (self.config.allow_bare_metal or not _is_bare_metal(o)) + ), ) return [ offer.with_availability(availability=InstanceAvailability.AVAILABLE) @@ -81,11 +87,24 @@ def create_instance( offer_backend_data = validate_extra_ignore( HotAisleOfferBackendData, instance_offer.backend_data ) - vm_data = self.api_client.create_virtual_machine(offer_backend_data.vm_specs) + if offer_backend_data.bare_metal_specs is not None: + server_data = self.api_client.reserve_bare_metal_server( + specs=offer_backend_data.bare_metal_specs, + description=instance_config.instance_name, + ) + # The deployment ID identifies this reservation, the name identifies the server. + instance_id = server_data["deployment_id"] + ip_address = server_data["ip_address"] + else: + vm_data = self.api_client.create_virtual_machine( + get_or_error(offer_backend_data.vm_specs) + ) + instance_id = vm_data["name"] + ip_address = vm_data["ip_address"] return JobProvisioningData( backend=instance_offer.backend, instance_type=instance_offer.instance, - instance_id=vm_data["name"], + instance_id=instance_id, hostname=None, internal_ip=None, region=instance_offer.region, @@ -95,7 +114,8 @@ def create_instance( dockerized=True, ssh_proxy=None, backend_data=HotAisleInstanceBackendData( - ip_address=vm_data["ip_address"] + ip_address=ip_address, + bare_metal=offer_backend_data.bare_metal_specs is not None, ).model_dump_json(), ) @@ -105,38 +125,56 @@ def update_provisioning_data( project_ssh_public_key: str, project_ssh_private_key: str, ): - vm_state = self.api_client.get_vm_state(provisioning_data.instance_id) - if vm_state == "running": - if provisioning_data.hostname is None and provisioning_data.backend_data: - backend_data = HotAisleInstanceBackendData.load(provisioning_data.backend_data) - provisioning_data.hostname = backend_data.ip_address - commands = get_shim_commands(arch=provisioning_data.instance_type.resources.cpu_arch) - launch_command = "sudo sh -c " + shlex.quote(" && ".join(commands)) - thread = Thread( - target=_start_runner, - kwargs={ - "hostname": provisioning_data.hostname, - "project_ssh_private_key": project_ssh_private_key, - "launch_command": launch_command, - }, - daemon=True, - ) - thread.start() + backend_data = HotAisleInstanceBackendData.load(provisioning_data.backend_data) + hostname = backend_data.ip_address + port = 22 + if backend_data.bare_metal: + server_data = self.api_client.get_bare_metal_server(provisioning_data.instance_id) + os_install_status = (server_data.get("os_status") or {}).get("os_install_status") + if os_install_status == "failed": + raise ProvisioningError("Hot Aisle bare metal server OS installation failed") + if os_install_status != "installed": + return + # The bare metal server's ip_address is private, SSH is exposed via ssh_access. + ssh_access = server_data.get("ssh_access") or {} + hostname = ssh_access.get("ip_address") or hostname + port = ssh_access.get("port") or port + elif self.api_client.get_vm_state(provisioning_data.instance_id) != "running": + return + # Retried on the next check until the shim starts. + if not _start_runner( + hostname=hostname, + port=port, + project_ssh_private_key=project_ssh_private_key, + arch=provisioning_data.instance_type.resources.cpu_arch, + ): + return + provisioning_data.hostname = hostname + provisioning_data.ssh_port = port def terminate_instance( self, instance_id: str, region: str, backend_data: Optional[str] = None ): + if backend_data is not None and HotAisleInstanceBackendData.load(backend_data).bare_metal: + self.api_client.release_bare_metal_server( + instance_id, force=not FeatureFlags.HOTAISLE_BARE_METAL_NO_FORCE_RELEASE + ) + return vm_name = instance_id self.api_client.terminate_virtual_machine(vm_name) def _start_runner( hostname: str, + port: int, project_ssh_private_key: str, - launch_command: str, -): - _launch_runner( + arch: Optional[str], +) -> bool: + commands = get_shim_commands(arch=arch) + launch_command = "sudo sh -c " + shlex.quote(" && ".join(commands)) + return _launch_runner( hostname=hostname, + port=port, ssh_private_key=project_ssh_private_key, launch_command=launch_command, ) @@ -144,36 +182,68 @@ def _start_runner( def _launch_runner( hostname: str, + port: int, ssh_private_key: str, launch_command: str, -): - daemonized_command = f"{launch_command.rstrip('&')} >/tmp/dstack-shim.log 2>&1 & disown" - _run_ssh_command( +) -> bool: + # nohup instead of disown, which isn't available in all shells and would fail the exit code. + daemonized_command = ( + f"nohup {launch_command.rstrip('&')} >/tmp/dstack-shim.log 2>&1 bool: with tempfile.NamedTemporaryFile("w+", 0o600) as f: f.write(ssh_private_key) f.flush() - subprocess.run( - [ - "ssh", - "-F", - "none", - "-o", - "StrictHostKeyChecking=no", - "-i", - f.name, - f"hotaisle@{hostname}", - command, - ], - stdout=subprocess.DEVNULL, - stderr=subprocess.DEVNULL, + try: + proc = subprocess.run( + [ + "ssh", + "-F", + "none", + "-o", + "BatchMode=yes", + "-o", + f"ConnectTimeout={SSH_CONNECT_TIMEOUT_SECONDS}", + "-o", + "ConnectionAttempts=1", + "-o", + "StrictHostKeyChecking=no", + "-o", + "UserKnownHostsFile=/dev/null", + "-o", + "LogLevel=ERROR", + "-i", + f.name, + "-p", + str(port), + f"hotaisle@{hostname}", + command, + ], + stdout=subprocess.DEVNULL, + stderr=subprocess.PIPE, + text=True, + timeout=SSH_LAUNCH_TIMEOUT_SECONDS, + ) + except subprocess.TimeoutExpired: + logger.debug("Timed out running SSH command on Hot Aisle instance %s", hostname) + return False + if proc.returncode != 0: + logger.debug( + "SSH command failed on Hot Aisle instance %s: exit_code=%s stderr=%r", + hostname, + proc.returncode, + proc.stderr[-1000:], ) + return False + return True def _supported_instances(offer: InstanceOffer) -> bool: @@ -182,8 +252,14 @@ def _supported_instances(offer: InstanceOffer) -> bool: ) +def _is_bare_metal(offer: InstanceOffer) -> bool: + offer_backend_data = validate_extra_ignore(HotAisleOfferBackendData, offer.backend_data) + return offer_backend_data.bare_metal_specs is not None + + class HotAisleInstanceBackendData(CoreModel): ip_address: str + bare_metal: bool = False @classmethod def load(cls, raw: Optional[str]) -> "HotAisleInstanceBackendData": @@ -192,4 +268,5 @@ def load(cls, raw: Optional[str]) -> "HotAisleInstanceBackendData": class HotAisleOfferBackendData(CoreModel): - vm_specs: dict[str, Any] + vm_specs: Optional[dict[str, Any]] = None + bare_metal_specs: Optional[dict[str, Any]] = None diff --git a/src/dstack/_internal/core/backends/hotaisle/models.py b/src/dstack/_internal/core/backends/hotaisle/models.py index efee6b4e93..7bc730a98e 100644 --- a/src/dstack/_internal/core/backends/hotaisle/models.py +++ b/src/dstack/_internal/core/backends/hotaisle/models.py @@ -4,6 +4,8 @@ from dstack._internal.core.models.common import CoreModel +HOTAISLE_BARE_METAL_DEFAULT = False + class HotAisleAPIKeyCreds(CoreModel): type: Annotated[Literal["api_key"], Field(description="The type of credentials")] = "api_key" @@ -24,6 +26,15 @@ class HotAisleBackendConfig(CoreModel): Optional[List[str]], Field(description="The list of Hot Aisle regions. Omit to use all regions"), ] = None + bare_metal: Annotated[ + Optional[bool], + Field( + description=( + "Whether bare metal offers can be suggested in addition to VMs (experimental)." + f" Defaults to `{str(HOTAISLE_BARE_METAL_DEFAULT).lower()}`" + ) + ), + ] = None class HotAisleBackendConfigWithCreds(HotAisleBackendConfig): @@ -43,3 +54,9 @@ class HotAisleStoredConfig(HotAisleBackendConfig): class HotAisleConfig(HotAisleStoredConfig): creds: AnyHotAisleCreds + + @property + def allow_bare_metal(self) -> bool: + if self.bare_metal is not None: + return self.bare_metal + return HOTAISLE_BARE_METAL_DEFAULT diff --git a/src/dstack/_internal/core/compatibility/backends.py b/src/dstack/_internal/core/compatibility/backends.py new file mode 100644 index 0000000000..a4bb7f21a5 --- /dev/null +++ b/src/dstack/_internal/core/compatibility/backends.py @@ -0,0 +1,15 @@ +from dstack._internal.core.backends.hotaisle.models import HotAisleBackendConfigWithCreds +from dstack._internal.core.backends.models import AnyBackendConfigWithCreds +from dstack._internal.core.models.common import IncludeExcludeDictType + + +def get_backend_config_excludes(config: AnyBackendConfigWithCreds) -> IncludeExcludeDictType: + """ + Returns `config` exclude mapping to exclude certain fields from the create/update backend + request. Use this method to exclude new fields when they are not set to keep + clients backward-compatibility with older servers. + """ + excludes: IncludeExcludeDictType = {} + if isinstance(config, HotAisleBackendConfigWithCreds) and config.bare_metal is None: + excludes["bare_metal"] = True + return excludes diff --git a/src/dstack/_internal/settings.py b/src/dstack/_internal/settings.py index b04b94b856..1d890f1643 100644 --- a/src/dstack/_internal/settings.py +++ b/src/dstack/_internal/settings.py @@ -50,3 +50,13 @@ class FeatureFlags: """If DSTACK_FF_CLI_PRINT_JOB_CONNECTION_INFO enabled, `dstack apply` command prints server-provided IDE URL(s) and SSH command(s) before job logs (for dev-environments only). """ + + # TODO: Drop once offers carry the minimum reservation period and `dstack` keeps such instances + # idle until it ends. + HOTAISLE_BARE_METAL_NO_FORCE_RELEASE = ( + os.getenv("DSTACK_FF_HOTAISLE_BARE_METAL_NO_FORCE_RELEASE", "1") != "0" + ) + """Enabled unless set to `0`. If enabled, Hot Aisle bare metal servers are deleted without `force`, + so a server still within its minimum reservation period isn't deleted and must be released + manually. This prevents accidentally losing a prepaid server. + """ diff --git a/src/dstack/api/server/_backends.py b/src/dstack/api/server/_backends.py index 37afa04e05..0b180a42cc 100644 --- a/src/dstack/api/server/_backends.py +++ b/src/dstack/api/server/_backends.py @@ -6,6 +6,7 @@ AnyBackendConfigWithCreds, AnyBackendConfigWithCredsTagged, ) +from dstack._internal.core.compatibility.backends import get_backend_config_excludes from dstack._internal.core.models.backends.base import BackendType from dstack._internal.core.models.common import validate_extra_ignore from dstack._internal.server.schemas.backends import DeleteBackendsRequest @@ -27,7 +28,8 @@ def create( self, project_name: str, config: AnyBackendConfigWithCreds ) -> AnyBackendConfigWithCreds: resp = self._request( - f"/api/project/{project_name}/backends/create", body=config.model_dump_json() + f"/api/project/{project_name}/backends/create", + body=config.model_dump_json(exclude=get_backend_config_excludes(config)), ) return validate_extra_ignore(AnyBackendConfigWithCredsTagged, resp.json()) @@ -35,7 +37,8 @@ def update( self, project_name: str, config: AnyBackendConfigWithCreds ) -> AnyBackendConfigWithCreds: resp = self._request( - f"/api/project/{project_name}/backends/update", body=config.model_dump_json() + f"/api/project/{project_name}/backends/update", + body=config.model_dump_json(exclude=get_backend_config_excludes(config)), ) return validate_extra_ignore(AnyBackendConfigWithCredsTagged, resp.json()) diff --git a/src/tests/_internal/core/backends/hotaisle/__init__.py b/src/tests/_internal/core/backends/hotaisle/__init__.py new file mode 100644 index 0000000000..e69de29bb2 diff --git a/src/tests/_internal/core/backends/hotaisle/test_api_client.py b/src/tests/_internal/core/backends/hotaisle/test_api_client.py new file mode 100644 index 0000000000..c69c969d32 --- /dev/null +++ b/src/tests/_internal/core/backends/hotaisle/test_api_client.py @@ -0,0 +1,76 @@ +import pytest +import requests + +from dstack._internal.core.backends.hotaisle.api_client import API_URL, HotAisleAPIClient +from dstack._internal.core.errors import BackendError, NoCapacityError + +BARE_METAL_URL = f"{API_URL}/teams/test-team/bare_metal/" +SERVER_URL = f"{BARE_METAL_URL}deployment-id/" + + +def _client() -> HotAisleAPIClient: + return HotAisleAPIClient(api_key="test-key", team_handle="test-team") + + +class TestReserveBareMetalServer: + def test_posts_specs_and_description(self, requests_mock): + requests_mock.post(BARE_METAL_URL, json={"deployment_id": "deployment-id"}) + + server_data = _client().reserve_bare_metal_server( + specs={"cpu_cores": 104}, description="test-instance" + ) + + assert server_data == {"deployment_id": "deployment-id"} + assert requests_mock.last_request.json() == { + "specs": {"cpu_cores": 104}, + "description": "test-instance", + } + + @pytest.mark.parametrize( + ("status_code", "text"), + [(403, "tenant limit exceeded"), (404, "no available servers")], + ) + def test_raises_no_capacity(self, requests_mock, status_code, text): + requests_mock.post(BARE_METAL_URL, status_code=status_code, text=text) + + with pytest.raises(NoCapacityError, match=text): + _client().reserve_bare_metal_server(specs={}, description="test-instance") + + def test_raises_on_other_errors(self, requests_mock): + requests_mock.post(BARE_METAL_URL, status_code=402) + + with pytest.raises(requests.HTTPError): + _client().reserve_bare_metal_server(specs={}, description="test-instance") + + +class TestReleaseBareMetalServer: + def test_forces_release(self, requests_mock): + requests_mock.delete(SERVER_URL, status_code=204) + + _client().release_bare_metal_server("deployment-id") + + assert requests_mock.last_request.qs == {"force": ["true"]} + + def test_releases_without_force(self, requests_mock): + requests_mock.delete(SERVER_URL, status_code=204) + + _client().release_bare_metal_server("deployment-id", force=False) + + assert requests_mock.last_request.qs == {} + + def test_raises_if_refused_without_force(self, requests_mock): + requests_mock.delete(SERVER_URL, status_code=400, text="minimum usage requirement") + + with pytest.raises(BackendError, match="minimum usage requirement"): + _client().release_bare_metal_server("deployment-id", force=False) + + def test_ignores_missing_server(self, requests_mock): + requests_mock.delete(SERVER_URL, status_code=404) + + _client().release_bare_metal_server("deployment-id") + + def test_raises_on_other_errors(self, requests_mock): + requests_mock.delete(SERVER_URL, status_code=400) + + with pytest.raises(requests.HTTPError): + _client().release_bare_metal_server("deployment-id") diff --git a/src/tests/_internal/core/backends/hotaisle/test_compute.py b/src/tests/_internal/core/backends/hotaisle/test_compute.py new file mode 100644 index 0000000000..562014770a --- /dev/null +++ b/src/tests/_internal/core/backends/hotaisle/test_compute.py @@ -0,0 +1,396 @@ +import subprocess +from typing import Optional +from unittest.mock import MagicMock, call, patch + +import pytest +from gpuhunt.providers.hotaisle import API_URL + +from dstack._internal.core.backends.hotaisle.compute import ( + SSH_CONNECT_TIMEOUT_SECONDS, + SSH_LAUNCH_TIMEOUT_SECONDS, + HotAisleCompute, + HotAisleInstanceBackendData, + _launch_runner, + _run_ssh_command, +) +from dstack._internal.core.backends.hotaisle.models import HotAisleAPIKeyCreds, HotAisleConfig +from dstack._internal.core.errors import ProvisioningError +from dstack._internal.core.models.backends.base import BackendType +from dstack._internal.core.models.instances import ( + Disk, + Gpu, + InstanceAvailability, + InstanceConfiguration, + InstanceOfferWithAvailability, + InstanceType, + Resources, + SSHKey, +) +from dstack._internal.core.models.runs import JobProvisioningData +from dstack._internal.settings import FeatureFlags + +VM_SPECS = { + "cpu_cores": 13, + "ram_capacity": 224 * 1024**3, + "disk_capacity": 12288 * 1024**3, + "cpus": {"count": 1, "manufacturer": "Intel", "model": "Xeon Platinum 8470"}, + "gpus": [{"count": 1, "manufacturer": "AMD", "model": "MI300X"}], +} + +BARE_METAL_SPECS = { + "cpu_cores": 104, + "ram_capacity": 2048 * 1024**3, + "disk_capacity": 123839994396672, + "cpus": [{"count": 2, "manufacturer": "Intel", "model": "Xeon Platinum 8470", "cores": 52}], + "gpus": [{"count": 8, "manufacturer": "AMD", "model": "MI300X"}], +} + + +def _compute(bare_metal: Optional[bool] = None) -> HotAisleCompute: + return HotAisleCompute( + HotAisleConfig( + team_handle="test-team", + creds=HotAisleAPIKeyCreds(api_key="test-key"), + bare_metal=bare_metal, + ) + ) + + +def _compute_with_mocked_api_client() -> HotAisleCompute: + compute = _compute() + compute.api_client = MagicMock() + return compute + + +def _instance_type(name: str) -> InstanceType: + return InstanceType( + name=name, + resources=Resources( + cpus=13, + memory_mib=224 * 1024, + gpus=[Gpu(name="MI300X", memory_mib=192 * 1024)], + spot=False, + disk=Disk(size_mib=12288 * 1024), + ), + ) + + +def _offer(name: str, backend_data: dict) -> InstanceOfferWithAvailability: + return InstanceOfferWithAvailability( + backend=BackendType.HOTAISLE, + instance=_instance_type(name), + region="us-michigan-1", + price=1.99, + backend_data=backend_data, + availability=InstanceAvailability.AVAILABLE, + ) + + +def _instance_config() -> InstanceConfiguration: + return InstanceConfiguration( + project_name="test-project", + instance_name="test-instance", + user="test-user", + ssh_keys=[SSHKey(public="ssh-rsa AAAA test")], + ) + + +def _provisioning_data( + backend_data: Optional[HotAisleInstanceBackendData], +) -> JobProvisioningData: + return JobProvisioningData( + backend=BackendType.HOTAISLE, + instance_type=_instance_type("vm-mi300x-1"), + instance_id="instance-id", + hostname=None, + internal_ip=None, + region="us-michigan-1", + price=1.99, + username="hotaisle", + ssh_port=22, + dockerized=True, + ssh_proxy=None, + backend_data=backend_data.model_dump_json() if backend_data is not None else None, + ) + + +class TestGetAllOffersWithAvailability: + @pytest.mark.parametrize( + ("bare_metal", "expected_instance_types"), + [ + (None, ["vm-mi300x-1"]), + (False, ["vm-mi300x-1"]), + (True, ["vm-mi300x-1", "bm-mi300x-8"]), + ], + ids=["default", "disabled", "enabled"], + ) + def test_returns_bare_metal_offers_only_if_enabled( + self, requests_mock, bare_metal, expected_instance_types + ): + requests_mock.get( + f"{API_URL}/teams/test-team/virtual_machines/available/", + json=[{"OnDemandPrice": 199, "Specs": VM_SPECS}], + ) + requests_mock.get( + f"{API_URL}/teams/test-team/bare_metal/available/", + json=[{"OnDemandPrice": 2712, "Specs": BARE_METAL_SPECS}], + ) + + offers = _compute(bare_metal=bare_metal).get_all_offers_with_availability( + unallocated_resources=False + ) + + assert [offer.instance.name for offer in offers] == expected_instance_types + + def test_keeps_specs_in_backend_data(self, requests_mock): + requests_mock.get( + f"{API_URL}/teams/test-team/virtual_machines/available/", + json=[{"OnDemandPrice": 199, "Specs": VM_SPECS}], + ) + requests_mock.get( + f"{API_URL}/teams/test-team/bare_metal/available/", + json=[{"OnDemandPrice": 2712, "Specs": BARE_METAL_SPECS}], + ) + + offers = _compute(bare_metal=True).get_all_offers_with_availability( + unallocated_resources=False + ) + + assert [offer.backend_data for offer in offers] == [ + {"vm_specs": VM_SPECS}, + {"bare_metal_specs": BARE_METAL_SPECS}, + ] + + +class TestCreateInstance: + def test_creates_vm(self): + compute = _compute_with_mocked_api_client() + compute.api_client.create_virtual_machine.return_value = { + "name": "vm-name", + "ip_address": "10.0.0.1", + } + + provisioning_data = compute.create_instance( + _offer("vm-mi300x-1", {"vm_specs": VM_SPECS}), _instance_config(), None + ) + + compute.api_client.upload_ssh_key.assert_called_once_with("ssh-rsa AAAA test") + compute.api_client.create_virtual_machine.assert_called_once_with(VM_SPECS) + compute.api_client.reserve_bare_metal_server.assert_not_called() + assert provisioning_data.instance_id == "vm-name" + assert HotAisleInstanceBackendData.load( + provisioning_data.backend_data + ) == HotAisleInstanceBackendData(ip_address="10.0.0.1", bare_metal=False) + + def test_reserves_bare_metal_server(self): + compute = _compute_with_mocked_api_client() + compute.api_client.reserve_bare_metal_server.return_value = { + "deployment_id": "deployment-id", + "name": "server-01", + "ip_address": "10.0.0.2", + } + + provisioning_data = compute.create_instance( + _offer("bm-mi300x-8", {"bare_metal_specs": BARE_METAL_SPECS}), + _instance_config(), + None, + ) + + # The key must be uploaded before reserving, so the server accepts it. + assert compute.api_client.mock_calls == [ + call.upload_ssh_key("ssh-rsa AAAA test"), + call.reserve_bare_metal_server(specs=BARE_METAL_SPECS, description="test-instance"), + ] + assert provisioning_data.instance_id == "deployment-id" + assert provisioning_data.username == "hotaisle" + assert HotAisleInstanceBackendData.load( + provisioning_data.backend_data + ) == HotAisleInstanceBackendData(ip_address="10.0.0.2", bare_metal=True) + + +@patch("dstack._internal.core.backends.hotaisle.compute._run_ssh_command", return_value=True) +class TestUpdateProvisioningData: + def test_starts_shim_on_running_vm(self, ssh_mock): + compute = _compute_with_mocked_api_client() + compute.api_client.get_vm_state.return_value = "running" + provisioning_data = _provisioning_data(HotAisleInstanceBackendData(ip_address="10.0.0.1")) + + compute.update_provisioning_data(provisioning_data, "public-key", "private-key") + + compute.api_client.get_vm_state.assert_called_once_with("instance-id") + assert provisioning_data.hostname == "10.0.0.1" + assert provisioning_data.ssh_port == 22 + ssh_mock.assert_called_once() + assert ssh_mock.call_args.kwargs["hostname"] == "10.0.0.1" + + def test_retries_if_shim_fails_to_start(self, ssh_mock): + ssh_mock.return_value = False + compute = _compute_with_mocked_api_client() + compute.api_client.get_vm_state.return_value = "running" + provisioning_data = _provisioning_data(HotAisleInstanceBackendData(ip_address="10.0.0.1")) + + compute.update_provisioning_data(provisioning_data, "public-key", "private-key") + + ssh_mock.assert_called_once() + assert provisioning_data.hostname is None + + def test_waits_for_vm_to_run(self, ssh_mock): + compute = _compute_with_mocked_api_client() + compute.api_client.get_vm_state.return_value = "shut off" + provisioning_data = _provisioning_data(HotAisleInstanceBackendData(ip_address="10.0.0.1")) + + compute.update_provisioning_data(provisioning_data, "public-key", "private-key") + + assert provisioning_data.hostname is None + ssh_mock.assert_not_called() + + @pytest.mark.parametrize( + ("ssh_access", "hostname", "port"), + [ + ({"ip_address": "203.0.113.10", "port": 2222}, "203.0.113.10", 2222), + (None, "10.0.0.2", 22), + ], + ids=["ssh-access", "no-ssh-access"], + ) + def test_starts_shim_on_installed_bare_metal_server( + self, ssh_mock, ssh_access, hostname, port + ): + compute = _compute_with_mocked_api_client() + compute.api_client.get_bare_metal_server.return_value = { + "ip_address": "10.0.0.2", + "ssh_access": ssh_access, + "os_status": {"os_install_status": "installed"}, + } + provisioning_data = _provisioning_data( + HotAisleInstanceBackendData(ip_address="10.0.0.2", bare_metal=True) + ) + + compute.update_provisioning_data(provisioning_data, "public-key", "private-key") + + compute.api_client.get_bare_metal_server.assert_called_once_with("instance-id") + compute.api_client.get_vm_state.assert_not_called() + assert provisioning_data.hostname == hostname + assert provisioning_data.ssh_port == port + ssh_mock.assert_called_once() + assert ssh_mock.call_args.kwargs["hostname"] == hostname + assert ssh_mock.call_args.kwargs["port"] == port + + @pytest.mark.parametrize( + "server_data", + [ + {"os_status": {"os_install_status": "installing_os"}}, + {"os_status": {"os_install_status": "first_boot_tasks"}}, + {"os_status": None}, + {}, + ], + ) + def test_waits_for_bare_metal_os_installation(self, ssh_mock, server_data): + compute = _compute_with_mocked_api_client() + compute.api_client.get_bare_metal_server.return_value = server_data + provisioning_data = _provisioning_data( + HotAisleInstanceBackendData(ip_address="10.0.0.2", bare_metal=True) + ) + + compute.update_provisioning_data(provisioning_data, "public-key", "private-key") + + assert provisioning_data.hostname is None + ssh_mock.assert_not_called() + + def test_raises_on_failed_bare_metal_os_installation(self, ssh_mock): + compute = _compute_with_mocked_api_client() + compute.api_client.get_bare_metal_server.return_value = { + "os_status": {"os_install_status": "failed"} + } + provisioning_data = _provisioning_data( + HotAisleInstanceBackendData(ip_address="10.0.0.2", bare_metal=True) + ) + + with pytest.raises(ProvisioningError): + compute.update_provisioning_data(provisioning_data, "public-key", "private-key") + ssh_mock.assert_not_called() + + +class TestTerminateInstance: + def test_terminates_vm(self): + compute = _compute_with_mocked_api_client() + + compute.terminate_instance( + "vm-name", + "us-michigan-1", + HotAisleInstanceBackendData(ip_address="10.0.0.1").model_dump_json(), + ) + + compute.api_client.terminate_virtual_machine.assert_called_once_with("vm-name") + compute.api_client.release_bare_metal_server.assert_not_called() + + def test_terminates_vm_without_backend_data(self): + compute = _compute_with_mocked_api_client() + + compute.terminate_instance("vm-name", "us-michigan-1", None) + + compute.api_client.terminate_virtual_machine.assert_called_once_with("vm-name") + + @pytest.mark.parametrize( + ("no_force_release", "force"), + [(False, True), (True, False)], + ids=["force", "no-force-flag"], + ) + def test_releases_bare_metal_server(self, no_force_release, force): + compute = _compute_with_mocked_api_client() + + with patch.object(FeatureFlags, "HOTAISLE_BARE_METAL_NO_FORCE_RELEASE", no_force_release): + compute.terminate_instance( + "deployment-id", + "us-michigan-1", + HotAisleInstanceBackendData( + ip_address="10.0.0.2", bare_metal=True + ).model_dump_json(), + ) + + compute.api_client.release_bare_metal_server.assert_called_once_with( + "deployment-id", force=force + ) + compute.api_client.terminate_virtual_machine.assert_not_called() + + +class TestLaunchRunner: + @patch("dstack._internal.core.backends.hotaisle.compute._run_ssh_command", return_value=True) + def test_daemonizes_without_disown(self, ssh_mock): + assert _launch_runner("10.0.0.1", 22, "private-key", "sudo sh -c 'dstack-shim'") + + ssh_mock.assert_called_once_with( + hostname="10.0.0.1", + port=22, + ssh_private_key="private-key", + command="nohup sudo sh -c 'dstack-shim' >/tmp/dstack-shim.log 2>&1 dict: + return json.loads(config.model_dump_json(exclude=get_backend_config_excludes(config))) + + +class TestGetBackendConfigExcludes: + def test_excludes_unset_hotaisle_bare_metal(self): + config = HotAisleBackendConfigWithCreds( + team_handle="test-team", creds=HotAisleAPIKeyCreds(api_key="test-key") + ) + + assert "bare_metal" not in _request_body(config) + + def test_keeps_set_hotaisle_bare_metal(self): + config = HotAisleBackendConfigWithCreds( + team_handle="test-team", creds=HotAisleAPIKeyCreds(api_key="test-key"), bare_metal=True + ) + + assert _request_body(config)["bare_metal"] is True