From 42b0294ede27e9496be5e498f60b51d1ca820529 Mon Sep 17 00:00:00 2001 From: Simon Kelly Date: Wed, 5 Aug 2026 12:07:38 +0200 Subject: [PATCH 1/3] Add pluggable context provider hook for task error data Lets a provider attach extra context (e.g. a link to the originating issue) to a task's data when it errors. Wired into @track, the Celery signal handlers, and the Procrastinate integration, plus Task.error() for manual reporting. A start/reset snapshot pair lets providers that read back state from elsewhere tell a fresh capture from a stale one. --- taskbadger/_error_context.py | 59 ++++++++++++ taskbadger/celery.py | 22 ++++- taskbadger/context_providers/__init__.py | 32 +++++++ taskbadger/decorators.py | 6 +- taskbadger/mug.py | 2 + taskbadger/procrastinate.py | 9 +- taskbadger/sdk.py | 26 +++++- tests/test_error_context.py | 110 +++++++++++++++++++++++ 8 files changed, 260 insertions(+), 6 deletions(-) create mode 100644 taskbadger/_error_context.py create mode 100644 taskbadger/context_providers/__init__.py create mode 100644 tests/test_error_context.py diff --git a/taskbadger/_error_context.py b/taskbadger/_error_context.py new file mode 100644 index 0000000..221abb1 --- /dev/null +++ b/taskbadger/_error_context.py @@ -0,0 +1,59 @@ +"""Builds the `data` payload attached to a task when it errors, combining the +exception message with any registered `context_providers` (e.g. Sentry). Not +part of the public API. +""" + +import logging +from contextvars import ContextVar + +from taskbadger.mug import Badger + +log = logging.getLogger("taskbadger") + +_snapshots: ContextVar[dict] = ContextVar("taskbadger_context_snapshots", default=None) + + +def _providers(): + settings = Badger.current.settings + return settings.context_providers if settings else [] + + +def start_error_context(): + """Snapshot every configured context provider. Call this when a tracked task + starts, before user code runs, so `capture_error_data` can later tell a fresh + capture from a stale one left over from something unrelated. + + Returns a token that can be passed to `reset_error_context` to restore the + previous snapshot (for nested tracking within the same thread). + """ + snapshot = {} + for provider in _providers(): + try: + snapshot[provider.identifier] = provider.snapshot() + except Exception: + log.warning("Error snapshotting context provider '%s'", provider.identifier, exc_info=True) + return _snapshots.set(snapshot) + + +def reset_error_context(token) -> None: + _snapshots.reset(token) + + +def capture_error_data(exception: BaseException, message: str = None) -> dict: + """Arguments: + exception: The exception to report to context providers. + message: Text to store as `data["exception"]`. Defaults to `str(exception)`; + override when the caller has a more descriptive representation (e.g. Celery's + `ExceptionInfo`, which wraps the original exception). + """ + data = {"exception": message if message is not None else str(exception)} + snapshot = _snapshots.get() or {} + for provider in _providers(): + try: + extra = provider.capture_error_context(exception, snapshot.get(provider.identifier)) + except Exception: + log.warning("Error capturing context from provider '%s'", provider.identifier, exc_info=True) + extra = None + if extra: + data[provider.identifier] = extra + return data diff --git a/taskbadger/celery.py b/taskbadger/celery.py index 1804fd9..a0079bd 100644 --- a/taskbadger/celery.py +++ b/taskbadger/celery.py @@ -14,6 +14,7 @@ from kombu import serialization from . import sdk +from ._error_context import capture_error_data, reset_error_context, start_error_context from ._heartbeat import heartbeat from ._integrations import TERMINAL_STATES, resolve_heartbeat_options, safe_get_task, task_cache from .internal.models import StatusEnum @@ -30,6 +31,9 @@ # Marks a request whose signal handlers opened the Task Badger session, so that # only they close it again. TB_OWNS_SESSION = f"{KWARG_PREFIX}owns_session" +# Token returned by start_error_context(), stashed on the request so +# task_postrun_handler can restore the previous context (see comment there). +TB_ERROR_CTX_TOKEN = f"{KWARG_PREFIX}error_ctx_token" log = logging.getLogger("taskbadger") @@ -301,6 +305,16 @@ def task_prerun_handler(sender=None, **kwargs): _maybe_create_task(sender) _update_task(sender, StatusEnum.PROCESSING) _start_heartbeat(sender) + if _get_taskbadger_task_id(sender.request): + # Snapshotted here (same thread as the task body and the failure/retry + # signals below) so context providers can tell a fresh capture from a + # stale one if the task errors. The token is restored in + # task_postrun_handler rather than discarded: a task synchronously + # invoking another tracked task in its body (eager mode, `.apply()`, + # canvas primitives) would otherwise leave this task's context + # clobbered by the inner task's snapshot for the rest of its run. + token = start_error_context() + sender.request.update({TB_ERROR_CTX_TOKEN: token}) @task_postrun.connect @@ -309,6 +323,9 @@ def task_postrun_handler(sender=None, **kwargs): task_id = _get_taskbadger_task_id(sender.request) if task_id: heartbeat.stop(task_id) + token = sender.request.get(TB_ERROR_CTX_TOKEN) + if token is not None: + reset_error_context(token) @task_success.connect @@ -351,7 +368,10 @@ def _update_task(signal_sender, status, einfo=None): data = None if einfo: - data = DefaultMergeStrategy().merge(task.data, {"exception": str(einfo)}) + # `einfo.exception` wraps the real exception (see billiard.einfo.ExceptionWithTraceback); + # unwrap it so context providers (e.g. Sentry) see the original exception. + exc = getattr(einfo.exception, "exc", einfo.exception) + data = DefaultMergeStrategy().merge(task.data, capture_error_data(exc, message=str(einfo))) task = update_task_safe(task.id, status=status, data=data) if task: task_cache.set(task_id, task) diff --git a/taskbadger/context_providers/__init__.py b/taskbadger/context_providers/__init__.py new file mode 100644 index 0000000..c3ef166 --- /dev/null +++ b/taskbadger/context_providers/__init__.py @@ -0,0 +1,32 @@ +class ContextProvider: + """Base class for pluggable providers that attach extra context to a task's + ``data`` when it errors, e.g. so the TaskBadger UI can link out to an + external system (Sentry, Rollbar, etc.). + + Registered via `init(context_providers=[...])` and consulted whenever a + tracked task (via `@track`, the Celery/Procrastinate integrations) errors. + + Implementations that read back state some other system captured on its own + (rather than capturing it themselves, which risks duplicate reporting) + should override `snapshot` to record a baseline when the task starts, so + `capture_error_context` can tell a fresh capture from a stale one left + over from something unrelated. + """ + + identifier: str = None + + def snapshot(self): + """Called when a tracked task starts, before user code runs. Return an + opaque value to be passed back as *snapshot* to `capture_error_context`. + Default: `None` (no baseline tracking). + """ + return None + + def capture_error_context(self, exception: BaseException, snapshot=None) -> dict | None: + """Return extra context for *exception*, or `None` if there is nothing + to add. The result is stored under `data[self.identifier]`. + + *snapshot* is whatever this provider's `snapshot()` returned when the + task started. + """ + raise NotImplementedError diff --git a/taskbadger/decorators.py b/taskbadger/decorators.py index 1f722df..49174c4 100644 --- a/taskbadger/decorators.py +++ b/taskbadger/decorators.py @@ -1,6 +1,7 @@ import logging from functools import wraps +from ._error_context import capture_error_data, reset_error_context, start_error_context from .mug import Session from .safe_sdk import create_task_safe from .sdk import StatusEnum @@ -52,16 +53,19 @@ def _inner(*args, **kwargs): monitor_id=monitor_id, **task_kwargs, ) + token = start_error_context() try: result = func(*args, **kwargs) except Exception as e: _update_task( task, status=StatusEnum.ERROR, - data={"exception": str(e)}, + data=capture_error_data(e), data_merge_strategy="default", ) raise + finally: + reset_error_context(token) _update_task(task, status=StatusEnum.SUCCESS) return result diff --git a/taskbadger/mug.py b/taskbadger/mug.py index c027c61..6bf8b90 100644 --- a/taskbadger/mug.py +++ b/taskbadger/mug.py @@ -4,6 +4,7 @@ from contextvars import ContextVar from copy import deepcopy +from taskbadger.context_providers import ContextProvider from taskbadger.internal import AuthenticatedClient from taskbadger.systems import System @@ -21,6 +22,7 @@ class Settings: project_slug: str systems: dict[str, System] = dataclasses.field(default_factory=dict) before_create: Callback = None + context_providers: list[ContextProvider] = dataclasses.field(default_factory=list) def get_client(self): return AuthenticatedClient(self.base_url, self.token) diff --git a/taskbadger/procrastinate.py b/taskbadger/procrastinate.py index 3db2f43..3ce993b 100644 --- a/taskbadger/procrastinate.py +++ b/taskbadger/procrastinate.py @@ -17,6 +17,7 @@ import logging from contextvars import ContextVar +from ._error_context import capture_error_data, reset_error_context, start_error_context from ._heartbeat import heartbeat from ._integrations import ( TERMINAL_STATES, @@ -88,11 +89,14 @@ async def wrapped(*args, **kwargs): try: _update_status(tb_id, StatusEnum.PROCESSING) heartbeat.start(tb_id, _heartbeat_interval(task)) + ctx_token = start_error_context() try: result = await original_func(*args, **kwargs) except Exception as exc: _update_status(tb_id, StatusEnum.ERROR, exception=exc) raise + finally: + reset_error_context(ctx_token) _update_status(tb_id, StatusEnum.SUCCESS) return result finally: @@ -109,11 +113,14 @@ def wrapped(*args, **kwargs): try: _update_status(tb_id, StatusEnum.PROCESSING) heartbeat.start(tb_id, _heartbeat_interval(task)) + ctx_token = start_error_context() try: result = original_func(*args, **kwargs) except Exception as exc: _update_status(tb_id, StatusEnum.ERROR, exception=exc) raise + finally: + reset_error_context(ctx_token) _update_status(tb_id, StatusEnum.SUCCESS) return result finally: @@ -141,7 +148,7 @@ def _update_status(tb_id, status, exception=None): data = None if exception is not None and current is not None: base = dict(current.data) if current.data else None - data = DefaultMergeStrategy().merge(base, {"exception": str(exception)}) + data = DefaultMergeStrategy().merge(base, capture_error_data(exception)) if data is not None: updated = update_task_safe(tb_id, status=status, data=data) else: diff --git a/taskbadger/sdk.py b/taskbadger/sdk.py index 4389cf2..12c8fe6 100644 --- a/taskbadger/sdk.py +++ b/taskbadger/sdk.py @@ -5,6 +5,8 @@ import warnings from typing import Any +from taskbadger._error_context import capture_error_data +from taskbadger.context_providers import ContextProvider from taskbadger.exceptions import ( ConfigurationError, MissingConfiguration, @@ -63,6 +65,7 @@ def init( systems: list[System] = None, tags: dict[str, str] = None, before_create: Callback = None, + context_providers: list[ContextProvider] = None, ): """Initialize Task Badger client. @@ -73,9 +76,13 @@ def init( For legacy API keys, *organization_slug* and *project_slug* are required and a deprecation warning is emitted. + Arguments: + context_providers: Providers consulted when a tracked task errors, to attach extra + context (e.g. a Sentry issue link) to the task's `data`. See `taskbadger.context_providers`. + Call this function once per thread. """ - _init(_TB_HOST, organization_slug, project_slug, token, systems, tags, before_create) + _init(_TB_HOST, organization_slug, project_slug, token, systems, tags, before_create, context_providers) def _init( @@ -86,6 +93,7 @@ def _init( systems: list[System] = None, tags: dict[str, str] = None, before_create: Callback = None, + context_providers: list[ContextProvider] = None, ): host = host or os.environ.get("TASKBADGER_HOST", "https://taskbadger.net") organization_slug = organization_slug or os.environ.get("TASKBADGER_ORG") @@ -118,6 +126,7 @@ def _init( project_slug, systems={system.identifier: system for system in systems}, before_create=before_create, + context_providers=context_providers or [], ) Badger.current.bind(settings, tags) else: @@ -387,8 +396,19 @@ def success(self, value: int = None): """Update the task status to `success` and set the value.""" self.update(status=StatusEnum.SUCCESS, value=value) - def error(self, value: int = None, data: dict = None): - """Update the task status to `error` and set the value and data.""" + def error(self, value: int = None, data: dict = None, exception: BaseException = None): + """Update the task status to `error` and set the value and data. + + If `exception` is given, it's passed to any configured context providers + (e.g. Sentry, see [taskbadger.context_providers][]) and the result merged into `data`. + Called on its own (outside `@track` or the Celery/Procrastinate integrations), providers + have no baseline to compare against, so e.g. `SentryContextProvider` will report whatever + `sentry_sdk.last_event_id()` currently is. + """ + if exception is not None: + error_data = capture_error_data(exception) + error_data.update(data or {}) + data = error_data self.update(status=StatusEnum.ERROR, value=value, data=data) def canceled(self): diff --git a/tests/test_error_context.py b/tests/test_error_context.py new file mode 100644 index 0000000..2b4892a --- /dev/null +++ b/tests/test_error_context.py @@ -0,0 +1,110 @@ +from unittest import mock + +import pytest + +from taskbadger._error_context import capture_error_data, reset_error_context, start_error_context +from taskbadger.context_providers import ContextProvider +from taskbadger.mug import Badger, Settings + + +@pytest.fixture +def bind_settings_with_providers(): + def _bind(providers): + Badger.current.bind(Settings("https://taskbadger.net", "token", "org", "proj", context_providers=providers)) + + yield _bind + Badger.current.bind(None) + + +def test_capture_error_data_no_providers(): + assert capture_error_data(ValueError("boom")) == {"exception": "boom"} + + +def test_capture_error_data_message_override(): + data = capture_error_data(ValueError("boom"), message="custom message") + assert data == {"exception": "custom message"} + + +def test_capture_error_data_not_configured(): + Badger.current.bind(None) + assert capture_error_data(ValueError("boom")) == {"exception": "boom"} + + +def test_capture_error_data_with_provider(bind_settings_with_providers): + class FakeProvider(ContextProvider): + identifier = "fake" + + def capture_error_context(self, exception, snapshot=None): + return {"detail": str(exception)} + + bind_settings_with_providers([FakeProvider()]) + + data = capture_error_data(ValueError("boom")) + assert data == {"exception": "boom", "fake": {"detail": "boom"}} + + +def test_capture_error_data_provider_returning_none(bind_settings_with_providers): + class FakeProvider(ContextProvider): + identifier = "fake" + + def capture_error_context(self, exception, snapshot=None): + return None + + bind_settings_with_providers([FakeProvider()]) + + assert capture_error_data(ValueError("boom")) == {"exception": "boom"} + + +def test_capture_error_data_provider_raises(bind_settings_with_providers): + class BrokenProvider(ContextProvider): + identifier = "broken" + + def capture_error_context(self, exception, snapshot=None): + raise RuntimeError("provider failed") + + bind_settings_with_providers([BrokenProvider()]) + + with mock.patch("taskbadger._error_context.log") as log: + data = capture_error_data(ValueError("boom")) + assert data == {"exception": "boom"} + log.warning.assert_called_once() + + +def test_capture_error_data_uses_snapshot_from_start_error_context(bind_settings_with_providers): + """A provider that only reports a fresh capture (comparing against its + snapshot) should see the snapshot taken by start_error_context.""" + + calls = [] + + class TrackingProvider(ContextProvider): + identifier = "tracking" + + def snapshot(self): + return "baseline" + + def capture_error_context(self, exception, snapshot=None): + calls.append(snapshot) + return None if snapshot == "baseline" else {"unexpected": True} + + bind_settings_with_providers([TrackingProvider()]) + + token = start_error_context() + try: + capture_error_data(ValueError("boom")) + finally: + reset_error_context(token) + + assert calls == ["baseline"] + + +def test_capture_error_data_without_snapshot_passes_none(bind_settings_with_providers): + class FakeProvider(ContextProvider): + identifier = "fake" + + def capture_error_context(self, exception, snapshot=None): + return {"snapshot": snapshot} + + bind_settings_with_providers([FakeProvider()]) + + data = capture_error_data(ValueError("boom")) + assert data["fake"] == {"snapshot": None} From e0fcaa045fa4059005593c94c77f69f46189b67f Mon Sep 17 00:00:00 2001 From: Simon Kelly Date: Wed, 5 Aug 2026 12:07:49 +0200 Subject: [PATCH 2/3] Add Sentry context provider Links a failed task to its Sentry issue by reading back sentry_sdk.last_event_id() rather than capturing the exception itself, since the app is assumed to already report it via its own Sentry integration. A snapshot taken when the task starts filters out stale ids left over from an unrelated earlier capture. --- pyproject.toml | 4 ++ taskbadger/context_providers/sentry.py | 43 +++++++++++++++++++++ tests/test_context_providers_sentry.py | 52 ++++++++++++++++++++++++++ uv.lock | 24 ++++++++++-- 4 files changed, 120 insertions(+), 3 deletions(-) create mode 100644 taskbadger/context_providers/sentry.py create mode 100644 tests/test_context_providers_sentry.py diff --git a/pyproject.toml b/pyproject.toml index d8eb693..997a0be 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -47,6 +47,9 @@ cli = [ procrastinate = [ "procrastinate>=3.0", ] +sentry = [ + "sentry-sdk>=1.0", +] [tool.uv] package = true @@ -69,6 +72,7 @@ dev = [ "redis", "openapi-python-client", "taskbadger[cli]", + "taskbadger[sentry]", ] [project.scripts] diff --git a/taskbadger/context_providers/sentry.py b/taskbadger/context_providers/sentry.py new file mode 100644 index 0000000..06b888a --- /dev/null +++ b/taskbadger/context_providers/sentry.py @@ -0,0 +1,43 @@ +from taskbadger.context_providers import ContextProvider + + +class SentryContextProvider(ContextProvider): + """Links a failed task to the corresponding Sentry issue. + + Reads back `sentry_sdk.last_event_id()` rather than capturing the exception + itself, on the assumption the surrounding system already reports its own + exceptions to Sentry (e.g. via a framework integration). To avoid linking to + a stale event left over from something unrelated, a snapshot is taken when + the task starts and the event id is only reported if it changed by the time + the task errors. + + Requires the `sentry-sdk` package; a no-op if it isn't installed. + """ + + identifier = "sentry" + + def __init__(self, organization_slug: str = None, base_url: str = "https://sentry.io"): + self.organization_slug = organization_slug + self.base_url = base_url.rstrip("/") + + def snapshot(self): + try: + import sentry_sdk + except ImportError: + return None + return sentry_sdk.last_event_id() + + def capture_error_context(self, exception: BaseException, snapshot=None) -> dict | None: + try: + import sentry_sdk + except ImportError: + return None + + event_id = sentry_sdk.last_event_id() + if not event_id or event_id == snapshot: + return None + + context = {"event_id": event_id} + if self.organization_slug: + context["url"] = f"{self.base_url}/organizations/{self.organization_slug}/issues/?query={event_id}" + return context diff --git a/tests/test_context_providers_sentry.py b/tests/test_context_providers_sentry.py new file mode 100644 index 0000000..7617d77 --- /dev/null +++ b/tests/test_context_providers_sentry.py @@ -0,0 +1,52 @@ +import sys +from unittest import mock + +from taskbadger.context_providers.sentry import SentryContextProvider + + +def test_sentry_provider_not_installed(): + provider = SentryContextProvider() + with mock.patch.dict(sys.modules, {"sentry_sdk": None}): + assert provider.snapshot() is None + assert provider.capture_error_context(ValueError("boom")) is None + + +def test_sentry_provider_no_event_id(): + provider = SentryContextProvider() + with mock.patch("sentry_sdk.last_event_id", return_value=None): + assert provider.capture_error_context(ValueError("boom")) is None + + +def test_sentry_provider_stale_event_id_not_reported(): + """If last_event_id() hasn't changed since the snapshot, the exception + was never actually captured by Sentry -- don't report the stale id.""" + provider = SentryContextProvider() + with mock.patch("sentry_sdk.last_event_id", return_value="stale123"): + snapshot = provider.snapshot() + context = provider.capture_error_context(ValueError("boom"), snapshot) + assert context is None + + +def test_sentry_provider_event_id_only(): + provider = SentryContextProvider() + snapshot = None + with mock.patch("sentry_sdk.last_event_id", return_value="abc123"): + context = provider.capture_error_context(ValueError("boom"), snapshot) + assert context == {"event_id": "abc123"} + + +def test_sentry_provider_with_url(): + provider = SentryContextProvider(organization_slug="acme") + with mock.patch("sentry_sdk.last_event_id", return_value="abc123"): + context = provider.capture_error_context(ValueError("boom")) + assert context == { + "event_id": "abc123", + "url": "https://sentry.io/organizations/acme/issues/?query=abc123", + } + + +def test_sentry_provider_custom_base_url(): + provider = SentryContextProvider(organization_slug="acme", base_url="https://sentry.example.com/") + with mock.patch("sentry_sdk.last_event_id", return_value="abc123"): + context = provider.capture_error_context(ValueError("boom")) + assert context["url"] == "https://sentry.example.com/organizations/acme/issues/?query=abc123" diff --git a/uv.lock b/uv.lock index 637cbbf..ed08f6d 100644 --- a/uv.lock +++ b/uv.lock @@ -1156,6 +1156,19 @@ wheels = [ { url = "https://files.pythonhosted.org/packages/d7/2b/9555445e1201d92b3195f45cdb153a0b68f24e0a4273f6e3d5ab46e212bb/ruff-0.15.20-py3-none-win_arm64.whl", hash = "sha256:2f5b2a6d614e8700388806a14996c40fab2c47b819ef57d790a34878858ed9ca", size = 11343498, upload-time = "2026-06-25T17:20:35.03Z" }, ] +[[package]] +name = "sentry-sdk" +version = "2.66.1" +source = { registry = "https://pypi.org/simple" } +dependencies = [ + { name = "certifi" }, + { name = "urllib3" }, +] +sdist = { url = "https://files.pythonhosted.org/packages/7f/6f/d59cad0889d15fde85254cf58e701484de3f3f0406003b3197746910b19b/sentry_sdk-2.66.1.tar.gz", hash = "sha256:f882fb08710c5f8bfc603aafa3e901b384009a19cc3f76a572b863392ee81cdc", size = 940543, upload-time = "2026-07-22T12:26:54.553Z" } +wheels = [ + { url = "https://files.pythonhosted.org/packages/89/d3/726bd88f0eece09ddf431bea4c9191c18e7a8d070b854eb0014d447712ee/sentry_sdk-2.66.1-py3-none-any.whl", hash = "sha256:86002793161d9a95ef04bdd8d442e9bfece5d989b755f05d6360215094a7aff6", size = 505555, upload-time = "2026-07-22T12:26:52.71Z" }, +] + [[package]] name = "shellingham" version = "1.5.4" @@ -1176,7 +1189,7 @@ wheels = [ [[package]] name = "taskbadger" -version = "2.3.0" +version = "2.4.0" source = { editable = "." } dependencies = [ { name = "attrs" }, @@ -1197,6 +1210,9 @@ cli = [ procrastinate = [ { name = "procrastinate" }, ] +sentry = [ + { name = "sentry-sdk" }, +] [package.dev-dependencies] dev = [ @@ -1209,7 +1225,7 @@ dev = [ { name = "pytest-celery" }, { name = "pytest-httpx" }, { name = "redis" }, - { name = "taskbadger", extra = ["cli"] }, + { name = "taskbadger", extra = ["cli", "sentry"] }, ] [package.metadata] @@ -1220,10 +1236,11 @@ requires-dist = [ { name = "procrastinate", marker = "extra == 'procrastinate'", specifier = ">=3.0" }, { name = "python-dateutil", specifier = ">=2.8.0" }, { name = "rich", marker = "extra == 'cli'", specifier = ">=13.0" }, + { name = "sentry-sdk", marker = "extra == 'sentry'", specifier = ">=1.0" }, { name = "tomlkit", specifier = ">=0.12.5" }, { name = "typer", marker = "extra == 'cli'", specifier = ">=0.12" }, ] -provides-extras = ["celery", "cli", "procrastinate"] +provides-extras = ["celery", "cli", "procrastinate", "sentry"] [package.metadata.requires-dev] dev = [ @@ -1236,6 +1253,7 @@ dev = [ { name = "pytest-httpx" }, { name = "redis" }, { name = "taskbadger", extras = ["cli"] }, + { name = "taskbadger", extras = ["sentry"] }, ] [[package]] From 1462e7480101a073b7bb4cd137c9716d8298e1ac Mon Sep 17 00:00:00 2001 From: Simon Kelly Date: Wed, 5 Aug 2026 13:54:35 +0200 Subject: [PATCH 3/3] Add integration tests for context provider wiring Covers the Celery, Procrastinate, @track, and Task.error() entry points with a real provider, not just the low-level helpers. Co-Authored-By: Claude Sonnet 5 --- tests/test_celery_error.py | 111 +++++++++++++++++++++++++++++++++++- tests/test_decorators.py | 27 +++++++++ tests/test_procrastinate.py | 31 ++++++++++ tests/test_sdk.py | 58 +++++++++++++++++++ 4 files changed, 226 insertions(+), 1 deletion(-) diff --git a/tests/test_celery_error.py b/tests/test_celery_error.py index d3a6ea1..d1727d7 100644 --- a/tests/test_celery_error.py +++ b/tests/test_celery_error.py @@ -3,7 +3,16 @@ import pytest from taskbadger import StatusEnum -from taskbadger.celery import Task +from taskbadger._error_context import capture_error_data +from taskbadger.celery import ( + TB_ERROR_CTX_TOKEN, + TB_TASK_ID, + Task, + task_failure_handler, + task_postrun_handler, + task_prerun_handler, +) +from taskbadger.context_providers import ContextProvider from taskbadger.mug import Badger from tests.utils import task_for_test @@ -41,3 +50,103 @@ def add_error(self, a, b): data_kwarg = update.call_args_list[1][1]["data"] assert "Traceback" in data_kwarg["exception"] assert Badger.current.session().client is None + + +class _FakeRequest(dict): + """Minimal stand-in for Celery's request `Context`: dict-like plus a + `.headers` attribute, which is all `_get_taskbadger_task_id` needs.""" + + headers = None + + +class _FakeEinfo: + """Minimal stand-in for Celery's `ExceptionInfo`: wraps the exception and + renders a traceback-shaped string, mirroring what `task_failure`/`task_retry` + signals actually pass to `_update_task`.""" + + def __init__(self, exc): + self.exception = exc + + def __str__(self): + return f"Traceback (most recent call last):\n{self.exception!r}" + + +@pytest.mark.usefixtures("_bind_settings") +def test_signal_handlers_wire_context_provider_snapshot_to_failure(): + """`task_prerun_handler` snapshots providers and stashes the token on the + request; `task_failure_handler` should see that exact snapshot when + building error data; `task_postrun_handler` then resets it.""" + + seen_snapshots = [] + + class TrackingProvider(ContextProvider): + identifier = "tracking" + + def snapshot(self): + return "baseline" + + def capture_error_context(self, exception, snapshot=None): + seen_snapshots.append(snapshot) + return {"snapshot": snapshot} + + Badger.current.settings.context_providers = [TrackingProvider()] + + sender = mock.Mock() + sender.request = _FakeRequest({TB_TASK_ID: "tb-1"}) + task = task_for_test(id="tb-1", status=StatusEnum.PROCESSING) + sender.taskbadger_task = task + + with ( + mock.patch("taskbadger.celery.safe_get_task", return_value=task), + mock.patch("taskbadger.celery.update_task_safe", return_value=task) as update, + mock.patch("taskbadger.celery.enter_session"), + mock.patch("taskbadger.celery.exit_session"), + ): + task_prerun_handler(sender=sender) + assert TB_ERROR_CTX_TOKEN in sender.request + + task_failure_handler(sender=sender, einfo=_FakeEinfo(ValueError("boom"))) + task_postrun_handler(sender=sender) + + # The snapshot seen at failure time is the one taken at prerun, not a + # missing/None one -- proving the token round-trips through the request. + assert seen_snapshots == ["baseline"] + data_kwarg = update.call_args.kwargs["data"] + assert data_kwarg["tracking"] == {"snapshot": "baseline"} + + +@pytest.mark.usefixtures("_bind_settings") +def test_postrun_resets_context_after_error(): + """Once `task_postrun_handler` runs, a later error in the same thread with + no provider snapshot taken shouldn't see a stale one left over from the + previous task.""" + + class TrackingProvider(ContextProvider): + identifier = "tracking" + + def snapshot(self): + return "baseline" + + def capture_error_context(self, exception, snapshot=None): + return {"snapshot": snapshot} + + Badger.current.settings.context_providers = [TrackingProvider()] + + task = task_for_test(id="tb-2", status=StatusEnum.PROCESSING) + sender = mock.Mock() + sender.request = _FakeRequest({TB_TASK_ID: "tb-2"}) + sender.taskbadger_task = task + + with ( + mock.patch("taskbadger.celery.safe_get_task", return_value=task), + mock.patch("taskbadger.celery.update_task_safe", return_value=task), + mock.patch("taskbadger.celery.enter_session"), + mock.patch("taskbadger.celery.exit_session"), + ): + task_prerun_handler(sender=sender) + task_postrun_handler(sender=sender) + + # A failure reported outside of any tracked task's prerun/postrun window + # (e.g. directly via Task.error) has no snapshot to compare against. + data = capture_error_data(ValueError("boom")) + assert data["tracking"] == {"snapshot": None} diff --git a/tests/test_decorators.py b/tests/test_decorators.py index f7960bc..209f88c 100644 --- a/tests/test_decorators.py +++ b/tests/test_decorators.py @@ -1,6 +1,7 @@ from unittest import mock from taskbadger import track +from taskbadger.context_providers import ContextProvider from taskbadger.mug import Badger, Settings @@ -52,6 +53,32 @@ def test(arg): assert update.call_args.kwargs["data"]["exception"] == "test" +@mock.patch("taskbadger.decorators.create_task_safe") +@mock.patch("taskbadger.decorators._update_safe") +def test_track_decorator_error_consults_context_provider(update, create): + class FakeProvider(ContextProvider): + identifier = "fake" + + def capture_error_context(self, exception, snapshot=None): + return {"detail": str(exception)} + + Badger.current.bind(Settings("https://taskbadger.net", "token", "org", "proj", context_providers=[FakeProvider()])) + try: + + @track + def test(arg): + raise Exception("test") + + try: + test("test") + except Exception: + pass + finally: + Badger.current.bind(None) + + assert update.call_args.kwargs["data"] == {"exception": "test", "fake": {"detail": "test"}} + + @mock.patch("taskbadger.decorators._update_safe") def test_track_decorator_badger_not_configured(update): @track diff --git a/tests/test_procrastinate.py b/tests/test_procrastinate.py index 7854261..bc60a54 100644 --- a/tests/test_procrastinate.py +++ b/tests/test_procrastinate.py @@ -7,6 +7,8 @@ from procrastinate import testing from taskbadger import StatusEnum +from taskbadger.context_providers import ContextProvider +from taskbadger.mug import Badger from taskbadger.procrastinate import TB_TASK_ID_KWARG, _instrument_task, current_task, track from tests.utils import task_for_test @@ -91,6 +93,35 @@ def boom(): assert err_call.kwargs["data"] == {"x": 1, "exception": "nope"} +@pytest.mark.usefixtures("_bind_settings") +def test_worker_marks_error_consults_context_provider(app): + class FakeProvider(ContextProvider): + identifier = "fake" + + def capture_error_context(self, exception, snapshot=None): + return {"detail": str(exception)} + + Badger.current.settings.context_providers = [FakeProvider()] + + @app.task(name="boom_with_provider") + def boom(): + raise ValueError("nope") + + _instrument_task(boom, system=None, manual=True) + + with ( + mock.patch("taskbadger.procrastinate.update_task_safe") as update, + mock.patch("taskbadger.sdk.get_task") as get, + ): + get.return_value = task_for_test(status=StatusEnum.PROCESSING) + update.return_value = task_for_test(status=StatusEnum.PROCESSING) + with pytest.raises(ValueError, match="nope"): + boom.func(**{TB_TASK_ID_KWARG: "tb-provider"}) + + err_call = update.call_args_list[-1] + assert err_call.kwargs["data"] == {"exception": "nope", "fake": {"detail": "nope"}} + + @pytest.mark.usefixtures("_bind_settings") def test_worker_no_id_runs_clean(app): @app.task(name="add2") diff --git a/tests/test_sdk.py b/tests/test_sdk.py index 768cade..e4ef259 100644 --- a/tests/test_sdk.py +++ b/tests/test_sdk.py @@ -6,6 +6,7 @@ import pytest from taskbadger import Action, EmailIntegration, StatusEnum, WebhookIntegration, create_task +from taskbadger.context_providers import ContextProvider from taskbadger.exceptions import TaskbadgerException from taskbadger.internal.models import ( PatchedTaskRequest, @@ -207,6 +208,63 @@ def test_update_data(settings, patched_update): _verify_update(settings, patched_update, data={"a": 1}) +def test_error_with_exception(settings, patched_update): + api_task = task_for_test() + task = Task(api_task) + + patched_update.return_value = Response(HTTPStatus.OK, b"", {}, api_task) + task.error(exception=ValueError("boom")) + + _verify_update(settings, patched_update, status=StatusEnum.ERROR, data={"exception": "boom"}) + + +def test_error_with_exception_consults_context_provider(settings, patched_update): + class FakeProvider(ContextProvider): + identifier = "fake" + + def capture_error_context(self, exception, snapshot=None): + return {"detail": str(exception)} + + settings.context_providers = [FakeProvider()] + + api_task = task_for_test() + task = Task(api_task) + patched_update.return_value = Response(HTTPStatus.OK, b"", {}, api_task) + task.error(exception=ValueError("boom")) + + _verify_update( + settings, + patched_update, + status=StatusEnum.ERROR, + data={"exception": "boom", "fake": {"detail": "boom"}}, + ) + + +def test_error_explicit_data_overrides_provider_data(settings, patched_update): + """Explicit `data` passed to `error()` wins over provider-derived data for + overlapping keys.""" + + class FakeProvider(ContextProvider): + identifier = "fake" + + def capture_error_context(self, exception, snapshot=None): + return {"detail": "from provider"} + + settings.context_providers = [FakeProvider()] + + api_task = task_for_test() + task = Task(api_task) + patched_update.return_value = Response(HTTPStatus.OK, b"", {}, api_task) + task.error(exception=ValueError("boom"), data={"exception": "overridden", "extra": 1}) + + _verify_update( + settings, + patched_update, + status=StatusEnum.ERROR, + data={"exception": "overridden", "extra": 1, "fake": {"detail": "from provider"}}, + ) + + def test_increment_value(settings, patched_update): api_task = task_for_test() task = Task(api_task)