Skip to content
Open
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
48 changes: 44 additions & 4 deletions lemur/common/celery.py
Original file line number Diff line number Diff line change
Expand Up @@ -18,6 +18,7 @@
after_setup_logger,
after_setup_task_logger,
task_failure,
task_prerun,
task_received,
task_revoked,
task_success,
Expand Down Expand Up @@ -53,6 +54,7 @@
flask_app = create_app()

red = RedisHandler().redis()
_task_started_at = {}


def make_celery(app):
Expand Down Expand Up @@ -197,6 +199,16 @@ def report_celery_last_success_metrics():
metrics.send(f"{function}.success", "counter", 1)


@task_prerun.connect
def report_task_started(**kwargs):
"""
Record task start time so we can emit duration on completion/failure.
"""
task_id = kwargs.get("task_id")
if task_id:
_task_started_at[task_id] = time.monotonic()


@task_received.connect
def report_number_pending_tasks(**kwargs):
"""
Expand All @@ -213,6 +225,35 @@ def report_number_pending_tasks(**kwargs):
)


def _emit_task_duration(status, **kwargs):
"""
Emit a task duration metric if we previously recorded a start time for the task.

Returns the tags for the task (from get_celery_request_tags) so callers can continue
to use them for further metrics/logging.
"""
tags = get_celery_request_tags(**kwargs)
started_at = _task_started_at.pop(tags["task_id"], None)
if started_at is None:
return tags
duration_ms = int((time.monotonic() - started_at) * 1000)
metrics.send(
"celery.task_duration",
"TIMER",
duration_ms,
metric_tags={"task_name": tags["task_name"], "status": status},
)
return tags


def _task_status_from_failure(**kwargs):
einfo = kwargs.get("einfo")
exception = getattr(einfo, "exception", None)
if isinstance(exception, SoftTimeLimitExceeded):
return "timeout"
return "failure"


@task_success.connect
def report_successful_task(**kwargs):
"""
Expand All @@ -221,7 +262,7 @@ def report_successful_task(**kwargs):
https://docs.celeryproject.org/en/latest/userguide/signals.html#task-success
"""
with flask_app.app_context():
tags = get_celery_request_tags(**kwargs)
tags = _emit_task_duration("success", **kwargs)
red.set(f"{tags['task_name']}.last_success", int(time.time()))
metrics.send("celery.successful_task", "TIMER", 1, metric_tags=tags)
# Emit failed_task=0 on success so the counter stays dense (0 when healthy)
Expand All @@ -246,14 +287,13 @@ def report_failed_task(**kwargs):
"function": f"{__name__}.{sys._getframe().f_code.co_name}",
"Message": "Celery Task Failure",
}
error_tags = _emit_task_duration(_task_status_from_failure(**kwargs), **kwargs)

# Add traceback if exception info is in the kwargs
einfo = kwargs.get("einfo")
if einfo:
log_data["traceback"] = einfo.traceback

error_tags = get_celery_request_tags(**kwargs)

log_data.update(error_tags)
current_app.logger.error(log_data)
metrics.send("celery.failed_task", "counter", 1, metric_tags=error_tags)
Expand All @@ -272,7 +312,7 @@ def report_revoked_task(**kwargs):
"Message": "Celery Task Revoked",
}

error_tags = get_celery_request_tags(**kwargs)
error_tags = _emit_task_duration("revoked", **kwargs)

log_data.update(error_tags)
current_app.logger.error(log_data)
Expand Down
7 changes: 6 additions & 1 deletion lemur/tests/conftest.py
Original file line number Diff line number Diff line change
Expand Up @@ -5,7 +5,7 @@
from cryptography import x509
from cryptography.hazmat.backends import default_backend
from cryptography.hazmat.primitives import hashes
from flask import current_app
from flask import current_app, g
from flask_principal import identity_changed, Identity
from sqlalchemy.sql import text

Expand Down Expand Up @@ -104,6 +104,11 @@ def session(db, request):
db.session.begin_nested()
yield db.session
db.session.rollback()
# g is app-context scoped and persists across tests. Clear request-scoped
# state so a stale, detached User from an earlier test isn't reused by a
# later test (which raises DetachedInstanceError when its attributes are
# expired).
g.pop("current_user", None)


@pytest.fixture(scope="function")
Expand Down
81 changes: 81 additions & 0 deletions lemur/tests/test_celery_metrics.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,81 @@
"""Tests for centralized Celery task duration metrics."""
import sys
from unittest.mock import MagicMock, patch

if "lemur.common.celery" not in sys.modules:
with patch("redis.StrictRedis") as _mock_redis:
_mock_redis.return_value.set.return_value = True
import lemur.common.celery # noqa: F401

import lemur.common.celery as _celery_module # noqa: E402


class FakeEinfo:
def __init__(self, exc):
self.exception = exc
self.traceback = "fake traceback"


class FakeRequest:
def __init__(self, task_id="task-1", name="lemur.common.celery.fake_task"):
self.id = task_id
self.hostname = "worker-1"
self.name = name


class FakeTask:
def __init__(self, task_id="task-1", name="lemur.common.celery.fake_task"):
self.request = FakeRequest(task_id=task_id, name=name)
self.hostname = "worker-1"
self.name = name


@patch("lemur.common.celery.metrics")
def test_task_duration_emitted_on_success_and_clears_start_time(mock_metrics):
_celery_module._task_started_at.clear()
_celery_module._task_started_at["task-1"] = 100.0

with patch.object(_celery_module, "current_app", MagicMock()), patch(
"lemur.common.celery.time.monotonic", return_value=101.234
), patch("lemur.common.celery.time.time", return_value=2000):
fake_task = FakeTask()
_celery_module.report_successful_task(sender=fake_task, request=fake_task.request)

duration_calls = [c for c in mock_metrics.send.call_args_list if c.args[0] == "celery.task_duration"]
assert len(duration_calls) == 1
assert duration_calls[0].kwargs["metric_tags"] == {"task_name": "lemur.common.celery.fake_task", "status": "success"}
assert duration_calls[0].args[2] == 1233
assert "task-1" not in _celery_module._task_started_at


@patch("lemur.common.celery.metrics")
def test_task_duration_emitted_on_failure_with_timeout_status(mock_metrics):
_celery_module._task_started_at.clear()
_celery_module._task_started_at["task-2"] = 50.0

with patch.object(_celery_module, "current_app", MagicMock()), patch(
"lemur.common.celery.time.monotonic", return_value=51.5
), patch("lemur.common.celery.time.time", return_value=2000):
fake_task = FakeTask(task_id="task-2")
_celery_module.report_failed_task(
sender=fake_task,
request=fake_task.request,
einfo=FakeEinfo(_celery_module.SoftTimeLimitExceeded()),
)

duration_calls = [c for c in mock_metrics.send.call_args_list if c.args[0] == "celery.task_duration"]
assert len(duration_calls) == 1
assert duration_calls[0].kwargs["metric_tags"]["status"] == "timeout"
assert duration_calls[0].args[2] == 1500
assert "task-2" not in _celery_module._task_started_at


@patch("lemur.common.celery.metrics")
def test_task_duration_not_emitted_when_no_start_time(mock_metrics):
_celery_module._task_started_at.clear()

with patch.object(_celery_module, "current_app", MagicMock()):
fake_task = FakeTask(task_id="task-3")
_celery_module.report_successful_task(sender=fake_task, request=fake_task.request)

assert not [c for c in mock_metrics.send.call_args_list if c.args[0] == "celery.task_duration"]
Loading