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