From 999016150f62b6b93233176f0fd585897a991043 Mon Sep 17 00:00:00 2001 From: Rodrigo Antunes Date: Tue, 15 Sep 2026 09:30:33 -0300 Subject: [PATCH] fix(RHINENG-31023): keep consumer spans active MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit The Kafka consumer spans in listener and grouper closed prematurely at scheduling time (~50-300µs) because the span wrapped only the task scheduling. Descendant operations (DB queries and Kafka sends) ran in background tasks under the already-ended parent span. Move consumer spans into the awaited coroutines: - listener: wrap consume_inventory_msg and consume_advisor_msg - grouper: wrap _start_item_processing - await outgoing producer sends in listener and grouper so they finish before consumer spans close Assisted-by: Cursor:gemini-3.8-flash --- common/utils.py | 3 +- grouper/common.py | 1 + grouper/grouper.py | 50 +++--- grouper/queue.py | 115 ++++++++---- listener/advisor_processor.py | 12 +- listener/inventory_processor.py | 12 +- listener/listener.py | 93 +++++----- tests/grouper_tests/__init__.py | 1 + tests/grouper_tests/test_grouper_tracing.py | 125 +++++++++++++ tests/listener_tests/test_listener_tracing.py | 168 ++++++++++++++++++ 10 files changed, 462 insertions(+), 118 deletions(-) create mode 100644 tests/grouper_tests/__init__.py create mode 100644 tests/grouper_tests/test_grouper_tracing.py create mode 100644 tests/listener_tests/test_listener_tracing.py diff --git a/common/utils.py b/common/utils.py index f985d0299..5f3f87fa1 100644 --- a/common/utils.py +++ b/common/utils.py @@ -123,8 +123,9 @@ def send_msg_to_payload_tracker(producer, msg_dict, status, status_msg=None, loo } if status_msg: tracking_payload["status_msg"] = status_msg - producer.send(tracking_payload, loop=loop) + fut = producer.send(tracking_payload, loop=loop) LOGGER.debug("Sent message to topic %s: %s", producer.topic, str(tracking_payload)) + return fut def send_remediations_update(producer, inventory_id: str, cves: list, loop=None) -> None: diff --git a/grouper/common.py b/grouper/common.py index 8a23b3373..93cd5f871 100644 --- a/grouper/common.py +++ b/grouper/common.py @@ -63,6 +63,7 @@ class QueueItem: request_id: str otel_context: otel_context_api.Context | None = None + topic: str = "" second_upload_event: asyncio.Event = None def __post_init__(self): diff --git a/grouper/grouper.py b/grouper/grouper.py index 6dfb5b07e..6f8091551 100644 --- a/grouper/grouper.py +++ b/grouper/grouper.py @@ -21,7 +21,6 @@ from common.telemetry import instrument_outbound_http from common.telemetry import instrument_psycopg from common.telemetry import instrument_psycopg2 -from common.telemetry import use_otel_context from common.utils import create_task_and_log from .common import CFG @@ -85,35 +84,26 @@ async def consume_message(self, msg: ConsumerRecord, unlock: asyncio.BoundedSema parent_ctx = extract_context_from_msg_headers(msg.headers) - with use_otel_context(parent_ctx): - with TRACER.start_as_current_span( - f"process {msg.topic}", - attributes={ - "messaging.system": "kafka", - "messaging.operation.name": "process", - "messaging.destination.name": msg.topic, - "vulnerability.inventory_id": inventory_id, - "vulnerability.msg_type": msg_type.value, - }, - ): - if msg_type is GrouperMessageType.INVENTORY_UPLOAD: - self.queue.push_inventory_msg( - org_id, - inventory_id, - reporter, - changed, - (msg_dict.get("platform_metadata") or {}).get("request_id", ""), - otel_context=parent_ctx, - ) - elif msg_type is GrouperMessageType.ADVISOR_UPLOAD: - self.queue.push_advisor_msg( - org_id, - inventory_id, - reporter, - changed, - (msg_dict.get("platform_metadata") or {}).get("request_id", ""), - otel_context=parent_ctx, - ) + if msg_type is GrouperMessageType.INVENTORY_UPLOAD: + self.queue.push_inventory_msg( + org_id, + inventory_id, + reporter, + changed, + (msg_dict.get("platform_metadata") or {}).get("request_id", ""), + otel_context=parent_ctx, + topic=msg.topic, + ) + elif msg_type is GrouperMessageType.ADVISOR_UPLOAD: + self.queue.push_advisor_msg( + org_id, + inventory_id, + reporter, + changed, + (msg_dict.get("platform_metadata") or {}).get("request_id", ""), + otel_context=parent_ctx, + topic=msg.topic, + ) async def _start_grouping_inventory(self) -> None: """Start of the grouping inventory uploads""" diff --git a/grouper/queue.py b/grouper/queue.py index f8b52abc4..ea721f5ec 100644 --- a/grouper/queue.py +++ b/grouper/queue.py @@ -9,10 +9,13 @@ from typing import Dict from opentelemetry import context as otel_context_api +from opentelemetry import trace from common.constants import EvaluatorMessageType from common.logging import get_logger from common.mqueue import MQWriter +from common.telemetry import get_tracer +from common.telemetry import threadctx from common.telemetry import use_otel_context from common.utils import create_task_and_log from common.utils import send_msg_to_payload_tracker @@ -26,9 +29,11 @@ from .common import QUEUE_SIZE from .common import UNCHANGED_SYSTEM from .common import BoundedSemaphorePrometheus +from .common import GrouperMessageType from .common import QueueItem LOGGER = get_logger(__name__) +TRACER = get_tracer(__name__) class GrouperQueue: @@ -64,6 +69,7 @@ def push_inventory_msg( inventory_changed: bool, request_id: str, otel_context: otel_context_api.Context | None = None, + topic: str = "", ) -> None: """Push inventory upload message to queue""" is_updated = False @@ -71,7 +77,7 @@ def push_inventory_msg( item = self._queue.get(inventory_id) if not item: LOGGER.info("pushing listener upload to queue for system: %s, org_id: %s", inventory_id, org_id) - item = QueueItem(True, inventory_changed, False, False, request_id, otel_context=otel_context) + item = QueueItem(True, inventory_changed, False, False, request_id, otel_context=otel_context, topic=topic) self._queue[inventory_id] = item else: is_updated = True @@ -86,6 +92,8 @@ def push_inventory_msg( item.request_id = request_id if otel_context is not None: item.otel_context = otel_context + if topic and not item.topic: + item.topic = topic if item.inventory_upload and item.advisor_upload: LOGGER.info("obtained both uploads for system: %s, account: %s, releasing lock", inventory_id, org_id) @@ -102,6 +110,7 @@ def push_advisor_msg( advisor_changed: bool, request_id: str, otel_context: otel_context_api.Context | None = None, + topic: str = "", ) -> None: """Push advisor message to queue""" is_updated = False @@ -109,7 +118,7 @@ def push_advisor_msg( item = self._queue.get(inventory_id) if not item: LOGGER.info("pushing advisor upload to queue for system: %s, org_id: %s", inventory_id, org_id) - item = QueueItem(False, False, True, advisor_changed, request_id, otel_context=otel_context) + item = QueueItem(False, False, True, advisor_changed, request_id, otel_context=otel_context, topic=topic) self._queue[inventory_id] = item else: is_updated = True @@ -124,6 +133,8 @@ def push_advisor_msg( item.request_id = request_id if otel_context is not None and item.otel_context is None: item.otel_context = otel_context + if topic and not item.topic: + item.topic = topic if item.inventory_upload and item.advisor_upload: LOGGER.info("obtained both messages for system: %s, account: %s, releasing lock", inventory_id, org_id) @@ -135,40 +146,61 @@ def push_advisor_msg( async def _start_item_processing(self, org_id: str, inventory_id: str, reporter: str) -> None: """Single queue item waiting coroutine""" item = self._queue[inventory_id] + topic = item.topic or (CFG.grouper_inventory_topic if item.inventory_upload else CFG.grouper_advisor_topic) + msg_type = GrouperMessageType.INVENTORY_UPLOAD.value if item.inventory_upload else GrouperMessageType.ADVISOR_UPLOAD.value - # RHSM systems do not need to wait for advisor message - if reporter == "rhsm-system-profile-bridge": - LOGGER.debug( - "reporter %s skipped waiting for %s message for system: %s, account: %s", - reporter, - "inventory" if item.advisor_upload else "advisor", - inventory_id, - org_id, - ) - else: - LOGGER.debug( - "starting waiting for %s msg for system: %s, account: %s", - "inventory" if item.advisor_upload else "advisor", - inventory_id, - org_id, - ) - - try: - await asyncio.wait_for(item.second_upload_event.wait(), CFG.grouper_messages_timeout_sec) - PAIR_HIT.inc() - except asyncio.TimeoutError: - LOGGER.info("timing out while waiting for message for system: %s, account: %s", inventory_id, org_id) - PAIR_MISS.inc() - self._queue.pop(inventory_id, None) - - if item.inventory_upload: - self.max_inventory_msgs.release() - if item.advisor_upload: - self.max_advisor_msgs.release() - - await self._send_for_evaluation(item, org_id, inventory_id) - - QUEUE_SIZE.dec() + with use_otel_context(item.otel_context): + with TRACER.start_as_current_span( + f"process {topic}", + kind=trace.SpanKind.CONSUMER, + attributes={ + "messaging.system": "kafka", + "messaging.operation.name": "process", + "messaging.destination.name": topic, + "vulnerability.inventory_id": inventory_id, + "vulnerability.msg_type": msg_type, + }, + ) as span: + if org_id: + span.set_attribute("rh.org_id", org_id) + threadctx.org_id = org_id + if item.request_id: + span.set_attribute("rh.request_id", item.request_id) + threadctx.request_id = item.request_id + + # RHSM systems do not need to wait for advisor message + if reporter == "rhsm-system-profile-bridge": + LOGGER.debug( + "reporter %s skipped waiting for %s message for system: %s, account: %s", + reporter, + "inventory" if item.advisor_upload else "advisor", + inventory_id, + org_id, + ) + else: + LOGGER.debug( + "starting waiting for %s msg for system: %s, account: %s", + "inventory" if item.advisor_upload else "advisor", + inventory_id, + org_id, + ) + + try: + await asyncio.wait_for(item.second_upload_event.wait(), CFG.grouper_messages_timeout_sec) + PAIR_HIT.inc() + except asyncio.TimeoutError: + LOGGER.info("timing out while waiting for message for system: %s, account: %s", inventory_id, org_id) + PAIR_MISS.inc() + self._queue.pop(inventory_id, None) + + if item.inventory_upload: + self.max_inventory_msgs.release() + if item.advisor_upload: + self.max_advisor_msgs.release() + + await self._send_for_evaluation(item, org_id, inventory_id) + + QUEUE_SIZE.dec() async def _send_for_evaluation(self, item: QueueItem, org_id: str, inventory_id: str) -> None: """Sends message to evaluate a system""" @@ -184,15 +216,20 @@ async def _send_for_evaluation(self, item: QueueItem, org_id: str, inventory_id: if (not item.inventory_changed and not item.advisor_changed) and not CFG.disable_optimisation: UNCHANGED_SYSTEM.inc() LOGGER.info("skipping evaluation, system not changed: %s, org_id: %s", inventory_id, org_id) - send_msg_to_payload_tracker( + tracker_fut = send_msg_to_payload_tracker( self.payload_tracker, msg, "success", status_msg="unchanged system, not sending to evaluator", loop=self.loop ) + if tracker_fut: + await tracker_fut return CHANGED_SYSTEM.inc() LOGGER.info("sending upload message to evaluator: %s, org_id: %s", inventory_id, org_id) - send_msg_to_payload_tracker( + tracker_fut = send_msg_to_payload_tracker( self.payload_tracker, msg, "processing", status_msg="changed system, sending to evaluator", loop=self.loop ) - with use_otel_context(item.otel_context): - self.evaluator.send(msg) + evaluator_fut = self.evaluator.send(msg) + if tracker_fut: + await tracker_fut + if evaluator_fut: + await evaluator_fut diff --git a/listener/advisor_processor.py b/listener/advisor_processor.py index 1759f3615..13db3c3d4 100644 --- a/listener/advisor_processor.py +++ b/listener/advisor_processor.py @@ -185,11 +185,11 @@ def _send_for_evaluation( }, "timestamp": timestamp, } - self.grouper.send(msg, loop=self.loop, key=org_id) + return self.grouper.send(msg, loop=self.loop, key=org_id) def _send_to_payload_tracker(self, status: str, msg: AdvisorMsg, message=None): """Send payload tracker message""" - send_msg_to_payload_tracker(self.payload_tracker, msg.msg["input"], status, status_msg=message, loop=self.loop) + return send_msg_to_payload_tracker(self.payload_tracker, msg.msg["input"], status, status_msg=message, loop=self.loop) async def _process_upload(self, msg: AdvisorMsg): """Process message from advisor""" @@ -221,8 +221,12 @@ async def _process_upload(self, msg: AdvisorMsg): LOGGER.info( "advisor data inserted, system: %s, org_id: %s, reporter: %s, request_id: %s", inventory_id, org_id, reporter, request_id ) - self._send_for_evaluation(org_id, inventory_id, request_id, reporter, timestamp, import_status) - self._send_to_payload_tracker("received", msg, message="system received from advisor, sending to grouper") + eval_fut = self._send_for_evaluation(org_id, inventory_id, request_id, reporter, timestamp, import_status) + tracker_fut = self._send_to_payload_tracker("received", msg, message="system received from advisor, sending to grouper") + if eval_fut: + await eval_fut + if tracker_fut: + await tracker_fut async def process_msg(self, msg: AdvisorMsg): """Process single advisor msg""" diff --git a/listener/inventory_processor.py b/listener/inventory_processor.py index ae6c96072..81b0806cd 100644 --- a/listener/inventory_processor.py +++ b/listener/inventory_processor.py @@ -467,11 +467,11 @@ def _send_for_evaluation( }, "timestamp": timestamp, } - self.grouper.send(msg, loop=self.loop, key=org_id) + return self.grouper.send(msg, loop=self.loop, key=org_id) def _send_to_payload_tracker(self, status: str, msg: InventoryMsg, message=None): """Send payload tracker message""" - send_msg_to_payload_tracker(self.payload_tracker, msg.msg, status, status_msg=message, loop=self.loop) + return send_msg_to_payload_tracker(self.payload_tracker, msg.msg, status, status_msg=message, loop=self.loop) async def _process_upload(self, msg: InventoryMsg): """Process upload message defined by QueueItem""" @@ -547,8 +547,12 @@ async def _process_upload(self, msg: InventoryMsg): LOGGER.info( "Inventory data processed, system: %s, org_id: %s, reporter: %s, request_id: %s", inventory_id, org_id, reporter, request_id ) - self._send_for_evaluation(org_id, inventory_id, request_id, reporter, timestamp, import_status) - self._send_to_payload_tracker("received", msg, message="system received from inventory, sending to grouper") + eval_fut = self._send_for_evaluation(org_id, inventory_id, request_id, reporter, timestamp, import_status) + tracker_fut = self._send_to_payload_tracker("received", msg, message="system received from inventory, sending to grouper") + if eval_fut: + await eval_fut + if tracker_fut: + await tracker_fut async def _process_delete(self, msg: InventoryMsg): """Process inventory delete message""" diff --git a/listener/listener.py b/listener/listener.py index 1cb00b548..dd4789390 100644 --- a/listener/listener.py +++ b/listener/listener.py @@ -7,6 +7,7 @@ import signal from aiokafka.structs import ConsumerRecord +from opentelemetry import context as otel_context_api from opentelemetry import trace from psycopg_pool import AsyncConnectionPool @@ -182,12 +183,32 @@ async def _consume_inventory_msg(self, msg: dict) -> InventoryMsgType: await self.inventory_msg_processor.process_msg(InventoryMsg(msg_type, msg)) return msg_type - async def consume_inventory_msg(self, msg: dict): + async def consume_inventory_msg(self, msg: dict, topic: str = "", parent_ctx: otel_context_api.Context | None = None): """Consume inventory message, semaphore wrapper""" - try: - await self._consume_inventory_msg(msg) - finally: - self.max_messages_semaphore.release() + with use_otel_context(parent_ctx): + with TRACER.start_as_current_span( + f"process {topic}", + kind=trace.SpanKind.CONSUMER, + attributes={ + "messaging.system": "kafka", + "messaging.operation.name": "process", + "messaging.destination.name": topic, + }, + ) as span: + host_data = msg.get("host") if isinstance(msg.get("host"), dict) else {} + org_id = host_data.get("org_id") or msg.get("org_id") + if org_id: + span.set_attribute("rh.org_id", org_id) + threadctx.org_id = org_id + metadata = msg.get("platform_metadata") + request_id = metadata.get("request_id") if isinstance(metadata, dict) else None + if request_id: + span.set_attribute("rh.request_id", request_id) + threadctx.request_id = request_id + try: + await self._consume_inventory_msg(msg) + finally: + self.max_messages_semaphore.release() async def _consume_advisor_msg(self, msg: dict): """Consumes advisor message""" @@ -203,12 +224,32 @@ async def _consume_advisor_msg(self, msg: dict): await self.advisor_msg_processor.process_msg(AdvisorMsg(msg)) - async def consume_advisor_msg(self, msg: dict): + async def consume_advisor_msg(self, msg: dict, topic: str = "", parent_ctx: otel_context_api.Context | None = None): """Consume advisor message, semaphore wrapper""" - try: - await self._consume_advisor_msg(msg) - finally: - self.max_messages_semaphore.release() + with use_otel_context(parent_ctx): + with TRACER.start_as_current_span( + f"process {topic}", + kind=trace.SpanKind.CONSUMER, + attributes={ + "messaging.system": "kafka", + "messaging.operation.name": "process", + "messaging.destination.name": topic, + }, + ) as span: + input_data = msg.get("input") if isinstance(msg.get("input"), dict) else {} + org_id = input_data.get("host", {}).get("org_id") if isinstance(input_data.get("host"), dict) else None + if org_id: + span.set_attribute("rh.org_id", org_id) + threadctx.org_id = org_id + metadata = input_data.get("platform_metadata") + request_id = metadata.get("request_id") if isinstance(metadata, dict) else None + if request_id: + span.set_attribute("rh.request_id", request_id) + threadctx.request_id = request_id + try: + await self._consume_advisor_msg(msg) + finally: + self.max_messages_semaphore.release() def _consume_message(self, msg: ConsumerRecord): """Consumes single message from advisor or inventory""" @@ -222,38 +263,10 @@ def _consume_message(self, msg: ConsumerRecord): parent_ctx = extract_context_from_msg_headers(msg.headers) if msg_dict.get("host") or msg_dict.get("type") == "delete": - with use_otel_context(parent_ctx): - with TRACER.start_as_current_span( - f"process {msg.topic}", - kind=trace.SpanKind.CONSUMER, - attributes={ - "messaging.system": "kafka", - "messaging.operation.name": "process", - "messaging.destination.name": msg.topic, - }, - ) as span: - org_id = msg_dict.get("host", {}).get("org_id") or msg_dict.get("org_id") - if org_id: - span.set_attribute("rh.org_id", org_id) - threadctx.org_id = org_id - request_id = (msg_dict.get("platform_metadata") or {}).get("request_id") - if request_id: - span.set_attribute("rh.request_id", request_id) - threadctx.request_id = request_id - create_task_and_log(self.consume_inventory_msg(msg_dict), LOGGER, self.loop) + create_task_and_log(self.consume_inventory_msg(msg_dict, msg.topic, parent_ctx), LOGGER, self.loop) PROCESS_MESSAGES.inc() elif msg_dict.get("input"): - with use_otel_context(parent_ctx): - with TRACER.start_as_current_span( - f"process {msg.topic}", - kind=trace.SpanKind.CONSUMER, - attributes={ - "messaging.system": "kafka", - "messaging.operation.name": "process", - "messaging.destination.name": msg.topic, - }, - ): - create_task_and_log(self.consume_advisor_msg(msg_dict), LOGGER, self.loop) + create_task_and_log(self.consume_advisor_msg(msg_dict, msg.topic, parent_ctx), LOGGER, self.loop) PROCESS_MESSAGES.inc() else: LOGGER.exception("Unknown message obtained: %s", msg) diff --git a/tests/grouper_tests/__init__.py b/tests/grouper_tests/__init__.py new file mode 100644 index 000000000..40a96afc6 --- /dev/null +++ b/tests/grouper_tests/__init__.py @@ -0,0 +1 @@ +# -*- coding: utf-8 -*- diff --git a/tests/grouper_tests/test_grouper_tracing.py b/tests/grouper_tests/test_grouper_tracing.py new file mode 100644 index 000000000..ff7308dc5 --- /dev/null +++ b/tests/grouper_tests/test_grouper_tracing.py @@ -0,0 +1,125 @@ +# -*- coding: utf-8 -*- +# pylint: disable=no-self-use,redefined-outer-name +""" +Tests asserting that the OpenTelemetry consumer span in grouper outlives its children. +""" + +import asyncio +from unittest.mock import AsyncMock +from unittest.mock import patch + +import pytest +from opentelemetry import trace +from opentelemetry.sdk.trace import TracerProvider +from opentelemetry.sdk.trace.export import SimpleSpanProcessor +from opentelemetry.sdk.trace.export.in_memory_span_exporter import InMemorySpanExporter + +from common.telemetry import threadctx +from common.utils import create_task_and_log +from grouper.queue import TRACER +from grouper.queue import GrouperQueue + + +@pytest.fixture +def otel_exporter(): + """Sets up an in-memory span exporter with a fresh TracerProvider.""" + from common import mqueue as mq + from grouper import queue as gq + from listener import listener as ll + + provider = TracerProvider() + exporter = InMemorySpanExporter() + provider.add_span_processor(SimpleSpanProcessor(exporter)) + + orig_provider = trace._TRACER_PROVIDER + trace._TRACER_PROVIDER = provider + ll.TRACER._real_tracer = None + mq.TRACER._real_tracer = None + gq.TRACER._real_tracer = None + + yield exporter + + trace._TRACER_PROVIDER = orig_provider + ll.TRACER._real_tracer = None + mq.TRACER._real_tracer = None + gq.TRACER._real_tracer = None + + +def build_mock_grouper_queue(loop): + """Creates a GrouperQueue with mocked MQWriter.""" + with patch("grouper.queue.MQWriter"), patch("common.mqueue.AIOKafkaProducer"): + return GrouperQueue(loop) + + +@pytest.mark.asyncio(loop_scope="function") +async def test_grouper_consumer_span_outlives_children(otel_exporter): + """Assert that the consumer span wraps item processing, remains active during it, and outlives children.""" + loop = asyncio.get_running_loop() + queue = build_mock_grouper_queue(loop) + + active_checks = {} + captured_threadctx = {} + + async def fake_send(item, org_id, inventory_id): + curr = trace.get_current_span() + active_checks["is_recording"] = curr.is_recording() + active_checks["span_name"] = curr.name + active_checks["finished_before_child"] = len(otel_exporter.get_finished_spans()) + captured_threadctx["org_id"] = threadctx.org_id + captured_threadctx["request_id"] = threadctx.request_id + with TRACER.start_as_current_span("send.to.evaluator"): + pass + active_checks["finished_after_child"] = len(otel_exporter.get_finished_spans()) + + queue._send_for_evaluation = AsyncMock(side_effect=fake_send) + + task = None + orig_create_task_and_log = create_task_and_log + + def spy_create_task_and_log(coroutine, logger, loop): + nonlocal task + task = orig_create_task_and_log(coroutine, logger, loop) + return task + + with patch("grouper.queue.create_task_and_log", side_effect=spy_create_task_and_log): + await queue.max_inventory_msgs.acquire() + queue.push_inventory_msg( + org_id="org-123", + inventory_id="sys-1", + reporter="rhsm-system-profile-bridge", + inventory_changed=True, + request_id="req-789", + topic="platform.vulnerability.inventory-upload", + ) + + assert task is not None + await asyncio.wait_for(task, timeout=5.0) + + assert active_checks["is_recording"] is True + assert active_checks["span_name"] == "process platform.vulnerability.inventory-upload" + assert active_checks["finished_before_child"] == 0 + assert active_checks["finished_after_child"] == 1 + + spans = otel_exporter.get_finished_spans() + assert len(spans) == 2, f"Expected 2 spans, got {len(spans)}" + + child_span = next(s for s in spans if s.name == "send.to.evaluator") + parent_span = next(s for s in spans if s.name == "process platform.vulnerability.inventory-upload") + + # Child must be nested under the consumer span + assert child_span.parent.span_id == parent_span.context.span_id + # Parent span must start before child starts + assert parent_span.start_time <= child_span.start_time + # Parent span must outlive the child span + assert parent_span.end_time >= child_span.end_time + + # Span attributes + assert parent_span.attributes.get("rh.org_id") == "org-123" + assert parent_span.attributes.get("rh.request_id") == "req-789" + assert parent_span.attributes.get("vulnerability.inventory_id") == "sys-1" + assert parent_span.attributes.get("vulnerability.msg_type") == "inventory_upload" + assert parent_span.attributes.get("messaging.system") == "kafka" + assert parent_span.attributes.get("messaging.destination.name") == "platform.vulnerability.inventory-upload" + + assert captured_threadctx.get("org_id") == "org-123" + assert captured_threadctx.get("request_id") == "req-789" diff --git a/tests/listener_tests/test_listener_tracing.py b/tests/listener_tests/test_listener_tracing.py new file mode 100644 index 000000000..c117d46f9 --- /dev/null +++ b/tests/listener_tests/test_listener_tracing.py @@ -0,0 +1,168 @@ +# -*- coding: utf-8 -*- +# pylint: disable=no-self-use,redefined-outer-name +""" +Tests asserting that the OpenTelemetry consumer span in listener outlives its children. +""" + +import asyncio +from unittest.mock import AsyncMock +from unittest.mock import MagicMock +from unittest.mock import patch + +import pytest +from opentelemetry import trace +from opentelemetry.sdk.trace import TracerProvider +from opentelemetry.sdk.trace.export import SimpleSpanProcessor +from opentelemetry.sdk.trace.export.in_memory_span_exporter import InMemorySpanExporter + +from common.telemetry import threadctx +from listener.listener import TRACER +from listener.listener import Listener + + +@pytest.fixture +def otel_exporter(): + """Sets up an in-memory span exporter with a fresh TracerProvider.""" + from common import mqueue as mq + from grouper import queue as gq + from listener import listener as ll + + provider = TracerProvider() + exporter = InMemorySpanExporter() + provider.add_span_processor(SimpleSpanProcessor(exporter)) + + orig_provider = trace._TRACER_PROVIDER + trace._TRACER_PROVIDER = provider + ll.TRACER._real_tracer = None + mq.TRACER._real_tracer = None + gq.TRACER._real_tracer = None + + yield exporter + + trace._TRACER_PROVIDER = orig_provider + ll.TRACER._real_tracer = None + mq.TRACER._real_tracer = None + gq.TRACER._real_tracer = None + + +def build_mock_listener(loop): + """Creates a Listener instance with mocked MQ and DB dependencies.""" + with ( + patch("listener.listener.MQReader"), + patch("listener.listener.MQWriter"), + patch("listener.inventory_processor.MQWriter"), + patch("listener.advisor_processor.MQWriter"), + ): + return Listener(MagicMock(), loop) + + +@pytest.mark.asyncio(loop_scope="function") +async def test_inventory_consumer_span_outlives_children(otel_exporter): + """Assert that the consumer span wraps processing, remains active during it, and outlives children.""" + loop = asyncio.get_running_loop() + listener = build_mock_listener(loop) + + active_checks = {} + + async def fake_consume(msg): + curr = trace.get_current_span() + active_checks["is_recording"] = curr.is_recording() + active_checks["span_name"] = curr.name + active_checks["finished_before_child"] = len(otel_exporter.get_finished_spans()) + with TRACER.start_as_current_span("db.select"): + await asyncio.sleep(0.01) + active_checks["finished_after_child"] = len(otel_exporter.get_finished_spans()) + + listener._consume_inventory_msg = AsyncMock(side_effect=fake_consume) + + msg = { + "host": {"id": "sys-1", "org_id": "12345"}, + "platform_metadata": {"request_id": "req-999"}, + } + + await listener.max_messages_semaphore.acquire() + await listener.consume_inventory_msg(msg, topic="platform.inventory.events") + + # While processing was running, consumer span was active and not yet finished + assert active_checks["is_recording"] is True + assert active_checks["span_name"] == "process platform.inventory.events" + assert active_checks["finished_before_child"] == 0 + assert active_checks["finished_after_child"] == 1 + + spans = otel_exporter.get_finished_spans() + assert len(spans) == 2, f"Expected 2 spans, got {len(spans)}" + + child_span = next(s for s in spans if s.name == "db.select") + parent_span = next(s for s in spans if s.name == "process platform.inventory.events") + + # Child must be nested under the consumer span + assert child_span.parent.span_id == parent_span.context.span_id + # Parent span must start before or when child starts + assert parent_span.start_time <= child_span.start_time + # Parent span must outlive the child span + assert parent_span.end_time >= child_span.end_time + + # Span attributes + assert parent_span.attributes.get("rh.org_id") == "12345" + assert parent_span.attributes.get("rh.request_id") == "req-999" + assert parent_span.attributes.get("messaging.system") == "kafka" + assert parent_span.attributes.get("messaging.operation.name") == "process" + assert parent_span.attributes.get("messaging.destination.name") == "platform.inventory.events" + + # ContextVars threadctx should have been set + assert threadctx.org_id == "12345" + assert threadctx.request_id == "req-999" + + +@pytest.mark.asyncio(loop_scope="function") +async def test_advisor_consumer_span_outlives_children(otel_exporter): + """Assert that the consumer span wraps processing, remains active during it, and outlives children.""" + loop = asyncio.get_running_loop() + listener = build_mock_listener(loop) + + active_checks = {} + + async def fake_consume(msg): + curr = trace.get_current_span() + active_checks["is_recording"] = curr.is_recording() + active_checks["span_name"] = curr.name + active_checks["finished_before_child"] = len(otel_exporter.get_finished_spans()) + with TRACER.start_as_current_span("db.insert"): + await asyncio.sleep(0.01) + active_checks["finished_after_child"] = len(otel_exporter.get_finished_spans()) + + listener._consume_advisor_msg = AsyncMock(side_effect=fake_consume) + + msg = { + "input": { + "host": {"id": "sys-2", "org_id": "67890"}, + "platform_metadata": {"request_id": "req-888"}, + } + } + + await listener.max_messages_semaphore.acquire() + await listener.consume_advisor_msg(msg, topic="platform.advisor.results") + + # While processing was running, consumer span was active and not yet finished + assert active_checks["is_recording"] is True + assert active_checks["span_name"] == "process platform.advisor.results" + assert active_checks["finished_before_child"] == 0 + assert active_checks["finished_after_child"] == 1 + + spans = otel_exporter.get_finished_spans() + assert len(spans) == 2, f"Expected 2 spans, got {len(spans)}" + + child_span = next(s for s in spans if s.name == "db.insert") + parent_span = next(s for s in spans if s.name == "process platform.advisor.results") + + assert child_span.parent.span_id == parent_span.context.span_id + assert parent_span.start_time <= child_span.start_time + assert parent_span.end_time >= child_span.end_time + + assert parent_span.attributes.get("rh.org_id") == "67890" + assert parent_span.attributes.get("rh.request_id") == "req-888" + assert parent_span.attributes.get("messaging.system") == "kafka" + assert parent_span.attributes.get("messaging.destination.name") == "platform.advisor.results" + + assert threadctx.org_id == "67890" + assert threadctx.request_id == "req-888"