diff --git a/docs/getting-started.md b/docs/getting-started.md index b320e9e2..681a2f00 100644 --- a/docs/getting-started.md +++ b/docs/getting-started.md @@ -1,5 +1,92 @@ # Getting Started +## Configure Azure Managed Authentication + +Pass an Azure `TokenCredential` to the Azure Managed client or worker. Use +`resource_id` to override the resource audience for token requests: + +```python +from azure.identity import AzureAuthorityHosts, DefaultAzureCredential +from durabletask.azuremanaged import DurableTaskSchedulerClient + +credential = DefaultAzureCredential( + authority=AzureAuthorityHosts.AZURE_GOVERNMENT, +) + +client = DurableTaskSchedulerClient( + host_address="https://myaccount.usgovvirginia.durabletask.azure.us", + taskhub="my-task-hub", + token_credential=credential, + resource_id="https://durabletask.azure.us", +) +``` + +The same `resource_id` parameter is available on `DurableTaskSchedulerWorker`, +`AsyncDurableTaskSchedulerClient`, and the preview `SandboxActivitiesClient` and +`SandboxWorker`. The async client requires an async credential, such as +`azure.identity.aio.DefaultAzureCredential`. Sandbox workers continue to use +their runtime-injected endpoint and managed identity. + +An explicit resource ID takes precedence over `REGION_NAME`. If `resource_id` +is `None` or `""`, the SDK resolves the default when the client or worker is +constructed: + +| `REGION_NAME` | Default resource ID | +| --- | --- | +| Starts with `usgov` or `usdod` (case-insensitive) | `https://durabletask.azure.us` | +| Any other value, empty, or unset | `https://durabletask.io` | + +For example, to select the government default without passing `resource_id`: + +Bash: + +```bash +export REGION_NAME=usgovvirginia +``` + +PowerShell: + +```powershell +$env:REGION_NAME = "usgovvirginia" +``` + +The SDK trims surrounding whitespace and trailing slashes, removes an existing +`/.default` suffix (case-insensitive), and appends `/.default` for the token +request. For example, both `https://durabletask.azure.us/` and +`https://durabletask.azure.us/.default` request +`https://durabletask.azure.us/.default`. Nonempty inputs that become empty after +normalization, such as whitespace, `///`, or `/.default`, raise `ValueError`. + +### Credential Ownership and Authority Host + +`DurableTaskSchedulerClient`, `AsyncDurableTaskSchedulerClient`, +`SandboxActivitiesClient`, and `DurableTaskSchedulerWorker` require a +`token_credential` argument. They use the caller-created credential directly; +they do not construct `DefaultAzureCredential` or another credential on the +caller's behalf. Passing `None` disables SDK token authentication rather than +creating a default credential. A caller-supplied channel remains responsible +for its own authentication. + +Credentials may come from Azure Identity or implement the Azure Core +`TokenCredential` / `AsyncTokenCredential` protocol themselves. Token requests +use `get_token(scope)`; the protocol has no supported per-request authority +override. Configure `authority` when creating a compatible Azure Identity +credential, as in the example above. Omitting it preserves the credential's +default behavior, including `AZURE_AUTHORITY_HOST` where applicable. + +The preview `SandboxWorker` is the only Azure Managed runtime path that creates +an Azure Identity credential internally. It creates `ManagedIdentityCredential` +using the injected `DTS_UMI_CLIENT_ID`. Managed identities use their hosting +environment's identity endpoint and ignore the authority setting, so this path +does not require an authority-host option either. + +> [!NOTE] +> The resource audience, service endpoint, and credential authority/cloud are +> separate settings. Neither `resource_id` nor `REGION_NAME` changes the endpoint +> or the credential's authority. Configure the credential and any underlying +> developer tools for the target cloud separately. Other clouds or custom +> audiences require an explicit resource ID; they are not inferred from the endpoint. + ## Run the Order Processing Example - Check out the [Durable Task Scheduler diff --git a/durabletask-azuremanaged/CHANGELOG.md b/durabletask-azuremanaged/CHANGELOG.md index 18a4c9ec..895b254a 100644 --- a/durabletask-azuremanaged/CHANGELOG.md +++ b/durabletask-azuremanaged/CHANGELOG.md @@ -7,6 +7,23 @@ adheres to [Semantic Versioning](https://semver.org/spec/v2.0.0.html). ## Unreleased +ADDED + +- Added optional `resource_id` configuration for the token audience on +`DurableTaskSchedulerClient`, `AsyncDurableTaskSchedulerClient`, +`DurableTaskSchedulerWorker`, `SandboxActivitiesClient`, and `SandboxWorker`. +Explicit values override region defaults and support custom resource URIs. +Surrounding whitespace, trailing slashes, and an existing `/.default` suffix +are normalized before token requests; values that become empty are rejected. + +CHANGED + +- When `resource_id` is omitted or empty, Azure Managed clients and workers now +use `https://durabletask.azure.us` if `REGION_NAME` starts with `usgov` or `usdod` +(case-insensitive). Other regions, including an unset `REGION_NAME`, retain +`https://durabletask.io`. The endpoint and credential authority remain separately +configured. + FIXED - With the corresponding core SDK update, asynchronous Azure Blob payload diff --git a/durabletask-azuremanaged/durabletask/azuremanaged/client.py b/durabletask-azuremanaged/durabletask/azuremanaged/client.py index 9a553644..33f21475 100644 --- a/durabletask-azuremanaged/durabletask/azuremanaged/client.py +++ b/durabletask-azuremanaged/durabletask/azuremanaged/client.py @@ -25,10 +25,22 @@ # Client class used for Durable Task Scheduler (DTS) class DurableTaskSchedulerClient(TaskHubGrpcClient): + """A client for Azure Durable Task Scheduler. + + ``resource_id`` optionally overrides the token audience. If None or empty, + it defaults to ``https://durabletask.azure.us`` when ``REGION_NAME`` starts + with ``usgov`` or ``usdod`` (case-insensitive), or ``https://durabletask.io`` + otherwise. Surrounding whitespace, trailing slashes, and an existing + ``/.default`` suffix are removed before requesting the ``/.default`` scope. + Values that become empty raise ``ValueError``. This does not configure the + service endpoint or the credential's authority. + """ + def __init__(self, *, host_address: str, taskhub: str, token_credential: TokenCredential | None, + resource_id: str | None = None, channel: grpc.Channel | None = None, secure_channel: bool = True, interceptors: Sequence[shared.ClientInterceptor] | None = None, @@ -47,7 +59,8 @@ def __init__(self, *, resolved_interceptors: list[shared.ClientInterceptor] = ( list(interceptors) if interceptors is not None else [] ) - resolved_interceptors.append(DTSDefaultClientInterceptorImpl(token_credential, taskhub)) + resolved_interceptors.append(DTSDefaultClientInterceptorImpl( + token_credential, taskhub, resource_id=resource_id)) # We pass in None for the metadata so we don't construct an additional interceptor in the parent class # Since the parent class doesn't use anything metadata for anything else, we can set it as None @@ -81,6 +94,12 @@ class AsyncDurableTaskSchedulerClient(AsyncTaskHubGrpcClient): taskhub (str): The name of the task hub. Cannot be empty. token_credential (TokenCredential | None): Azure credential for authentication. If None, anonymous authentication will be used. + resource_id (str | None, optional): Token audience override. If None or empty, + defaults to ``https://durabletask.azure.us`` when ``REGION_NAME`` starts + with ``usgov`` or ``usdod`` (case-insensitive), or ``https://durabletask.io`` + otherwise. Surrounding whitespace, trailing slashes, and an existing + ``/.default`` suffix are removed before requesting the ``/.default`` scope. + Does not configure the service endpoint or the credential's authority. secure_channel (bool, optional): Whether to use a secure gRPC channel (TLS). Defaults to True. resiliency_options (GrpcClientResiliencyOptions | None, optional): Client-side @@ -98,6 +117,7 @@ class AsyncDurableTaskSchedulerClient(AsyncTaskHubGrpcClient): Raises: ValueError: If taskhub is empty or None. + ValueError: If resource_id becomes empty after normalization. Example: >>> from azure.identity.aio import DefaultAzureCredential @@ -116,6 +136,7 @@ def __init__(self, *, host_address: str, taskhub: str, token_credential: AsyncTokenCredential | None, + resource_id: str | None = None, channel: grpc.aio.Channel | None = None, secure_channel: bool = True, interceptors: Sequence[shared.AsyncClientInterceptor] | None = None, @@ -134,7 +155,8 @@ def __init__(self, *, resolved_interceptors: list[shared.AsyncClientInterceptor] = ( list(interceptors) if interceptors is not None else [] ) - resolved_interceptors.append(DTSAsyncDefaultClientInterceptorImpl(token_credential, taskhub)) + resolved_interceptors.append(DTSAsyncDefaultClientInterceptorImpl( + token_credential, taskhub, resource_id=resource_id)) # We pass in None for the metadata so we don't construct an additional interceptor in the parent class # Since the parent class doesn't use anything metadata for anything else, we can set it as None diff --git a/durabletask-azuremanaged/durabletask/azuremanaged/internal/access_token_manager.py b/durabletask-azuremanaged/durabletask/azuremanaged/internal/access_token_manager.py index fedda9d3..271f1bf8 100644 --- a/durabletask-azuremanaged/durabletask/azuremanaged/internal/access_token_manager.py +++ b/durabletask-azuremanaged/durabletask/azuremanaged/internal/access_token_manager.py @@ -1,6 +1,7 @@ # Copyright (c) Microsoft Corporation. # Licensed under the MIT License. import asyncio +import os from datetime import datetime, timedelta, timezone from threading import Lock @@ -10,14 +11,31 @@ import durabletask.internal.shared as shared +def resolve_resource_id(resource_id: str | None) -> str: + """Normalize an explicit token audience or resolve the current region's default.""" + if resource_id is None or resource_id == "": + region = os.getenv("REGION_NAME", "") + if region.lower().startswith(("usgov", "usdod")): + return "https://durabletask.azure.us" + return "https://durabletask.io" + + resource_id = resource_id.strip().rstrip("/") + if resource_id.lower().endswith("/.default"): + resource_id = resource_id[:-len("/.default")].rstrip("/") + if not resource_id: + raise ValueError("resource_id cannot be empty after normalization.") + return resource_id + + # By default, when there's 10minutes left before the token expires, refresh the token class AccessTokenManager: _token: AccessToken | None expiry_time: datetime | None - def __init__(self, token_credential: TokenCredential | None, refresh_interval_seconds: int = 600): - self._scope = "https://durabletask.io/.default" + def __init__(self, token_credential: TokenCredential | None, refresh_interval_seconds: int = 600, + *, resource_id: str | None = None): + self._scope = f"{resolve_resource_id(resource_id)}/.default" self._refresh_interval_seconds = refresh_interval_seconds self._logger = shared.get_logger("token_manager") @@ -63,8 +81,8 @@ class AsyncAccessTokenManager: _token: AccessToken | None def __init__(self, token_credential: AsyncTokenCredential | None, - refresh_interval_seconds: int = 600): - self._scope = "https://durabletask.io/.default" + refresh_interval_seconds: int = 600, *, resource_id: str | None = None): + self._scope = f"{resolve_resource_id(resource_id)}/.default" self._refresh_interval_seconds = refresh_interval_seconds self._logger = shared.get_logger("async_token_manager") diff --git a/durabletask-azuremanaged/durabletask/azuremanaged/internal/durabletask_grpc_interceptor.py b/durabletask-azuremanaged/durabletask/azuremanaged/internal/durabletask_grpc_interceptor.py index 8d9eebd3..f476934e 100644 --- a/durabletask-azuremanaged/durabletask/azuremanaged/internal/durabletask_grpc_interceptor.py +++ b/durabletask-azuremanaged/durabletask/azuremanaged/internal/durabletask_grpc_interceptor.py @@ -11,6 +11,7 @@ from durabletask.azuremanaged.internal.access_token_manager import ( AccessTokenManager, AsyncAccessTokenManager, + resolve_resource_id, ) from durabletask.internal.grpc_interceptor import ( DefaultAsyncClientInterceptorImpl, @@ -43,7 +44,10 @@ def __init__( self, token_credential: TokenCredential | None, taskhub_name: str, - worker_id: str | None = None): + worker_id: str | None = None, + *, resource_id: str | None = None): + if token_credential is None: + resolve_resource_id(resource_id) user_agent = f"durabletask-python/{_get_sdk_version()}" self._metadata = [ ("taskhub", taskhub_name), @@ -58,7 +62,8 @@ def __init__( self._token_manager = None if token_credential is not None: self._token_credential = token_credential - self._token_manager = AccessTokenManager(token_credential=self._token_credential) + self._token_manager = AccessTokenManager( + token_credential=self._token_credential, resource_id=resource_id) def _upsert_authorization_header(self, token: str) -> None: found = False @@ -91,7 +96,10 @@ class DTSAsyncDefaultClientInterceptorImpl(DefaultAsyncClientInterceptorImpl): This class implements async gRPC interceptors to add DTS-specific headers (task hub name, user agent, and authentication token) to all async calls.""" - def __init__(self, token_credential: AsyncTokenCredential | None, taskhub_name: str): + def __init__(self, token_credential: AsyncTokenCredential | None, taskhub_name: str, + *, resource_id: str | None = None): + if token_credential is None: + resolve_resource_id(resource_id) user_agent = f"durabletask-python/{_get_sdk_version()}" self._metadata = [ ("taskhub", taskhub_name), @@ -104,7 +112,8 @@ def __init__(self, token_credential: AsyncTokenCredential | None, taskhub_name: self._token_manager = None if token_credential is not None: self._token_credential = token_credential - self._token_manager = AsyncAccessTokenManager(token_credential=self._token_credential) + self._token_manager = AsyncAccessTokenManager( + token_credential=self._token_credential, resource_id=resource_id) def _upsert_authorization_header(self, token: str) -> None: found = False diff --git a/durabletask-azuremanaged/durabletask/azuremanaged/preview/sandboxes/client.py b/durabletask-azuremanaged/durabletask/azuremanaged/preview/sandboxes/client.py index e306fe79..c65c274f 100644 --- a/durabletask-azuremanaged/durabletask/azuremanaged/preview/sandboxes/client.py +++ b/durabletask-azuremanaged/durabletask/azuremanaged/preview/sandboxes/client.py @@ -18,13 +18,19 @@ class SandboxActivitiesClient: - """Client for Durable Task Scheduler sandbox activity management operations.""" + """Client for Durable Task Scheduler sandbox activity management operations. + + ``resource_id`` selects the token audience using the same normalization and + ``REGION_NAME`` defaults as ``DurableTaskSchedulerClient``. It does not change + the endpoint or credential authority. A supplied channel owns its authentication. + """ def __init__( self, *, host_address: str, taskhub: str, token_credential: Optional[TokenCredential], + resource_id: Optional[str] = None, channel: Optional[grpc.Channel] = None, secure_channel: bool = True, interceptors: Optional[Sequence[shared.ClientInterceptor]] = None, @@ -33,6 +39,7 @@ def __init__( host_address=host_address, taskhub=taskhub, token_credential=token_credential, + resource_id=resource_id, channel=channel, secure_channel=secure_channel, interceptors=interceptors, diff --git a/durabletask-azuremanaged/durabletask/azuremanaged/preview/sandboxes/transport.py b/durabletask-azuremanaged/durabletask/azuremanaged/preview/sandboxes/transport.py index e2b53cfb..c567b01a 100644 --- a/durabletask-azuremanaged/durabletask/azuremanaged/preview/sandboxes/transport.py +++ b/durabletask-azuremanaged/durabletask/azuremanaged/preview/sandboxes/transport.py @@ -40,6 +40,7 @@ def __init__( host_address: str, taskhub: str, token_credential: Optional[TokenCredential], + resource_id: Optional[str] = None, channel: Optional[grpc.Channel] = None, secure_channel: bool = True, interceptors: Optional[Sequence[shared.ClientInterceptor]] = None, @@ -52,7 +53,8 @@ def __init__( resolved_interceptors: list[shared.ClientInterceptor] = ( list(interceptors) if interceptors is not None else [] ) - resolved_interceptors.append(DTSDefaultClientInterceptorImpl(token_credential, taskhub)) + resolved_interceptors.append(DTSDefaultClientInterceptorImpl( + token_credential, taskhub, resource_id=resource_id)) channel = shared.get_grpc_channel( host_address=host_address, secure_channel=secure_channel, diff --git a/durabletask-azuremanaged/durabletask/azuremanaged/preview/sandboxes/worker.py b/durabletask-azuremanaged/durabletask/azuremanaged/preview/sandboxes/worker.py index c87d12da..384ca6b3 100644 --- a/durabletask-azuremanaged/durabletask/azuremanaged/preview/sandboxes/worker.py +++ b/durabletask-azuremanaged/durabletask/azuremanaged/preview/sandboxes/worker.py @@ -12,6 +12,7 @@ from azure.identity import ManagedIdentityCredential from durabletask.azuremanaged.internal import sandbox_service_pb2 as pb +from durabletask.azuremanaged.internal.access_token_manager import resolve_resource_id from durabletask.azuremanaged.preview.sandboxes.helpers import SandboxActivity from durabletask.azuremanaged.preview.sandboxes.helpers import resolve_activities from durabletask.azuremanaged.preview.sandboxes.worker_profiles import ( @@ -38,9 +39,15 @@ class SandboxWorker(DurableTaskSchedulerWorker): This worker registers a live worker session with Durable Task Scheduler and restricts dispatch to the activities registered on this worker. + + ``resource_id`` selects the token audience using the same normalization and + ``REGION_NAME`` defaults as ``DurableTaskSchedulerWorker``. It applies to both + activity execution and worker registration without changing the runtime endpoint + or managed identity configuration. """ - def __init__(self) -> None: + def __init__(self, *, resource_id: str | None = None) -> None: + resolved_resource_id = resource_id or resolve_resource_id(None) resolved_host_address = _resolve_host_address() resolved_taskhub = _resolve_taskhub() resolved_secure_channel = _resolve_secure_channel(resolved_host_address) @@ -54,12 +61,14 @@ def __init__(self) -> None: self._sandbox_host_address = resolved_host_address self._sandbox_secure_channel = resolved_secure_channel self._sandbox_token_credential = resolved_token_credential + self._sandbox_resource_id = resolved_resource_id self._sandbox_logger = shared.get_logger("worker") super().__init__( host_address=resolved_host_address, taskhub=resolved_taskhub, token_credential=resolved_token_credential, + resource_id=resolved_resource_id, secure_channel=resolved_secure_channel, concurrency_options=concurrency_options) @@ -139,6 +148,7 @@ def _run_sandbox_registration_loop(self) -> None: host_address=self._sandbox_host_address, taskhub=self._sandbox_taskhub, token_credential=self._sandbox_token_credential, + resource_id=self._sandbox_resource_id, secure_channel=self._sandbox_secure_channel) client.connect_sandbox_activity_worker(self._registration_messages()) retry_delay = 1.0 diff --git a/durabletask-azuremanaged/durabletask/azuremanaged/worker.py b/durabletask-azuremanaged/durabletask/azuremanaged/worker.py index fcc84e82..f4d1a9c4 100644 --- a/durabletask-azuremanaged/durabletask/azuremanaged/worker.py +++ b/durabletask-azuremanaged/durabletask/azuremanaged/worker.py @@ -39,6 +39,12 @@ class DurableTaskSchedulerWorker(TaskHubGrpcWorker): taskhub (str): The name of the task hub. Cannot be empty. token_credential (TokenCredential | None): Azure credential for authentication. If None, anonymous authentication will be used. + resource_id (str | None, optional): Token audience override. If None or empty, + defaults to ``https://durabletask.azure.us`` when ``REGION_NAME`` starts + with ``usgov`` or ``usdod`` (case-insensitive), or ``https://durabletask.io`` + otherwise. Surrounding whitespace, trailing slashes, and an existing + ``/.default`` suffix are removed before requesting the ``/.default`` scope. + Does not configure the service endpoint or the credential's authority. secure_channel (bool, optional): Whether to use a secure gRPC channel (TLS). Defaults to True. resiliency_options (GrpcWorkerResiliencyOptions | None, optional): Worker-side @@ -60,6 +66,7 @@ class DurableTaskSchedulerWorker(TaskHubGrpcWorker): Raises: ValueError: If taskhub is empty or None. + ValueError: If resource_id becomes empty after normalization. Example: >>> from azure.identity import DefaultAzureCredential @@ -86,6 +93,7 @@ def __init__(self, *, host_address: str, taskhub: str, token_credential: TokenCredential | None, + resource_id: str | None = None, channel: grpc.Channel | None = None, secure_channel: bool = True, interceptors: Sequence[shared.ClientInterceptor] | None = None, @@ -107,7 +115,8 @@ def __init__(self, *, list(interceptors) if interceptors is not None else [] ) resolved_interceptors.append( - DTSDefaultClientInterceptorImpl(token_credential, taskhub, worker_id=worker_id) + DTSDefaultClientInterceptorImpl( + token_credential, taskhub, worker_id=worker_id, resource_id=resource_id) ) # We pass in None for the metadata so we don't construct an additional interceptor in the parent class diff --git a/examples/sandboxes/README.md b/examples/sandboxes/README.md index 8d44f865..a9946fc7 100644 --- a/examples/sandboxes/README.md +++ b/examples/sandboxes/README.md @@ -58,6 +58,12 @@ the SDK. In a sandbox, `SandboxWorker()` reads `DTS_ENDPOINT`, `DTS_TASK_HUB`, Durable Task Scheduler. It also reads optional `DTS_SANDBOX_PROVIDER` metadata when present. The worker requires `DTS_AUTHENTICATION=ManagedIdentity` and reports its sandbox ID plus registered activity identities when it connects. + +The optional `SandboxWorker(resource_id=...)` argument overrides only the token +audience, for both activity execution and worker registration. When omitted, +`REGION_NAME` selects the audience as described in +[Azure Managed authentication](../../docs/getting-started.md#configure-azure-managed-authentication). +It does not override the runtime-injected endpoint or identity. Durable Task Scheduler validates they match the worker_profile before advertising worker capacity. diff --git a/tests/durabletask-azuremanaged/test_resource_id.py b/tests/durabletask-azuremanaged/test_resource_id.py new file mode 100644 index 00000000..8fe2b44c --- /dev/null +++ b/tests/durabletask-azuremanaged/test_resource_id.py @@ -0,0 +1,203 @@ +# Copyright (c) Microsoft Corporation. +# Licensed under the MIT License. + +from collections.abc import Callable +from datetime import datetime, timedelta, timezone +from unittest.mock import AsyncMock, Mock, call, patch + +import grpc +import pytest +from azure.core.credentials import AccessToken + +from durabletask.azuremanaged.client import ( + AsyncDurableTaskSchedulerClient, + DurableTaskSchedulerClient, +) +from durabletask.azuremanaged.internal.access_token_manager import ( + AccessTokenManager, + AsyncAccessTokenManager, +) +from durabletask.azuremanaged.preview.sandboxes.client import SandboxActivitiesClient +from durabletask.azuremanaged.worker import DurableTaskSchedulerWorker + + +_PUBLIC = "https://durabletask.io" +_GOVERNMENT = "https://durabletask.azure.us" +_RESOURCE_CASES = [ + (None, None, _PUBLIC), + ("", None, _PUBLIC), + ("westus2", None, _PUBLIC), + ("chinaeast2", None, _PUBLIC), + ("notusgov", None, _PUBLIC), + ("notusdod", None, _PUBLIC), + ("usgovvirginia", None, _GOVERNMENT), + ("USGOVARIZONA", None, _GOVERNMENT), + ("UsGovTexas", None, _GOVERNMENT), + ("usdodcentral", None, _GOVERNMENT), + ("USDODEAST", None, _GOVERNMENT), + ("UsDodCentral", None, _GOVERNMENT), + (None, "", _PUBLIC), + ("usgovvirginia", "", _GOVERNMENT), + ("usdodcentral", "", _GOVERNMENT), + ("usgovvirginia", _PUBLIC, _PUBLIC), + ("usdodcentral", _PUBLIC, _PUBLIC), + ("westus2", _GOVERNMENT, _GOVERNMENT), + ("chinaeast2", "https://durabletask.example", "https://durabletask.example"), + (None, _GOVERNMENT + "/", _GOVERNMENT), + (None, _GOVERNMENT + "/.default", _GOVERNMENT), + (None, _GOVERNMENT + "//.default//", _GOVERNMENT), + (None, " \t" + _GOVERNMENT + "/.default/ \t", _GOVERNMENT), + ("usgovvirginia", "api://CustomAudience/resource/.DEFAULT/", "api://CustomAudience/resource"), + (None, "api://custom/.default/.default", "api://custom/.default"), +] +_INVALID_RESOURCE_IDS = [" \t ", "///", "/.default", " /.DEFAULT/// "] +_EXPIRED = datetime.now(timezone.utc) - timedelta(hours=1) + + +def _set_region(monkeypatch: pytest.MonkeyPatch, region: str | None) -> None: + if region is None: + monkeypatch.delenv("REGION_NAME", raising=False) + else: + monkeypatch.setenv("REGION_NAME", region) + + +@pytest.mark.parametrize(("region", "resource_id", "expected_resource"), _RESOURCE_CASES) +def test_sync_token_scope_and_refresh( + monkeypatch: pytest.MonkeyPatch, region: str | None, + resource_id: str | None, expected_resource: str) -> None: + _set_region(monkeypatch, region) + token = AccessToken("test-token", 9999999999) + credential = Mock() + credential.get_token.return_value = token + manager = AccessTokenManager(credential, resource_id=resource_id) + + credential.get_token.assert_not_called() + assert manager.get_access_token() == token + assert manager.get_access_token() == token + credential.get_token.assert_called_once_with(f"{expected_resource}/.default") + + manager.expiry_time = _EXPIRED + assert manager.get_access_token() == token + assert credential.get_token.call_args_list == [call(f"{expected_resource}/.default")] * 2 + + +@pytest.mark.parametrize(("region", "resource_id", "expected_resource"), _RESOURCE_CASES) +async def test_async_token_scope_and_refresh( + monkeypatch: pytest.MonkeyPatch, region: str | None, + resource_id: str | None, expected_resource: str) -> None: + _set_region(monkeypatch, region) + token = AccessToken("test-token", 9999999999) + credential = Mock() + credential.get_token = AsyncMock(return_value=token) + manager = AsyncAccessTokenManager(credential, resource_id=resource_id) + + credential.get_token.assert_not_called() + assert await manager.get_access_token() == token + assert await manager.get_access_token() == token + credential.get_token.assert_awaited_once_with(f"{expected_resource}/.default") + + manager.expiry_time = _EXPIRED + assert await manager.get_access_token() == token + assert credential.get_token.await_args_list == [call(f"{expected_resource}/.default")] * 2 + + +@pytest.mark.parametrize("resource_id", _INVALID_RESOURCE_IDS) +@pytest.mark.parametrize("manager_type", [AccessTokenManager, AsyncAccessTokenManager]) +def test_invalid_resource_ids_are_rejected_without_credentials( + resource_id: str, manager_type: type[AccessTokenManager] | type[AsyncAccessTokenManager]) -> None: + with pytest.raises(ValueError, match="resource_id cannot be empty after normalization"): + manager_type(None, resource_id=resource_id) + + +async def test_region_default_is_resolved_per_manager_and_pinned_for_refresh( + monkeypatch: pytest.MonkeyPatch) -> None: + credential = Mock() + credential.get_token.return_value = AccessToken("sync-token", 9999999999) + async_credential = Mock() + async_credential.get_token = AsyncMock(return_value=AccessToken("async-token", 9999999999)) + + monkeypatch.setenv("REGION_NAME", "usgovvirginia") + government = AccessTokenManager(credential) + async_government = AsyncAccessTokenManager(async_credential) + monkeypatch.setenv("REGION_NAME", "westus2") + public = AccessTokenManager(credential) + async_public = AsyncAccessTokenManager(async_credential) + monkeypatch.setenv("REGION_NAME", "usdodcentral") + + for manager in (government, public): + manager.get_access_token() + manager.expiry_time = _EXPIRED + manager.get_access_token() + for async_manager in (async_government, async_public): + await async_manager.get_access_token() + async_manager.expiry_time = _EXPIRED + await async_manager.get_access_token() + + expected_calls = [call(f"{_GOVERNMENT}/.default")] * 2 + [call(f"{_PUBLIC}/.default")] * 2 + assert credential.get_token.call_args_list == expected_calls + assert async_credential.get_token.await_args_list == expected_calls + + +@pytest.mark.parametrize(("factory", "base_init", "is_async"), [ + (DurableTaskSchedulerClient, "durabletask.azuremanaged.client.TaskHubGrpcClient.__init__", False), + (DurableTaskSchedulerWorker, "durabletask.azuremanaged.worker.TaskHubGrpcWorker.__init__", False), + (AsyncDurableTaskSchedulerClient, "durabletask.azuremanaged.client.AsyncTaskHubGrpcClient.__init__", True), + (SandboxActivitiesClient, "durabletask.azuremanaged.preview.sandboxes.transport.shared.get_grpc_channel", False), +]) +@pytest.mark.parametrize(("region", "resource_id", "expected_resource"), [ + (None, None, _PUBLIC), + ("UsGovVirginia", None, _GOVERNMENT), + ("USDODEAST", "", _GOVERNMENT), + ("usgovvirginia", " \t" + _PUBLIC + "/.DEFAULT/ \t", _PUBLIC), + ("westus2", _GOVERNMENT, _GOVERNMENT), + ("usgovvirginia", "api://custom/resource", "api://custom/resource"), + ("westus2", "api://custom/.default/.default", "api://custom/.default"), +]) +async def test_public_clients_and_worker_request_configured_scope( + monkeypatch: pytest.MonkeyPatch, factory: Callable[..., object], base_init: str, + is_async: bool, region: str | None, resource_id: str | None, + expected_resource: str) -> None: + _set_region(monkeypatch, region) + token = AccessToken("configured-token", 9999999999) + credential = Mock() + credential.get_token = AsyncMock(return_value=token) if is_async else Mock(return_value=token) + + with patch(base_init) as init: + init.return_value = Mock() if factory is SandboxActivitiesClient else None + factory( + host_address="localhost:4001", + taskhub="test-hub", + token_credential=credential, + resource_id=resource_id, + ) + + credential.get_token.assert_not_called() + assert init.call_args.kwargs["host_address"] == "localhost:4001" + interceptor = init.call_args.kwargs["interceptors"][-1] + details = Mock( + spec=grpc.ClientCallDetails, method="/test", timeout=None, metadata=(), + credentials=None, wait_for_ready=False, compression=None, + ) + if is_async: + result = await interceptor._intercept_call(details) + credential.get_token.assert_awaited_once_with(f"{expected_resource}/.default") + else: + result = interceptor._intercept_call(details) + credential.get_token.assert_called_once_with(f"{expected_resource}/.default") + metadata = dict(result.metadata) + assert metadata["taskhub"] == "test-hub" + assert metadata["authorization"] == "Bearer configured-token" + if factory is DurableTaskSchedulerWorker: + assert metadata["workerid"] + + +@pytest.mark.parametrize("factory", [ + DurableTaskSchedulerClient, AsyncDurableTaskSchedulerClient, + DurableTaskSchedulerWorker, SandboxActivitiesClient, +]) +@pytest.mark.parametrize("resource_id", _INVALID_RESOURCE_IDS) +def test_public_constructors_reject_invalid_resource_without_credentials( + factory: Callable[..., object], resource_id: str) -> None: + with pytest.raises(ValueError, match="resource_id cannot be empty after normalization"): + factory(host_address="localhost:4001", taskhub="test-hub", + token_credential=None, resource_id=resource_id) diff --git a/tests/durabletask-azuremanaged/test_sandboxes_extension.py b/tests/durabletask-azuremanaged/test_sandboxes_extension.py index 8b66abe3..f4c1876a 100644 --- a/tests/durabletask-azuremanaged/test_sandboxes_extension.py +++ b/tests/durabletask-azuremanaged/test_sandboxes_extension.py @@ -700,7 +700,10 @@ def test_generated_stub_uses_sandbox_rpc_paths() -> None: def test_sandbox_worker_constructor_does_not_expose_runtime_contract() -> None: - assert list(inspect.signature(SandboxWorker).parameters) == [] + parameters = inspect.signature(SandboxWorker).parameters + assert list(parameters) == ["resource_id"] + assert parameters["resource_id"].kind == inspect.Parameter.KEYWORD_ONLY + assert parameters["resource_id"].default is None assert "_execute_activity" not in SandboxWorker.__dict__ assert "add_activity" not in SandboxWorker.__dict__ @@ -1042,6 +1045,56 @@ def connect(_messages): assert backoff.upper_bounds == [1.0, 2.0, 4.0] +@pytest.mark.parametrize(("region", "resource_id", "expected_resource"), [ + ("westus2", None, "https://durabletask.io"), + ("UsGovVirginia", None, "https://durabletask.azure.us"), + ("USDODEAST", "", "https://durabletask.azure.us"), + ("usgovvirginia", " https://durabletask.example/.default/ ", "https://durabletask.example"), + ("westus2", "api://custom/.default/.default", "api://custom/.default"), +]) +def test_sandbox_worker_uses_same_audience_for_execution_and_registration( + monkeypatch, region: str, resource_id: str | None, expected_resource: str) -> None: + monkeypatch.setenv("REGION_NAME", region) + worker = _build_registration_test_worker(monkeypatch, resource_id=resource_id) + requested_scopes: list[tuple[str, ...]] = [] + + def get_token(*scopes, **kwargs): + requested_scopes.append(scopes) + return AccessToken("sandbox-token", 9999999999) + + monkeypatch.setattr(worker._sandbox_token_credential, "get_token", get_token) + worker._interceptors[-1]._token_manager.get_access_token() + assert requested_scopes == [(f"{expected_resource}/.default",)] + + # Registration can begin or reconnect after the environment has changed. + monkeypatch.setenv("REGION_NAME", "usdodcentral" if region == "westus2" else "westus2") + attempts = 0 + + def connect(_messages): + nonlocal attempts + attempts += 1 + if attempts == 1: + raise _FakeRpcError(grpc.StatusCode.CANCELLED, "Channel closed!") + worker._sandbox_registration_stop.set() + return object() + + transports = _install_fake_registration_transport(monkeypatch, connect) + monkeypatch.setattr(sandbox_worker, "random", _StubBackoff()) + worker._run_sandbox_registration_loop() + + assert len(transports) == 2 + assert all(transport.kwargs["resource_id"] == (resource_id or expected_resource) + for transport in transports) + assert all(transport.kwargs["token_credential"] is worker._sandbox_token_credential + for transport in transports) + + +@pytest.mark.parametrize("resource_id", [" \t ", "///", "/.default", " /.DEFAULT/// "]) +def test_sandbox_worker_rejects_invalid_resource_id(monkeypatch, resource_id: str) -> None: + with pytest.raises(ValueError, match="resource_id cannot be empty after normalization"): + _build_registration_test_worker(monkeypatch, resource_id=resource_id) + + def test_sandbox_registration_rebuilds_transport_after_channel_shutdown(monkeypatch) -> None: worker = _build_registration_test_worker(monkeypatch) attempts = 0 @@ -1129,7 +1182,7 @@ def factory(**kwargs) -> _FakeRegistrationTransport: return transports -def _build_registration_test_worker(monkeypatch) -> SandboxWorker: +def _build_registration_test_worker(monkeypatch, resource_id: str | None = None) -> SandboxWorker: monkeypatch.setenv("DTS_ENDPOINT", "http://localhost:8080") monkeypatch.setenv("DTS_TASK_HUB", "env-hub") monkeypatch.setenv("DTS_WORKER_PROFILE_ID", "env-profile") @@ -1139,7 +1192,7 @@ def _build_registration_test_worker(monkeypatch) -> SandboxWorker: def RegistrationActivity(_ctx, value): return value - worker = SandboxWorker() + worker = SandboxWorker(resource_id=resource_id) worker.add_activity(RegistrationActivity) worker._configure_sandbox_activity_filters() worker._sandbox_registration_stop.clear()