From e82721f6e2664d48160c21d77c2ee6585401f371 Mon Sep 17 00:00:00 2001 From: Anto Subash Date: Thu, 21 May 2026 23:10:49 +0200 Subject: [PATCH] feat(background_tasks): propagate task_id + job_id into worker log context Adds `bind_task_context()` plus a `task_prerun` / `task_postrun` pair that binds Celery's `task_id` / `task_name` into a contextvars-backed log context. A `LogContextFilter` copies the bindings onto every LogRecord; hosting's `JsonFormatter` already emits the keys via `_EXTRA_KEYS`. App code wraps task bodies in `bind_task_context(job_id=...)` to attach domain identifiers. Closes #165. --- .../hosting/simple_module_hosting/logging.py | 3 + modules/background_tasks/README.md | 27 +++ .../background_tasks/__init__.py | 12 ++ .../background_tasks/celery_app.py | 3 + .../background_tasks/log_context.py | 123 ++++++++++++++ .../background_tasks/signals.py | 10 +- .../tests/test_log_context.py | 156 ++++++++++++++++++ 7 files changed, 331 insertions(+), 3 deletions(-) create mode 100644 modules/background_tasks/background_tasks/log_context.py create mode 100644 modules/background_tasks/tests/test_log_context.py diff --git a/framework/hosting/simple_module_hosting/logging.py b/framework/hosting/simple_module_hosting/logging.py index 8d116b32..8485687f 100644 --- a/framework/hosting/simple_module_hosting/logging.py +++ b/framework/hosting/simple_module_hosting/logging.py @@ -38,6 +38,9 @@ class JsonFormatter(logging.Formatter): "entity", "entity_id", "db_duration_ms", + # Bound by log filters on Celery / job-runner workers. + "task_id", + "task_name", ) def format(self, record: logging.LogRecord) -> str: diff --git a/modules/background_tasks/README.md b/modules/background_tasks/README.md index c8762877..31fd660c 100644 --- a/modules/background_tasks/README.md +++ b/modules/background_tasks/README.md @@ -52,6 +52,33 @@ Run a worker locally: uv run celery -A background_tasks.celery_app worker --loglevel=info ``` +## Worker log context + +Every worker log line automatically carries the Celery task identifiers +that fired it. A `LogContextFilter` is attached when the Celery app is +built (`build_celery`) and the `task_prerun` / `task_postrun` signals +bind `task_id` + `task_name` into `contextvars` for the task's duration: + +```jsonc +{"level": "INFO", "logger": "reports.tasks", "message": "ingest done", + "task_id": "9c2a…", "task_name": "reports.generate"} +``` + +Use `bind_task_context(...)` to attach app-level identifiers (the +domain `job_id` that named a Celery task is the canonical example): + +```python +from background_tasks import bind_task_context + +@celery_app.task +def process_dataset(job_id: int) -> None: + with bind_task_context(job_id=job_id): + logger.info("starting ingest") # now carries job_id too +``` + +Bindings nest cleanly and restore on exit. structlog users can mount the +same `contextvars` directly via `structlog.contextvars.merge_contextvars`. + ## Depends on - `simple_module_core`, `simple_module_db`, `simple_module_hosting` diff --git a/modules/background_tasks/background_tasks/__init__.py b/modules/background_tasks/background_tasks/__init__.py index d065990b..f1a81e72 100644 --- a/modules/background_tasks/background_tasks/__init__.py +++ b/modules/background_tasks/background_tasks/__init__.py @@ -1 +1,13 @@ """BackgroundTasks module — Celery + Redis task queue with admin UI.""" + +from background_tasks.log_context import ( + bind_task_context, + get_log_context, + install_log_filter, +) + +__all__ = [ + "bind_task_context", + "get_log_context", + "install_log_filter", +] diff --git a/modules/background_tasks/background_tasks/celery_app.py b/modules/background_tasks/background_tasks/celery_app.py index 1eaa5b40..bc15b1de 100644 --- a/modules/background_tasks/background_tasks/celery_app.py +++ b/modules/background_tasks/background_tasks/celery_app.py @@ -96,5 +96,8 @@ def build_celery(settings: BackgroundTasksSettings) -> Celery: # Side-effect import: connects signal handlers to this Celery instance's # ``celery.signals.*`` globals. Safe to import repeatedly. from background_tasks import signals # noqa: F401 + from background_tasks.log_context import install_log_filter + + install_log_filter() return celery diff --git a/modules/background_tasks/background_tasks/log_context.py b/modules/background_tasks/background_tasks/log_context.py new file mode 100644 index 00000000..a2c3941f --- /dev/null +++ b/modules/background_tasks/background_tasks/log_context.py @@ -0,0 +1,123 @@ +"""Contextvars-based log context for Celery tasks. + +Workers run outside the HTTP request lifecycle, so the hosting layer's +``correlation_id`` ContextVar doesn't reach them. This module gives task +code the equivalent affordance: ``task_id`` / ``task_name`` bound from +Celery signals plus arbitrary domain identifiers via +:func:`bind_task_context`. See ``modules/background_tasks/README.md``. +""" + +from __future__ import annotations + +import logging +from collections.abc import Iterator, Mapping +from contextlib import contextmanager +from contextvars import ContextVar, Token +from types import MappingProxyType +from typing import Any + +# Read-only so a future caller can't mutate the shared default in place. +_EMPTY: Mapping[str, Any] = MappingProxyType({}) + +current_log_context: ContextVar[Mapping[str, Any]] = ContextVar( + "current_log_context", default=_EMPTY +) + +# Stdlib LogRecord populates these attrs by default. Binding a key with +# the same name would be silently shadowed by the LogRecord attr (or +# rejected by stdlib's own makeRecord guard), so we reject at bind time +# instead of at log time. +_RESERVED_RECORD_KEYS: frozenset[str] = frozenset(logging.makeLogRecord({}).__dict__) | { + "message", + "asctime", +} + +# Pairs ``task_prerun`` → ``task_postrun`` across signal calls. Keyed by +# Celery task UUID so threaded / gevent pools — which can interleave +# prerun/postrun pairs within one process — stay isolated. +_signal_tokens: dict[str, Token[Mapping[str, Any]]] = {} + +_log = logging.getLogger(__name__) + + +def get_log_context() -> dict[str, Any]: + """Return a snapshot of every currently-bound key.""" + return dict(current_log_context.get()) + + +@contextmanager +def bind_task_context(**identifiers: Any) -> Iterator[None]: + """Layer ``identifiers`` onto the current task's log context. + + Nests cleanly. Raises ``ValueError`` if any key collides with a + stdlib :class:`LogRecord` attribute (``name``, ``module``, ...) — + those would be silently shadowed downstream. + """ + collisions = identifiers.keys() & _RESERVED_RECORD_KEYS + if collisions: + raise ValueError( + f"Cannot bind log-context keys that collide with LogRecord " + f"attributes: {sorted(collisions)}" + ) + merged = {**current_log_context.get(), **identifiers} + token = current_log_context.set(merged) + try: + yield + finally: + current_log_context.reset(token) + + +class LogContextFilter(logging.Filter): + """Copy bound log-context keys onto each :class:`LogRecord`. + + Attached via :func:`install_log_filter`; downstream formatters read + ``record.task_id`` etc. when emitting the line. + """ + + def filter(self, record: logging.LogRecord) -> bool: + ctx = current_log_context.get() + if not ctx: + return True + record_dict = record.__dict__ + for key, value in ctx.items(): + # An explicit ``extra={key: ...}`` on the logger call wins. + record_dict.setdefault(key, value) + return True + + +def install_log_filter(logger: logging.Logger | None = None) -> LogContextFilter: + """Attach a :class:`LogContextFilter` to ``logger`` (root if omitted). Idempotent.""" + target = logger if logger is not None else logging.getLogger() + for existing in target.filters: + if isinstance(existing, LogContextFilter): + return existing + log_filter = LogContextFilter() + target.addFilter(log_filter) + return log_filter + + +def signal_task_started(*, task_id: str | None, task_name: str | None) -> None: + """Bind ``task_id`` / ``task_name`` for a Celery task's duration. + + Paired with :func:`signal_task_finished` via the postrun signal. + """ + if not task_id: + return + merged = {**current_log_context.get(), "task_id": task_id, "task_name": task_name} + _signal_tokens[task_id] = current_log_context.set(merged) + + +def signal_task_finished(*, task_id: str | None) -> None: + """Reset the binding from :func:`signal_task_started`.""" + if not task_id: + return + token = _signal_tokens.pop(task_id, None) + if token is None: + return + try: + current_log_context.reset(token) + except ValueError: + # Token belongs to a different context — possible under exotic + # eventlet patching. The var falls back when the task's context + # exits, so the leak is bounded. + _log.debug("Log-context reset skipped for task_id=%s", task_id) diff --git a/modules/background_tasks/background_tasks/signals.py b/modules/background_tasks/background_tasks/signals.py index fecedfb5..db7e3970 100644 --- a/modules/background_tasks/background_tasks/signals.py +++ b/modules/background_tasks/background_tasks/signals.py @@ -38,6 +38,7 @@ ) from background_tasks.constants import DEFAULT_QUEUE, TaskStatus from background_tasks.contracts.events import TaskFailed +from background_tasks.log_context import signal_task_finished, signal_task_started from background_tasks.models import TaskExecution from background_tasks.sync_db import sync_session @@ -161,14 +162,15 @@ def on_task_prerun( kwargs: Any = None, **_k: Any, ) -> None: - """Flip the row to ``running`` and start the heartbeat.""" + """Flip the row to ``running``, start the heartbeat, bind log context.""" args_n, kwargs_n = coerce_args_kwargs(args, kwargs) now = now_utc() + name = task_name_of(sender, task) _apply( "on_task_prerun", celery_task_id=task_id, defaults={ - "task_name": task_name_of(sender, task), + "task_name": name, "status": TaskStatus.RUNNING, "args": args_n, "kwargs": kwargs_n, @@ -176,6 +178,7 @@ def on_task_prerun( "heartbeat_at": now, }, ) + signal_task_started(task_id=task_id, task_name=name) @signals.task_postrun.connect @@ -185,7 +188,7 @@ def on_task_postrun( task: Any = None, **_k: Any, ) -> None: - """Refresh the heartbeat on normal completion. + """Refresh the heartbeat, then unbind the log context. Terminal status is written by ``task_success`` / ``task_failure`` / ``task_retry`` which fire *before* postrun; postrun only refreshes the @@ -196,6 +199,7 @@ def on_task_postrun( celery_task_id=task_id, defaults={"task_name": task_name_of(sender, task), "heartbeat_at": now_utc()}, ) + signal_task_finished(task_id=task_id) @signals.task_success.connect diff --git a/modules/background_tasks/tests/test_log_context.py b/modules/background_tasks/tests/test_log_context.py new file mode 100644 index 00000000..c77ecdf4 --- /dev/null +++ b/modules/background_tasks/tests/test_log_context.py @@ -0,0 +1,156 @@ +"""Tests for the Celery log-context layer.""" + +from __future__ import annotations + +import logging +import uuid +from pathlib import Path +from types import SimpleNamespace + +import pytest +from background_tasks.log_context import ( + LogContextFilter, + _signal_tokens, + bind_task_context, + get_log_context, + install_log_filter, +) +from background_tasks.signals import on_task_postrun, on_task_prerun + + +def _make_record() -> logging.LogRecord: + return logging.LogRecord( + name="test", + level=logging.INFO, + pathname=__file__, + lineno=1, + msg="hello", + args=(), + exc_info=None, + ) + + +def _fake_task(name: str) -> SimpleNamespace: + return SimpleNamespace(name=name) + + +@pytest.fixture(autouse=True) +def _reset_signal_tokens() -> None: + """Clear the prerun/postrun bookkeeping in case a prior test panicked.""" + _signal_tokens.clear() + + +class TestBindTaskContext: + def test_layers_identifiers_and_restores_on_exit(self) -> None: + assert get_log_context() == {} + + with bind_task_context(job_id=42, tenant_id="acme"): + assert get_log_context() == {"job_id": 42, "tenant_id": "acme"} + + assert get_log_context() == {} + + def test_nests_cleanly(self) -> None: + with bind_task_context(job_id=1): + with bind_task_context(step="ingest"): + assert get_log_context() == {"job_id": 1, "step": "ingest"} + assert get_log_context() == {"job_id": 1} + + def test_inner_shadows_outer_then_restores(self) -> None: + with bind_task_context(job_id=1), bind_task_context(job_id=2): + assert get_log_context()["job_id"] == 2 + assert get_log_context() == {} + + def test_restores_even_when_block_raises(self) -> None: + with pytest.raises(RuntimeError, match="boom"), bind_task_context(job_id=99): + raise RuntimeError("boom") + assert get_log_context() == {} + + def test_rejects_keys_that_collide_with_logrecord_attrs(self) -> None: + """`record.name` is the logger name; binding `name=...` would be shadowed.""" + with pytest.raises(ValueError, match="name"), bind_task_context(name="oops"): + pass + + +class TestLogContextFilter: + def test_injects_bound_keys_onto_record(self) -> None: + log_filter = LogContextFilter() + record = _make_record() + + with bind_task_context(task_id="t-1", task_name="reports.ingest", job_id=99): + log_filter.filter(record) + + assert record.task_id == "t-1" # type: ignore[attr-defined] + assert record.task_name == "reports.ingest" # type: ignore[attr-defined] + assert record.job_id == 99 # type: ignore[attr-defined] + + def test_no_op_when_unbound(self) -> None: + log_filter = LogContextFilter() + record = _make_record() + log_filter.filter(record) + assert not hasattr(record, "task_id") + assert not hasattr(record, "job_id") + + def test_does_not_overwrite_explicit_extra(self) -> None: + log_filter = LogContextFilter() + record = _make_record() + record.job_id = "from-extra" # type: ignore[attr-defined] + + with bind_task_context(job_id="from-context"): + log_filter.filter(record) + + assert record.job_id == "from-extra" # type: ignore[attr-defined] + + +class TestInstallLogFilter: + def test_idempotent(self) -> None: + scratch = logging.getLogger("bg_tasks_test_install") + scratch.filters.clear() + try: + first = install_log_filter(scratch) + second = install_log_filter(scratch) + assert first is second + count = sum(1 for f in scratch.filters if isinstance(f, LogContextFilter)) + assert count == 1 + finally: + scratch.filters.clear() + + +class TestSignalBinding: + def test_prerun_binds_and_postrun_unbinds(self, sync_sqlite: Path) -> None: + task_id = str(uuid.uuid4()) + task = _fake_task("demo.ingest") + + assert get_log_context() == {} + + on_task_prerun(sender=task, task_id=task_id, task=task, args=[], kwargs={}) + + assert get_log_context() == {"task_id": task_id, "task_name": "demo.ingest"} + + on_task_postrun(sender=task, task_id=task_id, task=task) + + assert get_log_context() == {} + assert _signal_tokens == {} + + def test_postrun_without_prerun_is_noop(self, sync_sqlite: Path) -> None: + task = _fake_task("demo.detached") + on_task_postrun(sender=task, task_id="never-bound", task=task) + assert get_log_context() == {} + + def test_bind_task_context_layers_on_top_of_signal_binding(self, sync_sqlite: Path) -> None: + task_id = str(uuid.uuid4()) + task = _fake_task("demo.layered") + + on_task_prerun(sender=task, task_id=task_id, task=task, args=[], kwargs={}) + try: + with bind_task_context(job_id=7): + assert get_log_context() == { + "task_id": task_id, + "task_name": "demo.layered", + "job_id": 7, + } + assert get_log_context() == { + "task_id": task_id, + "task_name": "demo.layered", + } + finally: + on_task_postrun(sender=task, task_id=task_id, task=task)