Skip to content

Commit 65223b7

Browse files
Mohammed0tareks3rius
authored andcommitted
feat: add OTel metrics to OpenTelemetryMiddleware
Add counters and histograms to OpenTelemetryMiddleware: - tasks_sent: producer-side counter per task name - task_success / task_errors: consumer-side counters with retry_error attribute - task_execution_time: histogram using result.execution_time - task_wait_time: histogram measuring queue time from send to receive via UTC timestamps in labels Add tests covering all instruments, retry_error attribute paths, and queue time correctness.
1 parent eaed387 commit 65223b7

3 files changed

Lines changed: 262 additions & 1 deletion

File tree

‎taskiq/middlewares/opentelemetry_middleware.py‎

Lines changed: 87 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -1,5 +1,6 @@
11
import logging
22
from contextlib import AbstractContextManager
3+
from datetime import datetime, timezone
34
from importlib.metadata import version
45
from typing import Any, TypeVar
56

@@ -59,6 +60,9 @@
5960
_TASK_RETRY_REASON_KEY = "taskiq.retry.reason"
6061
_TASK_NAME_KEY = "taskiq.task_name"
6162

63+
_TASK_QUEUE_TIME_KEY = "_taskiq_queue_time"
64+
_TASK_RECEIVED_TIME_KEY = "_taskiq_broker_receive_time"
65+
6266

6367
def set_attributes_from_context(span: Span, context: dict[str, Any]) -> None:
6468
"""Helper to extract meta values from a Taskiq Context."""
@@ -170,6 +174,37 @@ def __init__(
170174
if meter is None
171175
else meter
172176
)
177+
# Create metrics
178+
# 1- Number of tasks sent. Producer (Counter)
179+
self.n_tasks_sent_counter = self._meter.create_counter(
180+
name="tasks_sent",
181+
unit="1",
182+
description="Number of tasks sent from the producer side",
183+
)
184+
# 2- Number of errors by task name. consumer (Counter)
185+
self.n_errors_counter = self._meter.create_counter(
186+
name="task_errors",
187+
unit="1",
188+
description="Number of errors raised",
189+
)
190+
# 3- Number of task successes. consumer (Counter)
191+
self.n_success_counter = self._meter.create_counter(
192+
name="task_success",
193+
unit="1",
194+
description="Number of tasks completed successfully",
195+
)
196+
# 4- Task execution time. consumer (Histogram)
197+
self.execution_time_hist = self._meter.create_histogram(
198+
"task_execution_time",
199+
unit="s",
200+
description="Time to finish executing tasks",
201+
)
202+
# 5- Task wait time. both (Histogram)
203+
self.task_wait_time = self._meter.create_histogram(
204+
"task_wait_time",
205+
unit="s",
206+
description="Time the tasks waited before executing",
207+
)
173208

174209
def pre_send(self, message: TaskiqMessage) -> TaskiqMessage:
175210
"""
@@ -193,7 +228,7 @@ def pre_send(self, message: TaskiqMessage) -> TaskiqMessage:
193228
activation.__enter__()
194229
attach_context(message, span, activation, None, is_publish=True)
195230
inject(message.labels)
196-
231+
message.labels[_TASK_QUEUE_TIME_KEY] = datetime.now(timezone.utc).timestamp()
197232
return message
198233

199234
def post_send(self, message: TaskiqMessage) -> None:
@@ -214,6 +249,7 @@ def post_send(self, message: TaskiqMessage) -> None:
214249

215250
activation.__exit__(None, None, None)
216251
detach_context(message, is_publish=True)
252+
self.n_tasks_sent_counter.add(1, attributes={"task_name": message.task_name})
217253

218254
def pre_execute(self, message: TaskiqMessage) -> TaskiqMessage:
219255
"""
@@ -236,6 +272,7 @@ def pre_execute(self, message: TaskiqMessage) -> TaskiqMessage:
236272
activation = trace.use_span(span, end_on_exit=True)
237273
activation.__enter__() # pylint: disable=E1101
238274
attach_context(message, span, activation, token)
275+
message.labels[_TASK_RECEIVED_TIME_KEY] = datetime.now(timezone.utc).timestamp()
239276
return message
240277

