diff --git a/sdk/ml/azure-ai-ml/CHANGELOG.md b/sdk/ml/azure-ai-ml/CHANGELOG.md index d3a3c3401922..631eb2f3a146 100644 --- a/sdk/ml/azure-ai-ml/CHANGELOG.md +++ b/sdk/ml/azure-ai-ml/CHANGELOG.md @@ -5,6 +5,7 @@ ### Features Added ### Bugs Fixed +- Fixed registry-backed `MLClient` sending online endpoint/deployment calls to the registry's resource group instead of the workspace's, causing `ResourceNotFound` for cross-resource-group workspaces. ## 1.35.0 (2026-09-08) diff --git a/sdk/ml/azure-ai-ml/azure/ai/ml/_ml_client.py b/sdk/ml/azure-ai-ml/azure/ai/ml/_ml_client.py index e227029d10e5..9a47bdfe16a7 100644 --- a/sdk/ml/azure-ai-ml/azure/ai/ml/_ml_client.py +++ b/sdk/ml/azure-ai-ml/azure/ai/ml/_ml_client.py @@ -273,6 +273,17 @@ def __init__( workspace_id, workspace_location, ) + self._online_operation_scope = ( + OperationScope( + self._ws_operation_scope.subscription_id, + self._ws_operation_scope.resource_group_name, + workspace_name, + workspace_id=workspace_id, + workspace_location=workspace_location, + ) + if registry_name or registry_reference + else self._operation_scope + ) # Cannot send multiple base_url as azure-cli sets the base_url automatically. kwargs.pop("base_url", None) @@ -319,6 +330,12 @@ def __init__( base_url=base_url, **kwargs, ) + self._service_client_02_2022_preview_online = ServiceClient022022Preview( + subscription_id=self._online_operation_scope._subscription_id, + credential=self._credential, + base_url=base_url, + **kwargs, + ) self._service_client_05_2022 = ServiceClient052022( credential=self._credential, @@ -420,6 +437,12 @@ def __init__( base_url=base_url, **kwargs, ) + self._service_client_04_2023_preview_online = ServiceClient042023Preview( + credential=self._credential, + subscription_id=self._online_operation_scope._subscription_id, + base_url=base_url, + **kwargs, + ) self._service_client_06_2023_preview = ServiceClient062023Preview( credential=self._credential, @@ -604,9 +627,9 @@ def __init__( self._local_endpoint_helper = _LocalEndpointHelper(requests_pipeline=self._requests_pipeline) self._local_deployment_helper = _LocalDeploymentHelper(self._operation_container) self._online_endpoints = OnlineEndpointOperations( - self._ws_operation_scope if registry_reference else self._operation_scope, + self._online_operation_scope, self._operation_config, - self._service_client_02_2022_preview, + self._service_client_02_2022_preview_online, self._operation_container, self._local_endpoint_helper, self._credential, @@ -625,9 +648,9 @@ def __init__( self._operation_container.add(AzureMLResourceType.BATCH_ENDPOINT, self._batch_endpoints) self._operation_container.add(AzureMLResourceType.ONLINE_ENDPOINT, self._online_endpoints) self._online_deployments = OnlineDeploymentOperations( - self._ws_operation_scope if registry_reference else self._operation_scope, + self._online_operation_scope, self._operation_config, - self._service_client_04_2023_preview, + self._service_client_04_2023_preview_online, self._operation_container, self._local_deployment_helper, self._credential, diff --git a/sdk/ml/azure-ai-ml/azure/ai/ml/operations/_online_deployment_operations.py b/sdk/ml/azure-ai-ml/azure/ai/ml/operations/_online_deployment_operations.py index 24510ea4ef25..f71204f05e32 100644 --- a/sdk/ml/azure-ai-ml/azure/ai/ml/operations/_online_deployment_operations.py +++ b/sdk/ml/azure-ai-ml/azure/ai/ml/operations/_online_deployment_operations.py @@ -373,6 +373,8 @@ def _get_workspace_location(self) -> str: :return: The workspace location :rtype: str """ + if self._operation_scope._workspace_location: + return self._operation_scope._workspace_location return str( self._all_operations.all_operations[AzureMLResourceType.WORKSPACE].get(self._workspace_name).location ) diff --git a/sdk/ml/azure-ai-ml/azure/ai/ml/operations/_online_endpoint_operations.py b/sdk/ml/azure-ai-ml/azure/ai/ml/operations/_online_endpoint_operations.py index 2b27e5774529..0e0761f7d051 100644 --- a/sdk/ml/azure-ai-ml/azure/ai/ml/operations/_online_endpoint_operations.py +++ b/sdk/ml/azure-ai-ml/azure/ai/ml/operations/_online_endpoint_operations.py @@ -367,6 +367,8 @@ def invoke( return str(response.text()) def _get_workspace_location(self) -> str: + if self._operation_scope._workspace_location: + return self._operation_scope._workspace_location return str( self._all_operations.all_operations[AzureMLResourceType.WORKSPACE].get(self._workspace_name).location ) diff --git a/sdk/ml/azure-ai-ml/tests/conftest.py b/sdk/ml/azure-ai-ml/tests/conftest.py index 1cb9e7db7665..37ce4c601d8a 100644 --- a/sdk/ml/azure-ai-ml/tests/conftest.py +++ b/sdk/ml/azure-ai-ml/tests/conftest.py @@ -493,6 +493,30 @@ def sdkv2_registry_client(e2e_ws_scope: OperationScope, auth: ClientSecretCreden ) +@pytest.fixture +def registry_backed_client( + e2e_ws_scope: OperationScope, auth: ClientSecretCredential, sdkv2_registry_client: MLClient +) -> MLClient: + """Return a registry-backed client whose online operations target the test workspace.""" + registry_clients = ( + sdkv2_registry_client._service_client_10_2021_dataplanepreview, + sdkv2_registry_client.resource_group_name, + sdkv2_registry_client.subscription_id, + sdkv2_registry_client._service_client_model_dataplane, + sdkv2_registry_client._service_client_registry_arm, + ) + with patch("azure.ai.ml._ml_client.get_registry_client", return_value=registry_clients): + return MLClient( + credential=auth, + subscription_id=e2e_ws_scope.subscription_id, + resource_group_name=e2e_ws_scope.resource_group_name, + workspace_name=e2e_ws_scope.workspace_name, + registry_reference="sdkv2-testFeed", + logging_enable=getenv(E2E_TEST_LOGGING_ENABLED), + cloud="AzureCloud", + ) + + @pytest.fixture def only_registry_client(e2e_ws_scope: OperationScope, auth: ClientSecretCredential) -> MLClient: """return a machine learning client using default e2e testing workspace""" diff --git a/sdk/ml/azure-ai-ml/tests/internal_utils/unittests/test_ml_client.py b/sdk/ml/azure-ai-ml/tests/internal_utils/unittests/test_ml_client.py index aa7f3a55942f..371ee9ac895e 100644 --- a/sdk/ml/azure-ai-ml/tests/internal_utils/unittests/test_ml_client.py +++ b/sdk/ml/azure-ai-ml/tests/internal_utils/unittests/test_ml_client.py @@ -1,5 +1,6 @@ import logging import os +from contextlib import ExitStack from unittest import mock from unittest.mock import Mock, patch @@ -471,6 +472,98 @@ def test_ml_client_with_both_workspace_registry_names_throws( message = exception.value.args[0] assert message == "Both workspace_name and registry_name cannot be used together, for the ml_client." + def test_registry_backed_online_operations_keep_workspace_scope(self, mock_credential) -> None: + workspace_sub = "workspace-sub" + workspace_rg = "workspace-rg" + workspace_name = "workspace-name" + registry_sub = "registry-sub" + registry_rg = "registry-rg" + captured = {} + workspace_details = Mock(location="eastus", _workspace_id="workspace-id") + + def make_fake_operation(name): + class FakeOperation: + def __init__(self, *args, **kwargs): + captured[name] = { + "args": args, + "kwargs": kwargs, + "scope": args[0] if args else kwargs.get("operation_scope"), + } + + def get(self, *args, **kwargs): + return workspace_details + + return FakeOperation + + operation_names = [ + "WorkspaceOperations", + "WorkspaceOutboundRuleOperations", + "RegistryOperations", + "WorkspaceConnectionsOperations", + "CapabilityHostsOperations", + "ComputeOperations", + "DatastoreOperations", + "ModelOperations", + "EvaluatorOperations", + "CodeOperations", + "EnvironmentOperations", + "OnlineEndpointOperations", + "BatchEndpointOperations", + "OnlineDeploymentOperations", + "BatchDeploymentOperations", + "DeploymentTemplateOperations", + "DataOperations", + "ComponentOperations", + "JobOperations", + "ScheduleOperations", + "IndexOperations", + "FeatureStoreOperations", + "FeatureSetOperations", + "FeatureStoreEntityOperations", + "AzureOpenAIDeploymentOperations", + "ServerlessEndpointOperations", + "MarketplaceSubscriptionOperations", + ] + with ExitStack() as stack: + for operation_name in operation_names: + stack.enter_context( + patch(f"azure.ai.ml._ml_client.{operation_name}", make_fake_operation(operation_name)) + ) + + stack.enter_context(patch("azure.ai.ml._ml_client.get_deployments_operation", return_value=Mock())) + stack.enter_context( + patch( + "azure.ai.ml._ml_client.get_registry_client", + return_value=(Mock(), registry_rg, registry_sub, Mock(), Mock()), + ) + ) + + MLClient( + credential=mock_credential, + subscription_id=workspace_sub, + resource_group_name=workspace_rg, + workspace_name=workspace_name, + registry_reference="test-registry", + ) + + online_endpoint_scope = captured["OnlineEndpointOperations"]["scope"] + online_deployment_scope = captured["OnlineDeploymentOperations"]["scope"] + online_endpoint_client = captured["OnlineEndpointOperations"]["args"][2] + online_deployment_client = captured["OnlineDeploymentOperations"]["args"][2] + model_scope = captured["ModelOperations"]["scope"] + + assert online_endpoint_client._config.subscription_id == workspace_sub + assert online_deployment_client._config.subscription_id == workspace_sub + assert online_endpoint_scope.subscription_id == workspace_sub + assert online_endpoint_scope.resource_group_name == workspace_rg + assert online_endpoint_scope.workspace_name == workspace_name + assert online_deployment_scope.subscription_id == workspace_sub + assert online_deployment_scope.resource_group_name == workspace_rg + assert online_deployment_scope.workspace_name == workspace_name + assert model_scope.subscription_id == registry_sub + assert model_scope.resource_group_name == registry_rg + assert online_endpoint_scope is not model_scope + def test_ml_client_with_cli_config(self, mock_credential): # This cloud config should not work and it should NOT overwrite the hardcoded AzureCloud kwargs = { diff --git a/sdk/ml/azure-ai-ml/tests/online_services/e2etests/test_online_deployment.py b/sdk/ml/azure-ai-ml/tests/online_services/e2etests/test_online_deployment.py index 6319cf7b92f3..9aa2181c689a 100644 --- a/sdk/ml/azure-ai-ml/tests/online_services/e2etests/test_online_deployment.py +++ b/sdk/ml/azure-ai-ml/tests/online_services/e2etests/test_online_deployment.py @@ -53,7 +53,7 @@ def test_online_deployment_create( def test_online_deployment_create_when_registry_assets( self, sdkv2_registry_client: MLClient, - client: MLClient, + registry_backed_client: MLClient, randstr: Callable[[], str], rand_online_name: Callable[[], str], rand_online_deployment_name: Callable[[], str], @@ -69,7 +69,7 @@ def test_online_deployment_create_when_registry_assets( endpoint = load_online_endpoint(endpoint_yaml) endpoint_name = rand_online_name("endpoint_name") endpoint.name = endpoint_name - endpoint = client.online_endpoints.begin_create_or_update(endpoint).result() + endpoint = registry_backed_client.online_endpoints.begin_create_or_update(endpoint).result() assert endpoint.name == endpoint_name # create a deployment @@ -80,23 +80,23 @@ def test_online_deployment_create_when_registry_assets( deployment.model = model try: - client.online_deployments.begin_create_or_update(deployment).result() - dep = client.online_deployments.get(name=deployment.name, endpoint_name=endpoint.name) + registry_backed_client.online_deployments.begin_create_or_update(deployment).result() + dep = registry_backed_client.online_deployments.get(name=deployment.name, endpoint_name=endpoint.name) assert dep.name == deployment.name - deps = client.online_deployments.list(endpoint_name=endpoint.name) + deps = registry_backed_client.online_deployments.list(endpoint_name=endpoint.name) assert len(list(deps)) > 0 endpoint.traffic = {deployment.name: 100} - client.online_endpoints.begin_create_or_update(endpoint).result() - endpoint_updated = client.online_endpoints.get(endpoint.name) + registry_backed_client.online_endpoints.begin_create_or_update(endpoint).result() + endpoint_updated = registry_backed_client.online_endpoints.get(endpoint.name) assert endpoint_updated.traffic[deployment.name] == 100 - client.online_endpoints.invoke( + registry_backed_client.online_endpoints.invoke( endpoint_name=endpoint.name, request_file="tests/test_configs/deployments/model-1/sample-request.json", ) finally: - client.online_endpoints.begin_delete(name=endpoint.name) + registry_backed_client.online_endpoints.begin_delete(name=endpoint.name) def test_online_deployment_update( self, client: MLClient, rand_online_name: Callable[[], str], rand_online_deployment_name: Callable[[], str] diff --git a/sdk/ml/azure-ai-ml/tests/online_services/unittests/test_online_deployments.py b/sdk/ml/azure-ai-ml/tests/online_services/unittests/test_online_deployments.py index 36ad1c70cc3d..bae9b7d117f9 100644 --- a/sdk/ml/azure-ai-ml/tests/online_services/unittests/test_online_deployments.py +++ b/sdk/ml/azure-ai-ml/tests/online_services/unittests/test_online_deployments.py @@ -140,6 +140,16 @@ def mock_online_deployment_operations( @pytest.mark.unittest @pytest.mark.production_experiences_test class TestOnlineDeploymentOperations: + def test_get_workspace_location_uses_cached_scope( + self, + mock_online_deployment_operations: OnlineDeploymentOperations, + mock_workspace_operations: WorkspaceOperations, + ) -> None: + mock_online_deployment_operations._operation_scope._workspace_location = "eastus" + + assert mock_online_deployment_operations._get_workspace_location() == "eastus" + mock_workspace_operations._operation.get.assert_not_called() + def test_online_deployment_k8s_create( self, mock_online_deployment_operations: OnlineDeploymentOperations, diff --git a/sdk/ml/azure-ai-ml/tests/online_services/unittests/test_online_endpoints.py b/sdk/ml/azure-ai-ml/tests/online_services/unittests/test_online_endpoints.py index 0b111fd26f97..a425a32d7d2b 100644 --- a/sdk/ml/azure-ai-ml/tests/online_services/unittests/test_online_endpoints.py +++ b/sdk/ml/azure-ai-ml/tests/online_services/unittests/test_online_endpoints.py @@ -232,6 +232,16 @@ def mock_online_endpoint_operations( @pytest.mark.unittest @pytest.mark.production_experiences_test class TestOnlineEndpointsOperations: + def test_get_workspace_location_uses_cached_scope( + self, + mock_online_endpoint_operations: OnlineEndpointOperations, + mock_workspace_operations: WorkspaceOperations, + ) -> None: + mock_online_endpoint_operations._operation_scope._workspace_location = "eastus" + + assert mock_online_endpoint_operations._get_workspace_location() == "eastus" + mock_workspace_operations._operation.get.assert_not_called() + def test_online_list(self, mock_online_endpoint_operations: OnlineEndpointOperations) -> None: mock_online_endpoint_operations.list() mock_online_endpoint_operations._online_operation.list.assert_called_once()