@@ -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
4367def 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