Skip to content

Commit 4a4fc3f

Browse files
snopokeclaude
andcommitted
Record celery/procrastinate task ID as external_id
Co-Authored-By: Claude Opus 4.8 (1M context) <noreply@anthropic.com>
1 parent 73eef62 commit 4a4fc3f

5 files changed

Lines changed: 178 additions & 18 deletions

File tree

taskbadger/celery.py

Lines changed: 6 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -137,7 +137,9 @@ def task_publish_handler(sender=None, headers=None, body=None, **kwargs):
137137
ctask = celery.current_app.tasks.get(sender)
138138

139139
# get kwargs from the task class (set via decorator)
140-
kwargs = getattr(ctask, TB_KWARGS_ARG, {})
140+
# copy: the raw attr is shared across invocations, and we mutate per-publish
141+
# values (external_id, status, name) into it below
142+
kwargs = dict(getattr(ctask, TB_KWARGS_ARG, {}))
141143
for attr in dir(ctask):
142144
if attr.startswith(KWARG_PREFIX) and attr not in IGNORE_ARGS:
143145
kwargs[attr.removeprefix(KWARG_PREFIX)] = getattr(ctask, attr)
@@ -147,6 +149,7 @@ def task_publish_handler(sender=None, headers=None, body=None, **kwargs):
147149
kwargs["status"] = StatusEnum.PENDING
148150
if routing_key and "queue" not in kwargs:
149151
kwargs["queue"] = routing_key
152+
kwargs.setdefault("external_id", headers["id"])
150153
name = kwargs.pop("name", headers["task"])
151154

152155
global_record_task_args = celery_system and celery_system.record_task_args
@@ -247,7 +250,8 @@ def _maybe_create_task(signal_sender):
247250

248251
delivery_info = getattr(signal_sender.request, "delivery_info", None) or {}
249252
queue = delivery_info.get("routing_key")
250-
task = create_task_safe(task_name, status=StatusEnum.PENDING, data=data, queue=queue)
253+
external_id = signal_sender.request.id
254+
task = create_task_safe(task_name, status=StatusEnum.PENDING, data=data, queue=queue, external_id=external_id)
251255
if task:
252256
# Store the task ID in the request so _update_task can find it
253257
signal_sender.request.update({TB_TASK_ID: task.id})

taskbadger/procrastinate.py

Lines changed: 24 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -149,12 +149,16 @@ def _wrap_defer(task):
149149
@functools.wraps(original_defer)
150150
def defer(**kwargs):
151151
kwargs = _maybe_create_pending(task, kwargs)
152-
return original_defer(**kwargs)
152+
job_id = original_defer(**kwargs)
153+
_record_external_id(kwargs, job_id)
154+
return job_id
153155

154156
@functools.wraps(original_defer_async)
155157
async def defer_async(**kwargs):
156158
kwargs = _maybe_create_pending(task, kwargs)
157-
return await original_defer_async(**kwargs)
159+
job_id = await original_defer_async(**kwargs)
160+
_record_external_id(kwargs, job_id)
161+
return job_id
158162

159163
task.defer = defer
160164
task.defer_async = defer_async
@@ -214,6 +218,18 @@ def _maybe_create_pending(task, kwargs):
214218
return new_kwargs
215219

216220

