Skip to content
Merged
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
Original file line number Diff line number Diff line change
Expand Up @@ -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(
Expand Down Expand Up @@ -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:
Expand Down
142 changes: 142 additions & 0 deletions tests/durabletask-azuremanaged/test_sandboxes_extension.py
Original file line number Diff line number Diff line change
Expand Up @@ -2,6 +2,7 @@
# Licensed under the MIT License.

import inspect
import threading

import grpc
from azure.core.credentials import AccessToken
Expand Down Expand Up @@ -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] = []
Expand Down
Loading