import json import unittest from typing import Any, Optional, Tuple import httpx from pwo import async_test from kaya.core import HttpContext, KayaApp from kaya.core._asgi import AsgiWebSocket from kaya.otel import OTelMixin from opentelemetry.sdk.metrics import MeterProvider from opentelemetry.sdk.metrics.export import InMemoryMetricReader from opentelemetry.sdk.resources import Resource 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 opentelemetry.trace import StatusCode TRACEPARENT = '00-4bf92f3577b34da6a3ce929d0e0e4736-00f067aa0ba902b7-01' def make_mixin(**kwargs) -> Tuple[OTelMixin, InMemorySpanExporter, InMemoryMetricReader]: resource = Resource.create({'service.name': 'test-service'}) span_exporter = InMemorySpanExporter() tracer_provider = TracerProvider(resource=resource) tracer_provider.add_span_processor(SimpleSpanProcessor(span_exporter)) metric_reader = InMemoryMetricReader() meter_provider = MeterProvider(resource=resource, metric_readers=[metric_reader]) mixin = OTelMixin( service_name='test-service', tracer_provider=tracer_provider, meter_provider=meter_provider, **kwargs, ) return mixin, span_exporter, metric_reader def make_app(mixin: OTelMixin) -> KayaApp: app = KayaApp(mixins=[mixin]) @app.GET('/hello') async def hello(ctx: HttpContext) -> None: await ctx.send_str(200, json.dumps({'ok': True})) @app.GET('/boom') async def boom(ctx: HttpContext) -> None: await ctx.send_str(500, 'boom') return app async def request(app: KayaApp, path: str = '/hello', headers: Optional[dict[str, str]] = None) -> httpx.Response: transport = httpx.ASGITransport(app=app, client=('127.0.0.1', 123)) async with httpx.AsyncClient(transport=transport, base_url="http://127.0.0.1:80") as http_client: return await http_client.get(path, headers=headers) class HttpTracingTest(unittest.TestCase): @async_test async def test_request_produces_server_span(self): mixin, span_exporter, _ = make_mixin() app = make_app(mixin) r = await request(app) self.assertEqual(200, r.status_code) spans = span_exporter.get_finished_spans() self.assertEqual(1, len(spans)) span = spans[0] self.assertEqual('GET /hello', span.name) assert span.attributes is not None self.assertEqual('GET', span.attributes['http.request.method']) self.assertEqual('/hello', span.attributes['url.path']) self.assertEqual(200, span.attributes['http.response.status_code']) self.assertEqual(StatusCode.UNSET, span.status.status_code) @async_test async def test_5xx_marks_span_as_error(self): mixin, span_exporter, _ = make_mixin() app = make_app(mixin) r = await request(app, '/boom') self.assertEqual(500, r.status_code) spans = span_exporter.get_finished_spans() self.assertEqual(1, len(spans)) self.assertEqual(500, spans[0].attributes['http.response.status_code']) # type: ignore[index] self.assertEqual(StatusCode.ERROR, spans[0].status.status_code) @async_test async def test_traceparent_header_propagates(self): mixin, span_exporter, _ = make_mixin() app = make_app(mixin) await request(app, headers={'traceparent': TRACEPARENT}) spans = span_exporter.get_finished_spans() self.assertEqual(1, len(spans)) span = spans[0] self.assertEqual(0x4bf92f3577b34da6a3ce929d0e0e4736, span.context.trace_id) # type: ignore[union-attr] assert span.parent is not None self.assertEqual(0x00f067aa0ba902b7, span.parent.span_id) @async_test async def test_unmatched_route_still_traced(self): mixin, span_exporter, _ = make_mixin() app = make_app(mixin) r = await request(app, '/nowhere') self.assertEqual(404, r.status_code) spans = span_exporter.get_finished_spans() self.assertEqual(1, len(spans)) self.assertEqual(404, spans[0].attributes['http.response.status_code']) # type: ignore[index] @async_test async def test_no_leaked_spans_after_requests(self): mixin, span_exporter, _ = make_mixin() app = make_app(mixin) await request(app) await request(app, '/boom') self.assertEqual(0, len(mixin._spans)) self.assertEqual(2, len(span_exporter.get_finished_spans())) @async_test async def test_exception_marks_span_as_error(self): mixin, span_exporter, _ = make_mixin() app = KayaApp(mixins=[mixin]) @app.GET('/raises') async def raises(ctx: HttpContext) -> None: raise RuntimeError('boom') transport = httpx.ASGITransport(app=app, client=('127.0.0.1', 123)) async with httpx.AsyncClient(transport=transport, base_url="http://127.0.0.1:80") as http_client: with self.assertRaises(RuntimeError): await http_client.get('/raises') spans = span_exporter.get_finished_spans() self.assertEqual(1, len(spans)) span = spans[0] self.assertEqual(StatusCode.ERROR, span.status.status_code) self.assertTrue(any(event.name == 'exception' for event in span.events)) self.assertEqual(0, len(mixin._spans)) @async_test async def test_excluded_path_skips_tracing_and_metrics(self): mixin, span_exporter, metric_reader = make_mixin(excluded_paths=('/health',)) app = KayaApp(mixins=[mixin]) @app.GET('/health') async def health(ctx: HttpContext) -> None: await ctx.send_str(200, 'ok') @app.GET('/hello') async def hello(ctx: HttpContext) -> None: await ctx.send_str(200, 'hello') await request(app, '/health') await request(app, '/hello') spans = span_exporter.get_finished_spans() self.assertEqual(1, len(spans)) self.assertEqual('GET /hello', spans[0].name) metrics = MetricsTest._metric_names(metric_reader) duration_points = list(metrics['http.server.request.duration'].data.data_points) self.assertEqual(1, len(duration_points)) self.assertEqual(1, duration_points[0].count) @async_test async def test_route_template_used_for_span_name_and_metrics(self): mixin, span_exporter, metric_reader = make_mixin() app = KayaApp(mixins=[mixin]) @app.GET('/items/${item_id:int}') async def item(ctx: HttpContext, item_id: int) -> None: await ctx.send_str(200, str(item_id)) r = await request(app, '/items/123') self.assertEqual(200, r.status_code) spans = span_exporter.get_finished_spans() self.assertEqual(1, len(spans)) span = spans[0] self.assertEqual('GET /items/${item_id:int}', span.name) assert span.attributes is not None self.assertEqual('/items/${item_id:int}', span.attributes['http.route']) self.assertEqual('/items/123', span.attributes['url.path']) metrics = MetricsTest._metric_names(metric_reader) duration_points = list(metrics['http.server.request.duration'].data.data_points) self.assertEqual(1, len(duration_points)) duration_attributes = dict(duration_points[0].attributes) self.assertEqual('GET', duration_attributes['http.request.method']) self.assertEqual('/items/${item_id:int}', duration_attributes['http.route']) self.assertEqual(200, duration_attributes['http.response.status_code']) self.assertNotIn('url.path', duration_attributes) active_points = list(metrics['http.server.active_requests'].data.data_points) self.assertEqual(1, len(active_points)) self.assertEqual({'http.request.method': 'GET'}, dict(active_points[0].attributes)) @async_test async def test_header_capture_and_sanitization(self): mixin, span_exporter, _ = make_mixin( capture_request_headers=('X-Request-Id', 'Authorization'), capture_response_headers=('X-Response-Id', 'Set-Cookie'), ) app = KayaApp(mixins=[mixin]) @app.GET('/headers') async def headers(ctx: HttpContext) -> None: await ctx.send_str(200, 'ok', { 'X-Response-Id': 'res-1', 'Set-Cookie': 'sid=secret', }) r = await request(app, '/headers', headers={ 'X-Request-Id': 'req-1', 'Authorization': 'Bearer secret', }) self.assertEqual(200, r.status_code) spans = span_exporter.get_finished_spans() self.assertEqual(1, len(spans)) attributes = spans[0].attributes assert attributes is not None self.assertEqual(['req-1'], list(attributes['http.request.header.x_request_id'])) self.assertEqual(['REDACTED'], list(attributes['http.request.header.authorization'])) self.assertEqual(['res-1'], list(attributes['http.response.header.x_response_id'])) self.assertEqual(['REDACTED'], list(attributes['http.response.header.set_cookie'])) @async_test async def test_http_hooks_are_called(self): calls = [] def request_hook(span: Any, ctx: HttpContext) -> None: calls.append(('request', ctx.path)) span.set_attribute('test.request_hook', True) def response_hook(span: Any, ctx: HttpContext) -> None: calls.append(('response', ctx.path)) span.set_attribute('test.response_hook', True) mixin, span_exporter, _ = make_mixin( server_request_hook=request_hook, server_response_hook=response_hook, ) app = make_app(mixin) await request(app) self.assertEqual([('request', '/hello'), ('response', '/hello')], calls) attributes = span_exporter.get_finished_spans()[0].attributes assert attributes is not None self.assertTrue(attributes['test.request_hook']) self.assertTrue(attributes['test.response_hook']) class MetricsTest(unittest.TestCase): @staticmethod def _metric_names(metric_reader: InMemoryMetricReader) -> dict: data = metric_reader.get_metrics_data() return {m.name: m for rm in data.resource_metrics for m in rm.scope_metrics[0].metrics if rm.scope_metrics} @async_test async def test_duration_histogram_and_active_requests(self): mixin, _, metric_reader = make_mixin() app = make_app(mixin) await request(app) metrics = self._metric_names(metric_reader) self.assertIn('http.server.request.duration', metrics) self.assertIn('http.server.active_requests', metrics) duration_points = list(metrics['http.server.request.duration'].data.data_points) self.assertEqual(1, len(duration_points)) self.assertEqual(1, duration_points[0].count) self.assertGreaterEqual(duration_points[0].sum, 0) active_points = list(metrics['http.server.active_requests'].data.data_points) self.assertEqual(1, len(active_points)) self.assertEqual(0, active_points[0].value) class WebSocketTracingTest(unittest.TestCase): @staticmethod def _make_ws(headers=()): async def send(message): pass async def receive(): return {'type': 'websocket.connect'} scope = { 'type': 'websocket', 'path': '/ws/games/abc', 'query_string': b'', 'scheme': 'ws', 'client': ('127.0.0.1', 12345), 'server': ('127.0.0.1', 80), 'headers': headers, } return AsgiWebSocket(scope, receive, send) @async_test async def test_websocket_lifecycle_produces_span(self): mixin, span_exporter, _ = make_mixin() ws = self._make_ws() wrapped = await mixin._before_websocket(ws) self.assertIsNotNone(wrapped) self.assertEqual(1, len(mixin._spans)) await mixin._after_websocket(wrapped) # type: ignore[arg-type] self.assertEqual(0, len(mixin._spans)) spans = span_exporter.get_finished_spans() self.assertEqual(1, len(spans)) self.assertEqual('WS /ws/games/abc', spans[0].name) @async_test async def test_websocket_traceparent_propagates(self): mixin, span_exporter, _ = make_mixin() ws = self._make_ws(headers=[(b'traceparent', TRACEPARENT.encode())]) wrapped = await mixin._before_websocket(ws) await mixin._after_websocket(wrapped) # type: ignore[arg-type] spans = span_exporter.get_finished_spans() self.assertEqual(1, len(spans)) self.assertEqual(0x4bf92f3577b34da6a3ce929d0e0e4736, spans[0].context.trace_id) # type: ignore[union-attr] @async_test async def test_wrapped_websocket_delegates(self): mixin, _, _ = make_mixin() ws = self._make_ws() wrapped = await mixin._before_websocket(ws) assert wrapped is not None self.assertEqual('/ws/games/abc', wrapped.path) self.assertEqual(('127.0.0.1', 12345), wrapped.client) await mixin._after_websocket(wrapped) @async_test async def test_websocket_error_close_code_marks_span_error(self): mixin, span_exporter, _ = make_mixin() ws = self._make_ws() wrapped = await mixin._before_websocket(ws) assert wrapped is not None await wrapped.close(1011) await mixin._after_websocket(wrapped) spans = span_exporter.get_finished_spans() self.assertEqual(1, len(spans)) self.assertEqual(StatusCode.ERROR, spans[0].status.status_code) assert spans[0].attributes is not None self.assertEqual(1011, spans[0].attributes['kaya.websocket.close_code']) @async_test async def test_websocket_exception_marks_span_error(self): mixin, span_exporter, _ = make_mixin() ws = self._make_ws() wrapped = await mixin._before_websocket(ws) assert wrapped is not None wrapped.exception = RuntimeError('ws boom') await mixin._after_websocket(wrapped) spans = span_exporter.get_finished_spans() self.assertEqual(1, len(spans)) self.assertEqual(StatusCode.ERROR, spans[0].status.status_code) self.assertTrue(any(event.name == 'exception' for event in spans[0].events)) @async_test async def test_websocket_hooks_are_called(self): calls = [] def connect_hook(span: Any, ws: Any) -> None: calls.append(('connect', ws.path)) span.set_attribute('test.websocket_connect_hook', True) def close_hook(span: Any, ws: Any) -> None: calls.append(('close', ws.path)) span.set_attribute('test.websocket_close_hook', True) mixin, span_exporter, _ = make_mixin( websocket_connect_hook=connect_hook, websocket_close_hook=close_hook, ) ws = self._make_ws() wrapped = await mixin._before_websocket(ws) assert wrapped is not None await wrapped.close(1000) await mixin._after_websocket(wrapped) self.assertEqual([('connect', '/ws/games/abc'), ('close', '/ws/games/abc')], calls) spans = span_exporter.get_finished_spans() self.assertEqual(1, len(spans)) self.assertEqual(StatusCode.UNSET, spans[0].status.status_code) assert spans[0].attributes is not None self.assertTrue(spans[0].attributes['test.websocket_connect_hook']) self.assertTrue(spans[0].attributes['test.websocket_close_hook']) if __name__ == '__main__': unittest.main()