Skip to content

Commit c0678f2

Browse files
Preserve worker poll identity across backend outages (#73)
* Preserve Python worker poll identity during backend outages * Fail closed on ambiguous backend refusal contracts --------- Co-authored-by: Durable Workflow <support@durable-workflow.com>
1 parent 98c9e53 commit c0678f2

2 files changed

Lines changed: 237 additions & 7 deletions

File tree

‎src/durable_workflow/retry_policy.py‎

Lines changed: 72 additions & 7 deletions
Original file line numberDiff line numberDiff line change
@@ -65,6 +65,61 @@ def _storage_refusal(exc: Exception) -> tuple[ServerError, str | None] | None:
6565
return error, poll_id
6666

6767

68+
def _backend_unavailable_refusal(exc: Exception) -> tuple[bool, int | None]:
69+
if not isinstance(exc, httpx.HTTPStatusError) or exc.response.status_code != 503:
70+
return False, None
71+
try:
72+
body = exc.response.json()
73+
except ValueError:
74+
return False, None
75+
if not isinstance(body, dict) or body.get("reason") != "backend_unavailable":
76+
return False, None
77+
request = exc.request
78+
if request.method != "POST" or "X-Durable-Workflow-Protocol-Version" not in request.headers:
79+
return True, None
80+
operations = {
81+
"/api/worker/workflow-tasks/poll": "poll_workflow_task",
82+
"/api/worker/activity-tasks/poll": "poll_activity_task",
83+
"/api/worker/query-tasks/poll": "poll_query_task",
84+
"/api/worker/register": "register_worker",
85+
"/api/worker/heartbeat": "heartbeat_worker",
86+
}
87+
operation = next((name for path, name in operations.items() if request.url.path.endswith(path)), None)
88+
if operation is None:
89+
return True, None
90+
try:
91+
submitted = json.loads(request.content)
92+
except ValueError:
93+
return True, None
94+
if not isinstance(submitted, dict):
95+
return True, None
96+
worker_id = submitted.get("worker_id")
97+
queue = submitted.get("task_queue")
98+
delay = body.get("retry_after_seconds")
99+
if (
100+
not isinstance(worker_id, str) or not worker_id
101+
or body.get("operation") != operation
102+
or body.get("outcome") != "unknown"
103+
or body.get("worker_id") != worker_id
104+
or body.get("task_queue") != queue
105+
or body.get("retryable") is not True
106+
or type(delay) is not int or delay <= 0
107+
or (operation != "heartbeat_worker" and (not isinstance(queue, str) or not queue))
108+
):
109+
return True, None
110+
if operation.startswith("poll_"):
111+
poll_id = submitted.get("poll_request_id")
112+
if (
113+
not isinstance(poll_id, str) or not poll_id
114+
or "task" not in body or body["task"] is not None
115+
or body.get("poll_status") != "backend_unavailable"
116+
or body.get("poll_request_id") != poll_id
117+
or body.get("retry_same_poll_request_id") is not True
118+
):
119+
return True, None
120+
return True, delay
121+
122+
68123
@dataclass
69124
class TransportRetryPolicy:
70125
"""
@@ -123,7 +178,7 @@ async def execute(self, fn: Callable[[], Awaitable[T]]) -> T:
123178
Raises the last exception if all retries are exhausted.
124179
"""
125180
attempt = 0
126-
storage_attempt = 0
181+
worker_pause_attempt = 0
127182
last_exc: Exception | None = None
128183

129184
while attempt < self.max_attempts:
@@ -134,20 +189,30 @@ async def execute(self, fn: Callable[[], Awaitable[T]]) -> T:
134189
last_exc = exc
135190
stop = _worker_storage_admission_stop.get()
136191
refusal = _storage_refusal(exc) if stop is not None else None
137-
if refusal is not None and stop is not None:
192+
backend_refusal, backend_delay = _backend_unavailable_refusal(exc)
193+
if backend_refusal and backend_delay is None:
194+
raise
195+
pause: tuple[str, int] | None = None
196+
if refusal is not None:
138197
error, poll_id = refusal
139-
if not error.is_storage_admission_failure(poll_id) or stop():
198+
if not error.is_storage_admission_failure(poll_id):
140199
raise
141-
storage_attempt += 1
142200
assert isinstance(error.body, dict)
201+
pause = ("storage admission paused", error.body["retry_after_seconds"])
202+
elif backend_delay is not None:
203+
pause = ("worker backend unavailable", backend_delay)
204+
if pause is not None and stop is not None:
205+
if stop():
206+
raise
207+
worker_pause_attempt += 1
143208
delay = min(
144209
5.0,
145210
max(
146-
self.backoff_seconds(min(storage_attempt - 1, 6)),
147-
error.body["retry_after_seconds"],
211+
self.backoff_seconds(min(worker_pause_attempt - 1, 6)),
212+
pause[1],
148213
),
149214
)
150-
log.warning("storage admission paused; retrying the same worker request in %.2fs", delay)
215+
log.warning("%s; retrying the same worker request in %.2fs", pause[0], delay)
151216
# Do not consume the finite transport budget or repeat serialization/uploads.
152217
while delay > 0:
153218
if stop():

‎tests/test_storage_admission.py‎

Lines changed: 165 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -39,6 +39,30 @@ def pressure(poll_id: str | None = None, *, reason: str = "storage_pressure", en
3939
return body
4040

4141

42+
def backend_unavailable(request: httpx.Request) -> dict[str, Any]:
43+
submitted = json.loads(request.content)
44+
path = request.url.path
45+
operation = {
46+
"/api/worker/workflow-tasks/poll": "poll_workflow_task",
47+
"/api/worker/activity-tasks/poll": "poll_activity_task",
48+
"/api/worker/query-tasks/poll": "poll_query_task",
49+
"/api/worker/register": "register_worker",
50+
"/api/worker/heartbeat": "heartbeat_worker",
51+
}[path]
52+
body: dict[str, Any] = {
53+
"reason": "backend_unavailable", "operation": operation, "outcome": "unknown",
54+
"worker_id": submitted["worker_id"], "task_queue": submitted.get("task_queue"),
55+
"retryable": True, "retry_after_seconds": 1,
56+
}
57+
if path.endswith("/poll"):
58+
body.update({
59+
"task": None, "poll_status": "backend_unavailable",
60+
"poll_request_id": submitted["poll_request_id"],
61+
"retry_same_poll_request_id": True,
62+
})
63+
return body
64+
65+
4266
@contextmanager
4367
def worker_scope(stop: Callable[[], bool] = lambda: False) -> Iterator[None]:
4468
token = _worker_storage_admission_stop.set(stop)
@@ -104,6 +128,147 @@ def handler(request: httpx.Request) -> httpx.Response:
104128
assert sum(retry_sleeps) == pytest.approx(5)
105129

106130

131+
@pytest.mark.parametrize("kind", ["workflow", "activity", "query", "multiplexed"])
132+
async def test_backend_outage_preserves_logical_poll_after_transport_budget(
133+
kind: str, retry_sleeps: list[float],
134+
) -> None:
135+
requests: list[bytes] = []
136+
137+
def handler(request: httpx.Request) -> httpx.Response:
138+
requests.append(request.content)
139+
if len(requests) == 1:
140+
raise httpx.ReadTimeout("Response lost after a possible claim.", request=request)
141+
if len(requests) <= 5:
142+
return httpx.Response(503, json=backend_unavailable(request))
143+
return httpx.Response(200, json={"task": {"task_id": "reconciled-claim"}})
144+
145+
async with client_for(handler) as client:
146+
with worker_scope():
147+
if kind == "multiplexed":
148+
task = await client.poll_workflow_task(
149+
worker_id="backend-worker", task_queue="orders",
150+
task_kinds=("workflow", "update_validation"),
151+
)
152+
else:
153+
task = await getattr(client, f"poll_{kind}_task")(
154+
worker_id="backend-worker", task_queue="orders",
155+
)
156+
assert task == {"task_id": "reconciled-claim"}
157+
assert len(requests) == 6
158+
assert len(set(requests)) == 1
159+
assert sum(retry_sleeps) == pytest.approx(4)
160+
161+
162+
@pytest.mark.parametrize("method,kwargs", [
163+
("register_worker", {
164+
"worker_id": "backend-worker", "task_queue": "orders",
165+
"capability_manifest": PORTABLE_WORKER_AFFINITY_CAPABILITY_MANIFEST,
166+
}),
167+
("heartbeat_worker", {"worker_id": "backend-worker"}),
168+
])
169+
async def test_backend_outage_retries_worker_registration_and_heartbeat(
170+
method: str, kwargs: dict[str, Any], retry_sleeps: list[float],
171+
) -> None:
172+
requests: list[bytes] = []
173+
174+
def handler(request: httpx.Request) -> httpx.Response:
175+
requests.append(request.content)
176+
if len(requests) <= 4:
177+
return httpx.Response(503, json=backend_unavailable(request))
178+
return httpx.Response(200, json={"recorded": True})
179+
180+
async with client_for(handler) as client:
181+
with worker_scope():
182+
result = await getattr(client, method)(**kwargs)
183+
assert result == {"recorded": True}
184+
assert len(requests) == 5
185+
assert len(set(requests)) == 1
186+
assert sum(retry_sleeps) == pytest.approx(4)
187+
188+
189+
@pytest.mark.parametrize("override", [
190+
{"operation": "poll_activity_task"},
191+
{"outcome": "failed"}, {"worker_id": "other-worker"},
192+
{"task_queue": "other-queue"}, {"retryable": False},
193+
{"retry_after_seconds": 0}, {"retry_after_seconds": True},
194+
{"retry_after_seconds": "1"}, {"task": {"task_id": "claimed"}},
195+
{"poll_status": "empty"}, {"poll_request_id": "another-poll"},
196+
{"retry_same_poll_request_id": False},
197+
])
198+
async def test_invalid_backend_outage_contract_remains_bounded(
199+
override: dict[str, Any], retry_sleeps: list[float],
200+
) -> None:
201+
calls = 0
202+
203+
def handler(request: httpx.Request) -> httpx.Response:
204+
nonlocal calls
205+
calls += 1
206+
return httpx.Response(503, json={**backend_unavailable(request), **override})
207+
208+
async with client_for(handler) as client:
209+
with worker_scope(), pytest.raises(ServerError):
210+
await client.poll_workflow_task(worker_id="backend-worker", task_queue="orders")
211+
assert calls == 1
212+
assert not retry_sleeps
213+
214+
215+
async def test_legacy_validation_backend_refusal_does_not_repeat_ambiguous_poll(
216+
retry_sleeps: list[float],
217+
) -> None:
218+
calls = 0
219+
220+
def handler(request: httpx.Request) -> httpx.Response:
221+
nonlocal calls
222+
calls += 1
223+
return httpx.Response(503, json={
224+
"reason": "backend_unavailable", "operation": "poll_update_validation_task",
225+
"outcome": "unknown", "worker_id": "backend-worker", "task_queue": "orders",
226+
"retryable": False, "retry_after_seconds": 1, "task": None,
227+
"poll_status": "backend_unavailable", "poll_request_id": None,
228+
"retry_same_poll_request_id": True,
229+
})
230+
231+
async with client_for(handler) as client:
232+
with worker_scope(), pytest.raises(ServerError):
233+
await client.poll_update_validation_task(worker_id="backend-worker", task_queue="orders")
234+
assert calls == 1
235+
assert not retry_sleeps
236+
237+
238+
async def test_direct_backend_outage_poll_remains_bounded(retry_sleeps: list[float]) -> None:
239+
calls = 0
240+
241+
def handler(request: httpx.Request) -> httpx.Response:
242+
nonlocal calls
243+
calls += 1
244+
return httpx.Response(503, json=backend_unavailable(request))
245+
246+
async with client_for(handler) as client:
247+
with pytest.raises(ServerError):
248+
await client.poll_workflow_task(worker_id="backend-worker", task_queue="orders")
249+
assert calls == 2
250+
251+
252+
async def test_shutdown_interrupts_backend_outage_without_new_poll(monkeypatch: pytest.MonkeyPatch) -> None:
253+
calls = 0
254+
stopped = False
255+
256+
def handler(request: httpx.Request) -> httpx.Response:
257+
nonlocal calls
258+
calls += 1
259+
return httpx.Response(503, json=backend_unavailable(request))
260+
261+
async def stop_on_sleep(delay: float) -> None:
262+
nonlocal stopped
263+
stopped = True
264+
265+
monkeypatch.setattr(retry_module, "asyncio", SimpleNamespace(sleep=stop_on_sleep))
266+
async with client_for(handler) as client:
267+
with worker_scope(lambda: stopped), pytest.raises(ServerError):
268+
await client.poll_workflow_task(worker_id="backend-worker", task_queue="orders")
269+
assert calls == 1
270+
271+
107272
@pytest.mark.parametrize("method,kwargs", [
108273
("register_worker", {
109274
"worker_id": "storage-worker", "task_queue": "orders",

0 commit comments

Comments
 (0)