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
30 changes: 30 additions & 0 deletions src/durable_workflow/retry_policy.py
Original file line number Diff line number Diff line change
Expand Up @@ -80,6 +80,18 @@ def _backend_unavailable_refusal(exc: Exception) -> tuple[bool, int | None]:
"/api/worker/heartbeat": "heartbeat_worker",
}
operation = next((name for path, name in operations.items() if request.url.path.endswith(path)), None)
_, task_heartbeat_marker, task_heartbeat_tail = request.url.path.rpartition(
"/api/worker/workflow-tasks/"
)
task_id = task_heartbeat_tail.removesuffix("/heartbeat")
task_heartbeat = (
bool(task_heartbeat_marker)
and task_heartbeat_tail.endswith("/heartbeat")
and bool(task_id)
and "/" not in task_id
)
if task_heartbeat:
operation = "heartbeat_workflow_task"
if operation is None:
return False, None
try:
Expand All @@ -97,6 +109,24 @@ def _backend_unavailable_refusal(exc: Exception) -> tuple[bool, int | None]:
worker_id = submitted.get("worker_id")
queue = submitted.get("task_queue")
delay = body.get("retry_after_seconds")
if task_heartbeat:
lease_owner = submitted.get("lease_owner")
attempt = submitted.get("workflow_task_attempt")
if (
not isinstance(lease_owner, str) or not lease_owner
or type(attempt) is not int or attempt <= 0
or body.get("operation") != operation
or body.get("outcome") != "unknown"
or body.get("worker_id") != lease_owner
or body.get("task_queue") is not None
or body.get("task_id") != task_id
or body.get("lease_owner") != lease_owner
or body.get("workflow_task_attempt") != attempt
or body.get("retryable") is not True
or type(delay) is not int or delay <= 0
):
return True, None
return True, delay
if (
not isinstance(worker_id, str) or not worker_id
or body.get("operation") != operation
Expand Down
67 changes: 63 additions & 4 deletions tests/test_storage_admission.py
Original file line number Diff line number Diff line change
Expand Up @@ -42,7 +42,8 @@ def pressure(poll_id: str | None = None, *, reason: str = "storage_pressure", en
def backend_unavailable(request: httpx.Request) -> dict[str, Any]:
submitted = json.loads(request.content)
path = request.url.path
operation = {
task_heartbeat = "/api/worker/workflow-tasks/" in path and path.endswith("/heartbeat")
operation = "heartbeat_workflow_task" if task_heartbeat else {
"/api/worker/workflow-tasks/poll": "poll_workflow_task",
"/api/worker/activity-tasks/poll": "poll_activity_task",
"/api/worker/query-tasks/poll": "poll_query_task",
Expand All @@ -51,9 +52,16 @@ def backend_unavailable(request: httpx.Request) -> dict[str, Any]:
}[path]
body: dict[str, Any] = {
"reason": "backend_unavailable", "operation": operation, "outcome": "unknown",
"worker_id": submitted["worker_id"], "task_queue": submitted.get("task_queue"),
"worker_id": submitted.get("lease_owner") if task_heartbeat else submitted["worker_id"],
"task_queue": submitted.get("task_queue"),
"retryable": True, "retry_after_seconds": 1,
}
if task_heartbeat:
body.update({
"task_id": path.rpartition("/api/worker/workflow-tasks/")[2].removesuffix("/heartbeat"),
"lease_owner": submitted["lease_owner"],
"workflow_task_attempt": submitted["workflow_task_attempt"],
})
if path.endswith("/poll"):
body.update({
"task": None, "poll_status": "backend_unavailable",
Expand Down Expand Up @@ -85,9 +93,9 @@ async def sleep(delay: float) -> None:
return sleeps


def client_for(handler: Callable[..., Any]) -> Client:
def client_for(handler: Callable[..., Any], base_url: str = "https://runtime.example") -> Client:
client = Client(
"https://runtime.example", token="test-runtime-token",
base_url, token="test-runtime-token",
retry_policy=TransportRetryPolicy(max_attempts=2, initial_backoff_seconds=0, jitter=False),
)
client._http = httpx.AsyncClient(base_url=client.base_url, transport=httpx.MockTransport(handler))
Expand Down Expand Up @@ -186,6 +194,57 @@ def handler(request: httpx.Request) -> httpx.Response:
assert sum(retry_sleeps) == pytest.approx(4)


@pytest.mark.parametrize("base_url", ["https://runtime.example", "https://runtime.example/managed/namespace"])
async def test_backend_outage_retries_workflow_task_heartbeat_with_same_fence(
base_url: str, retry_sleeps: list[float],
) -> None:
requests: list[bytes] = []

def handler(request: httpx.Request) -> httpx.Response:
requests.append(request.content)
if len(requests) <= 4:
return httpx.Response(503, json=backend_unavailable(request))
return httpx.Response(200, json={
"task_id": "task-1", "lease_owner": "worker-1",
"workflow_task_attempt": 3, "renewed": True,
})

async with client_for(handler, base_url) as client:
with worker_scope():
result = await client.heartbeat_workflow_task(
task_id="task-1", lease_owner="worker-1", workflow_task_attempt=3,
)
assert result["renewed"] is True
assert len(requests) == 5
assert len(set(requests)) == 1
assert sum(retry_sleeps) == pytest.approx(4)


@pytest.mark.parametrize("override", [
{"task_id": "other-task"}, {"lease_owner": "other-worker"},
{"worker_id": "other-worker"}, {"workflow_task_attempt": 4},
{"operation": "heartbeat_worker"}, {"outcome": "failed"},
{"retryable": False}, {"retry_after_seconds": 0},
])
async def test_invalid_workflow_task_heartbeat_backend_response_is_not_retried(
override: dict[str, Any], retry_sleeps: list[float],
) -> None:
calls = 0

def handler(request: httpx.Request) -> httpx.Response:
nonlocal calls
calls += 1
return httpx.Response(503, json={**backend_unavailable(request), **override})

async with client_for(handler) as client:
with worker_scope(), pytest.raises(ServerError):
await client.heartbeat_workflow_task(
task_id="task-1", lease_owner="worker-1", workflow_task_attempt=3,
)
assert calls == 1
assert not retry_sleeps


@pytest.mark.parametrize("override", [
{"operation": "poll_activity_task"},
{"outcome": "failed"}, {"worker_id": "other-worker"},
Expand Down
Loading