diff --git a/microbootstrap/bootstrappers/faststream.py b/microbootstrap/bootstrappers/faststream.py index 97d0b37..ecef4c9 100644 --- a/microbootstrap/bootstrappers/faststream.py +++ b/microbootstrap/bootstrappers/faststream.py @@ -24,6 +24,7 @@ from microbootstrap.instruments.prometheus_instrument import FastStreamPrometheusConfig, PrometheusInstrument from microbootstrap.instruments.pyroscope_instrument import PyroscopeInstrument from microbootstrap.instruments.sentry_instrument import SentryInstrument +from microbootstrap.middlewares.faststream import FastStreamOpenTelemetryBaggageMiddleware from microbootstrap.settings import FastStreamSettings @@ -101,6 +102,11 @@ def bootstrap_after(self, application: AsgiFastStream) -> AsgiFastStream: # typ application.broker.add_middleware( self.instrument_config.opentelemetry_middleware_cls(tracer_provider=self.tracer_provider), ) + application.broker.add_middleware( + FastStreamOpenTelemetryBaggageMiddleware( + baggage_span_attributes=self.instrument_config.opentelemetry_baggage_span_attributes, + ), + ) return application @classmethod diff --git a/microbootstrap/middlewares/faststream.py b/microbootstrap/middlewares/faststream.py new file mode 100644 index 0000000..a18680d --- /dev/null +++ b/microbootstrap/middlewares/faststream.py @@ -0,0 +1,66 @@ +from __future__ import annotations +import typing + +from faststream._internal.middlewares import BaseMiddleware +from opentelemetry import baggage, propagate + +from microbootstrap.instruments.opentelemetry_instrument import opentelemetry_baggage_scope + + +if typing.TYPE_CHECKING: + from faststream._internal.basic_types import AsyncFuncAny + from faststream._internal.context import ContextRepo + from faststream.message import StreamMessage + + +class FastStreamOpenTelemetryBaggageMiddleware: + def __init__(self, *, baggage_span_attributes: typing.Mapping[str, str]) -> None: + self.baggage_span_attributes = baggage_span_attributes + + def __call__( + self, + msg: typing.Any, # noqa: ANN401 + /, + *, + context: ContextRepo, + ) -> _FastStreamOpenTelemetryBaggageMiddleware: + return _FastStreamOpenTelemetryBaggageMiddleware( + msg, + context=context, + baggage_span_attributes=self.baggage_span_attributes, + ) + + +class _FastStreamOpenTelemetryBaggageMiddleware(BaseMiddleware[typing.Any, typing.Any]): + def __init__( + self, + msg: typing.Any, # noqa: ANN401 + /, + *, + context: ContextRepo, + baggage_span_attributes: typing.Mapping[str, str], + ) -> None: + super().__init__(msg, context=context) + self.baggage_span_attributes = baggage_span_attributes + + async def consume_scope( + self, + call_next: AsyncFuncAny, + msg: StreamMessage[typing.Any], + ) -> typing.Any: # noqa: ANN401 + extracted_baggage: typing.Final = baggage.get_all(propagate.extract(msg.headers)) + with opentelemetry_baggage_scope( + extracted_baggage, + current_span_attributes=self.baggage_span_attributes, + ): + return await call_next(msg) + + async def publish_scope( + self, + call_next: typing.Callable[[typing.Any], typing.Awaitable[typing.Any]], + cmd: typing.Any, # noqa: ANN401 + ) -> typing.Any: # noqa: ANN401 + for field in propagate.get_global_textmap().fields: + cmd.headers.pop(field, None) + propagate.inject(cmd.headers) + return await call_next(cmd) diff --git a/tests/bootstrappers/test_faststream.py b/tests/bootstrappers/test_faststream.py index a3d5e46..35f14fe 100644 --- a/tests/bootstrappers/test_faststream.py +++ b/tests/bootstrappers/test_faststream.py @@ -12,6 +12,7 @@ from faststream.redis import RedisBroker, TestRedisBroker from faststream.redis.opentelemetry import RedisTelemetryMiddleware from faststream.redis.prometheus import RedisPrometheusMiddleware +from opentelemetry import baggage, trace from microbootstrap import opentelemetry_baggage_scope from microbootstrap.bootstrappers.faststream import FastStreamBootstrapper @@ -103,26 +104,57 @@ async def test_ok(self, broker: RedisBroker) -> None: assert response.status_code == status.HTTP_200_OK +@pytest.mark.parametrize("conversation_id", ["authoritative-value", None]) async def test_faststream_opentelemetry( monkeypatch: pytest.MonkeyPatch, faker: faker.Faker, broker: RedisBroker, minimal_opentelemetry_config: OpentelemetryConfig, + conversation_id: str | None, ) -> None: monkeypatch.setattr("opentelemetry.sdk.trace.TracerProvider.shutdown", mock.Mock()) + input_channel: typing.Final = faker.pystr() + output_channel: typing.Final = faker.pystr() + conversation_id_span_attribute: typing.Final = "conversation.id" + observed_context: list[tuple[object | None, object | None, object | None]] = [] + minimal_opentelemetry_config.opentelemetry_baggage_span_attributes = { + "conversation_id": conversation_id_span_attribute + } + + @broker.subscriber(input_channel) + async def handler(_: str) -> None: + with opentelemetry_baggage_scope( + {"conversation_id": conversation_id}, + current_span_attributes={"conversation_id": conversation_id_span_attribute}, + ): + await broker.publish(faker.pystr(), channel=output_channel) + + @broker.subscriber(output_channel) + async def capture_context(_: str) -> None: + current_span: typing.Final = trace.get_current_span() + observed_context.append( + ( + baggage.get_baggage("conversation_id"), + baggage.get_baggage("existing_key"), + current_span.attributes.get(conversation_id_span_attribute), # type: ignore[attr-defined] + ) + ) FastStreamBootstrapper(FastStreamSettings()).configure_application( FastStreamConfig(broker=broker) ).configure_instruments( FastStreamOpentelemetryConfig( - opentelemetry_middleware_cls=RedisTelemetryMiddleware, **minimal_opentelemetry_config.model_dump() + opentelemetry_middleware_cls=RedisTelemetryMiddleware, + **minimal_opentelemetry_config.model_dump(), ) ).bootstrap() async with TestRedisBroker(broker): - with mock.patch("opentelemetry.trace.use_span") as mock_capture_event: - await broker.publish(faker.pystr(), channel=faker.pystr()) - assert mock_capture_event.called + with opentelemetry_baggage_scope({"conversation_id": "stale-value", "existing_key": "existing-value"}): + await broker.publish(faker.pystr(), channel=input_channel) + + assert observed_context == [(conversation_id, "existing-value", conversation_id)] + assert baggage.get_baggage("conversation_id") is None async def test_faststream_logging(broker: RedisBroker, minimal_logging_config: LoggingConfig) -> None: