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