diff --git a/src/dstack/_internal/core/backends/azure/compute.py b/src/dstack/_internal/core/backends/azure/compute.py index 66ebb0a718..0a8036fdcc 100644 --- a/src/dstack/_internal/core/backends/azure/compute.py +++ b/src/dstack/_internal/core/backends/azure/compute.py @@ -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, @@ -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: @@ -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: @@ -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 diff --git a/src/tests/_internal/core/backends/azure/test_compute.py b/src/tests/_internal/core/backends/azure/test_compute.py index a0ce0afaef..039984ef26 100644 --- a/src/tests/_internal/core/backends/azure/test_compute.py +++ b/src/tests/_internal/core/backends/azure/test_compute.py @@ -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 @@ -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", + )