221+
def _record_external_id(kwargs, job_id):
222+
"""Record the Procrastinate job id as the TaskBadger task's ``external_id``.
223+
224+
``defer`` only returns the DB-assigned job id once the job is enqueued, so this
225+
runs as a follow-up update after both ids are known. No-op if the defer wasn't
226+
tracked (no injected id) or no job id came back."""
227+
tb_id = kwargs.get(TB_TASK_ID_KWARG)
228+
if tb_id is None or job_id is None:
229+
return
230+
update_task_safe(tb_id, external_id=str(job_id))
231+
232+
217233
def _serialize_kwargs(kwargs):
218234
"""Return a JSON-roundtrippable copy of the defer kwargs.
219235
@@ -296,14 +312,19 @@ def _patch_job_manager(app, system):
296312
@functools.wraps(original)
297313
async def patched(*, job, periodic_id, defer_timestamp):
298314
task = app.tasks.get(job.task_name)
315+
tb_id = None
299316
if task is not None:
300317
tb_task = _create_pending_task(task, job.task_kwargs, queue=job.queue)
301318
if tb_task is not None:
302319
new_kwargs = {**job.task_kwargs, TB_TASK_ID_KWARG: tb_task.id}
303320
job = job.evolve(task_kwargs=new_kwargs)
304-
return await jm._taskbadger_original_defer_periodic_job(
321+
tb_id = tb_task.id
322+
job_id = await jm._taskbadger_original_defer_periodic_job(
305323
job=job, periodic_id=periodic_id, defer_timestamp=defer_timestamp
306324
)
325+
if tb_id is not None and job_id is not None:
326+
update_task_safe(tb_id, external_id=str(job_id))
327+
return job_id
307328

308329
jm.defer_periodic_job = patched
309330

tests/test_celery.py

Lines changed: 59 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -98,7 +98,13 @@ def add_with_task_args(self, a, b):
9898
assert result.get(timeout=10, propagate=True) == 4
9999

100100
create.assert_called_once_with(
101-
"new_name", value_max=10, data={"foo": "bar"}, tags={"bar": "baz"}, status=StatusEnum.PENDING, queue="celery"
101+
"new_name",
102+
value_max=10,
103+
data={"foo": "bar"},
104+
tags={"bar": "baz"},
105+
status=StatusEnum.PENDING,
106+
queue="celery",
107+
external_id=mock.ANY,
102108
)
103109

104110

@@ -128,7 +134,9 @@ def add_with_task_args(self, a, b):
128134
)
129135
assert result.get(timeout=10, propagate=True) == 4
130136

131-
create.assert_called_once_with("new_name", value_max=10, actions=actions, status=StatusEnum.PENDING, queue="celery")
137+
create.assert_called_once_with(
138+
"new_name", value_max=10, actions=actions, status=StatusEnum.PENDING, queue="celery", external_id=mock.ANY
139+
)
132140

133141

134142
@pytest.mark.usefixtures("_bind_settings")
@@ -162,6 +170,7 @@ def add_with_task_args(self, a, b):
162170
data={"foo": "bar", "celery_task_args": [2, 2], "celery_task_kwargs": {}},
163171
status=StatusEnum.PENDING,
164172
queue="celery",
173+
external_id=mock.ANY,
165174
)
166175

167176

@@ -200,6 +209,7 @@ def add_with_task_kwargs(self, a, b, c=0):
200209
actions=actions,
201210
status=StatusEnum.PENDING,
202211
queue="celery",
212+
external_id=mock.ANY,
203213
)
204214

205215

@@ -237,6 +247,7 @@ def add_task_custom_serialization(self, a):
237247
data={"celery_task_args": [{"__type__": "A", "__value__": [2, 2]}], "celery_task_kwargs": {}},
238248
status=StatusEnum.PENDING,
239249
queue="celery",
250+
external_id=mock.ANY,
240251
)
241252

242253

@@ -264,7 +275,9 @@ def add_with_task_args_in_decorator(self, a, b):
264275
result = add_with_task_args_in_decorator.delay(2, 2)
265276
assert result.get(timeout=10, propagate=True) == 4
266277

267-
create.assert_called_once_with(mock.ANY, status=StatusEnum.PENDING, monitor_id="123", value_max=10, queue="celery")
278+
create.assert_called_once_with(
279+
mock.ANY, status=StatusEnum.PENDING, monitor_id="123", value_max=10, queue="celery", external_id=mock.ANY
280+
)
268281

269282

270283
@pytest.mark.usefixtures("_bind_settings")
@@ -288,6 +301,49 @@ def add_custom_queue(self, a, b):
288301
assert create.call_args.kwargs["queue"] == "high_priority"
289302

290303

304+
@pytest.mark.usefixtures("_bind_settings")
305+
def test_celery_records_task_id_as_external_id(celery_session_app, celery_session_worker):
306+
@celery_session_app.task(bind=True, base=Task)
307+
def add_external_id(self, a, b):
308+
return a + b
309+
310+
celery_session_worker.reload()
311+
312+
with (
313+
mock.patch("taskbadger.celery.create_task_safe") as create,
314+
mock.patch("taskbadger.celery.update_task_safe"),
315+
mock.patch("taskbadger.sdk.get_task"),
316+
):
317+
create.return_value = task_for_test()
318+
result = add_external_id.apply_async((2, 2), queue="high_priority")
319+
320+
assert create.call_args.kwargs["external_id"] == result.id
321+
322+
323+
@pytest.mark.usefixtures("_bind_settings")
324+
def test_celery_external_id_not_cached_across_invocations(celery_session_app, celery_session_worker):
325+
# Tasks defined with taskbadger_kwargs share a class-level dict; external_id
326+
# must not leak from one publish to the next.
327+
@celery_session_app.task(bind=True, base=Task, taskbadger_kwargs={"value_max": 10})
328+
def add_repeat(self, a, b):
329+
return a + b
330+
331+
celery_session_worker.reload()
332+
333+
with (
334+
mock.patch("taskbadger.celery.create_task_safe") as create,
335+
mock.patch("taskbadger.celery.update_task_safe"),
336+
mock.patch("taskbadger.sdk.get_task"),
337+
):
338+
create.return_value = task_for_test()
339+
result1 = add_repeat.apply_async((2, 2), queue="high_priority")
340+
result2 = add_repeat.apply_async((3, 3), queue="high_priority")
341+
342+
external_ids = [c.kwargs["external_id"] for c in create.call_args_list]
343+
assert external_ids == [result1.id, result2.id]
344+
assert result1.id != result2.id
345+
346+
291347
@pytest.mark.usefixtures("_bind_settings")
292348
def test_celery_task_retry(celery_session_app, celery_session_worker):
293349
"""Note: When a task is retried, the celery task ID remains the same but a new TB task

