diff --git a/mkdocs/docs/concepts/backends.md b/mkdocs/docs/concepts/backends.md index 110c1b78a5..6442075724 100644 --- a/mkdocs/docs/concepts/backends.md +++ b/mkdocs/docs/concepts/backends.md @@ -976,20 +976,18 @@ 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 "Instance types" + `dstack` supports Hot Aisle VMs (`vm-mi300x-1`, `vm-mi300x-2`, `vm-mi300x-4`, `vm-mi300x-8`) and bare metal servers (`bm-mi300x-8`). To use a specific type, set [`instance_types`](../reference/dstack.yml/fleet.md#instance_types) in the fleet configuration. -??? 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: + Some instances are prepaid for a minimum period (8 hours for bare metal). To avoid releasing them early, 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 the period. To check it for instance types in stock:
```shell - $ curl -H "Authorization: Token $API_KEY" https://admin.hotaisle.app/api/teams/$TEAM_HANDLE/virtual_machines/available/ | jq ".[] | {gpus: .Specs.gpus, MinimumReservationMinutes}" + $ curl -sS -H "Authorization: Token $API_KEY" https://admin.hotaisle.app/api/teams/$TEAM_HANDLE/virtual_machines/available/ | jq ".[]? | {gpus: .Specs.gpus, MinimumReservationMinutes}" + $ curl -sS -H "Authorization: Token $API_KEY" https://admin.hotaisle.app/api/teams/$TEAM_HANDLE/bare_metal/available/ | jq ".[]? | {gpus: .Specs.gpus, MinimumReservationMinutes}" ```
diff --git a/pyproject.toml b/pyproject.toml index 4c623306c4..8bbde6f21d 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -33,7 +33,8 @@ dependencies = [ "python-multipart>=0.0.16", "filelock", "psutil", - "gpuhunt==0.1.30", + # TODO: Release gpuhunt with Hot Aisle bare metal support and pin it. + "gpuhunt @ git+https://github.com/dstackai/gpuhunt.git@pr_hotaisle_bare_metal", "argcomplete>=3.5.0", "ignore-python>=0.2.0", "apscheduler<4", @@ -57,6 +58,10 @@ dstack = "dstack._internal.cli.main:main" [tool.hatch.version] path = "src/dstack/version.py" +[tool.hatch.metadata] +# TODO: Remove after pinning the released gpuhunt package. +allow-direct-references = true + [tool.hatch.build.targets.sdist] artifacts = [ "src/dstack/_internal/server/statics/**", 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..9a919e0d2a 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( @@ -81,11 +85,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 +112,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 +123,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 +180,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: @@ -184,6 +252,7 @@ def _supported_instances(offer: InstanceOffer) -> bool: class HotAisleInstanceBackendData(CoreModel): ip_address: str + bare_metal: bool = False @classmethod def load(cls, raw: Optional[str]) -> "HotAisleInstanceBackendData": @@ -192,4 +261,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/settings.py b/src/dstack/_internal/settings.py index b04b94b856..4e06a20c5b 100644 --- a/src/dstack/_internal/settings.py +++ b/src/dstack/_internal/settings.py @@ -50,3 +50,12 @@ 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: Change the default to "0" before merging https://github.com/dstackai/dstack/pull/4325. + 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/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..2ad0b61e52 --- /dev/null +++ b/src/tests/_internal/core/backends/hotaisle/test_compute.py @@ -0,0 +1,363 @@ +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() -> HotAisleCompute: + return HotAisleCompute( + HotAisleConfig(team_handle="test-team", creds=HotAisleAPIKeyCreds(api_key="test-key")) + ) + + +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: + def test_returns_vm_and_bare_metal_offers(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().get_all_offers_with_availability(unallocated_resources=False) + + assert [(offer.instance.name, offer.backend_data) for offer in offers] == [ + ("vm-mi300x-1", {"vm_specs": VM_SPECS}), + ("bm-mi300x-8", {"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