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
4 changes: 4 additions & 0 deletions pyproject.toml
Original file line number Diff line number Diff line change
Expand Up @@ -47,6 +47,9 @@ cli = [
procrastinate = [
"procrastinate>=3.0",
]
sentry = [
"sentry-sdk>=1.0",
]

[tool.uv]
package = true
Expand All @@ -69,6 +72,7 @@ dev = [
"redis",
"openapi-python-client",
"taskbadger[cli]",
"taskbadger[sentry]",
]

[project.scripts]
Expand Down
59 changes: 59 additions & 0 deletions taskbadger/_error_context.py
Original file line number Diff line number Diff line change
@@ -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
22 changes: 21 additions & 1 deletion taskbadger/celery.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand All @@ -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")

Expand Down Expand Up @@ -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
Expand All @@ -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
Expand Down Expand Up @@ -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)
Expand Down
32 changes: 32 additions & 0 deletions taskbadger/context_providers/__init__.py
Original file line number Diff line number Diff line change
@@ -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
43 changes: 43 additions & 0 deletions taskbadger/context_providers/sentry.py
Original file line number Diff line number Diff line change
@@ -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
6 changes: 5 additions & 1 deletion taskbadger/decorators.py
Original file line number Diff line number Diff line change
@@ -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
Expand Down Expand Up @@ -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
Expand Down
2 changes: 2 additions & 0 deletions taskbadger/mug.py
Original file line number Diff line number Diff line change
Expand Up @@ -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

Expand All @@ -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)
Expand Down
9 changes: 8 additions & 1 deletion taskbadger/procrastinate.py
Original file line number Diff line number Diff line change
Expand Up @@ -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,
Expand Down Expand Up @@ -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:
Expand All @@ -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:
Expand Down Expand Up @@ -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:
Expand Down
26 changes: 23 additions & 3 deletions taskbadger/sdk.py
Original file line number Diff line number Diff line change
Expand Up @@ -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,
Expand Down Expand Up @@ -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.

Expand All @@ -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(
Expand All @@ -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")
Expand Down Expand Up @@ -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:
Expand Down Expand Up @@ -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):
Expand Down
Loading