tests/test_celery_system_integration.py

Lines changed: 6 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -123,6 +123,7 @@ def add_normal(self, a, b):
123123
status=StatusEnum.PENDING,
124124
data={"celery_task_args": [2, 2], "celery_task_kwargs": {}},
125125
queue="celery",
126+
external_id=mock.ANY,
126127
)
127128
assert get_task.call_count == 1
128129
assert update.call_count == 2
@@ -154,7 +155,10 @@ def add_normal_with_override(a, b):
154155
assert result.get(timeout=10, propagate=True) == 4
155156

156157
create.assert_called_once_with(
157-
"tests.test_celery_system_integration.add_normal_with_override", status=StatusEnum.PENDING, queue="celery"
158+
"tests.test_celery_system_integration.add_normal_with_override",
159+
status=StatusEnum.PENDING,
160+
queue="celery",
161+
external_id=mock.ANY,
158162
)
159163

160164

@@ -189,6 +193,7 @@ def add_with_tags(a, b):
189193
status=StatusEnum.PENDING,
190194
tags=TaskRequestTags.from_dict({"tag1": "value1", "tag2": "override"}),
191195
queue="celery",
196+
external_id=result.id,
192197
)
193198
create.assert_called_with(
194199
client=mock.ANY,

tests/test_procrastinate.py

Lines changed: 83 additions & 9 deletions
Original file line numberDiff line numberDiff line change
@@ -115,7 +115,10 @@ def add3(a, b):
115115
_instrument_task(add3, system=None, manual=True)
116116

117117
tb = task_for_test()
118-
with mock.patch("taskbadger.procrastinate.create_task_safe", return_value=tb) as create:
118+
with (
119+
mock.patch("taskbadger.procrastinate.create_task_safe", return_value=tb) as create,
120+
mock.patch("taskbadger.procrastinate.update_task_safe"),
121+
):
119122
add3.defer(a=1, b=2)
120123

121124
create.assert_called_once()
@@ -137,7 +140,10 @@ def add_queued(a, b):
137140
_instrument_task(add_queued, system=None, manual=True)
138141

139142
tb = task_for_test()
140-
with mock.patch("taskbadger.procrastinate.create_task_safe", return_value=tb) as create:
143+
with (
144+
mock.patch("taskbadger.procrastinate.create_task_safe", return_value=tb) as create,
145+
mock.patch("taskbadger.procrastinate.update_task_safe"),
146+
):
141147
add_queued.defer(a=1, b=2)
142148

143149
assert create.call_args.kwargs["queue"] == "high_priority"
@@ -168,13 +174,69 @@ async def add5(a, b):
168174
_instrument_task(add5, system=None, manual=True)
169175

170176
tb = task_for_test()
171-
with mock.patch("taskbadger.procrastinate.create_task_safe", return_value=tb):
177+
with (
178+
mock.patch("taskbadger.procrastinate.create_task_safe", return_value=tb),
179+
mock.patch("taskbadger.procrastinate.update_task_safe"),
180+
):
172181
asyncio.run(add5.defer_async(a=1, b=2))
173182

174183
jobs = list(app.connector.jobs.values())
175184
assert jobs[0]["args"][TB_TASK_ID_KWARG] == tb.id
176185

177186

187+
@pytest.mark.usefixtures("_bind_settings")
188+
def test_defer_records_job_id_as_external_id(app):
189+
@app.task(name="add_ext")
190+
def add_ext(a, b):
191+
return a + b
192+
193+
_instrument_task(add_ext, system=None, manual=True)
194+
195+
tb = task_for_test()
196+
with (
197+
mock.patch("taskbadger.procrastinate.create_task_safe", return_value=tb),
198+
mock.patch("taskbadger.procrastinate.update_task_safe") as update,
199+
):
200+
job_id = add_ext.defer(a=1, b=2)
201+
202+
update.assert_called_once_with(tb.id, external_id=str(job_id))
203+
204+
205+
@pytest.mark.usefixtures("_bind_settings")
206+
def test_defer_async_records_job_id_as_external_id(app):
207+
@app.task(name="add_ext_async")
208+
async def add_ext_async(a, b):
209+
return a + b
210+
211+
_instrument_task(add_ext_async, system=None, manual=True)
212+
213+
tb = task_for_test()
214+
with (
215+
mock.patch("taskbadger.procrastinate.create_task_safe", return_value=tb),
216+
mock.patch("taskbadger.procrastinate.update_task_safe") as update,
217+
):
218+
job_id = asyncio.run(add_ext_async.defer_async(a=1, b=2))
219+
220+
update.assert_called_once_with(tb.id, external_id=str(job_id))
221+
222+
223+
def test_defer_no_external_id_when_untracked(app):
224+
@app.task(name="add_untracked")
225+
def add_untracked(a, b):
226+
return a + b
227+
228+
_instrument_task(add_untracked, system=None, manual=True)
229+
230+
# Badger is not configured, so no pending task is created and nothing to update.
231+
with (
232+
mock.patch("taskbadger.procrastinate.create_task_safe", return_value=None),
233+
mock.patch("taskbadger.procrastinate.update_task_safe") as update,
234+
):
235+
add_untracked.defer(a=1, b=2)
236+
237+
update.assert_not_called()
238+
239+
178240
@pytest.mark.usefixtures("_bind_settings")
179241
def test_end_to_end_via_worker(app):
180242
@app.task(name="add6")
@@ -194,7 +256,7 @@ def add6(a, b):
194256
app.run_worker(wait=False, install_signal_handlers=False, listen_notify=False)
195257

196258
create.assert_called_once()
197-
statuses = [c.kwargs["status"] for c in update.call_args_list]
259+
statuses = [c.kwargs["status"] for c in update.call_args_list if "status" in c.kwargs]
198260
assert statuses == [StatusEnum.PROCESSING, StatusEnum.SUCCESS]
199261

200262

@@ -206,7 +268,10 @@ def bare(a):
206268
return a
207269

208270
tb = task_for_test()
209-
with mock.patch("taskbadger.procrastinate.create_task_safe", return_value=tb):
271+
with (
272+
mock.patch("taskbadger.procrastinate.create_task_safe", return_value=tb),
273+
mock.patch("taskbadger.procrastinate.update_task_safe"),
274+
):
210275
bare.defer(a=1)
211276

212277
assert getattr(bare, "_taskbadger_manual") is True
@@ -223,7 +288,10 @@ def raw(a):
223288
return a
224289

225290
tb = task_for_test()
226-
with mock.patch("taskbadger.procrastinate.create_task_safe", return_value=tb) as create:
291+
with (
292+
mock.patch("taskbadger.procrastinate.create_task_safe", return_value=tb) as create,
293+
mock.patch("taskbadger.procrastinate.update_task_safe"),
294+
):
227295
raw.defer(a=1)
228296

229297
create.assert_called_once()
@@ -248,7 +316,10 @@ def dup(a):
248316
# Two @track applications must not double-wrap; defer once still creates one
249317
# PENDING task and injects one id.
250318
tb = task_for_test()
251-
with mock.patch("taskbadger.procrastinate.create_task_safe", return_value=tb) as create:
319+
with (
320+
mock.patch("taskbadger.procrastinate.create_task_safe", return_value=tb) as create,
321+
mock.patch("taskbadger.procrastinate.update_task_safe"),
322+
):
252323
dup.defer(a=1)
253324
assert create.call_count == 1
254325
jobs = list(app.connector.jobs.values())
@@ -312,7 +383,7 @@ def self_complete():
312383

313384
# The wrapper's post-call SUCCESS update is skipped because the cached
314385
# task is already SUCCESS. PROCESSING update is still allowed (early path).
315-
statuses = [c.kwargs["status"] for c in update.call_args_list]
386+
statuses = [c.kwargs["status"] for c in update.call_args_list if "status" in c.kwargs]
316387
assert StatusEnum.PROCESSING in statuses
317388
# Last attempted SUCCESS call should be suppressed
318389
assert statuses.count(StatusEnum.SUCCESS) == 0
@@ -326,7 +397,10 @@ def recorder(a, b):
326397
return a + b
327398

328399
tb = task_for_test()
329-
with mock.patch("taskbadger.procrastinate.create_task_safe", return_value=tb) as create:
400+
with (
401+
mock.patch("taskbadger.procrastinate.create_task_safe", return_value=tb) as create,
402+
mock.patch("taskbadger.procrastinate.update_task_safe"),
403+
):
330404
recorder.defer(a=5, b=6)
331405

332406
assert create.call_args.kwargs["data"] == {

0 commit comments

Comments
 (0)