Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
22 changes: 15 additions & 7 deletions src/dstack/_internal/core/backends/azure/compute.py
Original file line number Diff line number Diff line change
Expand Up @@ -186,7 +186,6 @@ def create_instance(
managed_identity_resource_group=managed_identity_resource_group,
image_reference=_get_image_ref(
compute_client=self._compute_client,
location=location,
variant=VMImageVariant.from_instance_type(instance_offer.instance),
),
vm_size=instance_offer.instance.name,
Expand Down Expand Up @@ -527,9 +526,12 @@ def _vm_type_available(vm_resource: ResourceSku) -> bool:
return False


# Public name Azure assigned to the gallery that scripts/publish_azure_image.sh publishes to
_COMMUNITY_GALLERY_NAME = "dstack-ebac134d-04b9-4c2b-8b6c-ad3e73904aa7" # Gen2


def _get_image_ref(
compute_client: compute_mgmt.ComputeManagementClient,
location: str,
variant: VMImageVariant,
) -> ImageReference:
if settings.DSTACK_VM_BASE_IMAGE_PREFIX:
Expand All @@ -539,12 +541,12 @@ def _get_image_ref(
image_name=variant.get_image_name(),
)
return ImageReference(id=image.id)
image = compute_client.community_gallery_images.get(
location=location,
public_gallery_name="dstack-ebac134d-04b9-4c2b-8b6c-ad3e73904aa7", # Gen2
gallery_image_name=variant.get_image_name(),
# Not looked up: the lookup fails with azure-mgmt-compute>=38.2.0 (#4333)
return ImageReference(
community_gallery_image_id=(
f"/CommunityGalleries/{_COMMUNITY_GALLERY_NAME}/Images/{variant.get_image_name()}"
)
)
return ImageReference(community_gallery_image_id=image.unique_id)


def _get_gateway_image_ref() -> ImageReference:
Expand Down Expand Up @@ -677,6 +679,12 @@ def _begin_create_instance(
message = e.error.message if e.error.message is not None else ""
raise NoCapacityError(message)
raise e
except ResourceNotFoundError as e:
# The image is not replicated to the location or does not exist
if e.error is not None and e.error.code == "GalleryImageNotFound":
image_id = image_reference.community_gallery_image_id or image_reference.id
raise ComputeError(f"VM image {image_id} is not available in {location}")
raise e
return poller


Expand Down
78 changes: 77 additions & 1 deletion src/tests/_internal/core/backends/azure/test_compute.py
Original file line number Diff line number Diff line change
@@ -1,7 +1,16 @@
from unittest.mock import Mock

import pytest
from azure.core.exceptions import ODataV4Format, ResourceNotFoundError
from azure.mgmt.compute.models import ImageReference

from dstack._internal import settings
from dstack._internal.core.backends.azure.compute import VMImageVariant
from dstack._internal.core.backends.azure.compute import (
VMImageVariant,
_begin_create_instance,
_get_image_ref,
)
from dstack._internal.core.errors import ComputeError
from dstack._internal.core.models.instances import Gpu, InstanceType, Resources


Expand Down Expand Up @@ -71,3 +80,70 @@ def test_from_instance_type(
)
def test_get_image_name(self, variant: VMImageVariant, expected_name: str):
assert variant.get_image_name() == expected_name


class TestGetImageRef:
def test_community_gallery_image(self, monkeypatch: pytest.MonkeyPatch):
monkeypatch.setattr(settings, "DSTACK_VM_BASE_IMAGE_PREFIX", "")
compute_client = Mock()
image_ref = _get_image_ref(compute_client=compute_client, variant=VMImageVariant.GRID)
assert image_ref.community_gallery_image_id == (
"/CommunityGalleries/dstack-ebac134d-04b9-4c2b-8b6c-ad3e73904aa7/Images/"
f"dstack-grid-{settings.DSTACK_VM_BASE_IMAGE_VERSION}"
)
assert compute_client.mock_calls == []


class TestBeginCreateInstance:
def test_raises_compute_error_if_image_not_available(self):
compute_client = Mock()
compute_client.virtual_machines.begin_create_or_update.side_effect = _not_found_error(
"GalleryImageNotFound"
)
image_id = "/CommunityGalleries/g/Images/dstack-0.14"
with pytest.raises(ComputeError, match=f"{image_id} is not available in westeurope"):
_begin_create_instance(
**_create_instance_kwargs(
compute_client, ImageReference(community_gallery_image_id=image_id)
)
)

def test_reraises_other_not_found_errors(self):
compute_client = Mock()
error = _not_found_error("ResourceGroupNotFound")
compute_client.virtual_machines.begin_create_or_update.side_effect = error
with pytest.raises(ResourceNotFoundError) as exc_info:
_begin_create_instance(
**_create_instance_kwargs(
compute_client, ImageReference(community_gallery_image_id="/img")
)
)
assert exc_info.value is error


def _not_found_error(code: str) -> ResourceNotFoundError:
error = ResourceNotFoundError(code)
error.error = ODataV4Format({"code": code, "message": code})
return error


def _create_instance_kwargs(compute_client: Mock, image_reference: ImageReference) -> dict:
return dict(
compute_client=compute_client,
subscription_id="subscription",
location="westeurope",
resource_group="resource-group",
network_security_group="security-group",
network="network",
subnet="subnet",
managed_identity_name=None,
managed_identity_resource_group=None,
image_reference=image_reference,
vm_size="Standard_NV6ads_A10_v5",
instance_name="instance",
user_data="",
ssh_pub_keys=["ssh-ed25519 AAAA"],
spot=True,
disk_size=100,
computer_name="runnervm",
)
Loading