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
3 changes: 3 additions & 0 deletions framework/hosting/simple_module_hosting/logging.py
Original file line number Diff line number Diff line change
Expand Up @@ -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:
Expand Down
27 changes: 27 additions & 0 deletions modules/background_tasks/README.md
Original file line number Diff line number Diff line change
Expand Up @@ -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`
Expand Down
12 changes: 12 additions & 0 deletions modules/background_tasks/background_tasks/__init__.py
Original file line number Diff line number Diff line change
@@ -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",
]
3 changes: 3 additions & 0 deletions modules/background_tasks/background_tasks/celery_app.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
123 changes: 123 additions & 0 deletions modules/background_tasks/background_tasks/log_context.py
Original file line number Diff line number Diff line change
@@ -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)
10 changes: 7 additions & 3 deletions modules/background_tasks/background_tasks/signals.py
Original file line number Diff line number Diff line change
Expand Up @@ -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

Expand Down Expand Up @@ -161,21 +162,23 @@ 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,
"started_at": now,
"heartbeat_at": now,
},
)
signal_task_started(task_id=task_id, task_name=name)


@signals.task_postrun.connect
Expand All @@ -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
Expand All @@ -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
Expand Down
156 changes: 156 additions & 0 deletions modules/background_tasks/tests/test_log_context.py
Original file line number Diff line number Diff line change
@@ -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)
Loading