241278
def post_save( # pylint: disable=R6301
@@ -313,3 +350,52 @@ def on_error(
313350
}
314351
span.record_exception(exception)
315352
span.set_status(Status(**status_kwargs)) # type: ignore[arg-type]
353+
354+
def post_execute(
355+
self,
356+
message: "TaskiqMessage",
357+
result: "TaskiqResult[Any]",
358+
) -> None:
359+
"""
360+
This function tracks number of errors and success executions.
361+
362+
:param message: received message.
363+
:param result: result of the execution.
364+
"""
365+
if result.is_err:
366+
retry_on_error = message.labels.get("retry_on_error")
367+
if isinstance(retry_on_error, str):
368+
retry_on_error = retry_on_error.lower() == "true"
369+
370+
if retry_on_error is None:
371+
retry_on_error = False
372+
373+
if retry_on_error:
374+
# Add retry reason metadata to span
375+
self.n_errors_counter.add(
376+
1,
377+
attributes={"retry_error": True, "task_name": message.task_name},
378+
)
379+
else:
380+
self.n_errors_counter.add(
381+
1,
382+
attributes={"retry_error": False, "task_name": message.task_name},
383+
)
384+
else:
385+
self.n_success_counter.add(
386+
1,
387+
attributes={"task_name": message.task_name},
388+
)
389+
self.execution_time_hist.record(
390+
result.execution_time,
391+
attributes={
392+
"task_name": message.task_name,
393+
},
394+
)
395+
task_receive_time = message.labels.get(_TASK_RECEIVED_TIME_KEY)
396+
task_send_time = message.labels.get(_TASK_QUEUE_TIME_KEY)
397+
if task_receive_time is not None and task_send_time is not None:
398+
self.task_wait_time.record(
399+
amount=task_receive_time - task_send_time,
400+
attributes={"task_name": message.task_name},
401+
)

‎tests/opentelemetry/taskiq_test_tasks.py‎

Lines changed: 6 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -1,3 +1,4 @@
1+
import asyncio
12
from typing import Any
23

34
from opentelemetry import baggage
@@ -26,3 +27,8 @@ async def task_raises() -> None:
2627
@broker.task
2728
async def task_returns_baggage() -> Any:
2829
return dict(baggage.get_all())
30+
31+
32+
@broker.task
33+
async def task_does_processing(wait_time: float) -> None:
34+
await asyncio.sleep(wait_time)
Lines changed: 169 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,169 @@
1+
import asyncio
2+
from typing import Any
3+
4+
from opentelemetry.sdk.metrics import MeterProvider
5+
from opentelemetry.sdk.metrics.export import InMemoryMetricReader
6+
from opentelemetry.test.test_base import TestBase
7+
8+
from taskiq.instrumentation import TaskiqInstrumentor
9+
10+
from .taskiq_test_tasks import (
11+
broker,
12+
task_add,
13+
task_does_processing,
14+
task_raises,
15+
)
16+
17+
18+
class TestTaskiqOTelMetrics(TestBase):
19+
def setUp(self) -> None:
20+
super().setUp()
21+
self.reader = InMemoryMetricReader()
22+
self.meter_provider = MeterProvider(metric_readers=[self.reader])
23+
TaskiqInstrumentor().instrument_broker(
24+
broker,
25+
meter_provider=self.meter_provider,
26+
)
27+
28+
def tearDown(self) -> None:
29+
super().tearDown()
30+
TaskiqInstrumentor().uninstrument_broker(broker)
31+
32+
def _get_data_points(self, metric_name: str) -> list[Any]:
33+
metrics = self.reader.get_metrics_data()
34+
if metrics is None:
35+
return []
36+
return [
37+
point
38+
for rm in metrics.resource_metrics
39+
for sm in rm.scope_metrics
40+
for metric in sm.metrics
41+
if metric.name == metric_name
42+
for point in metric.data.data_points
43+
]
44+
45+
def test_metrics_exist(self) -> None:
46+
async def test() -> None:
47+
await task_add.kiq(1, 2)
48+
await task_raises.kiq()
49+
await broker.wait_all()
50+
51+
asyncio.run(test())
52+
53+
metrics = self.reader.get_metrics_data()
54+
self.assertIsNotNone(metrics)
55+
expected = {
56+
"task_errors",
57+
"tasks_sent",
58+
"task_success",
59+
"task_execution_time",
60+
"task_wait_time",
61+
}
62+
found = {
63+
metric.name
64+
for rm in metrics.resource_metrics # type: ignore[union-attr]
65+
for sm in rm.scope_metrics
66+
for metric in sm.metrics
67+
}
68+
self.assertSetEqual(found.intersection(expected), expected)
69+
70+
def test_success_counter(self) -> None:
71+
async def test() -> None:
72+
for _ in range(3):
73+
await task_add.kiq(1, 2)
74+
await broker.wait_all()
75+
76+
asyncio.run(test())
77+
78+
points = self._get_data_points("task_success")
79+
self.assertEqual(len(points), 1)
80+
self.assertEqual(points[0].value, 3)
81+
82+
def test_error_counter_no_retry(self) -> None:
83+
async def test() -> None:
84+
for _ in range(3):
85+
await task_raises.kiq()
86+
await broker.wait_all()
87+
88+
asyncio.run(test())
89+
90+
points = self._get_data_points("task_errors")
91+
no_retry_points = [
92+
p for p in points if p.attributes.get("retry_error") is False
93+
]
94+
self.assertEqual(len(no_retry_points), 1)
95+
self.assertEqual(no_retry_points[0].value, 3)
96+
97+
def test_error_counter_with_retry(self) -> None:
98+
async def test() -> None:
99+
for _ in range(3):
100+
await task_raises.kicker().with_labels(retry_on_error="true").kiq()
101+
await broker.wait_all()
102+
103+
asyncio.run(test())
104+
105+
points = self._get_data_points("task_errors")
106+
retry_points = [p for p in points if p.attributes.get("retry_error") is True]
107+
self.assertEqual(len(retry_points), 1)
108+
self.assertEqual(retry_points[0].value, 3)
109+
110+
def test_execution_time_histogram(self) -> None:
111+
async def test() -> None:
112+
for _ in range(3):
113+
await task_does_processing.kiq(0.01)
114+
await broker.wait_all()
115+
116+
asyncio.run(test())
117+
118+
points = self._get_data_points("task_execution_time")
119+
self.assertEqual(len(points), 1)
120+
self.assertEqual(points[0].count, 3)
121+
self.assertGreater(points[0].sum, 0)
122+
123+
def test_task_wait_time_histogram(self) -> None:
124+
async def test() -> None:
125+
await task_does_processing.kiq(0.01)
126+
await broker.wait_all()
127+
128+
asyncio.run(test())
129+
130+
points = self._get_data_points("task_wait_time")
131+
self.assertEqual(len(points), 1)
132+
self.assertEqual(points[0].count, 1)
133+
self.assertGreaterEqual(points[0].sum, 0)
134+
135+
def test_queue_time(self) -> None:
136+
async def test() -> None:
137+
for _ in range(3):
138+
await task_add.kiq(1, 2)
139+
await broker.wait_all()
140+
141+
asyncio.run(test())
142+
143+
points = self._get_data_points("task_wait_time")
144+
# all 3 tasks share the same task_name so they aggregate into one data point
145+
self.assertEqual(len(points), 1)
146+
point = points[0]
147+
# 3 tasks recorded
148+
self.assertEqual(point.count, 3)
149+
# queue time must be non-negative — a negative value means timestamps
150+
# were not written/read correctly
151+
self.assertGreaterEqual(point.sum, 0)
152+
self.assertGreaterEqual(point.min, 0)
153+
# task_name attribute must be present and correct
154+
self.assertEqual(
155+
point.attributes.get("task_name"),
156+
"tests.opentelemetry.taskiq_test_tasks:task_add",
157+
)
158+
159+
def test_tasks_sent_counter(self) -> None:
160+
async def test() -> None:
161+
for _ in range(3):
162+
await task_add.kiq(1, 2)
163+
await broker.wait_all()
164+
165+
asyncio.run(test())
166+
167+
points = self._get_data_points("tasks_sent")
168+
self.assertEqual(len(points), 1)
169+
self.assertEqual(points[0].value, 3)

0 commit comments

Comments
 (0)