diff --git a/durabletask-azuremanaged/durabletask/azuremanaged/preview/sandboxes/worker.py b/durabletask-azuremanaged/durabletask/azuremanaged/preview/sandboxes/worker.py index 21cb7161..c87d12da 100644 --- a/durabletask-azuremanaged/durabletask/azuremanaged/preview/sandboxes/worker.py +++ b/durabletask-azuremanaged/durabletask/azuremanaged/preview/sandboxes/worker.py @@ -130,30 +130,50 @@ def _stop_sandbox_registration(self) -> None: def _run_sandbox_registration_loop(self) -> None: retry_delay = 1.0 - while not self._sandbox_registration_stop.is_set(): - try: - client = SandboxActivitiesGrpcTransport( - host_address=self._sandbox_host_address, - taskhub=self._sandbox_taskhub, - token_credential=self._sandbox_token_credential, - secure_channel=self._sandbox_secure_channel) + client: Optional[SandboxActivitiesGrpcTransport] = None + try: + while not self._sandbox_registration_stop.is_set(): try: + if client is None: + client = SandboxActivitiesGrpcTransport( + host_address=self._sandbox_host_address, + taskhub=self._sandbox_taskhub, + token_credential=self._sandbox_token_credential, + secure_channel=self._sandbox_secure_channel) client.connect_sandbox_activity_worker(self._registration_messages()) retry_delay = 1.0 - finally: - client.close() - except Exception as ex: - if self._sandbox_registration_stop.is_set(): - break - if not _is_retriable_registration_failure(ex): - self._sandbox_logger.error( - "Sandbox activity worker registration failed permanently: %s", ex) - self._sandbox_registration_stop.set() - break - self._sandbox_logger.warning("Sandbox activity worker registration failed: %s", ex) - delay = random.uniform(0, retry_delay) - self._sandbox_registration_stop.wait(delay) - retry_delay = min(retry_delay * 2, 30.0) + except Exception as ex: + # gRPC channels reconnect on their own after transient stream + # failures, so the transport (and its channel, stub, and cached + # access token) is reused instead of being rebuilt per attempt. + if _requires_new_registration_transport(ex): + self._close_sandbox_registration_transport(client) + client = None + if self._sandbox_registration_stop.is_set(): + break + if not _is_retriable_registration_failure(ex): + self._sandbox_logger.error( + "Sandbox activity worker registration failed permanently: %s", ex) + self._sandbox_registration_stop.set() + break + self._sandbox_logger.warning( + "Sandbox activity worker registration failed: %s", ex) + delay = random.uniform(0, retry_delay) + self._sandbox_registration_stop.wait(delay) + retry_delay = min(retry_delay * 2, 30.0) + finally: + self._close_sandbox_registration_transport(client) + + def _close_sandbox_registration_transport( + self, + client: Optional[SandboxActivitiesGrpcTransport]) -> None: + if client is None: + return + try: + client.close() + except Exception as ex: + self._sandbox_logger.debug( + "Sandbox activity worker registration transport failed to close: %s", ex) def _registration_messages(self) -> Iterator[pb.SandboxActivityWorkerMessage]: yield build_sandbox_worker_start( @@ -264,6 +284,26 @@ def _is_retriable_registration_failure(ex: Exception) -> bool: return isinstance(ex, OSError) +def _requires_new_registration_transport(ex: Exception) -> bool: + """Returns whether a failure leaves the registration transport unusable. + + A gRPC channel transparently reconnects after a stream fails, so the same + transport can serve the next registration attempt. A channel that has been + shut down rejects every later RPC, and a failure raised outside of gRPC (for + example while the channel itself is being created) can leave the transport + only partially initialized, so both cases need a replacement transport. + """ + if not isinstance(ex, grpc.RpcError): + return True + + if ex.code() != grpc.StatusCode.CANCELLED: + return False + + details = getattr(ex, "details", None) + resolved_details = details() if callable(details) else None + return isinstance(resolved_details, str) and "channel closed" in resolved_details.lower() + + def _resolve_max_concurrent_activities() -> int: value = os.getenv("DTS_SANDBOX_MAX_ACTIVITIES") if value is None: diff --git a/tests/durabletask-azuremanaged/test_sandboxes_extension.py b/tests/durabletask-azuremanaged/test_sandboxes_extension.py index caf96571..15991eeb 100644 --- a/tests/durabletask-azuremanaged/test_sandboxes_extension.py +++ b/tests/durabletask-azuremanaged/test_sandboxes_extension.py @@ -2,6 +2,7 @@ # Licensed under the MIT License. import inspect +import threading import grpc from azure.core.credentials import AccessToken @@ -766,6 +767,147 @@ def test_sandbox_registration_retries_failed_precondition_until_structured_reaso _FakeRpcError(grpc.StatusCode.FAILED_PRECONDITION, "worker profile does not match")) +def test_sandbox_registration_recreates_transport_only_when_channel_is_unusable() -> None: + assert not sandbox_worker._requires_new_registration_transport( + _FakeRpcError(grpc.StatusCode.UNAVAILABLE)) + assert not sandbox_worker._requires_new_registration_transport( + _FakeRpcError(grpc.StatusCode.FAILED_PRECONDITION, "sandbox not ready")) + assert not sandbox_worker._requires_new_registration_transport( + _FakeRpcError(grpc.StatusCode.CANCELLED, "Locally cancelled by application!")) + assert sandbox_worker._requires_new_registration_transport( + _FakeRpcError(grpc.StatusCode.CANCELLED, "Channel closed!")) + assert sandbox_worker._requires_new_registration_transport(OSError("connection reset")) + + +def test_sandbox_registration_reuses_one_transport_across_retriable_failures(monkeypatch) -> None: + worker = _build_registration_test_worker(monkeypatch) + attempts = 0 + + def connect(_messages): + nonlocal attempts + attempts += 1 + if attempts <= 3: + raise _FakeRpcError(grpc.StatusCode.UNAVAILABLE) + worker._sandbox_registration_stop.set() + return object() + + transports = _install_fake_registration_transport(monkeypatch, connect) + backoff = _StubBackoff() + monkeypatch.setattr(sandbox_worker, "random", backoff) + + worker._run_sandbox_registration_loop() + + assert attempts == 4 + assert len(transports) == 1 + assert transports[0].close_count == 1 + # The existing exponential backoff schedule must be preserved. + assert backoff.upper_bounds == [1.0, 2.0, 4.0] + + +def test_sandbox_registration_rebuilds_transport_after_channel_shutdown(monkeypatch) -> None: + worker = _build_registration_test_worker(monkeypatch) + 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 attempts == 2 + assert len(transports) == 2 + assert [transport.close_count for transport in transports] == [1, 1] + + +def test_sandbox_registration_loop_exits_promptly_on_stop_signal(monkeypatch) -> None: + worker = _build_registration_test_worker(monkeypatch) + worker._sandbox_heartbeat_interval_seconds = 0.05 + connected = threading.Event() + + def connect(messages): + connected.set() + for _ in messages: + pass + return object() + + transports = _install_fake_registration_transport(monkeypatch, connect) + + thread = threading.Thread(target=worker._run_sandbox_registration_loop, daemon=True) + thread.start() + try: + assert connected.wait(10) + worker._sandbox_registration_stop.set() + thread.join(timeout=10) + assert not thread.is_alive() + finally: + worker._sandbox_registration_stop.set() + thread.join(timeout=10) + + assert len(transports) == 1 + assert transports[0].close_count == 1 + + +class _StubBackoff: + """Records the backoff window used for each retry and never actually sleeps.""" + + def __init__(self) -> None: + self.upper_bounds: list[float] = [] + + def uniform(self, lower: float, upper: float) -> float: + self.upper_bounds.append(upper) + return 0.0 + + +class _FakeRegistrationTransport: + def __init__(self, connect, kwargs) -> None: + self.kwargs = kwargs + self.close_count = 0 + self._connect = connect + + def connect_sandbox_activity_worker(self, messages): + return self._connect(messages) + + def close(self) -> None: + self.close_count += 1 + + +def _install_fake_registration_transport( + monkeypatch, connect) -> list[_FakeRegistrationTransport]: + transports: list[_FakeRegistrationTransport] = [] + + def factory(**kwargs) -> _FakeRegistrationTransport: + transport = _FakeRegistrationTransport(connect, kwargs) + transports.append(transport) + return transport + + monkeypatch.setattr(sandbox_worker, "SandboxActivitiesGrpcTransport", factory) + return transports + + +def _build_registration_test_worker(monkeypatch) -> SandboxWorker: + monkeypatch.setenv("DTS_ENDPOINT", "http://localhost:8080") + monkeypatch.setenv("DTS_TASK_HUB", "env-hub") + monkeypatch.setenv("DTS_WORKER_PROFILE_ID", "env-profile") + monkeypatch.setenv("DTS_SANDBOX_PROVIDER", "Sandbox") + _configure_sandbox_worker_auth(monkeypatch) + + def RegistrationActivity(_ctx, value): + return value + + worker = SandboxWorker() + worker.add_activity(RegistrationActivity) + worker._configure_sandbox_activity_filters() + worker._sandbox_registration_stop.clear() + return worker + + class _RecordingChannel: def __init__(self) -> None: self.methods: list[str] = []