From abb2667922f59aef63ef9294ff7b338e5b7d6e1f Mon Sep 17 00:00:00 2001 From: Anthony Volk <14987227+anth-volk@users.noreply.github.com> Date: Thu, 17 Sep 2026 20:11:24 +0400 Subject: [PATCH] Refactor observability runtime by responsibility --- changelog.d/17.changed.md | 1 + .../engineering/skills/repository-guidance.md | 21 +- policyengine_observability/_metrics.py | 188 ++ policyengine_observability/_operations.py | 428 ++++ policyengine_observability/_requests.py | 291 +++ policyengine_observability/_state.py | 65 + policyengine_observability/_tracing.py | 434 ++++ policyengine_observability/logging.py | 363 +++ policyengine_observability/runtime.py | 2011 ++------------- policyengine_observability/segments.py | 401 ++- tests/conftest.py | 20 + tests/runtime_helpers.py | 198 ++ tests/test_google_destination.py | 334 +++ tests/test_runtime.py | 2217 ----------------- tests/test_runtime_emission.py | 425 ++++ tests/test_runtime_operations.py | 295 +++ tests/test_runtime_request_failures.py | 342 +++ tests/test_runtime_requests.py | 331 +++ tests/test_runtime_segments.py | 303 +++ tests/test_runtime_tracing.py | 420 ++++ ...nations.py => test_stdout_destinations.py} | 328 --- 21 files changed, 5062 insertions(+), 4354 deletions(-) create mode 100644 changelog.d/17.changed.md create mode 100644 policyengine_observability/_metrics.py create mode 100644 policyengine_observability/_operations.py create mode 100644 policyengine_observability/_requests.py create mode 100644 policyengine_observability/_state.py create mode 100644 policyengine_observability/_tracing.py create mode 100644 tests/runtime_helpers.py create mode 100644 tests/test_google_destination.py delete mode 100644 tests/test_runtime.py create mode 100644 tests/test_runtime_emission.py create mode 100644 tests/test_runtime_operations.py create mode 100644 tests/test_runtime_request_failures.py create mode 100644 tests/test_runtime_requests.py create mode 100644 tests/test_runtime_segments.py create mode 100644 tests/test_runtime_tracing.py rename tests/{test_destinations.py => test_stdout_destinations.py} (55%) diff --git a/changelog.d/17.changed.md b/changelog.d/17.changed.md new file mode 100644 index 0000000..e83c040 --- /dev/null +++ b/changelog.d/17.changed.md @@ -0,0 +1 @@ +Organize the observability runtime by request, operation, segment, logging, metrics, and tracing responsibilities without changing its public API. diff --git a/docs/engineering/skills/repository-guidance.md b/docs/engineering/skills/repository-guidance.md index ca44734..872d237 100644 --- a/docs/engineering/skills/repository-guidance.md +++ b/docs/engineering/skills/repository-guidance.md @@ -27,8 +27,17 @@ uv run --extra dev towncrier check --compare-with origin/main configuration. - `policyengine_observability/context.py` defines request and operation log payload structures. -- `policyengine_observability/runtime.py` owns context management, segments, - structured logs, metrics, traces, events, and fail-open behavior. +- `policyengine_observability/runtime.py` preserves the public runtime API, + configures the components, and coordinates their shutdown. +- `policyengine_observability/_state.py` owns shared context variables. +- `policyengine_observability/_operations.py` and `_requests.py` manage + operation and request lifecycles, respectively. +- `policyengine_observability/segments.py` manages segment naming, nesting, + and timing. +- `policyengine_observability/logging.py` emits structured logs and records + observability failures without interrupting application operations. +- `policyengine_observability/_metrics.py` and `_tracing.py` record metrics + and manage OpenTelemetry traces, respectively. - `policyengine_observability/adapters/` contains framework adapters such as Flask and FastAPI. - `policyengine_observability/integrations/` contains optional integrations @@ -57,9 +66,11 @@ uv run --extra dev towncrier check --compare-with origin/main ## Testing -Add focused tests for runtime context behavior and failure paths whenever -changing `runtime.py`. Adapter changes should include framework-level tests that -exercise request setup, response headers, error paths, and teardown behavior. +Add focused tests for context behavior and failure paths whenever changing +the runtime or its components. The corresponding `tests/test_runtime_*.py` +modules cover operations, requests, segments, log emission, and tracing. +Adapter changes should include framework-level tests that exercise request +setup, response headers, error paths, and teardown behavior. Release automation changes should include tests for the helper scripts when the logic is non-trivial. diff --git a/policyengine_observability/_metrics.py b/policyengine_observability/_metrics.py new file mode 100644 index 0000000..ac2d420 --- /dev/null +++ b/policyengine_observability/_metrics.py @@ -0,0 +1,188 @@ +"""Metric instrument creation and recording.""" + +from __future__ import annotations + +from typing import TYPE_CHECKING + +if TYPE_CHECKING: + from .runtime import ObservabilityRuntime + + +class _NoOpInstrument: + def add(self, *_args, **_kwargs) -> None: + return None + + def record(self, *_args, **_kwargs) -> None: + return None + + +class MetricRecorder: + def __init__(self, runtime: ObservabilityRuntime) -> None: + self.runtime = runtime + + def record_operation_metric( + self, + duration_seconds: float, + attributes: dict[str, str], + ) -> None: + try: + self.runtime.operation_duration.record( + duration_seconds, attributes + ) + self.runtime.operations.add(1, attributes) + except BaseException as exc: + self.runtime.log_observability_failure( + "metrics.record_operation", exc + ) + + def record_request_metric( + self, + duration_seconds: float, + attributes: dict[str, str], + ) -> None: + try: + self.runtime.http_duration.record(duration_seconds, attributes) + self.runtime.requests.add(1, attributes) + except BaseException as exc: + self.runtime.log_observability_failure( + "metrics.record_request", exc + ) + + def record_segment_metric( + self, + segment: str, + duration_seconds: float, + attributes: dict[str, str], + *, + backend_segment: bool = False, + ) -> None: + try: + segment_attributes = {**attributes, "segment": segment} + self.runtime.segment_duration.record( + duration_seconds, segment_attributes + ) + if segment == "calculation": + self.runtime.calculate_duration.record( + duration_seconds, attributes + ) + if backend_segment: + self.runtime.backend_duration.record( + duration_seconds, + segment_attributes, + ) + except BaseException as exc: + self.runtime.log_observability_failure( + "metrics.record_segment", + exc, + segment=segment, + ) + + def record_error_metric(self, attributes: dict[str, str]) -> None: + try: + self.runtime.errors.add(1, attributes) + except BaseException as exc: + self.runtime.log_observability_failure("metrics.record_error", exc) + + def record_rate_limited_metric(self, attributes: dict[str, str]) -> None: + try: + self.runtime.rate_limited.add(1, attributes) + except BaseException as exc: + self.runtime.log_observability_failure( + "metrics.record_rate_limited", exc + ) + + def record_failover_event_metric(self, attributes: dict[str, str]) -> None: + try: + self.runtime.failover_events.add(1, attributes) + except BaseException as exc: + self.runtime.log_observability_failure( + "metrics.record_failover_event", + exc, + ) + + def record_active_request( + self, + delta: int, + attributes: dict[str, str], + ) -> None: + try: + self.runtime.active_requests.add(delta, attributes) + except BaseException as exc: + self.runtime.log_observability_failure( + "metrics.add_active_request", exc + ) + + def _configure_instruments(self) -> None: + self.runtime.operation_duration = self.runtime._instrument( + getattr(self.runtime.meter, "create_histogram", None), + "policyengine.operation.duration", + unit="s", + description="PolicyEngine operation duration.", + ) + self.runtime.http_duration = self.runtime._instrument( + getattr(self.runtime.meter, "create_histogram", None), + "http.server.request.duration", + unit="s", + description="HTTP server request duration.", + ) + self.runtime.segment_duration = self.runtime._instrument( + getattr(self.runtime.meter, "create_histogram", None), + "policyengine.segment.duration", + unit="s", + description="PolicyEngine operation segment duration.", + ) + self.runtime.calculate_duration = self.runtime._instrument( + getattr(self.runtime.meter, "create_histogram", None), + "policyengine.calculate.duration", + unit="s", + description="PolicyEngine calculate operation duration.", + ) + self.runtime.backend_duration = self.runtime._instrument( + getattr(self.runtime.meter, "create_histogram", None), + "policyengine.backend.duration", + unit="s", + description="PolicyEngine backend call duration.", + ) + self.runtime.operations = self.runtime._instrument( + getattr(self.runtime.meter, "create_counter", None), + "policyengine.operations", + description="PolicyEngine operation count.", + ) + self.runtime.requests = self.runtime._instrument( + getattr(self.runtime.meter, "create_counter", None), + "policyengine.requests", + description="PolicyEngine request count.", + ) + self.runtime.errors = self.runtime._instrument( + getattr(self.runtime.meter, "create_counter", None), + "policyengine.errors", + description="PolicyEngine error count.", + ) + self.runtime.rate_limited = self.runtime._instrument( + getattr(self.runtime.meter, "create_counter", None), + "policyengine.rate_limited_requests", + description="PolicyEngine rate-limited request count.", + ) + self.runtime.failover_events = self.runtime._instrument( + getattr(self.runtime.meter, "create_counter", None), + "policyengine.failover.events", + description="PolicyEngine failover event count.", + ) + self.runtime.active_requests = self.runtime._instrument( + getattr(self.runtime.meter, "create_up_down_counter", None), + "http.server.active_requests", + description="Active HTTP server requests.", + ) + + def _instrument(self, factory, *args, **kwargs): + if factory is None: + return _NoOpInstrument() + try: + return factory(*args, **kwargs) + except BaseException as exc: + self.runtime.log_observability_failure( + "metrics.create_instrument", + exc, + instrument=args[0] if args else None, + ) + return _NoOpInstrument() diff --git a/policyengine_observability/_operations.py b/policyengine_observability/_operations.py new file mode 100644 index 0000000..9cefa7b --- /dev/null +++ b/policyengine_observability/_operations.py @@ -0,0 +1,428 @@ +"""Operation lifecycle and timing-scope behavior.""" + +from __future__ import annotations + +import inspect +import time +from contextlib import contextmanager +from functools import wraps +from typing import TYPE_CHECKING, Any + +from . import _state +from .context import ( + ErrorRecord, + OperationObservabilityContext, +) + +if TYPE_CHECKING: + from .runtime import ObservabilityRuntime + + +class OperationLifecycle: + def __init__(self, runtime: ObservabilityRuntime) -> None: + self.runtime = runtime + + def operation( + self, + name: str, + *, + flavor: str | None = None, + **attrs: Any, + ): + return _OperationManager( + self.runtime, name, flavor=flavor, attrs=attrs + ) + + def entrypoint( + self, + name: str | None = None, + *, + flavor: str | None = None, + **attrs: Any, + ): + def decorator(func): + operation_name = name or getattr(func, "__name__", "operation") + return self.runtime.operation( + operation_name, + flavor=flavor, + **attrs, + )(func) + + return decorator + + def start_operation( + self, + name: str, + *, + flavor: str | None = None, + parent_context: Any = None, + timings: dict[str, float] | None = None, + emit_log: bool = True, + record_metric: bool = True, + **attrs: Any, + ) -> dict[str, Any]: + handle = { + "operation": None, + "operation_token": None, + "timings_token": None, + "start_token": None, + "context_token": None, + } + if not self.runtime.enabled: + return handle + try: + operation = OperationObservabilityContext( + config=self.runtime.config, + name=self.runtime._safe_str(name), + flavor=flavor, + attributes={ + key: value + for key, value in attrs.items() + if value is not None + }, + timings_ms={}, + emit_log=emit_log, + record_metric=record_metric, + ) + operation.context_token = _state._OPERATION_CONTEXT.set(operation) + handle["operation"] = operation + handle["operation_token"] = operation.context_token + if timings is not None: + handle["timings_token"] = _state._TIMINGS.set(timings) + handle["start_token"] = _state._TURN_START.set(time.perf_counter()) + if parent_context is not None and self.runtime.tracer is not None: + try: + from opentelemetry import context as otel_context + + handle["context_token"] = otel_context.attach( + parent_context + ) + except BaseException as exc: + self.runtime.log_observability_failure( + "operation.context_attach", + exc, + ) + if self.runtime.tracer is not None: + operation.span_handle = self.runtime._start_span( + self.runtime._span_name(operation.name), + operation.span_attributes(), + ) + except BaseException as exc: + self.runtime.log_observability_failure( + "operation.start", exc, name=name + ) + return handle + + def end_operation( + self, + handle: dict[str, Any] | None, + error: BaseException | None = None, + ) -> None: + if not handle: + return + operation = handle.get("operation") + try: + if operation is not None and error is not None: + operation.error = ErrorRecord( + type=type(error).__name__, + message=self.runtime._safe_str(error), + handled=False, + stack=self.runtime._safe_traceback(error), + ) + self.runtime.record_error_metric( + operation.metric_attributes( + error_type=type(error).__name__ + ) + ) + if operation is not None: + self.runtime.complete_operation(operation) + if operation is not None: + self.runtime._end_span(operation.span_handle, error) + except BaseException as exc: + self.runtime.log_observability_failure("operation.end", exc) + finally: + context_token = handle.get("context_token") + if context_token is not None: + try: + from opentelemetry import context as otel_context + + otel_context.detach(context_token) + except BaseException as exc: + self.runtime.log_observability_failure( + "operation.context_detach", + exc, + ) + for var, key in ( + (_state._TIMINGS, "timings_token"), + (_state._TURN_START, "start_token"), + (_state._OPERATION_CONTEXT, "operation_token"), + ): + token = handle.get(key) + if token is not None: + try: + var.reset(token) + except BaseException as exc: + self.runtime.log_observability_failure( + "operation.context_reset", + exc, + token=key, + ) + + def complete_operation( + self, + operation: OperationObservabilityContext, + ) -> None: + if operation.metric_recorded: + return + operation.metric_recorded = True + if operation.record_metric: + self.runtime.record_operation_metric( + operation.duration_seconds(), + operation.metric_attributes(), + ) + if operation.emit_log: + self.runtime.emit_operation_log(operation) + + @contextmanager + def collect_timings(self, name: str = "operation", **attrs: Any): + timings: dict[str, float] = {} + handle = self.runtime.start_scope(timings, name=name, **attrs) + error: BaseException | None = None + try: + yield timings + except BaseException as exc: + error = exc + raise + finally: + self.runtime.end_scope(handle, error) + + def start_scope( + self, + timings: dict[str, float], + *, + name: str = "operation", + parent_context: Any = None, + **attrs: Any, + ) -> dict[str, Any]: + if self.runtime.current_operation() is None: + return { + "operation_handle": self.runtime.start_operation( + name, + parent_context=parent_context, + timings=timings, + **attrs, + ) + } + handle = { + "operation_handle": None, + "timings_token": None, + "start_token": None, + "context_token": None, + "span": None, + } + try: + handle["timings_token"] = _state._TIMINGS.set(timings) + except BaseException as exc: + self.runtime.log_observability_failure("scope.timings_set", exc) + try: + handle["start_token"] = _state._TURN_START.set(time.perf_counter()) + except BaseException as exc: + self.runtime.log_observability_failure("scope.start_set", exc) + if parent_context is not None and self.runtime.tracer is not None: + try: + from opentelemetry import context as otel_context + + handle["context_token"] = otel_context.attach(parent_context) + except BaseException as exc: + self.runtime.log_observability_failure( + "scope.context_attach", exc + ) + try: + if self.runtime.tracer is not None: + handle["span"] = self.runtime._start_span(name, attrs) + except BaseException as exc: + self.runtime.log_observability_failure( + "scope.span_start", exc, span=name + ) + handle["span"] = None + return handle + + def annotate( + self, + handle: dict[str, Any] | None = None, + **attrs: Any, + ) -> None: + try: + if handle: + span_handle = handle.get("span") + if span_handle is not None: + _cm, span = span_handle + for key, value in attrs.items(): + if value is not None: + span.set_attribute(key, value) + context = self.runtime.current_context() + if context is not None: + for key, value in attrs.items(): + context.set_attribute(key, value) + operation = self.runtime.current_operation() + if operation is not None: + for key, value in attrs.items(): + operation.set_attribute(key, value) + self.runtime._set_current_span_attributes( + operation.span_attributes() + ) + except BaseException as exc: + self.runtime.log_observability_failure("scope.annotate", exc) + + def end_scope( + self, + handle: dict[str, Any] | None, + error: BaseException | None = None, + ) -> None: + if not handle: + return + operation_handle = handle.get("operation_handle") + if operation_handle is not None: + self.runtime.end_operation(operation_handle, error) + return + try: + self.runtime._end_span(handle.get("span"), error) + except BaseException as exc: + self.runtime.log_observability_failure("scope.span_end", exc) + context_token = handle.get("context_token") + if context_token is not None: + try: + from opentelemetry import context as otel_context + + otel_context.detach(context_token) + except BaseException as exc: + self.runtime.log_observability_failure( + "scope.context_detach", exc + ) + for var, key in ( + (_state._TIMINGS, "timings_token"), + (_state._TURN_START, "start_token"), + ): + token = handle.get(key) + if token is not None: + try: + var.reset(token) + except BaseException as exc: + self.runtime.log_observability_failure( + "scope.context_reset", + exc, + token=key, + ) + + def mark(self, key: str, ms: float) -> None: + try: + timings = _state._TIMINGS.get() + if timings is not None: + timings[key] = round(float(ms), 1) + except BaseException as exc: + self.runtime.log_observability_failure("scope.mark", exc, key=key) + + def mark_ttft(self, key: str = "ttft_ms") -> None: + try: + start = _state._TURN_START.get() + if start is not None: + self.runtime.mark(key, (time.perf_counter() - start) * 1000.0) + except BaseException as exc: + self.runtime.log_observability_failure("scope.mark_ttft", exc) + + def mark_ttft_attribute(self, key: str = "ttft_ms") -> None: + try: + start = _state._TURN_START.get() + if start is None: + return + self.runtime.annotate( + **{key: round((time.perf_counter() - start) * 1000.0, 1)} + ) + except BaseException as exc: + self.runtime.log_observability_failure( + "scope.mark_ttft_attribute", exc + ) + + def _start_implicit_operation( + self, + segment_name: str, + attrs: dict[str, Any], + ) -> dict[str, Any] | None: + if ( + self.runtime.current_operation() is not None + or self.runtime.current_context() is not None + ): + return None + operation_name = attrs.get("operation") or segment_name + flavor = attrs.get("flavor") + operation_attrs = { + key: value + for key, value in attrs.items() + if key not in {"operation", "flavor"} and value is not None + } + return self.runtime.start_operation( + self.runtime._safe_str(operation_name), + flavor=self.runtime._safe_str(flavor) + if flavor is not None + else None, + **operation_attrs, + ) + + +class _OperationManager: + def __init__( + self, + runtime: ObservabilityRuntime, + name: str, + *, + flavor: str | None, + attrs: dict[str, Any], + ) -> None: + self.runtime = runtime + self.name = name + self.flavor = flavor + self.attrs = attrs + self.handle: dict[str, Any] | None = None + + def __enter__(self): + self.handle = self.runtime.start_operation( + self.name, + flavor=self.flavor, + **self.attrs, + ) + return self.runtime.current_operation() + + def __exit__(self, exc_type, exc, _traceback) -> bool: + self.runtime.end_operation(self.handle, exc) + return False + + async def __aenter__(self): + return self.__enter__() + + async def __aexit__(self, exc_type, exc, traceback) -> bool: + return self.__exit__(exc_type, exc, traceback) + + def __call__(self, func): + if inspect.iscoroutinefunction(func): + + @wraps(func) + async def async_wrapper(*args, **kwargs): + async with self.runtime.operation( + self.name, + flavor=self.flavor, + **self.attrs, + ): + return await func(*args, **kwargs) + + return async_wrapper + + @wraps(func) + def wrapper(*args, **kwargs): + with self.runtime.operation( + self.name, + flavor=self.flavor, + **self.attrs, + ): + return func(*args, **kwargs) + + return wrapper diff --git a/policyengine_observability/_requests.py b/policyengine_observability/_requests.py new file mode 100644 index 0000000..b326e91 --- /dev/null +++ b/policyengine_observability/_requests.py @@ -0,0 +1,291 @@ +"""Request lifecycle, response metadata, and cleanup.""" + +from __future__ import annotations + +from typing import TYPE_CHECKING, Any + +from . import _state +from ._state import REQUEST_ID_HEADER, TRACEPARENT_HEADER +from .context import ( + OperationObservabilityContext, + RequestObservabilityContext, +) + +if TYPE_CHECKING: + from .runtime import ObservabilityRuntime + + +class RequestLifecycle: + def __init__(self, runtime: ObservabilityRuntime) -> None: + self.runtime = runtime + + def begin_request( + self, + context: RequestObservabilityContext, + *, + carrier: Any = None, + ) -> None: + if not self.runtime.enabled: + return + try: + context.context_token = _state._REQUEST_CONTEXT.set(context) + context.set_attribute("endpoint", context.endpoint) + self.runtime._begin_request_operation(context) + self.runtime._start_request_span(context, carrier=carrier) + self.runtime.record_active_request(1, context.metric_attributes()) + except BaseException as exc: + self.runtime.log_observability_failure("request.begin", exc) + + def _begin_request_operation( + self, + context: RequestObservabilityContext, + ) -> None: + try: + parent_operation = _state._OPERATION_CONTEXT.get() + timings = context.timings_ms + timing_counts = context.timing_counts + segment_tree = context.segment_tree + segment_sequence = context.segment_sequence + if context.internal_dispatch and parent_operation is not None: + timings = parent_operation.timings_ms + timing_counts = parent_operation.timing_counts + segment_tree = parent_operation.segment_tree + segment_sequence = parent_operation.segment_sequence + context.timings_ms = timings + context.timing_counts = timing_counts + context.segment_tree = segment_tree + context.segment_sequence = segment_sequence + operation = OperationObservabilityContext( + config=context.config, + name=context.route, + flavor="http", + attributes={ + "route": context.route, + "method": context.method, + "endpoint": context.endpoint, + "path": context.path, + }, + timings_ms=timings, + timing_counts=timing_counts, + segment_tree=segment_tree, + segment_sequence=segment_sequence, + emit_log=False, + record_metric=False, + ) + operation.context_token = _state._OPERATION_CONTEXT.set(operation) + context.operation_context = operation + context.operation_token = operation.context_token + except BaseException as exc: + self.runtime.log_observability_failure( + "request.operation_begin", + exc, + request_id=getattr(context, "request_id", None), + ) + + def finish_request(self, status_code: int) -> dict[str, str]: + headers = self.runtime.prepare_response(status_code) + self.runtime.complete_request(status_code) + return headers + + def prepare_response(self, status_code: int) -> dict[str, str]: + if not self.runtime.enabled: + return {} + headers: dict[str, str] = {} + try: + context = self.runtime.current_context() + if context is None: + return headers + context.status_code = status_code + self.runtime._set_current_span_attributes( + context.span_attributes() + ) + if context.operation_context is not None: + context.operation_context.set_attribute( + "status_code", + str(status_code), + ) + headers[REQUEST_ID_HEADER] = context.request_id + traceparent = self.runtime.traceparent_header() + if traceparent: + headers[TRACEPARENT_HEADER] = traceparent + if status_code == 429: + context.set_attribute("rate_limited", True) + return headers + except BaseException as exc: + self.runtime.log_observability_failure( + "request.prepare_response", exc + ) + return headers + + def complete_request(self, status_code: int | None = None) -> None: + if not self.runtime.enabled: + return + try: + context = self.runtime.current_context() + if context is None: + return + if status_code is not None: + context.status_code = status_code + self.runtime._set_current_span_attributes( + context.span_attributes() + ) + if context.request_metric_recorded: + return + context.request_metric_recorded = True + if context.status_code == 429: + self.runtime.record_rate_limited_metric( + context.metric_attributes() + ) + self.runtime.record_request_metric( + context.duration_seconds(), + context.metric_attributes(), + ) + self.runtime._close_active_request(context) + except BaseException as exc: + self.runtime.log_observability_failure("request.complete", exc) + + def update_request_route( + self, + *, + route: str | None = None, + endpoint: str | None = None, + ) -> None: + if not self.runtime.enabled: + return + try: + context = self.runtime.current_context() + if context is None: + return + route_changed = bool(route and route != context.route) + old_active_attributes = ( + context.metric_attributes() + if route_changed and not context.active_closed + else None + ) + if route: + context.route = route + if context.operation_context is not None: + context.operation_context.name = route + context.operation_context.set_attribute("route", route) + if endpoint: + context.endpoint = endpoint + context.set_attribute("endpoint", endpoint) + if context.operation_context is not None: + context.operation_context.set_attribute( + "endpoint", + endpoint, + ) + self.runtime._set_current_span_attributes( + context.span_attributes() + ) + span = context.server_span + update_name = getattr(span, "update_name", None) + if route and update_name is not None: + update_name(route) + if old_active_attributes is not None: + self.runtime.record_active_request(-1, old_active_attributes) + self.runtime.record_active_request( + 1, context.metric_attributes() + ) + except BaseException as exc: + self.runtime.log_observability_failure("request.update_route", exc) + + def teardown_request(self, exc: BaseException | None = None) -> None: + if not self.runtime.enabled: + return + context = self.runtime.current_context() + if context is None: + return + try: + if exc is not None: + self.runtime.record_error( + exc, + handled=False, + status_code=context.status_code or 500, + ) + self.runtime._close_active_request(context) + self.runtime.emit_request_log(context) + except BaseException as observability_exc: + self.runtime.log_observability_failure( + "request.teardown", + observability_exc, + ) + finally: + self.runtime._close_request_span(context, exc) + self.runtime._reset_request_operation_context(context) + self.runtime._reset_request_context(context) + + def set_attribute(self, key: str, value: Any) -> None: + if not self.runtime.enabled: + return + try: + context = self.runtime.current_context() + if context is not None: + context.set_attribute(key, value) + if context.operation_context is not None: + context.operation_context.set_attribute(key, value) + self.runtime._set_current_span_attributes( + context.span_attributes(**{f"policyengine.{key}": value}) + ) + operation = self.runtime.current_operation() + if operation is not None and operation is not getattr( + context, "operation_context", None + ): + operation.set_attribute(key, value) + self.runtime._set_current_span_attributes( + operation.span_attributes(**{f"policyengine.{key}": value}) + ) + except BaseException as exc: + self.runtime.log_observability_failure( + "request.set_attribute", + exc, + attribute=key, + ) + + def _close_active_request( + self, + context: RequestObservabilityContext, + ) -> None: + try: + if context.active_closed: + return + context.active_closed = True + self.runtime.record_active_request(-1, context.metric_attributes()) + except BaseException as exc: + self.runtime.log_observability_failure( + "request.close_active", + exc, + request_id=getattr(context, "request_id", None), + ) + + def _reset_request_operation_context( + self, + context: RequestObservabilityContext, + ) -> None: + token = context.operation_token + if token is None: + return + try: + _state._OPERATION_CONTEXT.reset(token) + except BaseException as exc: + self.runtime.log_observability_failure( + "request.operation_context_reset", + exc, + request_id=getattr(context, "request_id", None), + ) + + def _reset_request_context( + self, + context: RequestObservabilityContext, + ) -> None: + token = context.context_token + if token is None: + return + try: + _state._REQUEST_CONTEXT.reset(token) + except BaseException as exc: + self.runtime.log_observability_failure( + "request.context_reset", + exc, + request_id=getattr(context, "request_id", None), + ) diff --git a/policyengine_observability/_state.py b/policyengine_observability/_state.py new file mode 100644 index 0000000..a601d3b --- /dev/null +++ b/policyengine_observability/_state.py @@ -0,0 +1,65 @@ +"""Shared context variables for requests, operations, and segments.""" + +from __future__ import annotations + +from contextvars import ContextVar +from typing import TYPE_CHECKING + +from .context import ( + OperationObservabilityContext, + RequestObservabilityContext, + SegmentTimingNode, +) + +if TYPE_CHECKING: + from .runtime import ObservabilityRuntime + +OBSERVABILITY_INTERNAL_DISPATCH_HEADER = "X-PolicyEngine-Internal-Dispatch" +REQUEST_ID_HEADER = "X-PolicyEngine-Request-Id" +TRACEPARENT_HEADER = "traceparent" + +_REQUEST_CONTEXT: ContextVar[RequestObservabilityContext | None] = ContextVar( + "policyengine_request_observability_context", + default=None, +) +_OPERATION_CONTEXT: ContextVar[OperationObservabilityContext | None] = ( + ContextVar( + "policyengine_operation_observability_context", + default=None, + ) +) +_TIMINGS: ContextVar[dict[str, float] | None] = ContextVar( + "policyengine_observability_timings", + default=None, +) +_TURN_START: ContextVar[float | None] = ContextVar( + "policyengine_observability_turn_start", + default=None, +) +_SEGMENT_STACK: ContextVar[tuple[tuple[int, SegmentTimingNode], ...]] = ( + ContextVar( + "policyengine_observability_segment_stack", + default=(), + ) +) + + +class ContextState: + def __init__(self, runtime: ObservabilityRuntime) -> None: + self.runtime = runtime + + def current_context(self) -> RequestObservabilityContext | None: + try: + return _REQUEST_CONTEXT.get() + except BaseException as exc: + self.runtime.log_observability_failure("context.current", exc) + return None + + def current_operation( + self, + ) -> OperationObservabilityContext | None: + try: + return _OPERATION_CONTEXT.get() + except BaseException as exc: + self.runtime.log_observability_failure("operation.current", exc) + return None diff --git a/policyengine_observability/_tracing.py b/policyengine_observability/_tracing.py new file mode 100644 index 0000000..b8dfc6c --- /dev/null +++ b/policyengine_observability/_tracing.py @@ -0,0 +1,434 @@ +"""Trace initialization, propagation, and span lifecycle.""" + +from __future__ import annotations + +from collections.abc import Iterator +from contextlib import contextmanager +from typing import TYPE_CHECKING, Any + +from ._state import TRACEPARENT_HEADER +from .context import ( + RequestObservabilityContext, +) + +if TYPE_CHECKING: + from .runtime import ObservabilityRuntime + + +def _is_safe_span_value(value: Any) -> bool: + return isinstance(value, str | bool | int | float) + + +class TraceRecorder: + def __init__(self, runtime: ObservabilityRuntime) -> None: + self.runtime = runtime + + def traceparent_header(self) -> str | None: + if not self.runtime.enabled or self.runtime.propagate is None: + return None + try: + carrier: dict[str, str] = {} + self.runtime.propagate.inject(carrier) + return carrier.get(TRACEPARENT_HEADER) + except BaseException as exc: + self.runtime.log_observability_failure( + "request.traceparent_header", exc + ) + return None + + def capture_context(self): + if self.runtime.tracer is None: + return None + try: + from opentelemetry import context as otel_context + + return otel_context.get_current() + except BaseException as exc: + self.runtime.log_observability_failure("otel.capture_context", exc) + return None + + def instrument_fastapi(self, app: Any) -> None: + if not self.runtime.enabled or not self.runtime.config.otel_enabled: + return + try: + from opentelemetry.instrumentation.fastapi import ( + FastAPIInstrumentor, + ) + + FastAPIInstrumentor.instrument_app(app) + except BaseException as exc: + self.runtime.log_observability_failure( + "fastapi.auto_instrument", + exc, + ) + + def instrument_httpx(self) -> None: + if ( + not self.runtime.enabled + or not self.runtime.config.otel_enabled + or self.runtime._httpx_instrumented + ): + return + try: + from opentelemetry.instrumentation.httpx import ( + HTTPXClientInstrumentor, + ) + + HTTPXClientInstrumentor().instrument() + self.runtime._httpx_instrumented = True + except BaseException as exc: + self.runtime.log_observability_failure( + "httpx.auto_instrument", exc + ) + + def _configure_otel(self) -> None: + try: + from opentelemetry import metrics, propagate, trace + from opentelemetry.sdk.metrics import MeterProvider + from opentelemetry.sdk.resources import ( + DEPLOYMENT_ENVIRONMENT, + SERVICE_NAME, + Resource, + ) + from opentelemetry.sdk.trace import TracerProvider + from opentelemetry.trace import SpanKind, Status, StatusCode + except BaseException as exc: + self.runtime.log_observability_failure( + "otel.configure_imports", exc + ) + return + + try: + resource = Resource.create( + { + SERVICE_NAME: self.runtime.config.service_name, + DEPLOYMENT_ENVIRONMENT: self.runtime.config.environment, + "service.role": self.runtime.config.service_role, + } + ) + tracer_provider = TracerProvider(resource=resource) + metric_readers = [] + if self.runtime.config.otlp_endpoint: + self.runtime._add_trace_exporter(tracer_provider) + metric_reader = self.runtime._metric_reader() + if metric_reader is not None: + metric_readers.append(metric_reader) + self.runtime.tracer_provider = tracer_provider + try: + trace.set_tracer_provider(tracer_provider) + except BaseException as exc: + self.runtime.log_observability_failure( + "otel.set_tracer_provider", + exc, + ) + try: + self.runtime.meter_provider = MeterProvider( + resource=resource, + metric_readers=metric_readers, + ) + metrics.set_meter_provider(self.runtime.meter_provider) + except BaseException as exc: + self.runtime.log_observability_failure( + "otel.set_meter_provider", + exc, + ) + self.runtime.trace = trace + self.runtime.propagate = propagate + self.runtime.SpanKind = SpanKind + self.runtime.Status = Status + self.runtime.StatusCode = StatusCode + tracer_name = ( + self.runtime.config.tracer_name + or self.runtime.config.service_name + ) + meter_name = ( + self.runtime.config.meter_name + or self.runtime.config.service_name + ) + self.runtime.tracer = trace.get_tracer(tracer_name) + self.runtime.meter = metrics.get_meter(meter_name) + self.runtime._configure_instruments() + except BaseException as exc: + self.runtime.log_observability_failure("otel.configure", exc) + + def _add_trace_exporter(self, tracer_provider) -> None: + try: + if self.runtime.config.otlp_protocol.startswith("http"): + from opentelemetry.exporter.otlp.proto.http.trace_exporter import ( + OTLPSpanExporter, + ) + else: + from opentelemetry.exporter.otlp.proto.grpc.trace_exporter import ( + OTLPSpanExporter, + ) + from opentelemetry.sdk.trace.export import BatchSpanProcessor + + tracer_provider.add_span_processor( + BatchSpanProcessor(OTLPSpanExporter()) + ) + except BaseException as exc: + self.runtime.log_observability_failure("otel.trace_exporter", exc) + + def _metric_reader(self): + try: + if self.runtime.config.otlp_protocol.startswith("http"): + from opentelemetry.exporter.otlp.proto.http.metric_exporter import ( + OTLPMetricExporter, + ) + else: + from opentelemetry.exporter.otlp.proto.grpc.metric_exporter import ( + OTLPMetricExporter, + ) + from opentelemetry.sdk.metrics.export import ( + PeriodicExportingMetricReader, + ) + + return PeriodicExportingMetricReader(OTLPMetricExporter()) + except BaseException as exc: + self.runtime.log_observability_failure("otel.metric_exporter", exc) + return None + + def _start_request_span( + self, + context: RequestObservabilityContext, + *, + carrier: Any = None, + ) -> None: + if self.runtime.tracer is None: + return + attrs = context.span_attributes() + parent_context = self.runtime._extract_context(carrier) + try: + context.server_span_cm = self.runtime.tracer.start_as_current_span( + context.route, + context=parent_context, + kind=self.runtime.SpanKind.SERVER + if self.runtime.SpanKind + else None, + attributes=attrs, + ) + context.server_span = context.server_span_cm.__enter__() + except BaseException as exc: + context.server_span_cm = None + context.server_span = None + self.runtime.log_observability_failure( + "otel.request_span_enter", exc + ) + + def _close_request_span( + self, + context: RequestObservabilityContext, + exc: BaseException | None, + ) -> None: + if context.span_closed: + return + context.span_closed = True + span_cm = context.server_span_cm + if span_cm is None: + return + try: + if exc is None: + span_cm.__exit__(None, None, None) + else: + span_cm.__exit__(type(exc), exc, exc.__traceback__) + except BaseException as observability_exc: + self.runtime.log_observability_failure( + "otel.request_span_exit", + observability_exc, + request_id=context.request_id, + ) + + @contextmanager + def _safe_span(self, name: str, attrs: dict[str, Any]) -> Iterator[Any]: + if self.runtime.tracer is None: + yield None + return + span_handle = self.runtime._start_span(name, attrs) + if span_handle is None: + yield None + return + _cm, span = span_handle + try: + yield span + except BaseException as exc: + try: + self.runtime._end_span(span_handle, exc) + except BaseException as observability_exc: + self.runtime.log_observability_failure( + "otel.span_exit", + observability_exc, + span=name, + ) + raise + else: + try: + self.runtime._end_span(span_handle) + except BaseException as exc: + self.runtime.log_observability_failure( + "otel.span_exit", + exc, + span=name, + ) + + def _start_span(self, name: str, attrs: dict[str, Any]): + try: + span_cm = self.runtime.tracer.start_as_current_span(name) + span = span_cm.__enter__() + except BaseException as exc: + self.runtime.log_observability_failure( + "otel.span_enter", exc, span=name + ) + return None + try: + for key, value in attrs.items(): + if value is not None: + span.set_attribute(key, value) + except BaseException as exc: + self.runtime.log_observability_failure( + "otel.span_attributes", + exc, + span=name, + ) + return span_cm, span + + def _end_span( + self, + span_handle, + error: BaseException | None = None, + ) -> None: + if span_handle is None: + return + span_cm, span = span_handle + try: + if error is not None: + self.runtime._record_exception_on_span( + span, + error, + handled=False, + status_code=500, + ) + except BaseException as exc: + self.runtime.log_observability_failure( + "otel.span_error_status", exc + ) + try: + span_cm.__exit__(None, None, None) + except BaseException as exc: + self.runtime.log_observability_failure("otel.span_exit", exc) + + def _segment_span_attributes( + self, + attrs: dict[str, Any], + ) -> dict[str, Any]: + context = self.runtime.current_context() + operation = self.runtime.current_operation() + span_attrs = { + key: value for key, value in attrs.items() if value is not None + } + if context is not None: + span_attrs = {**context.span_attributes(), **span_attrs} + elif operation is not None: + span_attrs = {**operation.span_attributes(), **span_attrs} + return span_attrs + + def _span_name(self, segment_name: str) -> str: + if not self.runtime.config.span_prefix: + return segment_name + return f"{self.runtime.config.span_prefix}.{segment_name}" + + def _set_current_span_attributes(self, attrs: dict[str, Any]) -> None: + span = self.runtime._current_span() + if span is None: + return + try: + for key, value in attrs.items(): + if value is not None: + span.set_attribute(key, value) + except BaseException as exc: + self.runtime.log_observability_failure( + "otel.set_span_attributes", exc + ) + + def _current_span(self): + if self.runtime.trace is None: + return None + try: + return self.runtime.trace.get_current_span() + except BaseException as exc: + self.runtime.log_observability_failure("otel.current_span", exc) + return None + + def _trace_ids(self) -> tuple[str | None, str | None]: + span = self.runtime._current_span() + if span is None: + return None, None + try: + context = span.get_span_context() + except BaseException as exc: + self.runtime.log_observability_failure("otel.span_context", exc) + return None, None + if not getattr(context, "is_valid", False): + return None, None + return f"{context.trace_id:032x}", f"{context.span_id:016x}" + + def _extract_context(self, carrier: Any): + if self.runtime.propagate is None or carrier is None: + return None + try: + return self.runtime.propagate.extract(carrier) + except BaseException as exc: + self.runtime.log_observability_failure("otel.extract_context", exc) + return None + + def _record_exception_on_span( + self, + span, + exc: BaseException, + *, + handled: bool, + status_code: int | None, + ) -> None: + try: + span.record_exception(exc) + span.set_attribute("error.type", type(exc).__name__) + span.set_attribute("error.handled", handled) + if ( + self.runtime.Status is not None + and self.runtime.StatusCode is not None + and ( + not handled + or (status_code is not None and status_code >= 500) + ) + ): + span.set_status( + self.runtime.Status( + self.runtime.StatusCode.ERROR, + self.runtime._safe_str(exc), + ) + ) + except BaseException as observability_exc: + self.runtime.log_observability_failure( + "otel.record_exception", + observability_exc, + original_error_type=type(exc).__name__, + ) + + def _add_span_event(self, event: str, fields: dict[str, Any]) -> None: + span = self.runtime._current_span() + if span is None: + return + try: + span.add_event( + event, + { + key: value + for key, value in fields.items() + if _is_safe_span_value(value) + }, + ) + except BaseException as exc: + self.runtime.log_observability_failure( + "otel.add_event", + exc, + event_name=event, + ) diff --git a/policyengine_observability/logging.py b/policyengine_observability/logging.py index d92e3bd..666d4ba 100644 --- a/policyengine_observability/logging.py +++ b/policyengine_observability/logging.py @@ -1,6 +1,23 @@ +"""Structured log emission and observability-failure handling.""" + from __future__ import annotations +import json import logging +import sys +import traceback +from datetime import UTC, datetime +from typing import TYPE_CHECKING, Any + +from .context import ( + ErrorRecord, + OperationObservabilityContext, + RequestObservabilityContext, + _metric_attrs, +) + +if TYPE_CHECKING: + from .runtime import ObservabilityRuntime class PlainMessageFormatter(logging.Formatter): @@ -15,3 +32,349 @@ def configure_plain_logger(logger: logging.Logger, level: int) -> None: handler = logging.StreamHandler() handler.setFormatter(PlainMessageFormatter()) logger.addHandler(handler) + + +REQUEST_LOGGER_NAME = "policyengine_observability.requests" +OPERATION_LOGGER_NAME = "policyengine_observability.operations" +EVENT_LOGGER_NAME = "policyengine_observability.events" +INTERNAL_LOGGER_NAME = "policyengine_observability.internal" + +REQUEST_LOGGER = logging.getLogger(REQUEST_LOGGER_NAME) +OPERATION_LOGGER = logging.getLogger(OPERATION_LOGGER_NAME) +EVENT_LOGGER = logging.getLogger(EVENT_LOGGER_NAME) +INTERNAL_LOGGER = logging.getLogger(INTERNAL_LOGGER_NAME) + + +class LogEmitter: + def __init__(self, runtime: ObservabilityRuntime) -> None: + self.runtime = runtime + + def record_error( + self, + exc: BaseException, + *, + handled: bool, + status_code: int | None = None, + include_stack: bool = True, + ) -> None: + if not self.runtime.enabled: + return + try: + context = self.runtime.current_context() + operation = self.runtime.current_operation() + error_record = ErrorRecord( + type=type(exc).__name__, + message=self.runtime._safe_str(exc), + handled=handled, + stack=( + self.runtime._safe_traceback(exc) + if include_stack + else None + ), + ) + if context is not None: + if status_code is not None: + context.status_code = status_code + context.error = error_record + self.runtime.record_error_metric( + context.metric_attributes(error_type=type(exc).__name__) + ) + elif operation is not None: + operation.error = error_record + self.runtime.record_error_metric( + operation.metric_attributes(error_type=type(exc).__name__) + ) + else: + return + span = self.runtime._current_span() + if span is not None: + self.runtime._record_exception_on_span( + span, + exc, + handled=handled, + status_code=status_code, + ) + except BaseException as observability_exc: + self.runtime.log_observability_failure( + "request.record_error", + observability_exc, + original_error_type=type(exc).__name__, + ) + + def record_event(self, event: str, **fields: Any) -> None: + if not self.runtime.enabled: + return + try: + context = self.runtime.current_context() + operation = self.runtime.current_operation() + base: dict[str, Any] = { + "schema_version": "policyengine.observability.event.v1", + "event": event, + "service_name": self.runtime.config.service_name, + "service_role": self.runtime.config.service_role, + "environment": self.runtime.config.environment, + "created_at": datetime.now(UTC).isoformat(), + } + if context is not None: + trace_id, span_id = self.runtime._trace_ids() + base.update( + { + "service_name": context.config.service_name, + "service_role": context.config.service_role, + "environment": context.config.environment, + "request_id": context.request_id, + "trace_id": trace_id, + "span_id": span_id, + "route": context.route, + "path": context.path, + } + ) + elif operation is not None: + trace_id, span_id = self.runtime._trace_ids() + base.update( + { + "service_name": operation.config.service_name, + "service_role": operation.config.service_role, + "environment": operation.config.environment, + "operation": operation.name, + "flavor": operation.flavor, + "trace_id": trace_id, + "span_id": span_id, + } + ) + clean_fields = { + key: value + for key, value in fields.items() + if value is not None + } + base.update(clean_fields) + self.runtime._emit_structured_log( + base, + log_type="event", + severity="INFO", + ) + self.runtime._add_span_event(event, clean_fields) + if event.startswith("modal_") or "fallback" in event: + attrs = ( + context.metric_attributes(event=event) + if context + else operation.metric_attributes(event=event) + if operation + else _metric_attrs( + {"event": event}, + self.runtime.config.metric_attribute_keys, + ) + ) + self.runtime.record_failover_event_metric(attrs) + except BaseException as exc: + self.runtime.log_observability_failure( + "request.record_event", + exc, + event_name=event, + ) + + def emit_request_log(self, context: RequestObservabilityContext) -> None: + if not self.runtime.enabled: + return + try: + if context.emitted: + return + context.emitted = True + if ( + context.internal_dispatch + or not context.config.request_logs_enabled + ): + return + trace_id, span_id = self.runtime._trace_ids() + payload = context.as_log_record( + trace_id=trace_id, + span_id=span_id, + ) + self.runtime._emit_structured_log( + payload, + log_type="request", + severity=self.runtime._severity_for_log_record(payload), + ) + except BaseException as exc: + self.runtime.log_observability_failure( + "request.emit_request_log", + exc, + request_id=getattr(context, "request_id", None), + ) + + def emit_operation_log( + self, + operation: OperationObservabilityContext, + ) -> None: + if not self.runtime.enabled: + return + try: + if operation.emitted: + return + operation.emitted = True + trace_id, span_id = self.runtime._trace_ids() + payload = operation.as_log_record( + trace_id=trace_id, + span_id=span_id, + ) + self.runtime._emit_structured_log( + payload, + log_type="operation", + severity=self.runtime._severity_for_log_record(payload), + ) + except BaseException as exc: + self.runtime.log_observability_failure( + "operation.emit_log", + exc, + operation=getattr(operation, "name", None), + ) + + def log_observability_failure( + self, + operation: str, + exc: BaseException, + **fields: Any, + ) -> None: + payload = self.runtime._internal_error_payload( + operation, exc, **fields + ) + if self.runtime._emitting_internal_error: + self.runtime._write_stderr(payload) + return + self.runtime._emitting_internal_error = True + try: + self.runtime._emit_structured_log( + payload, + log_type="internal", + severity="ERROR", + ) + except BaseException: + self.runtime._write_stderr(payload) + finally: + self.runtime._emitting_internal_error = False + + def _configure_loggers(self) -> None: + for logger in ( + REQUEST_LOGGER, + OPERATION_LOGGER, + EVENT_LOGGER, + INTERNAL_LOGGER, + ): + configure_plain_logger(logger, self.runtime.config.log_level) + + def _emit_structured_log( + self, + payload: dict[str, Any], + *, + log_type: str, + severity: str, + ) -> None: + try: + if not self.runtime.log_destination_manager.configured: + self.runtime._configure_loggers() + self.runtime.log_destination_manager.emit( + payload, + log_type=log_type, + severity=severity, + ) + except BaseException as exc: + if self.runtime._emitting_internal_error: + self.runtime._write_stderr(payload) + else: + self.runtime.log_observability_failure( + "logging.emit", + exc, + log_type=log_type, + ) + + def _handle_destination_failure( + self, + operation: str, + exc: BaseException, + **fields: Any, + ) -> None: + if self.runtime._emitting_internal_error: + self.runtime._write_stderr( + self.runtime._internal_error_payload(operation, exc, **fields) + ) + return + self.runtime.log_observability_failure(operation, exc, **fields) + + def _severity_for_log_record(self, payload: dict[str, Any]) -> str: + status_code = self.runtime._int_or_none(payload.get("status_code")) + error = payload.get("error") + handled = error.get("handled") if isinstance(error, dict) else None + if status_code is not None and status_code >= 500: + return "ERROR" + if error is not None and handled is False: + return "ERROR" + if error is not None or ( + status_code is not None and status_code >= 400 + ): + return "WARNING" + return "INFO" + + def _int_or_none(self, value: Any) -> int | None: + try: + return int(value) + except (TypeError, ValueError): + return None + + def _internal_error_payload( + self, + operation: str, + exc: BaseException, + **fields: Any, + ) -> dict[str, Any]: + payload = { + "schema_version": "policyengine.observability.internal_error.v1", + "event": "observability_internal_error", + "service_name": self.runtime.config.service_name, + "service_role": self.runtime.config.service_role, + "environment": self.runtime.config.environment, + "created_at": datetime.now(UTC).isoformat(), + "operation": operation, + "error": { + "type": type(exc).__name__, + "message": self.runtime._safe_str(exc), + "stack": self.runtime._safe_traceback(exc), + }, + } + payload.update( + {key: value for key, value in fields.items() if value is not None} + ) + return payload + + def _safe_str(self, value: Any) -> str: + try: + return str(value) + except BaseException: + return f"" + + def _safe_traceback(self, exc: BaseException) -> str: + try: + return "".join( + traceback.format_exception(type(exc), exc, exc.__traceback__) + ) + except BaseException: + return "" + + def _json(self, payload: dict[str, Any]) -> str: + try: + return json.dumps(payload, sort_keys=True, default=str) + except BaseException: + return json.dumps( + { + "schema_version": "policyengine.observability.internal_error.v1", + "event": "observability_internal_error", + "created_at": datetime.now(UTC).isoformat(), + "operation": "observability.failure_json", + }, + sort_keys=True, + ) + + def _write_stderr(self, payload: dict[str, Any]) -> None: + try: + sys.stderr.write(self.runtime._json(payload) + "\n") + except BaseException: + return diff --git a/policyengine_observability/runtime.py b/policyengine_observability/runtime.py index 051f79a..a2f7248 100644 --- a/policyengine_observability/runtime.py +++ b/policyengine_observability/runtime.py @@ -1,88 +1,42 @@ +"""Public runtime interface and component lifecycle.""" + from __future__ import annotations -import inspect -import json -import logging -import sys import threading import time -import traceback from collections.abc import AsyncIterator, Iterator -from contextlib import asynccontextmanager, contextmanager -from contextvars import ContextVar -from datetime import UTC, datetime from enum import Enum -from functools import wraps from typing import Any -from .config import ObservabilityConfig -from .context import ( - ErrorRecord, - OperationObservabilityContext, - RequestObservabilityContext, - SegmentTimingNode, - _metric_attrs, +from . import _state +from ._metrics import MetricRecorder, _NoOpInstrument +from ._operations import OperationLifecycle +from ._requests import RequestLifecycle +from ._state import ( + OBSERVABILITY_INTERNAL_DISPATCH_HEADER as OBSERVABILITY_INTERNAL_DISPATCH_HEADER, ) -from .destinations import LogDestinationManager -from .destinations.base import clamped -from .logging import configure_plain_logger -from .segments import coerce_segment_name - -OBSERVABILITY_INTERNAL_DISPATCH_HEADER = "X-PolicyEngine-Internal-Dispatch" -REQUEST_ID_HEADER = "X-PolicyEngine-Request-Id" -TRACEPARENT_HEADER = "traceparent" - -REQUEST_LOGGER_NAME = "policyengine_observability.requests" -OPERATION_LOGGER_NAME = "policyengine_observability.operations" -EVENT_LOGGER_NAME = "policyengine_observability.events" -INTERNAL_LOGGER_NAME = "policyengine_observability.internal" - -REQUEST_LOGGER = logging.getLogger(REQUEST_LOGGER_NAME) -OPERATION_LOGGER = logging.getLogger(OPERATION_LOGGER_NAME) -EVENT_LOGGER = logging.getLogger(EVENT_LOGGER_NAME) -INTERNAL_LOGGER = logging.getLogger(INTERNAL_LOGGER_NAME) - -_REQUEST_CONTEXT: ContextVar[RequestObservabilityContext | None] = ContextVar( - "policyengine_request_observability_context", - default=None, +from ._state import ( + REQUEST_ID_HEADER as REQUEST_ID_HEADER, ) -_OPERATION_CONTEXT: ContextVar[OperationObservabilityContext | None] = ( - ContextVar( - "policyengine_operation_observability_context", - default=None, - ) +from ._state import ( + TRACEPARENT_HEADER as TRACEPARENT_HEADER, ) -_TIMINGS: ContextVar[dict[str, float] | None] = ContextVar( - "policyengine_observability_timings", - default=None, +from ._state import ( + ContextState, ) -_TURN_START: ContextVar[float | None] = ContextVar( - "policyengine_observability_turn_start", - default=None, -) -_SEGMENT_STACK: ContextVar[tuple[tuple[int, SegmentTimingNode], ...]] = ( - ContextVar( - "policyengine_observability_segment_stack", - default=(), - ) -) - -MAX_SEGMENT_ATTR_LENGTH = 200 -SENSITIVE_SEGMENT_ATTR_PARTS = ( - "authorization", - "credential", - "password", - "secret", - "token", +from ._tracing import TraceRecorder +from .config import ObservabilityConfig +from .context import OperationObservabilityContext, RequestObservabilityContext +from .destinations import LogDestinationManager +from .destinations.base import clamped +from .logging import ( + EVENT_LOGGER, + INTERNAL_LOGGER, + OPERATION_LOGGER, + REQUEST_LOGGER, + LogEmitter, ) - - -class _NoOpInstrument: - def add(self, *_args, **_kwargs) -> None: - return None - - def record(self, *_args, **_kwargs) -> None: - return None +from .segments import SegmentRecorder class ObservabilityRuntime: @@ -128,6 +82,13 @@ def __init__( serializer=self._json, on_failure=self._handle_destination_failure, ) + self._context_state = ContextState(self) + self._operations = OperationLifecycle(self) + self._requests = RequestLifecycle(self) + self._segments = SegmentRecorder(self) + self._logging = LogEmitter(self) + self._metrics = MetricRecorder(self) + self._tracing = TraceRecorder(self) @classmethod def disabled(cls) -> ObservabilityRuntime: @@ -145,29 +106,13 @@ def configure(self) -> None: self.instrument_httpx() def current_context(self) -> RequestObservabilityContext | None: - try: - return _REQUEST_CONTEXT.get() - except BaseException as exc: - self.log_observability_failure("context.current", exc) - return None + return self._context_state.current_context() - def current_operation( - self, - ) -> OperationObservabilityContext | None: - try: - return _OPERATION_CONTEXT.get() - except BaseException as exc: - self.log_observability_failure("operation.current", exc) - return None + def current_operation(self) -> OperationObservabilityContext | None: + return self._context_state.current_operation() - def operation( - self, - name: str, - *, - flavor: str | None = None, - **attrs: Any, - ): - return _OperationManager(self, name, flavor=flavor, attrs=attrs) + def operation(self, name: str, *, flavor: str | None = None, **attrs: Any): + return self._operations.operation(name, flavor=flavor, **attrs) def entrypoint( self, @@ -176,15 +121,7 @@ def entrypoint( flavor: str | None = None, **attrs: Any, ): - def decorator(func): - operation_name = name or getattr(func, "__name__", "operation") - return self.operation( - operation_name, - flavor=flavor, - **attrs, - )(func) - - return decorator + return self._operations.entrypoint(name, flavor=flavor, **attrs) def start_operation( self, @@ -197,434 +134,67 @@ def start_operation( record_metric: bool = True, **attrs: Any, ) -> dict[str, Any]: - handle = { - "operation": None, - "operation_token": None, - "timings_token": None, - "start_token": None, - "context_token": None, - } - if not self.enabled: - return handle - try: - operation = OperationObservabilityContext( - config=self.config, - name=self._safe_str(name), - flavor=flavor, - attributes={ - key: value - for key, value in attrs.items() - if value is not None - }, - timings_ms={}, - emit_log=emit_log, - record_metric=record_metric, - ) - operation.context_token = _OPERATION_CONTEXT.set(operation) - handle["operation"] = operation - handle["operation_token"] = operation.context_token - if timings is not None: - handle["timings_token"] = _TIMINGS.set(timings) - handle["start_token"] = _TURN_START.set(time.perf_counter()) - if parent_context is not None and self.tracer is not None: - try: - from opentelemetry import context as otel_context - - handle["context_token"] = otel_context.attach( - parent_context - ) - except BaseException as exc: - self.log_observability_failure( - "operation.context_attach", - exc, - ) - if self.tracer is not None: - operation.span_handle = self._start_span( - self._span_name(operation.name), - operation.span_attributes(), - ) - except BaseException as exc: - self.log_observability_failure("operation.start", exc, name=name) - return handle + return self._operations.start_operation( + name, + flavor=flavor, + parent_context=parent_context, + timings=timings, + emit_log=emit_log, + record_metric=record_metric, + **attrs, + ) def end_operation( - self, - handle: dict[str, Any] | None, - error: BaseException | None = None, + self, handle: dict[str, Any] | None, error: BaseException | None = None ) -> None: - if not handle: - return - operation = handle.get("operation") - try: - if operation is not None and error is not None: - operation.error = ErrorRecord( - type=type(error).__name__, - message=self._safe_str(error), - handled=False, - stack=self._safe_traceback(error), - ) - self.record_error_metric( - operation.metric_attributes( - error_type=type(error).__name__ - ) - ) - if operation is not None: - self.complete_operation(operation) - if operation is not None: - self._end_span(operation.span_handle, error) - except BaseException as exc: - self.log_observability_failure("operation.end", exc) - finally: - context_token = handle.get("context_token") - if context_token is not None: - try: - from opentelemetry import context as otel_context - - otel_context.detach(context_token) - except BaseException as exc: - self.log_observability_failure( - "operation.context_detach", - exc, - ) - for var, key in ( - (_TIMINGS, "timings_token"), - (_TURN_START, "start_token"), - (_OPERATION_CONTEXT, "operation_token"), - ): - token = handle.get(key) - if token is not None: - try: - var.reset(token) - except BaseException as exc: - self.log_observability_failure( - "operation.context_reset", - exc, - token=key, - ) + return self._operations.end_operation(handle, error) def complete_operation( - self, - operation: OperationObservabilityContext, + self, operation: OperationObservabilityContext ) -> None: - if operation.metric_recorded: - return - operation.metric_recorded = True - if operation.record_metric: - self.record_operation_metric( - operation.duration_seconds(), - operation.metric_attributes(), - ) - if operation.emit_log: - self.emit_operation_log(operation) + return self._operations.complete_operation(operation) def begin_request( - self, - context: RequestObservabilityContext, - *, - carrier: Any = None, + self, context: RequestObservabilityContext, *, carrier: Any = None ) -> None: - if not self.enabled: - return - try: - context.context_token = _REQUEST_CONTEXT.set(context) - context.set_attribute("endpoint", context.endpoint) - self._begin_request_operation(context) - self._start_request_span(context, carrier=carrier) - self.record_active_request(1, context.metric_attributes()) - except BaseException as exc: - self.log_observability_failure("request.begin", exc) + return self._requests.begin_request(context, carrier=carrier) - def _begin_request_operation( - self, - context: RequestObservabilityContext, - ) -> None: - try: - parent_operation = _OPERATION_CONTEXT.get() - timings = context.timings_ms - timing_counts = context.timing_counts - segment_tree = context.segment_tree - segment_sequence = context.segment_sequence - if context.internal_dispatch and parent_operation is not None: - timings = parent_operation.timings_ms - timing_counts = parent_operation.timing_counts - segment_tree = parent_operation.segment_tree - segment_sequence = parent_operation.segment_sequence - context.timings_ms = timings - context.timing_counts = timing_counts - context.segment_tree = segment_tree - context.segment_sequence = segment_sequence - operation = OperationObservabilityContext( - config=context.config, - name=context.route, - flavor="http", - attributes={ - "route": context.route, - "method": context.method, - "endpoint": context.endpoint, - "path": context.path, - }, - timings_ms=timings, - timing_counts=timing_counts, - segment_tree=segment_tree, - segment_sequence=segment_sequence, - emit_log=False, - record_metric=False, - ) - operation.context_token = _OPERATION_CONTEXT.set(operation) - context.operation_context = operation - context.operation_token = operation.context_token - except BaseException as exc: - self.log_observability_failure( - "request.operation_begin", - exc, - request_id=getattr(context, "request_id", None), - ) + def _begin_request_operation(self, *args: Any, **kwargs: Any) -> Any: + return self._requests._begin_request_operation(*args, **kwargs) def finish_request(self, status_code: int) -> dict[str, str]: - headers = self.prepare_response(status_code) - self.complete_request(status_code) - return headers + return self._requests.finish_request(status_code) def prepare_response(self, status_code: int) -> dict[str, str]: - if not self.enabled: - return {} - headers: dict[str, str] = {} - try: - context = self.current_context() - if context is None: - return headers - context.status_code = status_code - self._set_current_span_attributes(context.span_attributes()) - if context.operation_context is not None: - context.operation_context.set_attribute( - "status_code", - str(status_code), - ) - headers[REQUEST_ID_HEADER] = context.request_id - traceparent = self.traceparent_header() - if traceparent: - headers[TRACEPARENT_HEADER] = traceparent - if status_code == 429: - context.set_attribute("rate_limited", True) - return headers - except BaseException as exc: - self.log_observability_failure("request.prepare_response", exc) - return headers + return self._requests.prepare_response(status_code) def complete_request(self, status_code: int | None = None) -> None: - if not self.enabled: - return - try: - context = self.current_context() - if context is None: - return - if status_code is not None: - context.status_code = status_code - self._set_current_span_attributes(context.span_attributes()) - if context.request_metric_recorded: - return - context.request_metric_recorded = True - if context.status_code == 429: - self.record_rate_limited_metric(context.metric_attributes()) - self.record_request_metric( - context.duration_seconds(), - context.metric_attributes(), - ) - self._close_active_request(context) - except BaseException as exc: - self.log_observability_failure("request.complete", exc) + return self._requests.complete_request(status_code) def update_request_route( - self, - *, - route: str | None = None, - endpoint: str | None = None, + self, *, route: str | None = None, endpoint: str | None = None ) -> None: - if not self.enabled: - return - try: - context = self.current_context() - if context is None: - return - route_changed = bool(route and route != context.route) - old_active_attributes = ( - context.metric_attributes() - if route_changed and not context.active_closed - else None - ) - if route: - context.route = route - if context.operation_context is not None: - context.operation_context.name = route - context.operation_context.set_attribute("route", route) - if endpoint: - context.endpoint = endpoint - context.set_attribute("endpoint", endpoint) - if context.operation_context is not None: - context.operation_context.set_attribute( - "endpoint", - endpoint, - ) - self._set_current_span_attributes(context.span_attributes()) - span = context.server_span - update_name = getattr(span, "update_name", None) - if route and update_name is not None: - update_name(route) - if old_active_attributes is not None: - self.record_active_request(-1, old_active_attributes) - self.record_active_request(1, context.metric_attributes()) - except BaseException as exc: - self.log_observability_failure("request.update_route", exc) + return self._requests.update_request_route( + route=route, endpoint=endpoint + ) def teardown_request(self, exc: BaseException | None = None) -> None: - if not self.enabled: - return - context = self.current_context() - if context is None: - return - try: - if exc is not None: - self.record_error( - exc, - handled=False, - status_code=context.status_code or 500, - ) - self._close_active_request(context) - self.emit_request_log(context) - except BaseException as observability_exc: - self.log_observability_failure( - "request.teardown", - observability_exc, - ) - finally: - self._close_request_span(context, exc) - self._reset_request_operation_context(context) - self._reset_request_context(context) + return self._requests.teardown_request(exc) def set_attribute(self, key: str, value: Any) -> None: - if not self.enabled: - return - try: - context = self.current_context() - if context is not None: - context.set_attribute(key, value) - if context.operation_context is not None: - context.operation_context.set_attribute(key, value) - self._set_current_span_attributes( - context.span_attributes(**{f"policyengine.{key}": value}) - ) - operation = self.current_operation() - if operation is not None and operation is not getattr( - context, "operation_context", None - ): - operation.set_attribute(key, value) - self._set_current_span_attributes( - operation.span_attributes(**{f"policyengine.{key}": value}) - ) - except BaseException as exc: - self.log_observability_failure( - "request.set_attribute", - exc, - attribute=key, - ) + return self._requests.set_attribute(key, value) def segment(self, name: Any, **attrs: Any) -> Iterator[Any]: - return _SegmentManager(self, name, attrs) + return self._segments.segment(name, **attrs) + + def _segment_context(self, *args: Any, **kwargs: Any) -> Any: + return self._segments._segment_context(*args, **kwargs) + + def asegment(self, name: Any, **attrs: Any) -> AsyncIterator[Any]: + return self._segments.asegment(name, **attrs) - @contextmanager - def _segment_context(self, name: Any, **attrs: Any) -> Iterator[Any]: - if not self.enabled: - yield None - return - segment_name = self._coerce_segment(name) - implicit_operation = self._start_implicit_operation( - segment_name, - attrs, - ) - start = self._safe_perf_counter(f"segment.{segment_name}.start") - segment_tree_handle = self._start_segment_tree_node( - segment_name, - attrs, - ) - span_attrs = self._segment_span_attributes(attrs) - span_name = self._span_name(segment_name) - error: BaseException | None = None - with self._safe_span(span_name, span_attrs) as span: - try: - yield span - except BaseException as exc: - error = exc - self._record_segment_safely( - segment_name, - start, - attrs, - segment_tree_handle=segment_tree_handle, - ) - raise - else: - self._record_segment_safely( - segment_name, - start, - attrs, - segment_tree_handle=segment_tree_handle, - ) - finally: - self._reset_segment_tree_stack(segment_tree_handle) - self.end_operation(implicit_operation, error) - - @asynccontextmanager - async def asegment(self, name: Any, **attrs: Any) -> AsyncIterator[Any]: - if not self.enabled: - yield None - return - segment_name = self._coerce_segment(name) - implicit_operation = self._start_implicit_operation( - segment_name, - attrs, - ) - start = self._safe_perf_counter(f"segment.{segment_name}.start") - segment_tree_handle = self._start_segment_tree_node( - segment_name, - attrs, - ) - span_attrs = self._segment_span_attributes(attrs) - span_name = self._span_name(segment_name) - error: BaseException | None = None - with self._safe_span(span_name, span_attrs) as span: - try: - yield span - except BaseException as exc: - error = exc - self._record_segment_safely( - segment_name, - start, - attrs, - segment_tree_handle=segment_tree_handle, - ) - raise - else: - self._record_segment_safely( - segment_name, - start, - attrs, - segment_tree_handle=segment_tree_handle, - ) - finally: - self._reset_segment_tree_stack(segment_tree_handle) - self.end_operation(implicit_operation, error) - - @contextmanager def collect_timings(self, name: str = "operation", **attrs: Any): - timings: dict[str, float] = {} - handle = self.start_scope(timings, name=name, **attrs) - error: BaseException | None = None - try: - yield timings - except BaseException as exc: - error = exc - raise - finally: - self.end_scope(handle, error) + return self._operations.collect_timings(name, **attrs) def start_scope( self, @@ -634,134 +204,28 @@ def start_scope( parent_context: Any = None, **attrs: Any, ) -> dict[str, Any]: - if self.current_operation() is None: - return { - "operation_handle": self.start_operation( - name, - parent_context=parent_context, - timings=timings, - **attrs, - ) - } - handle = { - "operation_handle": None, - "timings_token": None, - "start_token": None, - "context_token": None, - "span": None, - } - try: - handle["timings_token"] = _TIMINGS.set(timings) - except BaseException as exc: - self.log_observability_failure("scope.timings_set", exc) - try: - handle["start_token"] = _TURN_START.set(time.perf_counter()) - except BaseException as exc: - self.log_observability_failure("scope.start_set", exc) - if parent_context is not None and self.tracer is not None: - try: - from opentelemetry import context as otel_context - - handle["context_token"] = otel_context.attach(parent_context) - except BaseException as exc: - self.log_observability_failure("scope.context_attach", exc) - try: - if self.tracer is not None: - handle["span"] = self._start_span(name, attrs) - except BaseException as exc: - self.log_observability_failure("scope.span_start", exc, span=name) - handle["span"] = None - return handle + return self._operations.start_scope( + timings, name=name, parent_context=parent_context, **attrs + ) def annotate( - self, - handle: dict[str, Any] | None = None, - **attrs: Any, + self, handle: dict[str, Any] | None = None, **attrs: Any ) -> None: - try: - if handle: - span_handle = handle.get("span") - if span_handle is not None: - _cm, span = span_handle - for key, value in attrs.items(): - if value is not None: - span.set_attribute(key, value) - context = self.current_context() - if context is not None: - for key, value in attrs.items(): - context.set_attribute(key, value) - operation = self.current_operation() - if operation is not None: - for key, value in attrs.items(): - operation.set_attribute(key, value) - self._set_current_span_attributes(operation.span_attributes()) - except BaseException as exc: - self.log_observability_failure("scope.annotate", exc) + return self._operations.annotate(handle, **attrs) def end_scope( - self, - handle: dict[str, Any] | None, - error: BaseException | None = None, + self, handle: dict[str, Any] | None, error: BaseException | None = None ) -> None: - if not handle: - return - operation_handle = handle.get("operation_handle") - if operation_handle is not None: - self.end_operation(operation_handle, error) - return - try: - self._end_span(handle.get("span"), error) - except BaseException as exc: - self.log_observability_failure("scope.span_end", exc) - context_token = handle.get("context_token") - if context_token is not None: - try: - from opentelemetry import context as otel_context - - otel_context.detach(context_token) - except BaseException as exc: - self.log_observability_failure("scope.context_detach", exc) - for var, key in ( - (_TIMINGS, "timings_token"), - (_TURN_START, "start_token"), - ): - token = handle.get(key) - if token is not None: - try: - var.reset(token) - except BaseException as exc: - self.log_observability_failure( - "scope.context_reset", - exc, - token=key, - ) + return self._operations.end_scope(handle, error) def mark(self, key: str, ms: float) -> None: - try: - timings = _TIMINGS.get() - if timings is not None: - timings[key] = round(float(ms), 1) - except BaseException as exc: - self.log_observability_failure("scope.mark", exc, key=key) + return self._operations.mark(key, ms) def mark_ttft(self, key: str = "ttft_ms") -> None: - try: - start = _TURN_START.get() - if start is not None: - self.mark(key, (time.perf_counter() - start) * 1000.0) - except BaseException as exc: - self.log_observability_failure("scope.mark_ttft", exc) + return self._operations.mark_ttft(key) def mark_ttft_attribute(self, key: str = "ttft_ms") -> None: - try: - start = _TURN_START.get() - if start is None: - return - self.annotate( - **{key: round((time.perf_counter() - start) * 1000.0, 1)} - ) - except BaseException as exc: - self.log_observability_failure("scope.mark_ttft_attribute", exc) + return self._operations.mark_ttft_attribute(key) def record_error( self, @@ -771,217 +235,43 @@ def record_error( status_code: int | None = None, include_stack: bool = True, ) -> None: - if not self.enabled: - return - try: - context = self.current_context() - operation = self.current_operation() - error_record = ErrorRecord( - type=type(exc).__name__, - message=self._safe_str(exc), - handled=handled, - stack=(self._safe_traceback(exc) if include_stack else None), - ) - if context is not None: - if status_code is not None: - context.status_code = status_code - context.error = error_record - self.record_error_metric( - context.metric_attributes(error_type=type(exc).__name__) - ) - elif operation is not None: - operation.error = error_record - self.record_error_metric( - operation.metric_attributes(error_type=type(exc).__name__) - ) - else: - return - span = self._current_span() - if span is not None: - self._record_exception_on_span( - span, - exc, - handled=handled, - status_code=status_code, - ) - except BaseException as observability_exc: - self.log_observability_failure( - "request.record_error", - observability_exc, - original_error_type=type(exc).__name__, - ) + return self._logging.record_error( + exc, + handled=handled, + status_code=status_code, + include_stack=include_stack, + ) def record_event(self, event: str, **fields: Any) -> None: - if not self.enabled: - return - try: - context = self.current_context() - operation = self.current_operation() - base: dict[str, Any] = { - "schema_version": "policyengine.observability.event.v1", - "event": event, - "service_name": self.config.service_name, - "service_role": self.config.service_role, - "environment": self.config.environment, - "created_at": datetime.now(UTC).isoformat(), - } - if context is not None: - trace_id, span_id = self._trace_ids() - base.update( - { - "service_name": context.config.service_name, - "service_role": context.config.service_role, - "environment": context.config.environment, - "request_id": context.request_id, - "trace_id": trace_id, - "span_id": span_id, - "route": context.route, - "path": context.path, - } - ) - elif operation is not None: - trace_id, span_id = self._trace_ids() - base.update( - { - "service_name": operation.config.service_name, - "service_role": operation.config.service_role, - "environment": operation.config.environment, - "operation": operation.name, - "flavor": operation.flavor, - "trace_id": trace_id, - "span_id": span_id, - } - ) - clean_fields = { - key: value - for key, value in fields.items() - if value is not None - } - base.update(clean_fields) - self._emit_structured_log( - base, - log_type="event", - severity="INFO", - ) - self._add_span_event(event, clean_fields) - if event.startswith("modal_") or "fallback" in event: - attrs = ( - context.metric_attributes(event=event) - if context - else operation.metric_attributes(event=event) - if operation - else _metric_attrs( - {"event": event}, - self.config.metric_attribute_keys, - ) - ) - self.record_failover_event_metric(attrs) - except BaseException as exc: - self.log_observability_failure( - "request.record_event", - exc, - event_name=event, - ) + return self._logging.record_event(event, **fields) def traceparent_header(self) -> str | None: - if not self.enabled or self.propagate is None: - return None - try: - carrier: dict[str, str] = {} - self.propagate.inject(carrier) - return carrier.get(TRACEPARENT_HEADER) - except BaseException as exc: - self.log_observability_failure("request.traceparent_header", exc) - return None + return self._tracing.traceparent_header() def capture_context(self): - if self.tracer is None: - return None - try: - from opentelemetry import context as otel_context - - return otel_context.get_current() - except BaseException as exc: - self.log_observability_failure("otel.capture_context", exc) - return None + return self._tracing.capture_context() def emit_request_log(self, context: RequestObservabilityContext) -> None: - if not self.enabled: - return - try: - if context.emitted: - return - context.emitted = True - if ( - context.internal_dispatch - or not context.config.request_logs_enabled - ): - return - trace_id, span_id = self._trace_ids() - payload = context.as_log_record( - trace_id=trace_id, - span_id=span_id, - ) - self._emit_structured_log( - payload, - log_type="request", - severity=self._severity_for_log_record(payload), - ) - except BaseException as exc: - self.log_observability_failure( - "request.emit_request_log", - exc, - request_id=getattr(context, "request_id", None), - ) + return self._logging.emit_request_log(context) def emit_operation_log( - self, - operation: OperationObservabilityContext, + self, operation: OperationObservabilityContext ) -> None: - if not self.enabled: - return - try: - if operation.emitted: - return - operation.emitted = True - trace_id, span_id = self._trace_ids() - payload = operation.as_log_record( - trace_id=trace_id, - span_id=span_id, - ) - self._emit_structured_log( - payload, - log_type="operation", - severity=self._severity_for_log_record(payload), - ) - except BaseException as exc: - self.log_observability_failure( - "operation.emit_log", - exc, - operation=getattr(operation, "name", None), - ) + return self._logging.emit_operation_log(operation) def record_operation_metric( - self, - duration_seconds: float, - attributes: dict[str, str], + self, duration_seconds: float, attributes: dict[str, str] ) -> None: - try: - self.operation_duration.record(duration_seconds, attributes) - self.operations.add(1, attributes) - except BaseException as exc: - self.log_observability_failure("metrics.record_operation", exc) + return self._metrics.record_operation_metric( + duration_seconds, attributes + ) def record_request_metric( - self, - duration_seconds: float, - attributes: dict[str, str], + self, duration_seconds: float, attributes: dict[str, str] ) -> None: - try: - self.http_duration.record(duration_seconds, attributes) - self.requests.add(1, attributes) - except BaseException as exc: - self.log_observability_failure("metrics.record_request", exc) + return self._metrics.record_request_metric( + duration_seconds, attributes + ) def record_segment_metric( self, @@ -991,85 +281,32 @@ def record_segment_metric( *, backend_segment: bool = False, ) -> None: - try: - segment_attributes = {**attributes, "segment": segment} - self.segment_duration.record(duration_seconds, segment_attributes) - if segment == "calculation": - self.calculate_duration.record(duration_seconds, attributes) - if backend_segment: - self.backend_duration.record( - duration_seconds, - segment_attributes, - ) - except BaseException as exc: - self.log_observability_failure( - "metrics.record_segment", - exc, - segment=segment, - ) + return self._metrics.record_segment_metric( + segment, + duration_seconds, + attributes, + backend_segment=backend_segment, + ) def record_error_metric(self, attributes: dict[str, str]) -> None: - try: - self.errors.add(1, attributes) - except BaseException as exc: - self.log_observability_failure("metrics.record_error", exc) + return self._metrics.record_error_metric(attributes) def record_rate_limited_metric(self, attributes: dict[str, str]) -> None: - try: - self.rate_limited.add(1, attributes) - except BaseException as exc: - self.log_observability_failure("metrics.record_rate_limited", exc) + return self._metrics.record_rate_limited_metric(attributes) def record_failover_event_metric(self, attributes: dict[str, str]) -> None: - try: - self.failover_events.add(1, attributes) - except BaseException as exc: - self.log_observability_failure( - "metrics.record_failover_event", - exc, - ) + return self._metrics.record_failover_event_metric(attributes) def record_active_request( - self, - delta: int, - attributes: dict[str, str], + self, delta: int, attributes: dict[str, str] ) -> None: - try: - self.active_requests.add(delta, attributes) - except BaseException as exc: - self.log_observability_failure("metrics.add_active_request", exc) + return self._metrics.record_active_request(delta, attributes) def instrument_fastapi(self, app: Any) -> None: - if not self.enabled or not self.config.otel_enabled: - return - try: - from opentelemetry.instrumentation.fastapi import ( - FastAPIInstrumentor, - ) - - FastAPIInstrumentor.instrument_app(app) - except BaseException as exc: - self.log_observability_failure( - "fastapi.auto_instrument", - exc, - ) + return self._tracing.instrument_fastapi(app) def instrument_httpx(self) -> None: - if ( - not self.enabled - or not self.config.otel_enabled - or self._httpx_instrumented - ): - return - try: - from opentelemetry.instrumentation.httpx import ( - HTTPXClientInstrumentor, - ) - - HTTPXClientInstrumentor().instrument() - self._httpx_instrumented = True - except BaseException as exc: - self.log_observability_failure("httpx.auto_instrument", exc) + return self._tracing.instrument_httpx() def shutdown(self) -> None: budget = clamped( @@ -1146,969 +383,139 @@ def restart_log_destinations(self) -> None: self.log_destination_manager.configure() def log_observability_failure( - self, - operation: str, - exc: BaseException, - **fields: Any, - ) -> None: - payload = self._internal_error_payload(operation, exc, **fields) - if self._emitting_internal_error: - self._write_stderr(payload) - return - self._emitting_internal_error = True - try: - self._emit_structured_log( - payload, - log_type="internal", - severity="ERROR", - ) - except BaseException: - self._write_stderr(payload) - finally: - self._emitting_internal_error = False - - def _configure_loggers(self) -> None: - for logger in ( - REQUEST_LOGGER, - OPERATION_LOGGER, - EVENT_LOGGER, - INTERNAL_LOGGER, - ): - configure_plain_logger(logger, self.config.log_level) - - def _emit_structured_log( - self, - payload: dict[str, Any], - *, - log_type: str, - severity: str, + self, operation: str, exc: BaseException, **fields: Any ) -> None: - try: - if not self.log_destination_manager.configured: - self._configure_loggers() - self.log_destination_manager.emit( - payload, - log_type=log_type, - severity=severity, - ) - except BaseException as exc: - if self._emitting_internal_error: - self._write_stderr(payload) - else: - self.log_observability_failure( - "logging.emit", - exc, - log_type=log_type, - ) - - def _handle_destination_failure( - self, - operation: str, - exc: BaseException, - **fields: Any, - ) -> None: - if self._emitting_internal_error: - self._write_stderr( - self._internal_error_payload(operation, exc, **fields) - ) - return - self.log_observability_failure(operation, exc, **fields) - - def _severity_for_log_record(self, payload: dict[str, Any]) -> str: - status_code = self._int_or_none(payload.get("status_code")) - error = payload.get("error") - handled = error.get("handled") if isinstance(error, dict) else None - if status_code is not None and status_code >= 500: - return "ERROR" - if error is not None and handled is False: - return "ERROR" - if error is not None or ( - status_code is not None and status_code >= 400 - ): - return "WARNING" - return "INFO" - - def _int_or_none(self, value: Any) -> int | None: - try: - return int(value) - except (TypeError, ValueError): - return None - - def _internal_error_payload( - self, - operation: str, - exc: BaseException, - **fields: Any, - ) -> dict[str, Any]: - payload = { - "schema_version": "policyengine.observability.internal_error.v1", - "event": "observability_internal_error", - "service_name": self.config.service_name, - "service_role": self.config.service_role, - "environment": self.config.environment, - "created_at": datetime.now(UTC).isoformat(), - "operation": operation, - "error": { - "type": type(exc).__name__, - "message": self._safe_str(exc), - "stack": self._safe_traceback(exc), - }, - } - payload.update( - {key: value for key, value in fields.items() if value is not None} + return self._logging.log_observability_failure( + operation, exc, **fields ) - return payload - def _configure_otel(self) -> None: - try: - from opentelemetry import metrics, propagate, trace - from opentelemetry.sdk.metrics import MeterProvider - from opentelemetry.sdk.resources import ( - DEPLOYMENT_ENVIRONMENT, - SERVICE_NAME, - Resource, - ) - from opentelemetry.sdk.trace import TracerProvider - from opentelemetry.trace import SpanKind, Status, StatusCode - except BaseException as exc: - self.log_observability_failure("otel.configure_imports", exc) - return + def _configure_loggers(self, *args: Any, **kwargs: Any) -> Any: + return self._logging._configure_loggers(*args, **kwargs) - try: - resource = Resource.create( - { - SERVICE_NAME: self.config.service_name, - DEPLOYMENT_ENVIRONMENT: self.config.environment, - "service.role": self.config.service_role, - } - ) - tracer_provider = TracerProvider(resource=resource) - metric_readers = [] - if self.config.otlp_endpoint: - self._add_trace_exporter(tracer_provider) - metric_reader = self._metric_reader() - if metric_reader is not None: - metric_readers.append(metric_reader) - self.tracer_provider = tracer_provider - try: - trace.set_tracer_provider(tracer_provider) - except BaseException as exc: - self.log_observability_failure( - "otel.set_tracer_provider", - exc, - ) - try: - self.meter_provider = MeterProvider( - resource=resource, - metric_readers=metric_readers, - ) - metrics.set_meter_provider(self.meter_provider) - except BaseException as exc: - self.log_observability_failure( - "otel.set_meter_provider", - exc, - ) - self.trace = trace - self.propagate = propagate - self.SpanKind = SpanKind - self.Status = Status - self.StatusCode = StatusCode - tracer_name = self.config.tracer_name or self.config.service_name - meter_name = self.config.meter_name or self.config.service_name - self.tracer = trace.get_tracer(tracer_name) - self.meter = metrics.get_meter(meter_name) - self._configure_instruments() - except BaseException as exc: - self.log_observability_failure("otel.configure", exc) + def _emit_structured_log(self, *args: Any, **kwargs: Any) -> Any: + return self._logging._emit_structured_log(*args, **kwargs) - def _add_trace_exporter(self, tracer_provider) -> None: - try: - if self.config.otlp_protocol.startswith("http"): - from opentelemetry.exporter.otlp.proto.http.trace_exporter import ( - OTLPSpanExporter, - ) - else: - from opentelemetry.exporter.otlp.proto.grpc.trace_exporter import ( - OTLPSpanExporter, - ) - from opentelemetry.sdk.trace.export import BatchSpanProcessor - - tracer_provider.add_span_processor( - BatchSpanProcessor(OTLPSpanExporter()) - ) - except BaseException as exc: - self.log_observability_failure("otel.trace_exporter", exc) + def _handle_destination_failure(self, *args: Any, **kwargs: Any) -> Any: + return self._logging._handle_destination_failure(*args, **kwargs) - def _metric_reader(self): - try: - if self.config.otlp_protocol.startswith("http"): - from opentelemetry.exporter.otlp.proto.http.metric_exporter import ( - OTLPMetricExporter, - ) - else: - from opentelemetry.exporter.otlp.proto.grpc.metric_exporter import ( - OTLPMetricExporter, - ) - from opentelemetry.sdk.metrics.export import ( - PeriodicExportingMetricReader, - ) + def _severity_for_log_record(self, *args: Any, **kwargs: Any) -> Any: + return self._logging._severity_for_log_record(*args, **kwargs) - return PeriodicExportingMetricReader(OTLPMetricExporter()) - except BaseException as exc: - self.log_observability_failure("otel.metric_exporter", exc) - return None - - def _configure_instruments(self) -> None: - self.operation_duration = self._instrument( - getattr(self.meter, "create_histogram", None), - "policyengine.operation.duration", - unit="s", - description="PolicyEngine operation duration.", - ) - self.http_duration = self._instrument( - getattr(self.meter, "create_histogram", None), - "http.server.request.duration", - unit="s", - description="HTTP server request duration.", - ) - self.segment_duration = self._instrument( - getattr(self.meter, "create_histogram", None), - "policyengine.segment.duration", - unit="s", - description="PolicyEngine operation segment duration.", - ) - self.calculate_duration = self._instrument( - getattr(self.meter, "create_histogram", None), - "policyengine.calculate.duration", - unit="s", - description="PolicyEngine calculate operation duration.", - ) - self.backend_duration = self._instrument( - getattr(self.meter, "create_histogram", None), - "policyengine.backend.duration", - unit="s", - description="PolicyEngine backend call duration.", - ) - self.operations = self._instrument( - getattr(self.meter, "create_counter", None), - "policyengine.operations", - description="PolicyEngine operation count.", - ) - self.requests = self._instrument( - getattr(self.meter, "create_counter", None), - "policyengine.requests", - description="PolicyEngine request count.", - ) - self.errors = self._instrument( - getattr(self.meter, "create_counter", None), - "policyengine.errors", - description="PolicyEngine error count.", - ) - self.rate_limited = self._instrument( - getattr(self.meter, "create_counter", None), - "policyengine.rate_limited_requests", - description="PolicyEngine rate-limited request count.", - ) - self.failover_events = self._instrument( - getattr(self.meter, "create_counter", None), - "policyengine.failover.events", - description="PolicyEngine failover event count.", - ) - self.active_requests = self._instrument( - getattr(self.meter, "create_up_down_counter", None), - "http.server.active_requests", - description="Active HTTP server requests.", - ) + def _int_or_none(self, *args: Any, **kwargs: Any) -> Any: + return self._logging._int_or_none(*args, **kwargs) - def _instrument(self, factory, *args, **kwargs): - if factory is None: - return _NoOpInstrument() - try: - return factory(*args, **kwargs) - except BaseException as exc: - self.log_observability_failure( - "metrics.create_instrument", - exc, - instrument=args[0] if args else None, - ) - return _NoOpInstrument() + def _internal_error_payload(self, *args: Any, **kwargs: Any) -> Any: + return self._logging._internal_error_payload(*args, **kwargs) - def _start_request_span( - self, - context: RequestObservabilityContext, - *, - carrier: Any = None, - ) -> None: - if self.tracer is None: - return - attrs = context.span_attributes() - parent_context = self._extract_context(carrier) - try: - context.server_span_cm = self.tracer.start_as_current_span( - context.route, - context=parent_context, - kind=self.SpanKind.SERVER if self.SpanKind else None, - attributes=attrs, - ) - context.server_span = context.server_span_cm.__enter__() - except BaseException as exc: - context.server_span_cm = None - context.server_span = None - self.log_observability_failure("otel.request_span_enter", exc) + def _configure_otel(self, *args: Any, **kwargs: Any) -> Any: + return self._tracing._configure_otel(*args, **kwargs) - def _close_request_span( - self, - context: RequestObservabilityContext, - exc: BaseException | None, - ) -> None: - if context.span_closed: - return - context.span_closed = True - span_cm = context.server_span_cm - if span_cm is None: - return - try: - if exc is None: - span_cm.__exit__(None, None, None) - else: - span_cm.__exit__(type(exc), exc, exc.__traceback__) - except BaseException as observability_exc: - self.log_observability_failure( - "otel.request_span_exit", - observability_exc, - request_id=context.request_id, - ) + def _add_trace_exporter(self, *args: Any, **kwargs: Any) -> Any: + return self._tracing._add_trace_exporter(*args, **kwargs) - @contextmanager - def _safe_span(self, name: str, attrs: dict[str, Any]) -> Iterator[Any]: - if self.tracer is None: - yield None - return - span_handle = self._start_span(name, attrs) - if span_handle is None: - yield None - return - _cm, span = span_handle - try: - yield span - except BaseException as exc: - try: - self._end_span(span_handle, exc) - except BaseException as observability_exc: - self.log_observability_failure( - "otel.span_exit", - observability_exc, - span=name, - ) - raise - else: - try: - self._end_span(span_handle) - except BaseException as exc: - self.log_observability_failure( - "otel.span_exit", - exc, - span=name, - ) - - def _start_span(self, name: str, attrs: dict[str, Any]): - try: - span_cm = self.tracer.start_as_current_span(name) - span = span_cm.__enter__() - except BaseException as exc: - self.log_observability_failure("otel.span_enter", exc, span=name) - return None - try: - for key, value in attrs.items(): - if value is not None: - span.set_attribute(key, value) - except BaseException as exc: - self.log_observability_failure( - "otel.span_attributes", - exc, - span=name, - ) - return span_cm, span + def _metric_reader(self, *args: Any, **kwargs: Any) -> Any: + return self._tracing._metric_reader(*args, **kwargs) - def _end_span( - self, - span_handle, - error: BaseException | None = None, - ) -> None: - if span_handle is None: - return - span_cm, span = span_handle - try: - if error is not None: - self._record_exception_on_span( - span, - error, - handled=False, - status_code=500, - ) - except BaseException as exc: - self.log_observability_failure("otel.span_error_status", exc) - try: - span_cm.__exit__(None, None, None) - except BaseException as exc: - self.log_observability_failure("otel.span_exit", exc) + def _configure_instruments(self, *args: Any, **kwargs: Any) -> Any: + return self._metrics._configure_instruments(*args, **kwargs) - def _start_segment_tree_node( - self, - name: str, - attrs: dict[str, Any], - ) -> dict[str, Any] | None: - try: - owner = self._segment_tree_owner() - if owner is None: - return None - owner.segment_sequence[0] += 1 - node = SegmentTimingNode( - sequence=owner.segment_sequence[0], - name=name, - attrs=self._safe_segment_tree_attrs(attrs), - ) - owner_id = id(owner.segment_tree) - stack = _SEGMENT_STACK.get() - if stack and stack[-1][0] == owner_id: - stack[-1][1].children.append(node) - else: - owner.segment_tree.append(node) - token = _SEGMENT_STACK.set((*stack, (owner_id, node))) - return {"node": node, "token": token} - except BaseException as exc: - self.log_observability_failure( - "segment.tree_start", - exc, - segment=name, - ) - return None + def _instrument(self, *args: Any, **kwargs: Any) -> Any: + return self._metrics._instrument(*args, **kwargs) - def _finish_segment_tree_node( - self, - handle: dict[str, Any] | None, - duration_seconds: float, - ) -> None: - if not handle: - return - try: - node = handle.get("node") - if not isinstance(node, SegmentTimingNode): - return - node.duration_ms = duration_seconds * 1000 - except BaseException as exc: - self.log_observability_failure( - "segment.tree_finish", - exc, - ) + def _start_request_span(self, *args: Any, **kwargs: Any) -> Any: + return self._tracing._start_request_span(*args, **kwargs) - def _reset_segment_tree_stack( - self, - handle: dict[str, Any] | None, - ) -> None: - if not handle: - return - token = handle.get("token") - if token is None: - return - try: - _SEGMENT_STACK.reset(token) - except BaseException as exc: - self.log_observability_failure("segment.tree_reset", exc) + def _close_request_span(self, *args: Any, **kwargs: Any) -> Any: + return self._tracing._close_request_span(*args, **kwargs) - def _segment_tree_owner( - self, - ) -> RequestObservabilityContext | OperationObservabilityContext | None: - context = self.current_context() - if context is not None: - return context - return self.current_operation() + def _safe_span(self, *args: Any, **kwargs: Any) -> Any: + return self._tracing._safe_span(*args, **kwargs) - def _safe_segment_tree_attrs( - self, - attrs: dict[str, Any], - ) -> dict[str, Any]: - safe_attrs: dict[str, Any] = {} - for key, value in attrs.items(): - key_text = self._safe_str(key) - key_lower = key_text.lower() - if any(part in key_lower for part in SENSITIVE_SEGMENT_ATTR_PARTS): - continue - if value is None: - continue - if hasattr(value, "value"): - value = value.value - if isinstance(value, bool | int | float): - safe_attrs[key_text] = value - elif isinstance(value, str): - safe_attrs[key_text] = value[:MAX_SEGMENT_ATTR_LENGTH] - return safe_attrs - - def _record_segment_flat_timing( - self, - context: RequestObservabilityContext | None, - operation: OperationObservabilityContext | None, - name: str, - duration_ms: float, - ) -> None: - seen_timing_ids: set[int] = set() - seen_count_ids: set[int] = set() - for target in (context, operation): - if target is None: - continue - timings_id = id(target.timings_ms) - if timings_id not in seen_timing_ids: - target.timings_ms[name] = round( - target.timings_ms.get(name, 0.0) + duration_ms, - 3, - ) - seen_timing_ids.add(timings_id) - counts_id = id(target.timing_counts) - if counts_id not in seen_count_ids: - target.timing_counts[name] = ( - target.timing_counts.get(name, 0) + 1 - ) - seen_count_ids.add(counts_id) - - def _record_segment_safely( - self, - name: str, - start: float | None, - attrs: dict[str, Any], - *, - segment_tree_handle: dict[str, Any] | None = None, - ) -> None: - if start is None: - return - end = self._safe_perf_counter(f"segment.{name}.end") - if end is None: - return - try: - duration = end - start - self._finish_segment_tree_node(segment_tree_handle, duration) - self._record_timing(name, duration) - context = self.current_context() - operation = self.current_operation() - metric_extra = { - key: value - for key, value in attrs.items() - if ( - key in self.config.metric_attribute_keys - and value is not None - ) - } - duration_ms = duration * 1000 - self._record_segment_flat_timing( - context, - operation, - name, - duration_ms, - ) - if operation is not None: - metric_attributes = operation.metric_attributes( - segment=name, - **metric_extra, - ) - elif context is not None: - metric_attributes = context.metric_attributes( - segment=name, - **metric_extra, - ) - else: - metric_attributes = _metric_attrs( - { - "service.name": self.config.service_name, - "service.role": self.config.service_role, - "deployment.environment": self.config.environment, - "segment": name, - **metric_extra, - }, - self.config.metric_attribute_keys, - ) - self.record_segment_metric( - name, - duration, - metric_attributes, - backend_segment="backend" in metric_extra, - ) - except BaseException as exc: - self.log_observability_failure( - "request.record_segment", - exc, - segment=name, - ) + def _start_span(self, *args: Any, **kwargs: Any) -> Any: + return self._tracing._start_span(*args, **kwargs) - def _record_timing(self, name: str, duration_seconds: float) -> None: - try: - timings = _TIMINGS.get() - if timings is None: - return - key = f"{name}_ms" - duration_ms = duration_seconds * 1000.0 - timings[key] = round(timings.get(key, 0.0) + duration_ms, 1) - except BaseException as exc: - self.log_observability_failure( - "scope.record_timing", - exc, - segment=name, - ) + def _end_span(self, *args: Any, **kwargs: Any) -> Any: + return self._tracing._end_span(*args, **kwargs) - def _segment_span_attributes( - self, - attrs: dict[str, Any], - ) -> dict[str, Any]: - context = self.current_context() - operation = self.current_operation() - span_attrs = { - key: value for key, value in attrs.items() if value is not None - } - if context is not None: - span_attrs = {**context.span_attributes(), **span_attrs} - elif operation is not None: - span_attrs = {**operation.span_attributes(), **span_attrs} - return span_attrs - - def _span_name(self, segment_name: str) -> str: - if not self.config.span_prefix: - return segment_name - return f"{self.config.span_prefix}.{segment_name}" - - def _start_implicit_operation( - self, - segment_name: str, - attrs: dict[str, Any], - ) -> dict[str, Any] | None: - if ( - self.current_operation() is not None - or self.current_context() is not None - ): - return None - operation_name = attrs.get("operation") or segment_name - flavor = attrs.get("flavor") - operation_attrs = { - key: value - for key, value in attrs.items() - if key not in {"operation", "flavor"} and value is not None - } - return self.start_operation( - self._safe_str(operation_name), - flavor=self._safe_str(flavor) if flavor is not None else None, - **operation_attrs, - ) + def _start_segment_tree_node(self, *args: Any, **kwargs: Any) -> Any: + return self._segments._start_segment_tree_node(*args, **kwargs) - def _coerce_segment(self, name: Any) -> str: - segment, is_registered = coerce_segment_name( - name, - registry=self.segment_registry, - ) - if not is_registered: - self.log_observability_failure( - "segment.coerce", - ValueError("Unregistered observability segment."), - segment=segment, - segment_type=type(name).__name__, - ) - return segment + def _finish_segment_tree_node(self, *args: Any, **kwargs: Any) -> Any: + return self._segments._finish_segment_tree_node(*args, **kwargs) - def _set_current_span_attributes(self, attrs: dict[str, Any]) -> None: - span = self._current_span() - if span is None: - return - try: - for key, value in attrs.items(): - if value is not None: - span.set_attribute(key, value) - except BaseException as exc: - self.log_observability_failure("otel.set_span_attributes", exc) + def _reset_segment_tree_stack(self, *args: Any, **kwargs: Any) -> Any: + return self._segments._reset_segment_tree_stack(*args, **kwargs) - def _current_span(self): - if self.trace is None: - return None - try: - return self.trace.get_current_span() - except BaseException as exc: - self.log_observability_failure("otel.current_span", exc) - return None + def _segment_tree_owner(self, *args: Any, **kwargs: Any) -> Any: + return self._segments._segment_tree_owner(*args, **kwargs) - def _trace_ids(self) -> tuple[str | None, str | None]: - span = self._current_span() - if span is None: - return None, None - try: - context = span.get_span_context() - except BaseException as exc: - self.log_observability_failure("otel.span_context", exc) - return None, None - if not getattr(context, "is_valid", False): - return None, None - return f"{context.trace_id:032x}", f"{context.span_id:016x}" - - def _extract_context(self, carrier: Any): - if self.propagate is None or carrier is None: - return None - try: - return self.propagate.extract(carrier) - except BaseException as exc: - self.log_observability_failure("otel.extract_context", exc) - return None + def _safe_segment_tree_attrs(self, *args: Any, **kwargs: Any) -> Any: + return self._segments._safe_segment_tree_attrs(*args, **kwargs) - def _record_exception_on_span( - self, - span, - exc: BaseException, - *, - handled: bool, - status_code: int | None, - ) -> None: - try: - span.record_exception(exc) - span.set_attribute("error.type", type(exc).__name__) - span.set_attribute("error.handled", handled) - if ( - self.Status is not None - and self.StatusCode is not None - and ( - not handled - or (status_code is not None and status_code >= 500) - ) - ): - span.set_status( - self.Status( - self.StatusCode.ERROR, - self._safe_str(exc), - ) - ) - except BaseException as observability_exc: - self.log_observability_failure( - "otel.record_exception", - observability_exc, - original_error_type=type(exc).__name__, - ) + def _record_segment_flat_timing(self, *args: Any, **kwargs: Any) -> Any: + return self._segments._record_segment_flat_timing(*args, **kwargs) - def _add_span_event(self, event: str, fields: dict[str, Any]) -> None: - span = self._current_span() - if span is None: - return - try: - span.add_event( - event, - { - key: value - for key, value in fields.items() - if _is_safe_span_value(value) - }, - ) - except BaseException as exc: - self.log_observability_failure( - "otel.add_event", - exc, - event_name=event, - ) + def _record_segment_safely(self, *args: Any, **kwargs: Any) -> Any: + return self._segments._record_segment_safely(*args, **kwargs) - def _close_active_request( - self, - context: RequestObservabilityContext, - ) -> None: - try: - if context.active_closed: - return - context.active_closed = True - self.record_active_request(-1, context.metric_attributes()) - except BaseException as exc: - self.log_observability_failure( - "request.close_active", - exc, - request_id=getattr(context, "request_id", None), - ) - - def _reset_request_operation_context( - self, - context: RequestObservabilityContext, - ) -> None: - token = context.operation_token - if token is None: - return - try: - _OPERATION_CONTEXT.reset(token) - except BaseException as exc: - self.log_observability_failure( - "request.operation_context_reset", - exc, - request_id=getattr(context, "request_id", None), - ) - - def _reset_request_context( - self, - context: RequestObservabilityContext, - ) -> None: - token = context.context_token - if token is None: - return - try: - _REQUEST_CONTEXT.reset(token) - except BaseException as exc: - self.log_observability_failure( - "request.context_reset", - exc, - request_id=getattr(context, "request_id", None), - ) - - def _safe_perf_counter(self, operation: str) -> float | None: - try: - return time.perf_counter() - except BaseException as exc: - self.log_observability_failure(operation, exc) - return None - - def _safe_str(self, value: Any) -> str: - try: - return str(value) - except BaseException: - return f"" - - def _safe_traceback(self, exc: BaseException) -> str: - try: - return "".join( - traceback.format_exception(type(exc), exc, exc.__traceback__) - ) - except BaseException: - return "" - - def _json(self, payload: dict[str, Any]) -> str: - try: - return json.dumps(payload, sort_keys=True, default=str) - except BaseException: - return json.dumps( - { - "schema_version": "policyengine.observability.internal_error.v1", - "event": "observability_internal_error", - "created_at": datetime.now(UTC).isoformat(), - "operation": "observability.failure_json", - }, - sort_keys=True, - ) - - def _write_stderr(self, payload: dict[str, Any]) -> None: - try: - sys.stderr.write(self._json(payload) + "\n") - except BaseException: - return + def _record_timing(self, *args: Any, **kwargs: Any) -> Any: + return self._segments._record_timing(*args, **kwargs) + def _segment_span_attributes(self, *args: Any, **kwargs: Any) -> Any: + return self._tracing._segment_span_attributes(*args, **kwargs) -def _is_safe_span_value(value: Any) -> bool: - return isinstance(value, str | bool | int | float) + def _span_name(self, *args: Any, **kwargs: Any) -> Any: + return self._tracing._span_name(*args, **kwargs) + def _start_implicit_operation(self, *args: Any, **kwargs: Any) -> Any: + return self._operations._start_implicit_operation(*args, **kwargs) -class _OperationManager: - def __init__( - self, - runtime: ObservabilityRuntime, - name: str, - *, - flavor: str | None, - attrs: dict[str, Any], - ) -> None: - self.runtime = runtime - self.name = name - self.flavor = flavor - self.attrs = attrs - self.handle: dict[str, Any] | None = None - - def __enter__(self): - self.handle = self.runtime.start_operation( - self.name, - flavor=self.flavor, - **self.attrs, - ) - return self.runtime.current_operation() - - def __exit__(self, exc_type, exc, _traceback) -> bool: - self.runtime.end_operation(self.handle, exc) - return False + def _coerce_segment(self, *args: Any, **kwargs: Any) -> Any: + return self._segments._coerce_segment(*args, **kwargs) - async def __aenter__(self): - return self.__enter__() + def _set_current_span_attributes(self, *args: Any, **kwargs: Any) -> Any: + return self._tracing._set_current_span_attributes(*args, **kwargs) - async def __aexit__(self, exc_type, exc, traceback) -> bool: - return self.__exit__(exc_type, exc, traceback) + def _current_span(self, *args: Any, **kwargs: Any) -> Any: + return self._tracing._current_span(*args, **kwargs) - def __call__(self, func): - if inspect.iscoroutinefunction(func): + def _trace_ids(self, *args: Any, **kwargs: Any) -> Any: + return self._tracing._trace_ids(*args, **kwargs) - @wraps(func) - async def async_wrapper(*args, **kwargs): - async with self.runtime.operation( - self.name, - flavor=self.flavor, - **self.attrs, - ): - return await func(*args, **kwargs) + def _extract_context(self, *args: Any, **kwargs: Any) -> Any: + return self._tracing._extract_context(*args, **kwargs) - return async_wrapper + def _record_exception_on_span(self, *args: Any, **kwargs: Any) -> Any: + return self._tracing._record_exception_on_span(*args, **kwargs) - @wraps(func) - def wrapper(*args, **kwargs): - with self.runtime.operation( - self.name, - flavor=self.flavor, - **self.attrs, - ): - return func(*args, **kwargs) + def _add_span_event(self, *args: Any, **kwargs: Any) -> Any: + return self._tracing._add_span_event(*args, **kwargs) - return wrapper + def _close_active_request(self, *args: Any, **kwargs: Any) -> Any: + return self._requests._close_active_request(*args, **kwargs) + def _reset_request_operation_context( + self, *args: Any, **kwargs: Any + ) -> Any: + return self._requests._reset_request_operation_context(*args, **kwargs) -class _SegmentManager: - def __init__( - self, - runtime: ObservabilityRuntime, - name: Any, - attrs: dict[str, Any], - ) -> None: - self.runtime = runtime - self.name = name - self.attrs = attrs - self.context_manager = None - - def __enter__(self): - self.context_manager = self.runtime._segment_context( - self.name, - **self.attrs, - ) - return self.context_manager.__enter__() - - def __exit__(self, exc_type, exc, traceback) -> bool: - if self.context_manager is None: - return False - return bool(self.context_manager.__exit__(exc_type, exc, traceback)) - - async def __aenter__(self): - self.context_manager = self.runtime.asegment(self.name, **self.attrs) - return await self.context_manager.__aenter__() - - async def __aexit__(self, exc_type, exc, traceback) -> bool: - if self.context_manager is None: - return False - return bool( - await self.context_manager.__aexit__(exc_type, exc, traceback) - ) + def _reset_request_context(self, *args: Any, **kwargs: Any) -> Any: + return self._requests._reset_request_context(*args, **kwargs) - def __call__(self, func): - if inspect.iscoroutinefunction(func): + def _safe_perf_counter(self, *args: Any, **kwargs: Any) -> Any: + return self._segments._safe_perf_counter(*args, **kwargs) - @wraps(func) - async def async_wrapper(*args, **kwargs): - async with self.runtime.segment(self.name, **self.attrs): - return await func(*args, **kwargs) + def _safe_str(self, *args: Any, **kwargs: Any) -> Any: + return self._logging._safe_str(*args, **kwargs) - return async_wrapper + def _safe_traceback(self, *args: Any, **kwargs: Any) -> Any: + return self._logging._safe_traceback(*args, **kwargs) - @wraps(func) - def wrapper(*args, **kwargs): - with self.runtime.segment(self.name, **self.attrs): - return func(*args, **kwargs) + def _json(self, *args: Any, **kwargs: Any) -> Any: + return self._logging._json(*args, **kwargs) - return wrapper + def _write_stderr(self, *args: Any, **kwargs: Any) -> Any: + return self._logging._write_stderr(*args, **kwargs) _RUNTIME = ObservabilityRuntime(ObservabilityConfig()) @@ -2118,10 +525,10 @@ def set_observability_runtime(runtime: ObservabilityRuntime) -> None: global _RUNTIME _RUNTIME = runtime for context_var in ( - _REQUEST_CONTEXT, - _OPERATION_CONTEXT, - _TIMINGS, - _TURN_START, + _state._REQUEST_CONTEXT, + _state._OPERATION_CONTEXT, + _state._TIMINGS, + _state._TURN_START, ): try: context_var.set(None) diff --git a/policyengine_observability/segments.py b/policyengine_observability/segments.py index 977c524..588ed9d 100644 --- a/policyengine_observability/segments.py +++ b/policyengine_observability/segments.py @@ -1,8 +1,25 @@ +"""Segment names, nested timings, and context managers.""" + from __future__ import annotations -from collections.abc import Iterable +import inspect +import time +from collections.abc import AsyncIterator, Iterable, Iterator +from contextlib import asynccontextmanager, contextmanager from enum import Enum -from typing import Any +from functools import wraps +from typing import TYPE_CHECKING, Any + +from . import _state +from .context import ( + OperationObservabilityContext, + RequestObservabilityContext, + SegmentTimingNode, + _metric_attrs, +) + +if TYPE_CHECKING: + from .runtime import ObservabilityRuntime UNKNOWN_SEGMENT = "unknown_segment" @@ -33,3 +50,383 @@ def coerce_segment_name( except BaseException: return UNKNOWN_SEGMENT, False return segment, False if values else True + + +MAX_SEGMENT_ATTR_LENGTH = 200 +SENSITIVE_SEGMENT_ATTR_PARTS = ( + "authorization", + "credential", + "password", + "secret", + "token", +) + + +class SegmentRecorder: + def __init__(self, runtime: ObservabilityRuntime) -> None: + self.runtime = runtime + + def segment(self, name: Any, **attrs: Any) -> Iterator[Any]: + return _SegmentManager(self.runtime, name, attrs) + + @contextmanager + def _segment_context(self, name: Any, **attrs: Any) -> Iterator[Any]: + if not self.runtime.enabled: + yield None + return + segment_name = self.runtime._coerce_segment(name) + implicit_operation = self.runtime._start_implicit_operation( + segment_name, + attrs, + ) + start = self.runtime._safe_perf_counter( + f"segment.{segment_name}.start" + ) + segment_tree_handle = self.runtime._start_segment_tree_node( + segment_name, + attrs, + ) + span_attrs = self.runtime._segment_span_attributes(attrs) + span_name = self.runtime._span_name(segment_name) + error: BaseException | None = None + with self.runtime._safe_span(span_name, span_attrs) as span: + try: + yield span + except BaseException as exc: + error = exc + self.runtime._record_segment_safely( + segment_name, + start, + attrs, + segment_tree_handle=segment_tree_handle, + ) + raise + else: + self.runtime._record_segment_safely( + segment_name, + start, + attrs, + segment_tree_handle=segment_tree_handle, + ) + finally: + self.runtime._reset_segment_tree_stack(segment_tree_handle) + self.runtime.end_operation(implicit_operation, error) + + @asynccontextmanager + async def asegment(self, name: Any, **attrs: Any) -> AsyncIterator[Any]: + if not self.runtime.enabled: + yield None + return + segment_name = self.runtime._coerce_segment(name) + implicit_operation = self.runtime._start_implicit_operation( + segment_name, + attrs, + ) + start = self.runtime._safe_perf_counter( + f"segment.{segment_name}.start" + ) + segment_tree_handle = self.runtime._start_segment_tree_node( + segment_name, + attrs, + ) + span_attrs = self.runtime._segment_span_attributes(attrs) + span_name = self.runtime._span_name(segment_name) + error: BaseException | None = None + with self.runtime._safe_span(span_name, span_attrs) as span: + try: + yield span + except BaseException as exc: + error = exc + self.runtime._record_segment_safely( + segment_name, + start, + attrs, + segment_tree_handle=segment_tree_handle, + ) + raise + else: + self.runtime._record_segment_safely( + segment_name, + start, + attrs, + segment_tree_handle=segment_tree_handle, + ) + finally: + self.runtime._reset_segment_tree_stack(segment_tree_handle) + self.runtime.end_operation(implicit_operation, error) + + def _start_segment_tree_node( + self, + name: str, + attrs: dict[str, Any], + ) -> dict[str, Any] | None: + try: + owner = self.runtime._segment_tree_owner() + if owner is None: + return None + owner.segment_sequence[0] += 1 + node = SegmentTimingNode( + sequence=owner.segment_sequence[0], + name=name, + attrs=self.runtime._safe_segment_tree_attrs(attrs), + ) + owner_id = id(owner.segment_tree) + stack = _state._SEGMENT_STACK.get() + if stack and stack[-1][0] == owner_id: + stack[-1][1].children.append(node) + else: + owner.segment_tree.append(node) + token = _state._SEGMENT_STACK.set((*stack, (owner_id, node))) + return {"node": node, "token": token} + except BaseException as exc: + self.runtime.log_observability_failure( + "segment.tree_start", + exc, + segment=name, + ) + return None + + def _finish_segment_tree_node( + self, + handle: dict[str, Any] | None, + duration_seconds: float, + ) -> None: + if not handle: + return + try: + node = handle.get("node") + if not isinstance(node, SegmentTimingNode): + return + node.duration_ms = duration_seconds * 1000 + except BaseException as exc: + self.runtime.log_observability_failure( + "segment.tree_finish", + exc, + ) + + def _reset_segment_tree_stack( + self, + handle: dict[str, Any] | None, + ) -> None: + if not handle: + return + token = handle.get("token") + if token is None: + return + try: + _state._SEGMENT_STACK.reset(token) + except BaseException as exc: + self.runtime.log_observability_failure("segment.tree_reset", exc) + + def _segment_tree_owner( + self, + ) -> RequestObservabilityContext | OperationObservabilityContext | None: + context = self.runtime.current_context() + if context is not None: + return context + return self.runtime.current_operation() + + def _safe_segment_tree_attrs( + self, + attrs: dict[str, Any], + ) -> dict[str, Any]: + safe_attrs: dict[str, Any] = {} + for key, value in attrs.items(): + key_text = self.runtime._safe_str(key) + key_lower = key_text.lower() + if any(part in key_lower for part in SENSITIVE_SEGMENT_ATTR_PARTS): + continue + if value is None: + continue + if hasattr(value, "value"): + value = value.value + if isinstance(value, bool | int | float): + safe_attrs[key_text] = value + elif isinstance(value, str): + safe_attrs[key_text] = value[:MAX_SEGMENT_ATTR_LENGTH] + return safe_attrs + + def _record_segment_flat_timing( + self, + context: RequestObservabilityContext | None, + operation: OperationObservabilityContext | None, + name: str, + duration_ms: float, + ) -> None: + seen_timing_ids: set[int] = set() + seen_count_ids: set[int] = set() + for target in (context, operation): + if target is None: + continue + timings_id = id(target.timings_ms) + if timings_id not in seen_timing_ids: + target.timings_ms[name] = round( + target.timings_ms.get(name, 0.0) + duration_ms, + 3, + ) + seen_timing_ids.add(timings_id) + counts_id = id(target.timing_counts) + if counts_id not in seen_count_ids: + target.timing_counts[name] = ( + target.timing_counts.get(name, 0) + 1 + ) + seen_count_ids.add(counts_id) + + def _record_segment_safely( + self, + name: str, + start: float | None, + attrs: dict[str, Any], + *, + segment_tree_handle: dict[str, Any] | None = None, + ) -> None: + if start is None: + return + end = self.runtime._safe_perf_counter(f"segment.{name}.end") + if end is None: + return + try: + duration = end - start + self.runtime._finish_segment_tree_node( + segment_tree_handle, duration + ) + self.runtime._record_timing(name, duration) + context = self.runtime.current_context() + operation = self.runtime.current_operation() + metric_extra = { + key: value + for key, value in attrs.items() + if ( + key in self.runtime.config.metric_attribute_keys + and value is not None + ) + } + duration_ms = duration * 1000 + self.runtime._record_segment_flat_timing( + context, + operation, + name, + duration_ms, + ) + if operation is not None: + metric_attributes = operation.metric_attributes( + segment=name, + **metric_extra, + ) + elif context is not None: + metric_attributes = context.metric_attributes( + segment=name, + **metric_extra, + ) + else: + metric_attributes = _metric_attrs( + { + "service.name": self.runtime.config.service_name, + "service.role": self.runtime.config.service_role, + "deployment.environment": self.runtime.config.environment, + "segment": name, + **metric_extra, + }, + self.runtime.config.metric_attribute_keys, + ) + self.runtime.record_segment_metric( + name, + duration, + metric_attributes, + backend_segment="backend" in metric_extra, + ) + except BaseException as exc: + self.runtime.log_observability_failure( + "request.record_segment", + exc, + segment=name, + ) + + def _record_timing(self, name: str, duration_seconds: float) -> None: + try: + timings = _state._TIMINGS.get() + if timings is None: + return + key = f"{name}_ms" + duration_ms = duration_seconds * 1000.0 + timings[key] = round(timings.get(key, 0.0) + duration_ms, 1) + except BaseException as exc: + self.runtime.log_observability_failure( + "scope.record_timing", + exc, + segment=name, + ) + + def _coerce_segment(self, name: Any) -> str: + segment, is_registered = coerce_segment_name( + name, + registry=self.runtime.segment_registry, + ) + if not is_registered: + self.runtime.log_observability_failure( + "segment.coerce", + ValueError("Unregistered observability segment."), + segment=segment, + segment_type=type(name).__name__, + ) + return segment + + def _safe_perf_counter(self, operation: str) -> float | None: + try: + return time.perf_counter() + except BaseException as exc: + self.runtime.log_observability_failure(operation, exc) + return None + + +class _SegmentManager: + def __init__( + self, + runtime: ObservabilityRuntime, + name: Any, + attrs: dict[str, Any], + ) -> None: + self.runtime = runtime + self.name = name + self.attrs = attrs + self.context_manager = None + + def __enter__(self): + self.context_manager = self.runtime._segment_context( + self.name, + **self.attrs, + ) + return self.context_manager.__enter__() + + def __exit__(self, exc_type, exc, traceback) -> bool: + if self.context_manager is None: + return False + return bool(self.context_manager.__exit__(exc_type, exc, traceback)) + + async def __aenter__(self): + self.context_manager = self.runtime.asegment(self.name, **self.attrs) + return await self.context_manager.__aenter__() + + async def __aexit__(self, exc_type, exc, traceback) -> bool: + if self.context_manager is None: + return False + return bool( + await self.context_manager.__aexit__(exc_type, exc, traceback) + ) + + def __call__(self, func): + if inspect.iscoroutinefunction(func): + + @wraps(func) + async def async_wrapper(*args, **kwargs): + async with self.runtime.segment(self.name, **self.attrs): + return await func(*args, **kwargs) + + return async_wrapper + + @wraps(func) + def wrapper(*args, **kwargs): + with self.runtime.segment(self.name, **self.attrs): + return func(*args, **kwargs) + + return wrapper diff --git a/tests/conftest.py b/tests/conftest.py index dbc84a1..9114b37 100644 --- a/tests/conftest.py +++ b/tests/conftest.py @@ -3,12 +3,32 @@ import pytest from fakes import RecordingDestination +from policyengine_observability import _state from policyengine_observability.destinations.registry import ( _STRATEGIES, register_destination, ) +@pytest.fixture(autouse=True) +def isolated_observability_context(): + """Keep tests independent when request and operation tests run in separate files.""" + variables = ( + (_state._REQUEST_CONTEXT, None), + (_state._OPERATION_CONTEXT, None), + (_state._TIMINGS, None), + (_state._TURN_START, None), + (_state._SEGMENT_STACK, ()), + ) + for variable, default in variables: + variable.set(default) + try: + yield + finally: + for variable, default in variables: + variable.set(default) + + @pytest.fixture def fake_remote_strategy(): """Register a fake remote strategy; yields its destination name.""" diff --git a/tests/runtime_helpers.py b/tests/runtime_helpers.py new file mode 100644 index 0000000..5e287f5 --- /dev/null +++ b/tests/runtime_helpers.py @@ -0,0 +1,198 @@ +from __future__ import annotations + +from enum import StrEnum +from typing import Any + +from policyengine_observability import ( + ObservabilityConfig, + ObservabilityRuntime, +) + + +class SegmentName(StrEnum): + LOAD = "load" + SAVE = "save" + + +class RecordingSpan: + def __init__(self) -> None: + self.attributes = {} + self.exceptions = [] + self.events = [] + self.status = None + + def set_attribute(self, key, value) -> None: + self.attributes[key] = value + + def record_exception(self, exc) -> None: + self.exceptions.append(exc) + + def set_status(self, status) -> None: + self.status = status + + def add_event(self, event, fields) -> None: + self.events.append((event, fields)) + + def get_span_context(self): + return type( + "SpanContext", + (), + {"is_valid": False, "trace_id": 0, "span_id": 0}, + )() + + +class NamedRecordingSpan(RecordingSpan): + def __init__(self) -> None: + super().__init__() + self.names = [] + + def update_name(self, name: str) -> None: + self.names.append(name) + + +class ValidContextSpan(RecordingSpan): + def get_span_context(self): + return type( + "SpanContext", + (), + { + "is_valid": True, + "trace_id": 0x4BF92F3577B34DA6A3CE929D0E0E4736, + "span_id": 0x00F067AA0BA902B7, + }, + )() + + +class AttributeFailingSpan(RecordingSpan): + def set_attribute(self, key, value) -> None: + raise RuntimeError("attribute failed") + + +class ExceptionFailingSpan(RecordingSpan): + def record_exception(self, exc) -> None: + raise RuntimeError("record exception failed") + + +class RecordingSpanContextManager: + def __init__( + self, + span: RecordingSpan, + *, + fail_exit: bool = False, + ) -> None: + self.span = span + self.fail_exit = fail_exit + self.exited = False + + def __enter__(self): + return self.span + + def __exit__(self, *_args): + self.exited = True + if self.fail_exit: + raise RuntimeError("span exit failed") + return False + + +class RecordingTracer: + def __init__( + self, + span: RecordingSpan | None = None, + *, + fail_enter: bool = False, + fail_exit: bool = False, + ) -> None: + self.span = span or RecordingSpan() + self.fail_enter = fail_enter + self.fail_exit = fail_exit + self.calls = [] + self.last_context_manager = None + + def start_as_current_span(self, name, **kwargs): + self.calls.append((name, kwargs)) + if self.fail_enter: + raise RuntimeError("span enter failed") + self.last_context_manager = RecordingSpanContextManager( + self.span, + fail_exit=self.fail_exit, + ) + return self.last_context_manager + + +class RecordingMeter: + def __init__(self) -> None: + self.created = [] + + def create_histogram(self, name, **kwargs): + self.created.append(("histogram", name, kwargs)) + return RecordingInstrument() + + def create_counter(self, name, **kwargs): + self.created.append(("counter", name, kwargs)) + return RecordingInstrument() + + def create_up_down_counter(self, name, **kwargs): + self.created.append(("up_down_counter", name, kwargs)) + return RecordingInstrument() + + +class RecordingInstrument: + def __init__(self) -> None: + self.calls = [] + + def add(self, value, attributes=None) -> None: + self.calls.append(("add", value, attributes)) + + def record(self, value, attributes=None) -> None: + self.calls.append(("record", value, attributes)) + + +class FailingInstrument: + def add(self, *_args, **_kwargs) -> None: + raise RuntimeError("metric failed") + + def record(self, *_args, **_kwargs) -> None: + raise RuntimeError("metric failed") + + +class RecordingLogDestination: + def __init__(self, name: str = "recording") -> None: + self.name = name + self.calls = [] + + def emit( + self, + payload: dict[str, Any], + *, + log_type: str, + severity: str, + ) -> None: + self.calls.append((payload, log_type, severity)) + + +class FailingLogDestination: + name = "failing" + + def emit(self, *_args, **_kwargs) -> None: + raise RuntimeError("destination failed") + + +class RecordingPropagator: + def __init__(self) -> None: + self.extracted = None + + def inject(self, carrier) -> None: + carrier["traceparent"] = ( + "00-4bf92f3577b34da6a3ce929d0e0e4736-00f067aa0ba902b7-01" + ) + + def extract(self, carrier): + self.extracted = carrier + return {"parent": carrier} + + +def runtime(**kwargs) -> ObservabilityRuntime: + return ObservabilityRuntime( + ObservabilityConfig(service_name="svc", **kwargs), + segment_registry=SegmentName, + ) diff --git a/tests/test_google_destination.py b/tests/test_google_destination.py new file mode 100644 index 0000000..fffb9c0 --- /dev/null +++ b/tests/test_google_destination.py @@ -0,0 +1,334 @@ +from __future__ import annotations + +import json +import math + +import pytest + +from policyengine_observability.config import ObservabilityConfig +from policyengine_observability.destinations import ( + GoogleCloudLoggingDestination, + google_cloud_logging, + normalize_payload, +) +from policyengine_observability.destinations.base import ( + accepts_keyword, + clamped, +) + + +class Unprintable: + def __str__(self) -> str: + raise RuntimeError("cannot stringify") + + +class FakeLogger: + def __init__(self) -> None: + self.calls = [] + + def log_struct(self, payload, **kwargs) -> None: + self.calls.append((payload, kwargs)) + + +class FakeGapicApi: + def __init__(self) -> None: + self.calls = [] + + def write_log_entries(self, *args, **kwargs) -> None: + self.calls.append((args, kwargs)) + + +class FakeLoggingApi: + def __init__(self) -> None: + self._gapic_api = FakeGapicApi() + + +class FakeClient: + def __init__(self, *, gapic: bool = False) -> None: + self.project = "resolved-project" + self.fake_logger = FakeLogger() + self.log_names = [] + if gapic: + self.logging_api = FakeLoggingApi() + + def logger(self, log_name: str) -> FakeLogger: + self.log_names.append(log_name) + return self.fake_logger + + +def _google_destination(monkeypatch, client, **kwargs): + monkeypatch.setattr( + google_cloud_logging, + "load_google_credentials", + lambda *, prefer_workload_identity: None, + ) + monkeypatch.setattr( + google_cloud_logging, + "configure_google_application_credentials", + lambda: None, + ) + return GoogleCloudLoggingDestination( + project=None, + log_name="policyengine-observability", + client_factory=lambda _project, _credentials: client, + **kwargs, + ) + + +@pytest.mark.parametrize( + ("value", "expected"), + [ + (5.0, 5.0), + ("5", 5.0), + (0, 0.5), + (-3, 0.5), + (1000, 60.0), + (float("inf"), 10.0), + (float("nan"), 10.0), + (None, 10.0), + ("garbage", 10.0), + ], +) +def test_clamped_bounds_and_rejects_non_finite(value, expected) -> None: + result = clamped(value, low=0.5, high=60.0, default=10.0) + + assert result == expected + assert math.isfinite(result) + + +def test_accepts_keyword_covers_named_var_keyword_and_uninspectable() -> None: + def named(payload, *, timestamp=None): + pass + + def var_keyword(payload, **kwargs): + pass + + def blind(payload): + pass + + assert accepts_keyword(named, "timestamp") is True + assert accepts_keyword(var_keyword, "timestamp") is True + assert accepts_keyword(blind, "timestamp") is False + # Builtins without introspectable signatures degrade to False + # instead of raising at construction time. + assert accepts_keyword(min, "timestamp") is False + + +def test_normalize_payload_recursively_stringifies_unsafe_values() -> None: + normalized = normalize_payload( + { + "keep": "value", + "drop_none": None, + "bytes": b"value", + "list": [1, Unprintable()], + "nested": {"object": object()}, + } + ) + + assert normalized["keep"] == "value" + assert normalized["drop_none"] is None + assert normalized["bytes"] == "value" + assert normalized["list"] == [1, ""] + assert normalized["nested"]["object"].startswith(" None: + monkeypatch.setattr( + google_cloud_logging, + "load_google_credentials", + lambda *, prefer_workload_identity: None, + ) + monkeypatch.setattr( + google_cloud_logging, + "configure_google_application_credentials", + lambda: None, + ) + client = FakeClient() + destination = GoogleCloudLoggingDestination( + project=None, + log_name="policyengine-observability", + client_factory=lambda _project, _credentials: client, + ) + + destination.emit( + { + "schema_version": "policyengine.observability.request.v1", + "service_name": "svc", + "service_role": "api", + "environment": "production", + "request_id": "request-1", + "trace_id": "abc123", + "span_id": "def456", + "path": "/calculate", + "object": object(), + }, + log_type="request", + severity="ERROR", + ) + + payload, kwargs = client.fake_logger.calls[0] + assert client.log_names == ["policyengine-observability"] + assert payload["object"].startswith(" None: + client = FakeClient(gapic=True) + + destination = _google_destination( + monkeypatch, client, write_timeout_seconds=5.0 + ) + # The write path under log_struct funnels through this method; the + # rebinding must inject the bounded retry and per-call timeout. + client.logging_api._gapic_api.write_log_entries(request="sentinel") + + assert destination.write_timeout_seconds == 5.0 + ((args, kwargs),) = client.logging_api._gapic_api.calls + assert kwargs["request"] == "sentinel" + assert kwargs["timeout"] == 5.0 + assert kwargs["retry"].timeout == 5.0 + + +def test_google_destination_clamps_write_timeout(monkeypatch) -> None: + destination = _google_destination( + monkeypatch, FakeClient(gapic=True), write_timeout_seconds=0.0 + ) + + assert destination.write_timeout_seconds == 0.5 + + +def test_google_destination_without_gapic_transport_still_works( + monkeypatch, +) -> None: + client = FakeClient() + + destination = _google_destination(monkeypatch, client) + destination.emit({"event": "x"}, log_type="event", severity="INFO") + + assert len(client.fake_logger.calls) == 1 + + +def test_google_destination_forwards_enqueue_timestamp(monkeypatch) -> None: + from datetime import UTC, datetime + + client = FakeClient() + destination = _google_destination(monkeypatch, client) + stamp = datetime(2026, 7, 8, 12, 0, 0, tzinfo=UTC) + + destination.emit( + {"event": "x"}, log_type="event", severity="INFO", timestamp=stamp + ) + destination.emit({"event": "y"}, log_type="event", severity="INFO") + + (_, stamped_kwargs), (_, plain_kwargs) = client.fake_logger.calls + assert stamped_kwargs["timestamp"] is stamp + assert "timestamp" not in plain_kwargs + + +def test_google_destination_close_closes_client(monkeypatch) -> None: + class ClosableFakeClient(FakeClient): + def __init__(self) -> None: + super().__init__() + self.closed = 0 + + def close(self) -> None: + self.closed += 1 + + client = ClosableFakeClient() + destination = _google_destination(monkeypatch, client) + + destination.close() + + assert client.closed == 1 + + +def test_google_destination_close_tolerates_closeless_client( + monkeypatch, +) -> None: + destination = _google_destination(monkeypatch, FakeClient()) + + destination.close() # FakeClient has no close; must be a no-op + + +def test_google_destination_suppresses_instrumentation_entry( + monkeypatch, +) -> None: + logging_v2 = pytest.importorskip("google.cloud.logging_v2") + monkeypatch.setattr( + logging_v2, "_instrumentation_emitted", False, raising=False + ) + + _google_destination(monkeypatch, FakeClient()) + + assert logging_v2._instrumentation_emitted is True + + +def test_google_factory_reads_write_timeout_env(monkeypatch) -> None: + captured = {} + + class StubDestination: + def __init__(self, **kwargs) -> None: + captured.update(kwargs) + + monkeypatch.setattr( + google_cloud_logging, "GoogleCloudLoggingDestination", StubDestination + ) + monkeypatch.setenv("OBSERVABILITY_GOOGLE_WRITE_TIMEOUT_SECONDS", "2.5") + + from policyengine_observability.destinations.registry import ( + destination_strategy, + ) + + destination_strategy("google_cloud_logging").factory( + config=ObservabilityConfig(google_cloud_project="proj"), + loggers={}, + serializer=json.dumps, + ) + + assert captured["project"] == "proj" + assert captured["write_timeout_seconds"] == 2.5 + + +def test_google_factory_write_timeout_defaults_without_env( + monkeypatch, +) -> None: + captured = {} + + class StubDestination: + def __init__(self, **kwargs) -> None: + captured.update(kwargs) + + monkeypatch.setattr( + google_cloud_logging, "GoogleCloudLoggingDestination", StubDestination + ) + monkeypatch.delenv( + "OBSERVABILITY_GOOGLE_WRITE_TIMEOUT_SECONDS", raising=False + ) + + from policyengine_observability.destinations.registry import ( + destination_strategy, + ) + + destination_strategy("google_cloud_logging").factory( + config=ObservabilityConfig(google_cloud_project="proj"), + loggers={}, + serializer=json.dumps, + ) + + assert captured["write_timeout_seconds"] == 10.0 + + +# ── Stdout formatters ──────────────────────────────────────────────────── diff --git a/tests/test_runtime.py b/tests/test_runtime.py deleted file mode 100644 index e940594..0000000 --- a/tests/test_runtime.py +++ /dev/null @@ -1,2217 +0,0 @@ -from __future__ import annotations - -import asyncio -import builtins -import time -from enum import StrEnum -from types import SimpleNamespace -from typing import Any - -import pytest - -from policyengine_observability import ( - UNKNOWN_SEGMENT, - ObservabilityConfig, - ObservabilityRuntime, - RequestObservabilityContext, - coerce_segment_name, -) -from policyengine_observability import runtime as runtime_module -from policyengine_observability.config import DEFAULT_METRIC_ATTRIBUTE_KEYS -from policyengine_observability.destinations import ( - google_cloud_logging as google_cloud_logging_module, -) - - -class SegmentName(StrEnum): - LOAD = "load" - SAVE = "save" - - -class RecordingSpan: - def __init__(self) -> None: - self.attributes = {} - self.exceptions = [] - self.events = [] - self.status = None - - def set_attribute(self, key, value) -> None: - self.attributes[key] = value - - def record_exception(self, exc) -> None: - self.exceptions.append(exc) - - def set_status(self, status) -> None: - self.status = status - - def add_event(self, event, fields) -> None: - self.events.append((event, fields)) - - def get_span_context(self): - return type( - "SpanContext", - (), - {"is_valid": False, "trace_id": 0, "span_id": 0}, - )() - - -class NamedRecordingSpan(RecordingSpan): - def __init__(self) -> None: - super().__init__() - self.names = [] - - def update_name(self, name: str) -> None: - self.names.append(name) - - -class ValidContextSpan(RecordingSpan): - def get_span_context(self): - return type( - "SpanContext", - (), - { - "is_valid": True, - "trace_id": 0x4BF92F3577B34DA6A3CE929D0E0E4736, - "span_id": 0x00F067AA0BA902B7, - }, - )() - - -class AttributeFailingSpan(RecordingSpan): - def set_attribute(self, key, value) -> None: - raise RuntimeError("attribute failed") - - -class ExceptionFailingSpan(RecordingSpan): - def record_exception(self, exc) -> None: - raise RuntimeError("record exception failed") - - -class RecordingSpanContextManager: - def __init__( - self, - span: RecordingSpan, - *, - fail_exit: bool = False, - ) -> None: - self.span = span - self.fail_exit = fail_exit - self.exited = False - - def __enter__(self): - return self.span - - def __exit__(self, *_args): - self.exited = True - if self.fail_exit: - raise RuntimeError("span exit failed") - return False - - -class RecordingTracer: - def __init__( - self, - span: RecordingSpan | None = None, - *, - fail_enter: bool = False, - fail_exit: bool = False, - ) -> None: - self.span = span or RecordingSpan() - self.fail_enter = fail_enter - self.fail_exit = fail_exit - self.calls = [] - self.last_context_manager = None - - def start_as_current_span(self, name, **kwargs): - self.calls.append((name, kwargs)) - if self.fail_enter: - raise RuntimeError("span enter failed") - self.last_context_manager = RecordingSpanContextManager( - self.span, - fail_exit=self.fail_exit, - ) - return self.last_context_manager - - -class RecordingMeter: - def __init__(self) -> None: - self.created = [] - - def create_histogram(self, name, **kwargs): - self.created.append(("histogram", name, kwargs)) - return RecordingInstrument() - - def create_counter(self, name, **kwargs): - self.created.append(("counter", name, kwargs)) - return RecordingInstrument() - - def create_up_down_counter(self, name, **kwargs): - self.created.append(("up_down_counter", name, kwargs)) - return RecordingInstrument() - - -class RecordingInstrument: - def __init__(self) -> None: - self.calls = [] - - def add(self, value, attributes=None) -> None: - self.calls.append(("add", value, attributes)) - - def record(self, value, attributes=None) -> None: - self.calls.append(("record", value, attributes)) - - -class FailingInstrument: - def add(self, *_args, **_kwargs) -> None: - raise RuntimeError("metric failed") - - def record(self, *_args, **_kwargs) -> None: - raise RuntimeError("metric failed") - - -class RecordingLogDestination: - def __init__(self, name: str = "recording") -> None: - self.name = name - self.calls = [] - - def emit( - self, - payload: dict[str, Any], - *, - log_type: str, - severity: str, - ) -> None: - self.calls.append((payload, log_type, severity)) - - -class FailingLogDestination: - name = "failing" - - def emit(self, *_args, **_kwargs) -> None: - raise RuntimeError("destination failed") - - -class RecordingPropagator: - def __init__(self) -> None: - self.extracted = None - - def inject(self, carrier) -> None: - carrier["traceparent"] = ( - "00-4bf92f3577b34da6a3ce929d0e0e4736-00f067aa0ba902b7-01" - ) - - def extract(self, carrier): - self.extracted = carrier - return {"parent": carrier} - - -def runtime(**kwargs) -> ObservabilityRuntime: - return ObservabilityRuntime( - ObservabilityConfig(service_name="svc", **kwargs), - segment_registry=SegmentName, - ) - - -def test_segment_records_aggregated_timing() -> None: - observed = runtime() - - with observed.collect_timings("request") as timings: - with observed.segment(SegmentName.LOAD): - pass - with observed.segment(SegmentName.LOAD): - pass - - assert "load_ms" in timings - assert timings["load_ms"] >= 0 - - -def test_operation_log_accumulates_repeated_segment_timings() -> None: - observed = runtime() - handle = observed.start_operation("job") - operation = handle["operation"] - - try: - with observed.segment(SegmentName.LOAD): - pass - with observed.segment(SegmentName.LOAD): - pass - finally: - observed.end_operation(handle) - - assert operation.timings_ms["load"] >= 0 - assert operation.timing_counts["load"] == 2 - payload = operation.as_log_record(trace_id=None, span_id=None) - assert payload["timing_counts"]["load"] == 2 - - -def test_operation_log_records_ordered_nested_segment_tree() -> None: - observed = runtime() - handle = observed.start_operation("job") - operation = handle["operation"] - - try: - with observed.segment(SegmentName.LOAD): - with observed.segment( - SegmentName.SAVE, - simulation_kind="baseline", - token="SECRET", - payload={"not": "safe"}, - ): - pass - with observed.segment( - SegmentName.SAVE, - simulation_kind="reform", - ): - pass - finally: - observed.end_operation(handle) - - payload = operation.as_log_record(trace_id=None, span_id=None) - tree = payload["segment_tree"] - assert len(tree) == 1 - assert tree[0]["sequence"] == 1 - assert tree[0]["name"] == "load" - assert "duration_ms" in tree[0] - assert "self_ms" not in tree[0] - - children = tree[0]["children"] - assert [child["sequence"] for child in children] == [2, 3] - assert [child["name"] for child in children] == ["save", "save"] - assert children[0]["attrs"] == {"simulation_kind": "baseline"} - assert children[1]["attrs"] == {"simulation_kind": "reform"} - assert "token" not in children[0].get("attrs", {}) - assert "payload" not in children[0].get("attrs", {}) - assert payload["timing_counts"]["save"] == 2 - - -def test_operation_log_reserved_fields_override_attributes() -> None: - observed = runtime() - handle = observed.start_operation( - "job", - operation="attribute-operation", - duration_ms="attribute-duration", - timings_ms="attribute-timings", - timing_counts="attribute-counts", - segment_tree="attribute-tree", - error="attribute-error", - ) - operation = handle["operation"] - - try: - with observed.segment(SegmentName.LOAD): - pass - finally: - observed.end_operation(handle) - - payload = operation.as_log_record(trace_id=None, span_id=None) - assert payload["operation"] == "job" - assert isinstance(payload["duration_ms"], float) - assert isinstance(payload["timings_ms"], dict) - assert isinstance(payload["timing_counts"], dict) - assert isinstance(payload["segment_tree"], list) - assert payload["error"] is None - - -def test_async_segment_records_timing() -> None: - async def run() -> dict[str, float]: - observed = runtime() - with observed.collect_timings("request") as timings: - async with observed.asegment(SegmentName.SAVE): - pass - return timings - - timings = asyncio.run(run()) - - assert "save_ms" in timings - - -def test_async_segments_keep_independent_segment_tree_stacks() -> None: - async def run() -> list[dict[str, Any]]: - observed = runtime() - handle = observed.start_operation("job") - operation = handle["operation"] - - async def branch(branch_name: str) -> None: - async with observed.asegment(SegmentName.LOAD, branch=branch_name): - await asyncio.sleep(0) - async with observed.asegment( - SegmentName.SAVE, - branch=branch_name, - ): - await asyncio.sleep(0) - - try: - await asyncio.gather(branch("a"), branch("b")) - finally: - observed.end_operation(handle) - return operation.as_log_record(trace_id=None, span_id=None)[ - "segment_tree" - ] - - tree = asyncio.run(run()) - - assert [node["name"] for node in tree] == ["load", "load"] - assert [node["attrs"] for node in tree] == [ - {"branch": "a"}, - {"branch": "b"}, - ] - assert [node["children"][0]["attrs"] for node in tree] == [ - {"branch": "a"}, - {"branch": "b"}, - ] - - -def test_segment_preserves_business_exception_and_records_timing() -> None: - observed = runtime() - - with pytest.raises(ValueError, match="business failed"): - with observed.collect_timings("request") as timings: - with observed.segment(SegmentName.LOAD): - raise ValueError("business failed") - - assert "load_ms" in timings - - -def test_segment_tree_records_failed_segments_before_reraising() -> None: - observed = runtime() - handle = observed.start_operation("job") - operation = handle["operation"] - error = None - - try: - with observed.segment(SegmentName.LOAD): - raise ValueError("business failed") - except ValueError as exc: - error = exc - finally: - observed.end_operation(handle, error) - - payload = operation.as_log_record(trace_id=None, span_id=None) - assert payload["event"] == "operation_failed" - assert payload["segment_tree"][0]["name"] == "load" - assert "duration_ms" in payload["segment_tree"][0] - - -def test_unregistered_segment_falls_back_without_throwing() -> None: - class BrokenString: - def __str__(self) -> str: - raise RuntimeError("cannot stringify") - - observed = runtime() - - with observed.collect_timings("request") as timings: - with observed.segment(BrokenString()): - pass - - assert f"{UNKNOWN_SEGMENT}_ms" in timings - - -def test_segment_span_start_failure_does_not_skip_user_code() -> None: - observed = runtime() - observed.tracer = RecordingTracer(fail_enter=True) - executed = False - - with observed.collect_timings("request") as timings: - with observed.segment(SegmentName.LOAD): - executed = True - - assert executed - assert "load_ms" in timings - - -def test_segment_span_exit_failure_does_not_escape() -> None: - observed = runtime() - observed.tracer = RecordingTracer(fail_exit=True) - - with observed.collect_timings("request") as timings: - with observed.segment(SegmentName.LOAD): - pass - - assert "load_ms" in timings - - -def test_disabled_runtime_noops_across_public_methods() -> None: - observed = ObservabilityRuntime.disabled() - context = RequestObservabilityContext( - config=observed.config, - request_id="request-1", - method="GET", - route="/disabled", - path="/disabled", - endpoint="disabled", - query_keys=[], - content_length_bytes=None, - inbound={}, - ) - - handle = observed.start_operation("disabled") - observed.end_operation(handle) - observed.begin_request(context) - observed.complete_request(200) - observed.update_request_route(route="/other") - observed.teardown_request(None) - observed.set_attribute("key", "value") - observed.record_error(RuntimeError("ignored"), handled=True) - observed.record_event("ignored") - - with observed.segment(SegmentName.LOAD) as span: - assert span is None - - async def run() -> None: - async with observed.asegment(SegmentName.LOAD) as async_span: - assert async_span is None - - asyncio.run(run()) - assert observed.prepare_response(200) == {} - assert observed.current_context() is None - assert observed.current_operation() is None - - -def test_span_attribute_failure_does_not_drop_span_lifecycle() -> None: - observed = runtime() - span = AttributeFailingSpan() - observed.tracer = RecordingTracer(span=span) - failures = [] - observed.log_observability_failure = lambda operation, exc, **fields: ( - failures.append(operation) - ) - - with observed.segment(SegmentName.LOAD, tool="loader"): - pass - - assert "otel.span_attributes" in failures - assert observed.tracer.last_context_manager.exited - - -def test_collect_timings_records_block_exception_on_scope_span() -> None: - observed = runtime() - span = RecordingSpan() - observed.tracer = RecordingTracer(span=span) - - with pytest.raises(RuntimeError, match="scope failed"): - with observed.collect_timings("turn"): - raise RuntimeError("scope failed") - - assert len(span.exceptions) == 1 - assert isinstance(span.exceptions[0], RuntimeError) - - -def test_operation_context_manager_async_and_exception_paths() -> None: - async def run() -> None: - observed = runtime() - observed.errors = RecordingInstrument() - - with pytest.raises(RuntimeError, match="async failed"): - async with observed.operation("async_job", flavor="worker"): - raise RuntimeError("async failed") - - assert observed.errors.calls[0][2]["error_type"] == "RuntimeError" - assert observed.current_operation() is None - - asyncio.run(run()) - - -def test_start_operation_with_parent_context_attaches_and_detaches() -> None: - observed = runtime() - observed.tracer = RecordingTracer() - parent_context = object() - - handle = observed.start_operation( - "parented", - parent_context=parent_context, - ) - observed.end_operation(handle) - - assert observed.tracer.calls[0][0] == "parented" - assert observed.current_operation() is None - - -def test_operation_attach_detach_and_reset_failures_are_logged( - monkeypatch, -) -> None: - from opentelemetry import context as otel_context - - observed = runtime() - observed.tracer = RecordingTracer() - failures = [] - observed.log_observability_failure = lambda operation, exc, **fields: ( - failures.append((operation, fields.get("token"))) - ) - monkeypatch.setattr( - otel_context, - "attach", - lambda _context: (_ for _ in ()).throw(RuntimeError("attach failed")), - ) - - handle = observed.start_operation("job", parent_context=object()) - observed.end_operation(handle) - - monkeypatch.setattr( - otel_context, - "detach", - lambda _token: (_ for _ in ()).throw(RuntimeError("detach failed")), - ) - observed.end_operation( - { - "operation": None, - "context_token": object(), - "timings_token": object(), - "start_token": object(), - "operation_token": object(), - } - ) - - assert ("operation.context_attach", None) in failures - assert ("operation.context_detach", None) in failures - assert ("operation.context_reset", "timings_token") in failures - assert ("operation.context_reset", "start_token") in failures - assert ("operation.context_reset", "operation_token") in failures - - -def test_operation_end_and_start_failures_are_logged(monkeypatch) -> None: - class BrokenVar: - def set(self, _value): - raise RuntimeError("set failed") - - observed = runtime() - failures = [] - observed.log_observability_failure = lambda operation, exc, **fields: ( - failures.append(operation) - ) - operation_handle = observed.start_operation("job") - operation = operation_handle["operation"] - operation.metric_recorded = True - observed.complete_operation(operation) - observed.complete_operation = lambda _operation: (_ for _ in ()).throw( - RuntimeError("complete failed") - ) - observed.end_operation(operation_handle) - - monkeypatch.setattr(runtime_module, "_OPERATION_CONTEXT", BrokenVar()) - handle = observed.start_operation("job") - - assert handle["operation"] is None - assert failures == ["operation.end", "operation.start"] - - -def test_standalone_segment_creates_implicit_operation_metrics() -> None: - observed = runtime() - observed.segment_duration = RecordingInstrument() - observed.operation_duration = RecordingInstrument() - observed.operations = RecordingInstrument() - emitted_payloads = [] - observed.emit_operation_log = lambda operation: emitted_payloads.append( - operation.as_log_record(trace_id=None, span_id=None) - ) - - with observed.segment(SegmentName.LOAD, flavor="cli", tool="loader"): - pass - - _, _, segment_attributes = observed.segment_duration.calls[0] - _, _, operation_attributes = observed.operation_duration.calls[0] - assert segment_attributes["operation"] == "load" - assert segment_attributes["flavor"] == "cli" - assert segment_attributes["tool"] == "loader" - assert operation_attributes["operation"] == "load" - assert operation_attributes["flavor"] == "cli" - assert emitted_payloads[0]["segment_tree"][0]["name"] == "load" - assert emitted_payloads[0]["segment_tree"][0]["attrs"] == { - "flavor": "cli", - "tool": "loader", - } - assert observed.current_operation() is None - - -def test_segment_with_request_context_does_not_create_implicit_operation( - monkeypatch, -) -> None: - observed = runtime() - observed.segment_duration = RecordingInstrument() - observed.operation_duration = RecordingInstrument() - context = RequestObservabilityContext( - config=observed.config, - request_id="request-1", - method="GET", - route="/calculate", - path="/calculate", - endpoint="calculate", - query_keys=[], - content_length_bytes=None, - inbound={}, - ) - monkeypatch.setattr(observed, "current_context", lambda: context) - - with observed.segment(SegmentName.LOAD): - pass - - _, _, attributes = observed.segment_duration.calls[0] - assert attributes["route"] == "/calculate" - assert observed.operation_duration.calls == [] - - -def test_start_scope_outside_request_records_operation_segment_metrics() -> ( - None -): - observed = runtime() - observed.segment_duration = RecordingInstrument() - timings: dict[str, float] = {} - - handle = observed.start_scope( - timings, - name="chat_turn", - flavor="chat", - model="claude", - ) - with observed.segment(SegmentName.LOAD, tool="search"): - pass - observed.end_scope(handle) - - _, _, attributes = observed.segment_duration.calls[0] - assert "load_ms" in timings - assert attributes["operation"] == "chat_turn" - assert attributes["flavor"] == "chat" - assert attributes["model"] == "claude" - assert attributes["tool"] == "search" - - -def test_nested_scope_annotates_span_context_and_operation() -> None: - observed = runtime() - observed.tracer = RecordingTracer() - parent_context = object() - timings: dict[str, float] = {} - - with observed.operation("outer", flavor="chat"): - handle = observed.start_scope( - timings, - name="inner", - parent_context=parent_context, - ) - observed.annotate(handle, model="claude") - observed.mark("custom_ms", 1.23) - observed.mark_ttft() - observed.end_scope(handle) - - span = observed.tracer.span - assert span.attributes["model"] == "claude" - assert "custom_ms" in timings - - -def test_entrypoint_decorator_records_operation_metrics() -> None: - observed = runtime() - observed.operation_duration = RecordingInstrument() - observed.operations = RecordingInstrument() - - @observed.entrypoint("import_data", flavor="cli") - def run_import() -> str: - return "done" - - assert run_import() == "done" - _, _, attributes = observed.operation_duration.calls[0] - assert attributes["operation"] == "import_data" - assert attributes["flavor"] == "cli" - - -def test_async_segment_decorator_records_segment_metrics() -> None: - observed = runtime() - observed.segment_duration = RecordingInstrument() - - @observed.segment(SegmentName.SAVE, flavor="worker") - async def save() -> str: - return "saved" - - assert asyncio.run(save()) == "saved" - _, _, attributes = observed.segment_duration.calls[0] - assert attributes["operation"] == "save" - assert attributes["flavor"] == "worker" - - -def test_record_error_outside_request_uses_operation_context() -> None: - observed = runtime() - observed.errors = RecordingInstrument() - - with observed.operation("worker", flavor="queue"): - observed.record_error( - RuntimeError("failed"), - handled=True, - include_stack=False, - ) - - _, _, attributes = observed.errors.calls[0] - assert attributes["operation"] == "worker" - assert attributes["flavor"] == "queue" - assert attributes["error_type"] == "RuntimeError" - - -def test_record_error_on_request_updates_span_status() -> None: - observed = runtime() - span = RecordingSpan() - observed.trace = SimpleNamespace(get_current_span=lambda: span) - observed.StatusCode = SimpleNamespace(ERROR="ERROR") - observed.Status = lambda code, message: (code, message) - observed.errors = RecordingInstrument() - context = RequestObservabilityContext( - config=observed.config, - request_id="request-1", - method="GET", - route="/error", - path="/error", - endpoint="error", - query_keys=[], - content_length_bytes=None, - inbound={}, - ) - - observed.begin_request(context) - observed.record_error( - RuntimeError("failed"), - handled=True, - status_code=500, - ) - observed.teardown_request(None) - - assert span.exceptions - assert span.status == ("ERROR", "failed") - - -def test_request_lifecycle_records_headers_and_context_metrics() -> None: - observed = runtime() - observed.active_requests = RecordingInstrument() - observed.requests = RecordingInstrument() - observed.http_duration = RecordingInstrument() - context = RequestObservabilityContext( - config=observed.config, - request_id="request-1", - method="GET", - route="/calculate", - path="/calculate", - endpoint="calculate", - query_keys=["country"], - content_length_bytes=None, - inbound={"ip_source": "remote_addr", "client_ip": "127.0.0.1"}, - ) - - observed.begin_request(context) - with observed.segment(SegmentName.LOAD): - pass - headers = observed.finish_request(200) - observed.teardown_request(None) - - assert headers["X-PolicyEngine-Request-Id"] == "request-1" - assert context.status_code == 200 - assert "load" in context.timings_ms - assert observed.current_context() is None - assert observed.current_operation() is None - assert observed.active_requests.calls[0][1] == 1 - assert observed.active_requests.calls[-1][1] == -1 - assert observed.requests.calls[0][1] == 1 - - -def test_request_log_accumulates_repeated_segment_timings() -> None: - observed = runtime() - context = RequestObservabilityContext( - config=observed.config, - request_id="request-1", - method="GET", - route="/calculate", - path="/calculate", - endpoint="calculate", - query_keys=[], - content_length_bytes=None, - inbound={}, - ) - - observed.begin_request(context) - with observed.segment(SegmentName.LOAD): - pass - with observed.segment(SegmentName.LOAD): - pass - observed.finish_request(200) - observed.teardown_request(None) - - assert context.timings_ms["load"] >= 0 - assert context.timing_counts["load"] == 2 - payload = context.as_log_record(trace_id=None, span_id=None) - assert payload["timing_counts"]["load"] == 2 - assert [node["name"] for node in payload["segment_tree"]] == [ - "load", - "load", - ] - - -def test_internal_dispatch_segments_merge_into_parent_operation() -> None: - observed = runtime() - handle = observed.start_operation( - "modal_worker_dispatch", - flavor="modal_worker", - ) - parent_operation = handle["operation"] - context = RequestObservabilityContext( - config=observed.config, - request_id="request-1", - method="POST", - route="/calculate", - path="/calculate", - endpoint="calculate", - query_keys=[], - content_length_bytes=None, - inbound={}, - internal_dispatch=True, - ) - - try: - observed.begin_request(context) - with observed.segment(SegmentName.LOAD): - pass - observed.finish_request(200) - observed.teardown_request(None) - - assert context.timings_ms is parent_operation.timings_ms - assert context.timing_counts is parent_operation.timing_counts - assert context.segment_tree is parent_operation.segment_tree - assert "load" in parent_operation.timings_ms - assert parent_operation.timing_counts["load"] == 1 - assert parent_operation.segment_tree[0].name == "load" - assert observed.current_operation() is parent_operation - finally: - observed.end_operation(handle) - - assert observed.current_context() is None - assert observed.current_operation() is None - - -def test_non_internal_request_timings_do_not_leak_to_parent_operation() -> ( - None -): - observed = runtime() - handle = observed.start_operation("job", flavor="worker") - parent_operation = handle["operation"] - context = RequestObservabilityContext( - config=observed.config, - request_id="request-1", - method="POST", - route="/calculate", - path="/calculate", - endpoint="calculate", - query_keys=[], - content_length_bytes=None, - inbound={}, - ) - - try: - observed.begin_request(context) - with observed.segment(SegmentName.LOAD): - pass - observed.finish_request(200) - observed.teardown_request(None) - - assert context.timings_ms is not parent_operation.timings_ms - assert context.segment_tree is not parent_operation.segment_tree - assert "load" not in parent_operation.timings_ms - assert parent_operation.segment_tree == [] - assert context.segment_tree[0].name == "load" - assert observed.current_operation() is parent_operation - finally: - observed.end_operation(handle) - - assert observed.current_context() is None - assert observed.current_operation() is None - - -def test_set_attribute_updates_explicit_operation_inside_request() -> None: - observed = runtime() - context = RequestObservabilityContext( - config=observed.config, - request_id="request-1", - method="GET", - route="/chat", - path="/chat", - endpoint="chat", - query_keys=[], - content_length_bytes=None, - inbound={}, - ) - - observed.begin_request(context) - handle = observed.start_operation("chat.turn", flavor="chat") - operation = handle["operation"] - try: - observed.set_attribute("model", "claude") - finally: - observed.end_operation(handle) - observed.teardown_request(None) - - assert context.attributes["model"] == "claude" - assert context.operation_context.attributes["model"] == "claude" - assert operation.attributes["model"] == "claude" - assert observed.current_context() is None - assert observed.current_operation() is None - - -def test_mark_ttft_attribute_updates_current_operation() -> None: - observed = runtime() - handle = observed.start_operation("chat.turn", flavor="chat") - operation = handle["operation"] - - try: - observed.mark_ttft_attribute() - finally: - observed.end_operation(handle) - - assert operation.attributes["ttft_ms"] >= 0 - - -def test_request_methods_noop_without_current_context() -> None: - observed = runtime() - - assert observed.prepare_response(200) == {} - observed.complete_request(200) - observed.update_request_route(route="/missing") - observed.teardown_request(None) - - -def test_request_begin_operation_begin_and_lifecycle_failures_are_logged( - monkeypatch, -) -> None: - class BrokenVar: - def set(self, _value): - raise RuntimeError("set failed") - - observed = runtime() - failures = [] - observed.log_observability_failure = lambda operation, exc, **fields: ( - failures.append(operation) - ) - context = RequestObservabilityContext( - config=observed.config, - request_id="request-1", - method="GET", - route="/broken", - path="/broken", - endpoint="broken", - query_keys=[], - content_length_bytes=None, - inbound={}, - ) - - monkeypatch.setattr(runtime_module, "_REQUEST_CONTEXT", BrokenVar()) - observed.begin_request(context) - monkeypatch.setattr( - runtime_module, - "_REQUEST_CONTEXT", - runtime_module.ContextVar("request", default=None), - ) - monkeypatch.setattr(runtime_module, "_OPERATION_CONTEXT", BrokenVar()) - observed._begin_request_operation(context) - - assert failures == ["request.begin", "request.operation_begin"] - - -def test_request_prepare_complete_update_and_teardown_failures_are_logged() -> ( - None -): - observed = runtime() - failures = [] - observed.log_observability_failure = lambda operation, exc, **fields: ( - failures.append(operation) - ) - context = RequestObservabilityContext( - config=observed.config, - request_id="request-1", - method="GET", - route="/broken", - path="/broken", - endpoint="broken", - query_keys=[], - content_length_bytes=None, - inbound={}, - ) - observed.begin_request(context) - context.span_attributes = lambda **_extra: (_ for _ in ()).throw( - RuntimeError("span attrs failed") - ) - observed.prepare_response(200) - observed.complete_request(200) - observed.update_request_route(route="/other") - observed.emit_request_log = lambda _context: (_ for _ in ()).throw( - RuntimeError("emit failed") - ) - observed.teardown_request(None) - - assert "request.prepare_response" in failures - assert "request.complete" in failures - assert "request.update_route" in failures - assert "request.teardown" in failures - - -def test_set_attribute_failure_path_is_logged() -> None: - observed = runtime() - failures = [] - observed.log_observability_failure = lambda operation, exc, **fields: ( - failures.append(operation) - ) - observed.current_context = lambda: SimpleNamespace( - set_attribute=lambda *_args: (_ for _ in ()).throw( - RuntimeError("attribute failed") - ) - ) - - observed.set_attribute("tool", "loader") - - assert failures == ["request.set_attribute"] - - -def test_request_route_update_relabels_active_request_and_span() -> None: - observed = runtime() - observed.active_requests = RecordingInstrument() - span = NamedRecordingSpan() - context = RequestObservabilityContext( - config=observed.config, - request_id="request-1", - method="GET", - route="/initial", - path="/items/1", - endpoint="initial", - query_keys=[], - content_length_bytes=None, - inbound={}, - ) - context.server_span = span - - observed.begin_request(context) - observed.update_request_route(route="/items/", endpoint="item") - observed.teardown_request(None) - - assert context.route == "/items/" - assert context.endpoint == "item" - assert span.names == ["/items/"] - assert observed.active_requests.calls[1][1] == -1 - assert observed.active_requests.calls[2][1] == 1 - - -def test_prepare_response_includes_traceparent_when_available() -> None: - observed = runtime() - observed.propagate = RecordingPropagator() - context = RequestObservabilityContext( - config=observed.config, - request_id="request-1", - method="GET", - route="/trace", - path="/trace", - endpoint="trace", - query_keys=[], - content_length_bytes=None, - inbound={}, - ) - - observed.begin_request(context) - headers = observed.prepare_response(200) - observed.teardown_request(None) - - assert headers["traceparent"].startswith("00-4bf92f") - - -def test_rate_limited_request_records_rate_limit_metric() -> None: - observed = runtime() - observed.rate_limited = RecordingInstrument() - context = RequestObservabilityContext( - config=observed.config, - request_id="request-1", - method="GET", - route="/limited", - path="/limited", - endpoint="limited", - query_keys=[], - content_length_bytes=None, - inbound={}, - ) - - observed.begin_request(context) - headers = observed.finish_request(429) - observed.teardown_request(None) - - assert headers["X-PolicyEngine-Request-Id"] == "request-1" - assert context.attributes["rate_limited"] is True - assert observed.rate_limited.calls[0][0] == "add" - - -def test_teardown_request_records_unhandled_exception() -> None: - observed = runtime() - observed.errors = RecordingInstrument() - context = RequestObservabilityContext( - config=observed.config, - request_id="request-1", - method="GET", - route="/error", - path="/error", - endpoint="error", - query_keys=[], - content_length_bytes=None, - inbound={}, - ) - - observed.begin_request(context) - observed.teardown_request(RuntimeError("failed")) - - assert context.status_code == 500 - assert context.error is not None - assert context.error.handled is False - assert observed.errors.calls[0][0] == "add" - - -def test_from_env_invalid_shutdown_timeout_falls_back(monkeypatch) -> None: - monkeypatch.setenv("OBSERVABILITY_SHUTDOWN_TIMEOUT_SECONDS", "bad") - - config = ObservabilityConfig.from_env(service_name="svc") - - assert config.shutdown_timeout_seconds == 3.0 - - -def test_from_env_reads_stdout_format(monkeypatch) -> None: - monkeypatch.setenv("OBSERVABILITY_STDOUT_FORMAT", "google") - - config = ObservabilityConfig.from_env(service_name="svc") - - assert config.stdout_format == "google" - - -def test_from_env_reads_queue_knobs(monkeypatch) -> None: - monkeypatch.setenv("OBSERVABILITY_LOG_QUEUE_MAXSIZE", "50") - monkeypatch.setenv("OBSERVABILITY_LOG_QUEUE_CLOSE_TIMEOUT_SECONDS", "1.5") - - config = ObservabilityConfig.from_env(service_name="svc") - - assert config.log_queue_maxsize == 50 - assert config.log_queue_close_timeout_seconds == 1.5 - - -def test_from_env_queue_knobs_fall_back_on_garbage(monkeypatch) -> None: - monkeypatch.setenv("OBSERVABILITY_LOG_QUEUE_MAXSIZE", "many") - monkeypatch.setenv("OBSERVABILITY_LOG_QUEUE_CLOSE_TIMEOUT_SECONDS", "soon") - - config = ObservabilityConfig.from_env(service_name="svc") - - assert config.log_queue_maxsize == 1000 - assert config.log_queue_close_timeout_seconds == 2.0 - - -def test_from_env_enables_otel_by_default() -> None: - config = ObservabilityConfig.from_env(service_name="svc") - - assert config.otel_enabled is True - - -def test_from_env_allows_otel_opt_out(monkeypatch) -> None: - monkeypatch.setenv("OTEL_ENABLED", "false") - - config = ObservabilityConfig.from_env(service_name="svc") - - assert config.otel_enabled is False - - -def test_from_env_ignores_legacy_observability_otel_switch( - monkeypatch, -) -> None: - monkeypatch.setenv("OBSERVABILITY_OTEL_ENABLED", "false") - - config = ObservabilityConfig.from_env(service_name="svc") - - assert config.otel_enabled is True - - -def test_from_env_reads_boolean_csv_and_environment(monkeypatch) -> None: - monkeypatch.setenv("OBSERVABILITY_SERVICE_NAME", "env-svc") - monkeypatch.setenv("DEPLOYMENT_ENVIRONMENT", "production") - monkeypatch.setenv("OBSERVABILITY_ENABLED", "off") - monkeypatch.setenv("OBSERVABILITY_REQUEST_LOGS_ENABLED", "false") - monkeypatch.setenv("OBSERVABILITY_LOG_RAW_IP", "0") - monkeypatch.setenv("OBSERVABILITY_LOG_LEVEL", "warning") - monkeypatch.setenv("OTEL_ENABLED", "1") - monkeypatch.setenv("OTEL_EXPORTER_OTLP_ENDPOINT", "http://collector") - monkeypatch.setenv("OTEL_EXPORTER_OTLP_PROTOCOL", "http/protobuf") - monkeypatch.setenv("OBSERVABILITY_TRACER_NAME", "tracer") - monkeypatch.setenv("OBSERVABILITY_METER_NAME", "meter") - monkeypatch.setenv( - "OBSERVABILITY_METRIC_ATTRIBUTE_KEYS", - "service.name, custom", - ) - monkeypatch.setenv( - "OBSERVABILITY_EXTRA_METRIC_ATTRIBUTE_KEYS", - "custom, other", - ) - - config = ObservabilityConfig.from_env( - service_name="svc", - instrument_fastapi=True, - instrument_httpx=True, - ) - - assert config.service_name == "env-svc" - assert config.environment == "production" - assert config.enabled is False - assert config.request_logs_enabled is False - assert config.log_raw_ip is False - assert config.otel_enabled is True - assert config.otlp_endpoint == "http://collector" - assert config.otlp_protocol == "http/protobuf" - assert config.tracer_name == "tracer" - assert config.meter_name == "meter" - assert config.instrument_fastapi is True - assert config.instrument_httpx is True - assert config.metric_attribute_keys == ("service.name", "custom", "other") - - -def test_from_env_reads_log_destinations_and_google_config( - monkeypatch, -) -> None: - monkeypatch.setenv( - "OBSERVABILITY_LOG_DESTINATIONS", - "stdout, google-cloud-logging, stdout", - ) - monkeypatch.setenv("GCP_PROJECT", "fallback-project") - monkeypatch.setenv("OBSERVABILITY_GOOGLE_CLOUD_LOG_NAME", "custom-log") - - config = ObservabilityConfig.from_env(service_name="svc") - - assert config.log_destinations == ("stdout", "google-cloud-logging") - assert config.google_cloud_project == "fallback-project" - assert config.google_cloud_log_name == "custom-log" - - -def test_from_env_uses_default_log_destinations_without_env( - monkeypatch, -) -> None: - monkeypatch.delenv("OBSERVABILITY_LOG_DESTINATIONS", raising=False) - - config = ObservabilityConfig.from_env( - service_name="svc", - default_log_destinations=("google_cloud_logging",), - ) - - assert config.log_destinations == ("google_cloud_logging",) - - -def test_from_env_log_destinations_env_overrides_default( - monkeypatch, -) -> None: - monkeypatch.setenv("OBSERVABILITY_LOG_DESTINATIONS", "stdout") - - config = ObservabilityConfig.from_env( - service_name="svc", - default_log_destinations=("google_cloud_logging",), - ) - - assert config.log_destinations == ("stdout",) - - -def test_metric_attribute_keys_are_configurable() -> None: - config = ObservabilityConfig( - service_name="svc", - metric_attribute_keys=("service.name", "tool"), - ) - context = RequestObservabilityContext( - config=config, - request_id="request-1", - method="POST", - route="/chat", - path="/chat", - endpoint="chat", - query_keys=[], - content_length_bytes=None, - inbound={}, - ) - context.set_attribute("tool", "search") - context.set_attribute("model", "claude") - - assert context.metric_attributes() == { - "service.name": "svc", - "tool": "search", - } - - -def test_context_set_attribute_normalizes_enum_values() -> None: - observed = runtime() - operation = observed.start_operation("job")["operation"] - request = RequestObservabilityContext( - config=observed.config, - request_id="request-1", - method="GET", - route="/", - path="/", - endpoint="root", - query_keys=[], - content_length_bytes=None, - inbound={}, - ) - - operation.set_attribute("segment", SegmentName.LOAD) - request.set_attribute("segment", SegmentName.SAVE) - observed.end_operation({"operation": operation}) - - assert operation.attributes["segment"] == "load" - assert request.attributes["segment"] == "save" - - -def test_metric_attribute_keys_can_be_extended_from_env(monkeypatch) -> None: - monkeypatch.setenv("OBSERVABILITY_EXTRA_METRIC_ATTRIBUTE_KEYS", "custom") - - config = ObservabilityConfig.from_env(service_name="svc") - - assert config.metric_attribute_keys == ( - *DEFAULT_METRIC_ATTRIBUTE_KEYS, - "custom", - ) - - -def test_segment_metric_uses_configured_metric_attribute_keys() -> None: - observed = runtime( - metric_attribute_keys=( - "service.name", - "route", - "method", - "segment", - "tool", - ) - ) - observed.segment_duration = RecordingInstrument() - context = RequestObservabilityContext( - config=observed.config, - request_id="request-1", - method="POST", - route="/chat", - path="/chat", - endpoint="chat", - query_keys=[], - content_length_bytes=None, - inbound={}, - ) - - observed.begin_request(context) - with observed.segment(SegmentName.LOAD, tool="search", model="claude"): - pass - observed.finish_request(200) - observed.teardown_request(None) - - _, _, attributes = observed.segment_duration.calls[0] - assert attributes["tool"] == "search" - assert "model" not in attributes - - -def test_shutdown_calls_trace_and_metric_providers() -> None: - class Provider: - def __init__(self) -> None: - self.shutdown_called = False - - def shutdown(self) -> None: - self.shutdown_called = True - - observed = runtime(shutdown_timeout_seconds=1) - trace_provider = Provider() - meter_provider = Provider() - observed.tracer_provider = trace_provider - observed.meter_provider = meter_provider - - observed.shutdown() - - assert trace_provider.shutdown_called - assert meter_provider.shutdown_called - - -def test_shutdown_logs_provider_failures_and_timeout() -> None: - class FailingProvider: - def shutdown(self) -> None: - raise RuntimeError("shutdown failed") - - class SlowProvider: - def shutdown(self) -> None: - time.sleep(0.05) - - observed = runtime(shutdown_timeout_seconds=0.001) - observed.tracer_provider = FailingProvider() - observed.meter_provider = SlowProvider() - failures = [] - observed.log_observability_failure = lambda operation, exc, **fields: ( - failures.append(operation) - ) - - observed.shutdown() - - assert "otel.trace_shutdown" in failures - assert "otel.shutdown_timeout" in failures - - -def test_shutdown_closes_destinations_with_full_budget_and_no_thread( - monkeypatch, -) -> None: - observed = runtime(shutdown_timeout_seconds=2.0) - close_calls = [] - monkeypatch.setattr( - observed.log_destination_manager, - "close", - lambda deadline=None: close_calls.append(deadline), - ) - - def fail_thread(*args, **kwargs): - raise AssertionError( - "no watchdog thread should exist without providers" - ) - - monkeypatch.setattr(runtime_module.threading, "Thread", fail_thread) - - observed.shutdown() - - assert close_calls == [2.0] - - -def test_shutdown_destination_deadline_fits_inside_provider_budget( - monkeypatch, -) -> None: - """The prior design gave the log flush a deadline larger than the - join bounding it, starving provider shutdown; the deadline must be - derived from (and smaller than) the shutdown budget.""" - - class Provider: - def __init__(self) -> None: - self.shutdown_called = False - - def shutdown(self) -> None: - self.shutdown_called = True - - observed = runtime(shutdown_timeout_seconds=2.0) - provider = Provider() - observed.tracer_provider = provider - close_calls = [] - monkeypatch.setattr( - observed.log_destination_manager, - "close", - lambda deadline=None: close_calls.append(deadline), - ) - - observed.shutdown() - - assert close_calls == [1.0] - assert provider.shutdown_called - - -def test_shutdown_slow_destination_close_still_runs_providers() -> None: - class Provider: - def __init__(self) -> None: - self.shutdown_called = False - - def shutdown(self) -> None: - self.shutdown_called = True - - observed = runtime(shutdown_timeout_seconds=0.2) - provider = Provider() - observed.tracer_provider = provider - observed.log_destination_manager.close = lambda deadline=None: time.sleep( - 0.05 - ) - - observed.shutdown() - - assert provider.shutdown_called - - -def test_shutdown_clamps_pathological_budget(monkeypatch) -> None: - observed = runtime(shutdown_timeout_seconds=float("inf")) - close_calls = [] - monkeypatch.setattr( - observed.log_destination_manager, - "close", - lambda deadline=None: close_calls.append(deadline), - ) - - observed.shutdown() - - assert close_calls == [3.0] - - -def test_restart_log_destinations_rebuilds_from_config() -> None: - observed = runtime(otel_enabled=False) - observed.configure() - first = observed.log_destination_manager.destinations[0] - - observed.restart_log_destinations() - - rebuilt = observed.log_destination_manager.destinations - assert len(rebuilt) == 1 - assert rebuilt[0] is not first - assert observed.log_destination_manager.configured is True - - -def test_restart_log_destinations_noops_when_disabled(monkeypatch) -> None: - """The kill switch must hold across forks and snapshot restores: - a disabled runtime's restart must not build destinations.""" - observed = runtime(enabled=False) - configure_calls = [] - monkeypatch.setattr( - observed.log_destination_manager, - "configure", - lambda: configure_calls.append(True), - ) - - observed.restart_log_destinations() - - assert configure_calls == [] - - -def test_shutdown_survives_destination_close_failure(monkeypatch) -> None: - class Provider: - def __init__(self) -> None: - self.shutdown_called = False - - def shutdown(self) -> None: - self.shutdown_called = True - - observed = runtime(shutdown_timeout_seconds=1.0) - provider = Provider() - observed.tracer_provider = provider - failures = [] - observed.log_observability_failure = lambda operation, exc, **fields: ( - failures.append(operation) - ) - - def broken_close(deadline=None): - raise RuntimeError("close exploded") - - monkeypatch.setattr( - observed.log_destination_manager, "close", broken_close - ) - - observed.shutdown() - - assert provider.shutdown_called - assert "logging.destination_close" in failures - - -def test_configure_otel_creates_real_providers_and_instruments() -> None: - observed = runtime(otel_enabled=True) - - observed.configure() - - assert observed.tracer is not None - assert observed.meter is not None - assert observed.trace is not None - assert observed.propagate is not None - - -def test_configure_otel_with_exporters_does_not_throw() -> None: - observed = runtime( - otel_enabled=True, - otlp_endpoint="http://localhost:4318", - otlp_protocol="http/protobuf", - ) - failures = [] - observed.log_observability_failure = lambda operation, exc, **fields: ( - failures.append(operation) - ) - - observed.configure() - - assert observed.tracer_provider is not None - assert observed.meter_provider is not None - - -def test_configure_otel_import_failure_is_logged(monkeypatch) -> None: - observed = runtime(otel_enabled=True) - failures = [] - original_import = builtins.__import__ - - def failing_import(name, *args, **kwargs): - if name == "opentelemetry": - raise RuntimeError("otel missing") - return original_import(name, *args, **kwargs) - - monkeypatch.setattr(builtins, "__import__", failing_import) - observed.log_observability_failure = lambda operation, exc, **fields: ( - failures.append(operation) - ) - - observed.configure() - - assert failures == ["otel.configure_imports"] - - -def test_configure_instruments_and_instrument_failures() -> None: - observed = runtime() - meter = RecordingMeter() - observed.meter = meter - - observed._configure_instruments() - - assert len(meter.created) == 11 - - failures = [] - observed.log_observability_failure = lambda operation, exc, **fields: ( - failures.append((operation, fields.get("instrument"))) - ) - noop = observed._instrument( - lambda *_args, **_kwargs: (_ for _ in ()).throw( - RuntimeError("factory failed") - ), - "broken", - ) - - noop.add(1) - noop.record(1) - assert failures == [("metrics.create_instrument", "broken")] - - -def test_request_span_lifecycle_records_enter_and_exit_failures() -> None: - observed = runtime() - observed.tracer = RecordingTracer(fail_enter=True) - failures = [] - observed.log_observability_failure = lambda operation, exc, **fields: ( - failures.append(operation) - ) - context = RequestObservabilityContext( - config=observed.config, - request_id="request-1", - method="GET", - route="/span", - path="/span", - endpoint="span", - query_keys=[], - content_length_bytes=None, - inbound={}, - ) - - observed._start_request_span(context) - assert context.server_span is None - - observed.tracer = RecordingTracer(fail_exit=True) - observed._start_request_span(context) - observed._close_request_span(context, RuntimeError("failed")) - observed._close_request_span(context, None) - - assert failures == ["otel.request_span_enter", "otel.request_span_exit"] - - -def test_safe_span_records_exception_and_preserves_user_error() -> None: - observed = runtime() - observed.tracer = RecordingTracer() - - with pytest.raises(RuntimeError, match="business failed"): - with observed._safe_span("safe", {}): - raise RuntimeError("business failed") - - assert isinstance(observed.tracer.span.exceptions[0], RuntimeError) - - -def test_span_and_segment_failure_helpers_are_logged(monkeypatch) -> None: - observed = runtime() - failures = [] - observed.log_observability_failure = lambda operation, exc, **fields: ( - failures.append(operation) - ) - observed.trace = SimpleNamespace( - get_current_span=lambda: AttributeFailingSpan() - ) - observed._set_current_span_attributes({"key": "value"}) - - observed.trace = SimpleNamespace( - get_current_span=lambda: ExceptionFailingSpan() - ) - observed._record_exception_on_span( - ExceptionFailingSpan(), - RuntimeError("failed"), - handled=False, - status_code=500, - ) - observed._add_span_event("event", {"safe": "yes", "unsafe": object()}) - observed._record_segment_safely("missing_start", None, {}) - monkeypatch.setattr( - observed, - "_safe_perf_counter", - lambda _operation: None, - ) - observed._record_segment_safely("missing_end", 1.0, {}) - - assert "otel.set_span_attributes" in failures - assert "otel.record_exception" in failures - - -def test_segment_helpers_cover_operation_attrs_and_span_prefix() -> None: - observed = runtime(span_prefix="svc") - with observed.operation("job", flavor="cli"): - attrs = observed._segment_span_attributes({"tool": "loader"}) - - assert attrs["policyengine.operation"] == "job" - assert attrs["tool"] == "loader" - assert observed._span_name("load") == "svc.load" - - -def test_contextvar_failure_paths_are_logged(monkeypatch) -> None: - class BrokenVar: - def get(self): - raise RuntimeError("get failed") - - observed = runtime() - failures = [] - observed.log_observability_failure = lambda operation, exc, **fields: ( - failures.append(operation) - ) - monkeypatch.setattr(runtime_module, "_REQUEST_CONTEXT", BrokenVar()) - monkeypatch.setattr(runtime_module, "_OPERATION_CONTEXT", BrokenVar()) - - assert observed.current_context() is None - assert observed.current_operation() is None - assert failures == ["context.current", "operation.current"] - - -def test_runtime_owned_httpx_instrumentation_failure_does_not_throw( - monkeypatch, -) -> None: - observed = runtime(otel_enabled=True) - failures = [] - original_import = builtins.__import__ - - def failing_import(name, *args, **kwargs): - if name == "opentelemetry.instrumentation.httpx": - raise RuntimeError("instrumentation failed") - return original_import(name, *args, **kwargs) - - monkeypatch.setattr(builtins, "__import__", failing_import) - monkeypatch.setattr( - observed, - "log_observability_failure", - lambda operation, exc, **fields: failures.append(operation), - ) - - observed.instrument_httpx() - - assert failures == ["httpx.auto_instrument"] - - -def test_runtime_owned_httpx_instrumentation_success_and_wrapper() -> None: - from policyengine_observability.integrations.httpx import ( - instrument_httpx, - ) - - observed = runtime(otel_enabled=True) - - instrument_httpx(observed) - instrument_httpx(observed) - - assert observed._httpx_instrumented is True - - -def test_traceparent_capture_and_valid_trace_ids() -> None: - observed = runtime() - propagator = RecordingPropagator() - observed.propagate = propagator - span = ValidContextSpan() - observed.trace = SimpleNamespace(get_current_span=lambda: span) - - trace_id, span_id = observed._trace_ids() - - assert observed.traceparent_header().startswith("00-4bf92f") - assert observed._extract_context({"traceparent": "parent"}) == { - "parent": {"traceparent": "parent"} - } - assert propagator.extracted == {"traceparent": "parent"} - assert trace_id == "4bf92f3577b34da6a3ce929d0e0e4736" - assert span_id == "00f067aa0ba902b7" - - -def test_trace_helpers_log_failures_without_throwing() -> None: - observed = runtime() - failures = [] - observed.log_observability_failure = lambda operation, exc, **fields: ( - failures.append(operation) - ) - observed.propagate = SimpleNamespace( - inject=lambda _carrier: (_ for _ in ()).throw(RuntimeError("inject")), - extract=lambda _carrier: (_ for _ in ()).throw( - RuntimeError("extract") - ), - ) - observed.trace = SimpleNamespace( - get_current_span=lambda: (_ for _ in ()).throw(RuntimeError("span")) - ) - - assert observed.traceparent_header() is None - assert observed._extract_context({"traceparent": "parent"}) is None - assert observed._current_span() is None - assert failures == [ - "request.traceparent_header", - "otel.extract_context", - "otel.current_span", - ] - - -def test_record_event_covers_operation_context_and_no_context_metrics() -> ( - None -): - observed = runtime() - observed.failover_events = RecordingInstrument() - - with observed.operation("worker", flavor="queue"): - observed.record_event("modal_retry", attempt=1, ignored=None) - - observed.record_event("fallback_without_context") - - assert len(observed.failover_events.calls) == 2 - assert observed.failover_events.calls[0][2]["operation"] == "worker" - assert observed.failover_events.calls[1][2]["event"] == ( - "fallback_without_context" - ) - - -def test_record_event_request_context_and_emit_log_skip_paths() -> None: - observed = runtime() - observed.failover_events = RecordingInstrument() - context = RequestObservabilityContext( - config=observed.config, - request_id="request-1", - method="GET", - route="/event", - path="/event", - endpoint="event", - query_keys=[], - content_length_bytes=None, - inbound={}, - internal_dispatch=True, - ) - observed.begin_request(context) - observed.record_event("fallback_request", detail="request") - observed.emit_request_log(context) - observed.emit_request_log(context) - observed.teardown_request(None) - - operation = observed.start_operation("job")["operation"] - observed.emit_operation_log(operation) - observed.emit_operation_log(operation) - observed.end_operation({"operation": operation}) - - assert observed.failover_events.calls[0][2]["route"] == "/event" - assert context.emitted is True - assert operation.emitted is True - - -def test_operation_log_emits_to_configured_destinations_once() -> None: - observed = runtime() - destination = RecordingLogDestination() - observed.log_destination_manager.destinations = [destination] - observed.log_destination_manager.configured = True - - handle = observed.start_operation("job") - operation = handle["operation"] - observed.end_operation(handle) - observed.emit_operation_log(operation) - - assert len(destination.calls) == 1 - payload, log_type, severity = destination.calls[0] - assert payload["operation"] == "job" - assert payload["severity"] == "INFO" - assert log_type == "operation" - assert severity == "INFO" - - -def test_request_log_emits_to_configured_destination_once() -> None: - observed = runtime() - destination = RecordingLogDestination() - observed.log_destination_manager.destinations = [destination] - observed.log_destination_manager.configured = True - context = RequestObservabilityContext( - config=observed.config, - request_id="request-1", - method="GET", - route="/calculate", - path="/calculate", - endpoint="calculate", - query_keys=[], - content_length_bytes=None, - inbound={"client_ip": "203.0.113.1"}, - status_code=200, - ) - - observed.emit_request_log(context) - observed.emit_request_log(context) - - assert len(destination.calls) == 1 - payload, log_type, severity = destination.calls[0] - assert payload["request_id"] == "request-1" - assert payload["client_ip"] == "203.0.113.1" - assert payload["severity"] == "INFO" - assert log_type == "request" - assert severity == "INFO" - - -def test_request_log_reserved_fields_override_inbound_and_attributes() -> None: - observed = runtime() - context = RequestObservabilityContext( - config=observed.config, - request_id="request-1", - method="GET", - route="/calculate", - path="/calculate", - endpoint="calculate", - query_keys=[], - content_length_bytes=None, - inbound={ - "request_id": "inbound-request", - "status_code": "inbound-status", - "duration_ms": "inbound-duration", - "segment_tree": "inbound-tree", - }, - attributes={ - "request_id": "attribute-request", - "status_code": "attribute-status", - "duration_ms": "attribute-duration", - "timings_ms": "attribute-timings", - "timing_counts": "attribute-counts", - "segment_tree": "attribute-tree", - "error": "attribute-error", - }, - status_code=204, - ) - - payload = context.as_log_record(trace_id=None, span_id=None) - assert payload["request_id"] == "request-1" - assert payload["status_code"] == 204 - assert isinstance(payload["duration_ms"], float) - assert isinstance(payload["timings_ms"], dict) - assert isinstance(payload["timing_counts"], dict) - assert payload["segment_tree"] == [] - assert payload["error"] is None - - -def test_event_log_emits_to_configured_destination() -> None: - observed = runtime() - destination = RecordingLogDestination() - observed.log_destination_manager.destinations = [destination] - observed.log_destination_manager.configured = True - - observed.record_event("custom_event", detail="value") - - assert len(destination.calls) == 1 - payload, log_type, severity = destination.calls[0] - assert payload["event"] == "custom_event" - assert payload["detail"] == "value" - assert payload["service_name"] == "svc" - assert payload["severity"] == "INFO" - assert log_type == "event" - assert severity == "INFO" - - -def test_log_severity_maps_status_codes_and_errors() -> None: - observed = runtime() - - assert observed._severity_for_log_record({"status_code": 200}) == "INFO" - assert observed._severity_for_log_record({"status_code": 404}) == ( - "WARNING" - ) - assert observed._severity_for_log_record({"status_code": "500"}) == ( - "ERROR" - ) - assert ( - observed._severity_for_log_record({"error": {"handled": False}}) - == "ERROR" - ) - assert observed._severity_for_log_record({"error": {"handled": True}}) == ( - "WARNING" - ) - - -def test_destination_failure_logs_internal_error_without_throwing() -> None: - observed = runtime() - recording = RecordingLogDestination() - observed.log_destination_manager.destinations = [ - FailingLogDestination(), - recording, - ] - observed.log_destination_manager.configured = True - - with observed.operation("job"): - pass - - assert any( - payload["event"] == "observability_internal_error" - and payload["operation"] == "logging.destination_emit" - for payload, _log_type, _severity in recording.calls - ) - assert any( - payload.get("operation") == "job" - for payload, _log_type, _severity in recording.calls - ) - - -def test_all_destination_failures_fall_back_to_stderr(capsys) -> None: - observed = runtime() - observed.log_destination_manager.destinations = [FailingLogDestination()] - observed.log_destination_manager.configured = True - - with observed.operation("job"): - pass - - stderr = capsys.readouterr().err - assert "observability_internal_error" in stderr - assert "logging.destination_emit" in stderr - - -def test_unknown_destination_falls_back_to_stdout() -> None: - observed = ObservabilityRuntime( - ObservabilityConfig( - service_name="svc", - otel_enabled=False, - log_destinations=("missing",), - ) - ) - - observed.configure() - - assert [ - destination.name - for destination in observed.log_destination_manager.destinations - ] == ["stdout"] - - -def test_google_destination_init_failure_falls_back_to_stdout( - monkeypatch, -) -> None: - def fail_google_destination(**_kwargs): - raise ImportError("google-cloud-logging missing") - - monkeypatch.setattr( - google_cloud_logging_module, - "GoogleCloudLoggingDestination", - fail_google_destination, - ) - observed = ObservabilityRuntime( - ObservabilityConfig( - service_name="svc", - otel_enabled=False, - log_destinations=("google_cloud_logging",), - ) - ) - - observed.configure() - - assert [ - destination.name - for destination in observed.log_destination_manager.destinations - ] == ["stdout"] - - -def test_disabled_configure_does_not_initialize_log_destinations( - monkeypatch, -) -> None: - def fail_google_destination(**_kwargs): - raise AssertionError("google destination should not initialize") - - monkeypatch.setattr( - google_cloud_logging_module, - "GoogleCloudLoggingDestination", - fail_google_destination, - ) - observed = ObservabilityRuntime( - ObservabilityConfig( - service_name="svc", - enabled=False, - log_destinations=("google_cloud_logging",), - ) - ) - - observed.configure() - - assert observed.log_destination_manager.destinations == [] - assert observed.log_destination_manager.configured is False - - -def test_record_segment_metric_covers_calculation_and_backend() -> None: - observed = runtime() - observed.segment_duration = RecordingInstrument() - observed.calculate_duration = RecordingInstrument() - observed.backend_duration = RecordingInstrument() - - observed.record_segment_metric( - "calculation", - 0.1, - {"backend": "modal"}, - backend_segment=True, - ) - - assert observed.segment_duration.calls - assert observed.calculate_duration.calls - assert observed.backend_duration.calls - - -def test_metric_recording_failures_are_logged() -> None: - observed = runtime() - observed.operation_duration = FailingInstrument() - observed.http_duration = FailingInstrument() - observed.segment_duration = FailingInstrument() - observed.errors = FailingInstrument() - observed.rate_limited = FailingInstrument() - observed.failover_events = FailingInstrument() - observed.active_requests = FailingInstrument() - failures = [] - observed.log_observability_failure = lambda operation, exc, **fields: ( - failures.append(operation) - ) - - observed.record_operation_metric(0.1, {}) - observed.record_request_metric(0.1, {}) - observed.record_segment_metric("load", 0.1, {}) - observed.record_error_metric({}) - observed.record_rate_limited_metric({}) - observed.record_failover_event_metric({}) - observed.record_active_request(1, {}) - - assert failures == [ - "metrics.record_operation", - "metrics.record_request", - "metrics.record_segment", - "metrics.record_error", - "metrics.record_rate_limited", - "metrics.record_failover_event", - "metrics.add_active_request", - ] - - -def test_private_safety_helpers_cover_fallback_paths( - monkeypatch, - capsys, -) -> None: - observed = runtime() - - class Unprintable: - def __str__(self) -> str: - raise RuntimeError("cannot stringify") - - assert observed._safe_str(Unprintable()) == "" - assert observed._safe_traceback(Unprintable()) == "" - - original_dumps = __import__("json").dumps - - def failing_dumps(payload, *args, **kwargs): - if payload.get("event") == "bad": - raise RuntimeError("json failed") - return original_dumps(payload, *args, **kwargs) - - monkeypatch.setattr("json.dumps", failing_dumps) - assert "observability_internal_error" in observed._json({"event": "bad"}) - - monkeypatch.setattr( - "policyengine_observability.runtime.INTERNAL_LOGGER.error", - lambda _message: (_ for _ in ()).throw(RuntimeError("logger failed")), - ) - observed.log_observability_failure("test", RuntimeError("failed")) - assert "observability_internal_error" in capsys.readouterr().err - - -def test_coerce_segment_name_validates_registry() -> None: - assert coerce_segment_name(SegmentName.LOAD, registry=SegmentName) == ( - "load", - True, - ) - assert coerce_segment_name("other", registry=SegmentName) == ( - "other", - False, - ) - assert coerce_segment_name("load", registry=["load"]) == ("load", True) - assert coerce_segment_name(SegmentName.LOAD, registry=None) == ( - "load", - True, - ) diff --git a/tests/test_runtime_emission.py b/tests/test_runtime_emission.py new file mode 100644 index 0000000..8a57f91 --- /dev/null +++ b/tests/test_runtime_emission.py @@ -0,0 +1,425 @@ +from __future__ import annotations + +from types import SimpleNamespace + +from runtime_helpers import ( + FailingInstrument, + FailingLogDestination, + RecordingInstrument, + RecordingLogDestination, + SegmentName, + runtime, +) + +from policyengine_observability import ( + ObservabilityConfig, + ObservabilityRuntime, + RequestObservabilityContext, + coerce_segment_name, +) +from policyengine_observability.destinations import ( + google_cloud_logging as google_cloud_logging_module, +) + + +def test_trace_helpers_log_failures_without_throwing() -> None: + observed = runtime() + failures = [] + observed.log_observability_failure = lambda operation, exc, **fields: ( + failures.append(operation) + ) + observed.propagate = SimpleNamespace( + inject=lambda _carrier: (_ for _ in ()).throw(RuntimeError("inject")), + extract=lambda _carrier: (_ for _ in ()).throw( + RuntimeError("extract") + ), + ) + observed.trace = SimpleNamespace( + get_current_span=lambda: (_ for _ in ()).throw(RuntimeError("span")) + ) + + assert observed.traceparent_header() is None + assert observed._extract_context({"traceparent": "parent"}) is None + assert observed._current_span() is None + assert failures == [ + "request.traceparent_header", + "otel.extract_context", + "otel.current_span", + ] + + +def test_record_event_covers_operation_context_and_no_context_metrics() -> ( + None +): + observed = runtime() + observed.failover_events = RecordingInstrument() + + with observed.operation("worker", flavor="queue"): + observed.record_event("modal_retry", attempt=1, ignored=None) + + observed.record_event("fallback_without_context") + + assert len(observed.failover_events.calls) == 2 + assert observed.failover_events.calls[0][2]["operation"] == "worker" + assert observed.failover_events.calls[1][2]["event"] == ( + "fallback_without_context" + ) + + +def test_record_event_request_context_and_emit_log_skip_paths() -> None: + observed = runtime() + observed.failover_events = RecordingInstrument() + context = RequestObservabilityContext( + config=observed.config, + request_id="request-1", + method="GET", + route="/event", + path="/event", + endpoint="event", + query_keys=[], + content_length_bytes=None, + inbound={}, + internal_dispatch=True, + ) + observed.begin_request(context) + observed.record_event("fallback_request", detail="request") + observed.emit_request_log(context) + observed.emit_request_log(context) + observed.teardown_request(None) + + operation = observed.start_operation("job")["operation"] + observed.emit_operation_log(operation) + observed.emit_operation_log(operation) + observed.end_operation({"operation": operation}) + + assert observed.failover_events.calls[0][2]["route"] == "/event" + assert context.emitted is True + assert operation.emitted is True + + +def test_operation_log_emits_to_configured_destinations_once() -> None: + observed = runtime() + destination = RecordingLogDestination() + observed.log_destination_manager.destinations = [destination] + observed.log_destination_manager.configured = True + + handle = observed.start_operation("job") + operation = handle["operation"] + observed.end_operation(handle) + observed.emit_operation_log(operation) + + assert len(destination.calls) == 1 + payload, log_type, severity = destination.calls[0] + assert payload["operation"] == "job" + assert payload["severity"] == "INFO" + assert log_type == "operation" + assert severity == "INFO" + + +def test_request_log_emits_to_configured_destination_once() -> None: + observed = runtime() + destination = RecordingLogDestination() + observed.log_destination_manager.destinations = [destination] + observed.log_destination_manager.configured = True + context = RequestObservabilityContext( + config=observed.config, + request_id="request-1", + method="GET", + route="/calculate", + path="/calculate", + endpoint="calculate", + query_keys=[], + content_length_bytes=None, + inbound={"client_ip": "203.0.113.1"}, + status_code=200, + ) + + observed.emit_request_log(context) + observed.emit_request_log(context) + + assert len(destination.calls) == 1 + payload, log_type, severity = destination.calls[0] + assert payload["request_id"] == "request-1" + assert payload["client_ip"] == "203.0.113.1" + assert payload["severity"] == "INFO" + assert log_type == "request" + assert severity == "INFO" + + +def test_request_log_reserved_fields_override_inbound_and_attributes() -> None: + observed = runtime() + context = RequestObservabilityContext( + config=observed.config, + request_id="request-1", + method="GET", + route="/calculate", + path="/calculate", + endpoint="calculate", + query_keys=[], + content_length_bytes=None, + inbound={ + "request_id": "inbound-request", + "status_code": "inbound-status", + "duration_ms": "inbound-duration", + "segment_tree": "inbound-tree", + }, + attributes={ + "request_id": "attribute-request", + "status_code": "attribute-status", + "duration_ms": "attribute-duration", + "timings_ms": "attribute-timings", + "timing_counts": "attribute-counts", + "segment_tree": "attribute-tree", + "error": "attribute-error", + }, + status_code=204, + ) + + payload = context.as_log_record(trace_id=None, span_id=None) + assert payload["request_id"] == "request-1" + assert payload["status_code"] == 204 + assert isinstance(payload["duration_ms"], float) + assert isinstance(payload["timings_ms"], dict) + assert isinstance(payload["timing_counts"], dict) + assert payload["segment_tree"] == [] + assert payload["error"] is None + + +def test_event_log_emits_to_configured_destination() -> None: + observed = runtime() + destination = RecordingLogDestination() + observed.log_destination_manager.destinations = [destination] + observed.log_destination_manager.configured = True + + observed.record_event("custom_event", detail="value") + + assert len(destination.calls) == 1 + payload, log_type, severity = destination.calls[0] + assert payload["event"] == "custom_event" + assert payload["detail"] == "value" + assert payload["service_name"] == "svc" + assert payload["severity"] == "INFO" + assert log_type == "event" + assert severity == "INFO" + + +def test_log_severity_maps_status_codes_and_errors() -> None: + observed = runtime() + + assert observed._severity_for_log_record({"status_code": 200}) == "INFO" + assert observed._severity_for_log_record({"status_code": 404}) == ( + "WARNING" + ) + assert observed._severity_for_log_record({"status_code": "500"}) == ( + "ERROR" + ) + assert ( + observed._severity_for_log_record({"error": {"handled": False}}) + == "ERROR" + ) + assert observed._severity_for_log_record({"error": {"handled": True}}) == ( + "WARNING" + ) + + +def test_destination_failure_logs_internal_error_without_throwing() -> None: + observed = runtime() + recording = RecordingLogDestination() + observed.log_destination_manager.destinations = [ + FailingLogDestination(), + recording, + ] + observed.log_destination_manager.configured = True + + with observed.operation("job"): + pass + + assert any( + payload["event"] == "observability_internal_error" + and payload["operation"] == "logging.destination_emit" + for payload, _log_type, _severity in recording.calls + ) + assert any( + payload.get("operation") == "job" + for payload, _log_type, _severity in recording.calls + ) + + +def test_all_destination_failures_fall_back_to_stderr(capsys) -> None: + observed = runtime() + observed.log_destination_manager.destinations = [FailingLogDestination()] + observed.log_destination_manager.configured = True + + with observed.operation("job"): + pass + + stderr = capsys.readouterr().err + assert "observability_internal_error" in stderr + assert "logging.destination_emit" in stderr + + +def test_unknown_destination_falls_back_to_stdout() -> None: + observed = ObservabilityRuntime( + ObservabilityConfig( + service_name="svc", + otel_enabled=False, + log_destinations=("missing",), + ) + ) + + observed.configure() + + assert [ + destination.name + for destination in observed.log_destination_manager.destinations + ] == ["stdout"] + + +def test_google_destination_init_failure_falls_back_to_stdout( + monkeypatch, +) -> None: + def fail_google_destination(**_kwargs): + raise ImportError("google-cloud-logging missing") + + monkeypatch.setattr( + google_cloud_logging_module, + "GoogleCloudLoggingDestination", + fail_google_destination, + ) + observed = ObservabilityRuntime( + ObservabilityConfig( + service_name="svc", + otel_enabled=False, + log_destinations=("google_cloud_logging",), + ) + ) + + observed.configure() + + assert [ + destination.name + for destination in observed.log_destination_manager.destinations + ] == ["stdout"] + + +def test_disabled_configure_does_not_initialize_log_destinations( + monkeypatch, +) -> None: + def fail_google_destination(**_kwargs): + raise AssertionError("google destination should not initialize") + + monkeypatch.setattr( + google_cloud_logging_module, + "GoogleCloudLoggingDestination", + fail_google_destination, + ) + observed = ObservabilityRuntime( + ObservabilityConfig( + service_name="svc", + enabled=False, + log_destinations=("google_cloud_logging",), + ) + ) + + observed.configure() + + assert observed.log_destination_manager.destinations == [] + assert observed.log_destination_manager.configured is False + + +def test_record_segment_metric_covers_calculation_and_backend() -> None: + observed = runtime() + observed.segment_duration = RecordingInstrument() + observed.calculate_duration = RecordingInstrument() + observed.backend_duration = RecordingInstrument() + + observed.record_segment_metric( + "calculation", + 0.1, + {"backend": "modal"}, + backend_segment=True, + ) + + assert observed.segment_duration.calls + assert observed.calculate_duration.calls + assert observed.backend_duration.calls + + +def test_metric_recording_failures_are_logged() -> None: + observed = runtime() + observed.operation_duration = FailingInstrument() + observed.http_duration = FailingInstrument() + observed.segment_duration = FailingInstrument() + observed.errors = FailingInstrument() + observed.rate_limited = FailingInstrument() + observed.failover_events = FailingInstrument() + observed.active_requests = FailingInstrument() + failures = [] + observed.log_observability_failure = lambda operation, exc, **fields: ( + failures.append(operation) + ) + + observed.record_operation_metric(0.1, {}) + observed.record_request_metric(0.1, {}) + observed.record_segment_metric("load", 0.1, {}) + observed.record_error_metric({}) + observed.record_rate_limited_metric({}) + observed.record_failover_event_metric({}) + observed.record_active_request(1, {}) + + assert failures == [ + "metrics.record_operation", + "metrics.record_request", + "metrics.record_segment", + "metrics.record_error", + "metrics.record_rate_limited", + "metrics.record_failover_event", + "metrics.add_active_request", + ] + + +def test_private_safety_helpers_cover_fallback_paths( + monkeypatch, + capsys, +) -> None: + observed = runtime() + + class Unprintable: + def __str__(self) -> str: + raise RuntimeError("cannot stringify") + + assert observed._safe_str(Unprintable()) == "" + assert observed._safe_traceback(Unprintable()) == "" + + original_dumps = __import__("json").dumps + + def failing_dumps(payload, *args, **kwargs): + if payload.get("event") == "bad": + raise RuntimeError("json failed") + return original_dumps(payload, *args, **kwargs) + + monkeypatch.setattr("json.dumps", failing_dumps) + assert "observability_internal_error" in observed._json({"event": "bad"}) + + monkeypatch.setattr( + "policyengine_observability.runtime.INTERNAL_LOGGER.error", + lambda _message: (_ for _ in ()).throw(RuntimeError("logger failed")), + ) + observed.log_observability_failure("test", RuntimeError("failed")) + assert "observability_internal_error" in capsys.readouterr().err + + +def test_coerce_segment_name_validates_registry() -> None: + assert coerce_segment_name(SegmentName.LOAD, registry=SegmentName) == ( + "load", + True, + ) + assert coerce_segment_name("other", registry=SegmentName) == ( + "other", + False, + ) + assert coerce_segment_name("load", registry=["load"]) == ("load", True) + assert coerce_segment_name(SegmentName.LOAD, registry=None) == ( + "load", + True, + ) diff --git a/tests/test_runtime_operations.py b/tests/test_runtime_operations.py new file mode 100644 index 0000000..13f8e4a --- /dev/null +++ b/tests/test_runtime_operations.py @@ -0,0 +1,295 @@ +from __future__ import annotations + +import asyncio +from types import SimpleNamespace + +import pytest +from runtime_helpers import ( + RecordingInstrument, + RecordingSpan, + RecordingTracer, + SegmentName, + runtime, +) + +from policyengine_observability import ( + RequestObservabilityContext, +) +from policyengine_observability import _state as state_module + + +def test_operation_context_manager_async_and_exception_paths() -> None: + async def run() -> None: + observed = runtime() + observed.errors = RecordingInstrument() + + with pytest.raises(RuntimeError, match="async failed"): + async with observed.operation("async_job", flavor="worker"): + raise RuntimeError("async failed") + + assert observed.errors.calls[0][2]["error_type"] == "RuntimeError" + assert observed.current_operation() is None + + asyncio.run(run()) + + +def test_start_operation_with_parent_context_attaches_and_detaches() -> None: + observed = runtime() + observed.tracer = RecordingTracer() + parent_context = object() + + handle = observed.start_operation( + "parented", + parent_context=parent_context, + ) + observed.end_operation(handle) + + assert observed.tracer.calls[0][0] == "parented" + assert observed.current_operation() is None + + +def test_operation_attach_detach_and_reset_failures_are_logged( + monkeypatch, +) -> None: + from opentelemetry import context as otel_context + + observed = runtime() + observed.tracer = RecordingTracer() + failures = [] + observed.log_observability_failure = lambda operation, exc, **fields: ( + failures.append((operation, fields.get("token"))) + ) + monkeypatch.setattr( + otel_context, + "attach", + lambda _context: (_ for _ in ()).throw(RuntimeError("attach failed")), + ) + + handle = observed.start_operation("job", parent_context=object()) + observed.end_operation(handle) + + monkeypatch.setattr( + otel_context, + "detach", + lambda _token: (_ for _ in ()).throw(RuntimeError("detach failed")), + ) + observed.end_operation( + { + "operation": None, + "context_token": object(), + "timings_token": object(), + "start_token": object(), + "operation_token": object(), + } + ) + + assert ("operation.context_attach", None) in failures + assert ("operation.context_detach", None) in failures + assert ("operation.context_reset", "timings_token") in failures + assert ("operation.context_reset", "start_token") in failures + assert ("operation.context_reset", "operation_token") in failures + + +def test_operation_end_and_start_failures_are_logged(monkeypatch) -> None: + class BrokenVar: + def set(self, _value): + raise RuntimeError("set failed") + + observed = runtime() + failures = [] + observed.log_observability_failure = lambda operation, exc, **fields: ( + failures.append(operation) + ) + operation_handle = observed.start_operation("job") + operation = operation_handle["operation"] + operation.metric_recorded = True + observed.complete_operation(operation) + observed.complete_operation = lambda _operation: (_ for _ in ()).throw( + RuntimeError("complete failed") + ) + observed.end_operation(operation_handle) + + monkeypatch.setattr(state_module, "_OPERATION_CONTEXT", BrokenVar()) + handle = observed.start_operation("job") + + assert handle["operation"] is None + assert failures == ["operation.end", "operation.start"] + + +def test_standalone_segment_creates_implicit_operation_metrics() -> None: + observed = runtime() + observed.segment_duration = RecordingInstrument() + observed.operation_duration = RecordingInstrument() + observed.operations = RecordingInstrument() + emitted_payloads = [] + observed.emit_operation_log = lambda operation: emitted_payloads.append( + operation.as_log_record(trace_id=None, span_id=None) + ) + + with observed.segment(SegmentName.LOAD, flavor="cli", tool="loader"): + pass + + _, _, segment_attributes = observed.segment_duration.calls[0] + _, _, operation_attributes = observed.operation_duration.calls[0] + assert segment_attributes["operation"] == "load" + assert segment_attributes["flavor"] == "cli" + assert segment_attributes["tool"] == "loader" + assert operation_attributes["operation"] == "load" + assert operation_attributes["flavor"] == "cli" + assert emitted_payloads[0]["segment_tree"][0]["name"] == "load" + assert emitted_payloads[0]["segment_tree"][0]["attrs"] == { + "flavor": "cli", + "tool": "loader", + } + assert observed.current_operation() is None + + +def test_segment_with_request_context_does_not_create_implicit_operation( + monkeypatch, +) -> None: + observed = runtime() + observed.segment_duration = RecordingInstrument() + observed.operation_duration = RecordingInstrument() + context = RequestObservabilityContext( + config=observed.config, + request_id="request-1", + method="GET", + route="/calculate", + path="/calculate", + endpoint="calculate", + query_keys=[], + content_length_bytes=None, + inbound={}, + ) + monkeypatch.setattr(observed, "current_context", lambda: context) + + with observed.segment(SegmentName.LOAD): + pass + + _, _, attributes = observed.segment_duration.calls[0] + assert attributes["route"] == "/calculate" + assert observed.operation_duration.calls == [] + + +def test_start_scope_outside_request_records_operation_segment_metrics() -> ( + None +): + observed = runtime() + observed.segment_duration = RecordingInstrument() + timings: dict[str, float] = {} + + handle = observed.start_scope( + timings, + name="chat_turn", + flavor="chat", + model="claude", + ) + with observed.segment(SegmentName.LOAD, tool="search"): + pass + observed.end_scope(handle) + + _, _, attributes = observed.segment_duration.calls[0] + assert "load_ms" in timings + assert attributes["operation"] == "chat_turn" + assert attributes["flavor"] == "chat" + assert attributes["model"] == "claude" + assert attributes["tool"] == "search" + + +def test_nested_scope_annotates_span_context_and_operation() -> None: + observed = runtime() + observed.tracer = RecordingTracer() + parent_context = object() + timings: dict[str, float] = {} + + with observed.operation("outer", flavor="chat"): + handle = observed.start_scope( + timings, + name="inner", + parent_context=parent_context, + ) + observed.annotate(handle, model="claude") + observed.mark("custom_ms", 1.23) + observed.mark_ttft() + observed.end_scope(handle) + + span = observed.tracer.span + assert span.attributes["model"] == "claude" + assert "custom_ms" in timings + + +def test_entrypoint_decorator_records_operation_metrics() -> None: + observed = runtime() + observed.operation_duration = RecordingInstrument() + observed.operations = RecordingInstrument() + + @observed.entrypoint("import_data", flavor="cli") + def run_import() -> str: + return "done" + + assert run_import() == "done" + _, _, attributes = observed.operation_duration.calls[0] + assert attributes["operation"] == "import_data" + assert attributes["flavor"] == "cli" + + +def test_async_segment_decorator_records_segment_metrics() -> None: + observed = runtime() + observed.segment_duration = RecordingInstrument() + + @observed.segment(SegmentName.SAVE, flavor="worker") + async def save() -> str: + return "saved" + + assert asyncio.run(save()) == "saved" + _, _, attributes = observed.segment_duration.calls[0] + assert attributes["operation"] == "save" + assert attributes["flavor"] == "worker" + + +def test_record_error_outside_request_uses_operation_context() -> None: + observed = runtime() + observed.errors = RecordingInstrument() + + with observed.operation("worker", flavor="queue"): + observed.record_error( + RuntimeError("failed"), + handled=True, + include_stack=False, + ) + + _, _, attributes = observed.errors.calls[0] + assert attributes["operation"] == "worker" + assert attributes["flavor"] == "queue" + assert attributes["error_type"] == "RuntimeError" + + +def test_record_error_on_request_updates_span_status() -> None: + observed = runtime() + span = RecordingSpan() + observed.trace = SimpleNamespace(get_current_span=lambda: span) + observed.StatusCode = SimpleNamespace(ERROR="ERROR") + observed.Status = lambda code, message: (code, message) + observed.errors = RecordingInstrument() + context = RequestObservabilityContext( + config=observed.config, + request_id="request-1", + method="GET", + route="/error", + path="/error", + endpoint="error", + query_keys=[], + content_length_bytes=None, + inbound={}, + ) + + observed.begin_request(context) + observed.record_error( + RuntimeError("failed"), + handled=True, + status_code=500, + ) + observed.teardown_request(None) + + assert span.exceptions + assert span.status == ("ERROR", "failed") diff --git a/tests/test_runtime_request_failures.py b/tests/test_runtime_request_failures.py new file mode 100644 index 0000000..8a22d57 --- /dev/null +++ b/tests/test_runtime_request_failures.py @@ -0,0 +1,342 @@ +from __future__ import annotations + +from runtime_helpers import ( + RecordingInstrument, + RecordingPropagator, + SegmentName, + runtime, +) + +from policyengine_observability import ( + ObservabilityConfig, + RequestObservabilityContext, +) +from policyengine_observability.config import DEFAULT_METRIC_ATTRIBUTE_KEYS + + +def test_prepare_response_includes_traceparent_when_available() -> None: + observed = runtime() + observed.propagate = RecordingPropagator() + context = RequestObservabilityContext( + config=observed.config, + request_id="request-1", + method="GET", + route="/trace", + path="/trace", + endpoint="trace", + query_keys=[], + content_length_bytes=None, + inbound={}, + ) + + observed.begin_request(context) + headers = observed.prepare_response(200) + observed.teardown_request(None) + + assert headers["traceparent"].startswith("00-4bf92f") + + +def test_rate_limited_request_records_rate_limit_metric() -> None: + observed = runtime() + observed.rate_limited = RecordingInstrument() + context = RequestObservabilityContext( + config=observed.config, + request_id="request-1", + method="GET", + route="/limited", + path="/limited", + endpoint="limited", + query_keys=[], + content_length_bytes=None, + inbound={}, + ) + + observed.begin_request(context) + headers = observed.finish_request(429) + observed.teardown_request(None) + + assert headers["X-PolicyEngine-Request-Id"] == "request-1" + assert context.attributes["rate_limited"] is True + assert observed.rate_limited.calls[0][0] == "add" + + +def test_teardown_request_records_unhandled_exception() -> None: + observed = runtime() + observed.errors = RecordingInstrument() + context = RequestObservabilityContext( + config=observed.config, + request_id="request-1", + method="GET", + route="/error", + path="/error", + endpoint="error", + query_keys=[], + content_length_bytes=None, + inbound={}, + ) + + observed.begin_request(context) + observed.teardown_request(RuntimeError("failed")) + + assert context.status_code == 500 + assert context.error is not None + assert context.error.handled is False + assert observed.errors.calls[0][0] == "add" + + +def test_from_env_invalid_shutdown_timeout_falls_back(monkeypatch) -> None: + monkeypatch.setenv("OBSERVABILITY_SHUTDOWN_TIMEOUT_SECONDS", "bad") + + config = ObservabilityConfig.from_env(service_name="svc") + + assert config.shutdown_timeout_seconds == 3.0 + + +def test_from_env_reads_stdout_format(monkeypatch) -> None: + monkeypatch.setenv("OBSERVABILITY_STDOUT_FORMAT", "google") + + config = ObservabilityConfig.from_env(service_name="svc") + + assert config.stdout_format == "google" + + +def test_from_env_reads_queue_knobs(monkeypatch) -> None: + monkeypatch.setenv("OBSERVABILITY_LOG_QUEUE_MAXSIZE", "50") + monkeypatch.setenv("OBSERVABILITY_LOG_QUEUE_CLOSE_TIMEOUT_SECONDS", "1.5") + + config = ObservabilityConfig.from_env(service_name="svc") + + assert config.log_queue_maxsize == 50 + assert config.log_queue_close_timeout_seconds == 1.5 + + +def test_from_env_queue_knobs_fall_back_on_garbage(monkeypatch) -> None: + monkeypatch.setenv("OBSERVABILITY_LOG_QUEUE_MAXSIZE", "many") + monkeypatch.setenv("OBSERVABILITY_LOG_QUEUE_CLOSE_TIMEOUT_SECONDS", "soon") + + config = ObservabilityConfig.from_env(service_name="svc") + + assert config.log_queue_maxsize == 1000 + assert config.log_queue_close_timeout_seconds == 2.0 + + +def test_from_env_enables_otel_by_default() -> None: + config = ObservabilityConfig.from_env(service_name="svc") + + assert config.otel_enabled is True + + +def test_from_env_allows_otel_opt_out(monkeypatch) -> None: + monkeypatch.setenv("OTEL_ENABLED", "false") + + config = ObservabilityConfig.from_env(service_name="svc") + + assert config.otel_enabled is False + + +def test_from_env_ignores_legacy_observability_otel_switch( + monkeypatch, +) -> None: + monkeypatch.setenv("OBSERVABILITY_OTEL_ENABLED", "false") + + config = ObservabilityConfig.from_env(service_name="svc") + + assert config.otel_enabled is True + + +def test_from_env_reads_boolean_csv_and_environment(monkeypatch) -> None: + monkeypatch.setenv("OBSERVABILITY_SERVICE_NAME", "env-svc") + monkeypatch.setenv("DEPLOYMENT_ENVIRONMENT", "production") + monkeypatch.setenv("OBSERVABILITY_ENABLED", "off") + monkeypatch.setenv("OBSERVABILITY_REQUEST_LOGS_ENABLED", "false") + monkeypatch.setenv("OBSERVABILITY_LOG_RAW_IP", "0") + monkeypatch.setenv("OBSERVABILITY_LOG_LEVEL", "warning") + monkeypatch.setenv("OTEL_ENABLED", "1") + monkeypatch.setenv("OTEL_EXPORTER_OTLP_ENDPOINT", "http://collector") + monkeypatch.setenv("OTEL_EXPORTER_OTLP_PROTOCOL", "http/protobuf") + monkeypatch.setenv("OBSERVABILITY_TRACER_NAME", "tracer") + monkeypatch.setenv("OBSERVABILITY_METER_NAME", "meter") + monkeypatch.setenv( + "OBSERVABILITY_METRIC_ATTRIBUTE_KEYS", + "service.name, custom", + ) + monkeypatch.setenv( + "OBSERVABILITY_EXTRA_METRIC_ATTRIBUTE_KEYS", + "custom, other", + ) + + config = ObservabilityConfig.from_env( + service_name="svc", + instrument_fastapi=True, + instrument_httpx=True, + ) + + assert config.service_name == "env-svc" + assert config.environment == "production" + assert config.enabled is False + assert config.request_logs_enabled is False + assert config.log_raw_ip is False + assert config.otel_enabled is True + assert config.otlp_endpoint == "http://collector" + assert config.otlp_protocol == "http/protobuf" + assert config.tracer_name == "tracer" + assert config.meter_name == "meter" + assert config.instrument_fastapi is True + assert config.instrument_httpx is True + assert config.metric_attribute_keys == ("service.name", "custom", "other") + + +def test_from_env_reads_log_destinations_and_google_config( + monkeypatch, +) -> None: + monkeypatch.setenv( + "OBSERVABILITY_LOG_DESTINATIONS", + "stdout, google-cloud-logging, stdout", + ) + monkeypatch.setenv("GCP_PROJECT", "fallback-project") + monkeypatch.setenv("OBSERVABILITY_GOOGLE_CLOUD_LOG_NAME", "custom-log") + + config = ObservabilityConfig.from_env(service_name="svc") + + assert config.log_destinations == ("stdout", "google-cloud-logging") + assert config.google_cloud_project == "fallback-project" + assert config.google_cloud_log_name == "custom-log" + + +def test_from_env_uses_default_log_destinations_without_env( + monkeypatch, +) -> None: + monkeypatch.delenv("OBSERVABILITY_LOG_DESTINATIONS", raising=False) + + config = ObservabilityConfig.from_env( + service_name="svc", + default_log_destinations=("google_cloud_logging",), + ) + + assert config.log_destinations == ("google_cloud_logging",) + + +def test_from_env_log_destinations_env_overrides_default( + monkeypatch, +) -> None: + monkeypatch.setenv("OBSERVABILITY_LOG_DESTINATIONS", "stdout") + + config = ObservabilityConfig.from_env( + service_name="svc", + default_log_destinations=("google_cloud_logging",), + ) + + assert config.log_destinations == ("stdout",) + + +def test_metric_attribute_keys_are_configurable() -> None: + config = ObservabilityConfig( + service_name="svc", + metric_attribute_keys=("service.name", "tool"), + ) + context = RequestObservabilityContext( + config=config, + request_id="request-1", + method="POST", + route="/chat", + path="/chat", + endpoint="chat", + query_keys=[], + content_length_bytes=None, + inbound={}, + ) + context.set_attribute("tool", "search") + context.set_attribute("model", "claude") + + assert context.metric_attributes() == { + "service.name": "svc", + "tool": "search", + } + + +def test_context_set_attribute_normalizes_enum_values() -> None: + observed = runtime() + operation = observed.start_operation("job")["operation"] + request = RequestObservabilityContext( + config=observed.config, + request_id="request-1", + method="GET", + route="/", + path="/", + endpoint="root", + query_keys=[], + content_length_bytes=None, + inbound={}, + ) + + operation.set_attribute("segment", SegmentName.LOAD) + request.set_attribute("segment", SegmentName.SAVE) + observed.end_operation({"operation": operation}) + + assert operation.attributes["segment"] == "load" + assert request.attributes["segment"] == "save" + + +def test_metric_attribute_keys_can_be_extended_from_env(monkeypatch) -> None: + monkeypatch.setenv("OBSERVABILITY_EXTRA_METRIC_ATTRIBUTE_KEYS", "custom") + + config = ObservabilityConfig.from_env(service_name="svc") + + assert config.metric_attribute_keys == ( + *DEFAULT_METRIC_ATTRIBUTE_KEYS, + "custom", + ) + + +def test_segment_metric_uses_configured_metric_attribute_keys() -> None: + observed = runtime( + metric_attribute_keys=( + "service.name", + "route", + "method", + "segment", + "tool", + ) + ) + observed.segment_duration = RecordingInstrument() + context = RequestObservabilityContext( + config=observed.config, + request_id="request-1", + method="POST", + route="/chat", + path="/chat", + endpoint="chat", + query_keys=[], + content_length_bytes=None, + inbound={}, + ) + + observed.begin_request(context) + with observed.segment(SegmentName.LOAD, tool="search", model="claude"): + pass + observed.finish_request(200) + observed.teardown_request(None) + + _, _, attributes = observed.segment_duration.calls[0] + assert attributes["tool"] == "search" + assert "model" not in attributes + + +def test_shutdown_calls_trace_and_metric_providers() -> None: + class Provider: + def __init__(self) -> None: + self.shutdown_called = False + + def shutdown(self) -> None: + self.shutdown_called = True + + observed = runtime(shutdown_timeout_seconds=1) + trace_provider = Provider() + meter_provider = Provider() + observed.tracer_provider = trace_provider + observed.meter_provider = meter_provider + + observed.shutdown() + + assert trace_provider.shutdown_called + assert meter_provider.shutdown_called diff --git a/tests/test_runtime_requests.py b/tests/test_runtime_requests.py new file mode 100644 index 0000000..fb42d62 --- /dev/null +++ b/tests/test_runtime_requests.py @@ -0,0 +1,331 @@ +from __future__ import annotations + +from contextvars import ContextVar +from types import SimpleNamespace + +from runtime_helpers import ( + NamedRecordingSpan, + RecordingInstrument, + SegmentName, + runtime, +) + +from policyengine_observability import ( + RequestObservabilityContext, +) +from policyengine_observability import _state as state_module + + +def test_request_lifecycle_records_headers_and_context_metrics() -> None: + observed = runtime() + observed.active_requests = RecordingInstrument() + observed.requests = RecordingInstrument() + observed.http_duration = RecordingInstrument() + context = RequestObservabilityContext( + config=observed.config, + request_id="request-1", + method="GET", + route="/calculate", + path="/calculate", + endpoint="calculate", + query_keys=["country"], + content_length_bytes=None, + inbound={"ip_source": "remote_addr", "client_ip": "127.0.0.1"}, + ) + + observed.begin_request(context) + with observed.segment(SegmentName.LOAD): + pass + headers = observed.finish_request(200) + observed.teardown_request(None) + + assert headers["X-PolicyEngine-Request-Id"] == "request-1" + assert context.status_code == 200 + assert "load" in context.timings_ms + assert observed.current_context() is None + assert observed.current_operation() is None + assert observed.active_requests.calls[0][1] == 1 + assert observed.active_requests.calls[-1][1] == -1 + assert observed.requests.calls[0][1] == 1 + + +def test_request_log_accumulates_repeated_segment_timings() -> None: + observed = runtime() + context = RequestObservabilityContext( + config=observed.config, + request_id="request-1", + method="GET", + route="/calculate", + path="/calculate", + endpoint="calculate", + query_keys=[], + content_length_bytes=None, + inbound={}, + ) + + observed.begin_request(context) + with observed.segment(SegmentName.LOAD): + pass + with observed.segment(SegmentName.LOAD): + pass + observed.finish_request(200) + observed.teardown_request(None) + + assert context.timings_ms["load"] >= 0 + assert context.timing_counts["load"] == 2 + payload = context.as_log_record(trace_id=None, span_id=None) + assert payload["timing_counts"]["load"] == 2 + assert [node["name"] for node in payload["segment_tree"]] == [ + "load", + "load", + ] + + +def test_internal_dispatch_segments_merge_into_parent_operation() -> None: + observed = runtime() + handle = observed.start_operation( + "modal_worker_dispatch", + flavor="modal_worker", + ) + parent_operation = handle["operation"] + context = RequestObservabilityContext( + config=observed.config, + request_id="request-1", + method="POST", + route="/calculate", + path="/calculate", + endpoint="calculate", + query_keys=[], + content_length_bytes=None, + inbound={}, + internal_dispatch=True, + ) + + try: + observed.begin_request(context) + with observed.segment(SegmentName.LOAD): + pass + observed.finish_request(200) + observed.teardown_request(None) + + assert context.timings_ms is parent_operation.timings_ms + assert context.timing_counts is parent_operation.timing_counts + assert context.segment_tree is parent_operation.segment_tree + assert "load" in parent_operation.timings_ms + assert parent_operation.timing_counts["load"] == 1 + assert parent_operation.segment_tree[0].name == "load" + assert observed.current_operation() is parent_operation + finally: + observed.end_operation(handle) + + assert observed.current_context() is None + assert observed.current_operation() is None + + +def test_non_internal_request_timings_do_not_leak_to_parent_operation() -> ( + None +): + observed = runtime() + handle = observed.start_operation("job", flavor="worker") + parent_operation = handle["operation"] + context = RequestObservabilityContext( + config=observed.config, + request_id="request-1", + method="POST", + route="/calculate", + path="/calculate", + endpoint="calculate", + query_keys=[], + content_length_bytes=None, + inbound={}, + ) + + try: + observed.begin_request(context) + with observed.segment(SegmentName.LOAD): + pass + observed.finish_request(200) + observed.teardown_request(None) + + assert context.timings_ms is not parent_operation.timings_ms + assert context.segment_tree is not parent_operation.segment_tree + assert "load" not in parent_operation.timings_ms + assert parent_operation.segment_tree == [] + assert context.segment_tree[0].name == "load" + assert observed.current_operation() is parent_operation + finally: + observed.end_operation(handle) + + assert observed.current_context() is None + assert observed.current_operation() is None + + +def test_set_attribute_updates_explicit_operation_inside_request() -> None: + observed = runtime() + context = RequestObservabilityContext( + config=observed.config, + request_id="request-1", + method="GET", + route="/chat", + path="/chat", + endpoint="chat", + query_keys=[], + content_length_bytes=None, + inbound={}, + ) + + observed.begin_request(context) + handle = observed.start_operation("chat.turn", flavor="chat") + operation = handle["operation"] + try: + observed.set_attribute("model", "claude") + finally: + observed.end_operation(handle) + observed.teardown_request(None) + + assert context.attributes["model"] == "claude" + assert context.operation_context.attributes["model"] == "claude" + assert operation.attributes["model"] == "claude" + assert observed.current_context() is None + assert observed.current_operation() is None + + +def test_mark_ttft_attribute_updates_current_operation() -> None: + observed = runtime() + handle = observed.start_operation("chat.turn", flavor="chat") + operation = handle["operation"] + + try: + observed.mark_ttft_attribute() + finally: + observed.end_operation(handle) + + assert operation.attributes["ttft_ms"] >= 0 + + +def test_request_methods_noop_without_current_context() -> None: + observed = runtime() + + assert observed.prepare_response(200) == {} + observed.complete_request(200) + observed.update_request_route(route="/missing") + observed.teardown_request(None) + + +def test_request_begin_operation_begin_and_lifecycle_failures_are_logged( + monkeypatch, +) -> None: + class BrokenVar: + def set(self, _value): + raise RuntimeError("set failed") + + observed = runtime() + failures = [] + observed.log_observability_failure = lambda operation, exc, **fields: ( + failures.append(operation) + ) + context = RequestObservabilityContext( + config=observed.config, + request_id="request-1", + method="GET", + route="/broken", + path="/broken", + endpoint="broken", + query_keys=[], + content_length_bytes=None, + inbound={}, + ) + + monkeypatch.setattr(state_module, "_REQUEST_CONTEXT", BrokenVar()) + observed.begin_request(context) + monkeypatch.setattr( + state_module, + "_REQUEST_CONTEXT", + ContextVar("request", default=None), + ) + monkeypatch.setattr(state_module, "_OPERATION_CONTEXT", BrokenVar()) + observed._begin_request_operation(context) + + assert failures == ["request.begin", "request.operation_begin"] + + +def test_request_prepare_complete_update_and_teardown_failures_are_logged() -> ( + None +): + observed = runtime() + failures = [] + observed.log_observability_failure = lambda operation, exc, **fields: ( + failures.append(operation) + ) + context = RequestObservabilityContext( + config=observed.config, + request_id="request-1", + method="GET", + route="/broken", + path="/broken", + endpoint="broken", + query_keys=[], + content_length_bytes=None, + inbound={}, + ) + observed.begin_request(context) + context.span_attributes = lambda **_extra: (_ for _ in ()).throw( + RuntimeError("span attrs failed") + ) + observed.prepare_response(200) + observed.complete_request(200) + observed.update_request_route(route="/other") + observed.emit_request_log = lambda _context: (_ for _ in ()).throw( + RuntimeError("emit failed") + ) + observed.teardown_request(None) + + assert "request.prepare_response" in failures + assert "request.complete" in failures + assert "request.update_route" in failures + assert "request.teardown" in failures + + +def test_set_attribute_failure_path_is_logged() -> None: + observed = runtime() + failures = [] + observed.log_observability_failure = lambda operation, exc, **fields: ( + failures.append(operation) + ) + observed.current_context = lambda: SimpleNamespace( + set_attribute=lambda *_args: (_ for _ in ()).throw( + RuntimeError("attribute failed") + ) + ) + + observed.set_attribute("tool", "loader") + + assert failures == ["request.set_attribute"] + + +def test_request_route_update_relabels_active_request_and_span() -> None: + observed = runtime() + observed.active_requests = RecordingInstrument() + span = NamedRecordingSpan() + context = RequestObservabilityContext( + config=observed.config, + request_id="request-1", + method="GET", + route="/initial", + path="/items/1", + endpoint="initial", + query_keys=[], + content_length_bytes=None, + inbound={}, + ) + context.server_span = span + + observed.begin_request(context) + observed.update_request_route(route="/items/", endpoint="item") + observed.teardown_request(None) + + assert context.route == "/items/" + assert context.endpoint == "item" + assert span.names == ["/items/"] + assert observed.active_requests.calls[1][1] == -1 + assert observed.active_requests.calls[2][1] == 1 diff --git a/tests/test_runtime_segments.py b/tests/test_runtime_segments.py new file mode 100644 index 0000000..c92021a --- /dev/null +++ b/tests/test_runtime_segments.py @@ -0,0 +1,303 @@ +from __future__ import annotations + +import asyncio +from typing import Any + +import pytest +from runtime_helpers import ( + AttributeFailingSpan, + RecordingSpan, + RecordingTracer, + SegmentName, + runtime, +) + +from policyengine_observability import ( + UNKNOWN_SEGMENT, + ObservabilityRuntime, + RequestObservabilityContext, +) + + +def test_segment_records_aggregated_timing() -> None: + observed = runtime() + + with observed.collect_timings("request") as timings: + with observed.segment(SegmentName.LOAD): + pass + with observed.segment(SegmentName.LOAD): + pass + + assert "load_ms" in timings + assert timings["load_ms"] >= 0 + + +def test_operation_log_accumulates_repeated_segment_timings() -> None: + observed = runtime() + handle = observed.start_operation("job") + operation = handle["operation"] + + try: + with observed.segment(SegmentName.LOAD): + pass + with observed.segment(SegmentName.LOAD): + pass + finally: + observed.end_operation(handle) + + assert operation.timings_ms["load"] >= 0 + assert operation.timing_counts["load"] == 2 + payload = operation.as_log_record(trace_id=None, span_id=None) + assert payload["timing_counts"]["load"] == 2 + + +def test_operation_log_records_ordered_nested_segment_tree() -> None: + observed = runtime() + handle = observed.start_operation("job") + operation = handle["operation"] + + try: + with observed.segment(SegmentName.LOAD): + with observed.segment( + SegmentName.SAVE, + simulation_kind="baseline", + token="SECRET", + payload={"not": "safe"}, + ): + pass + with observed.segment( + SegmentName.SAVE, + simulation_kind="reform", + ): + pass + finally: + observed.end_operation(handle) + + payload = operation.as_log_record(trace_id=None, span_id=None) + tree = payload["segment_tree"] + assert len(tree) == 1 + assert tree[0]["sequence"] == 1 + assert tree[0]["name"] == "load" + assert "duration_ms" in tree[0] + assert "self_ms" not in tree[0] + + children = tree[0]["children"] + assert [child["sequence"] for child in children] == [2, 3] + assert [child["name"] for child in children] == ["save", "save"] + assert children[0]["attrs"] == {"simulation_kind": "baseline"} + assert children[1]["attrs"] == {"simulation_kind": "reform"} + assert "token" not in children[0].get("attrs", {}) + assert "payload" not in children[0].get("attrs", {}) + assert payload["timing_counts"]["save"] == 2 + + +def test_operation_log_reserved_fields_override_attributes() -> None: + observed = runtime() + handle = observed.start_operation( + "job", + operation="attribute-operation", + duration_ms="attribute-duration", + timings_ms="attribute-timings", + timing_counts="attribute-counts", + segment_tree="attribute-tree", + error="attribute-error", + ) + operation = handle["operation"] + + try: + with observed.segment(SegmentName.LOAD): + pass + finally: + observed.end_operation(handle) + + payload = operation.as_log_record(trace_id=None, span_id=None) + assert payload["operation"] == "job" + assert isinstance(payload["duration_ms"], float) + assert isinstance(payload["timings_ms"], dict) + assert isinstance(payload["timing_counts"], dict) + assert isinstance(payload["segment_tree"], list) + assert payload["error"] is None + + +def test_async_segment_records_timing() -> None: + async def run() -> dict[str, float]: + observed = runtime() + with observed.collect_timings("request") as timings: + async with observed.asegment(SegmentName.SAVE): + pass + return timings + + timings = asyncio.run(run()) + + assert "save_ms" in timings + + +def test_async_segments_keep_independent_segment_tree_stacks() -> None: + async def run() -> list[dict[str, Any]]: + observed = runtime() + handle = observed.start_operation("job") + operation = handle["operation"] + + async def branch(branch_name: str) -> None: + async with observed.asegment(SegmentName.LOAD, branch=branch_name): + await asyncio.sleep(0) + async with observed.asegment( + SegmentName.SAVE, + branch=branch_name, + ): + await asyncio.sleep(0) + + try: + await asyncio.gather(branch("a"), branch("b")) + finally: + observed.end_operation(handle) + return operation.as_log_record(trace_id=None, span_id=None)[ + "segment_tree" + ] + + tree = asyncio.run(run()) + + assert [node["name"] for node in tree] == ["load", "load"] + assert [node["attrs"] for node in tree] == [ + {"branch": "a"}, + {"branch": "b"}, + ] + assert [node["children"][0]["attrs"] for node in tree] == [ + {"branch": "a"}, + {"branch": "b"}, + ] + + +def test_segment_preserves_business_exception_and_records_timing() -> None: + observed = runtime() + + with pytest.raises(ValueError, match="business failed"): + with observed.collect_timings("request") as timings: + with observed.segment(SegmentName.LOAD): + raise ValueError("business failed") + + assert "load_ms" in timings + + +def test_segment_tree_records_failed_segments_before_reraising() -> None: + observed = runtime() + handle = observed.start_operation("job") + operation = handle["operation"] + error = None + + try: + with observed.segment(SegmentName.LOAD): + raise ValueError("business failed") + except ValueError as exc: + error = exc + finally: + observed.end_operation(handle, error) + + payload = operation.as_log_record(trace_id=None, span_id=None) + assert payload["event"] == "operation_failed" + assert payload["segment_tree"][0]["name"] == "load" + assert "duration_ms" in payload["segment_tree"][0] + + +def test_unregistered_segment_falls_back_without_throwing() -> None: + class BrokenString: + def __str__(self) -> str: + raise RuntimeError("cannot stringify") + + observed = runtime() + + with observed.collect_timings("request") as timings: + with observed.segment(BrokenString()): + pass + + assert f"{UNKNOWN_SEGMENT}_ms" in timings + + +def test_segment_span_start_failure_does_not_skip_user_code() -> None: + observed = runtime() + observed.tracer = RecordingTracer(fail_enter=True) + executed = False + + with observed.collect_timings("request") as timings: + with observed.segment(SegmentName.LOAD): + executed = True + + assert executed + assert "load_ms" in timings + + +def test_segment_span_exit_failure_does_not_escape() -> None: + observed = runtime() + observed.tracer = RecordingTracer(fail_exit=True) + + with observed.collect_timings("request") as timings: + with observed.segment(SegmentName.LOAD): + pass + + assert "load_ms" in timings + + +def test_disabled_runtime_noops_across_public_methods() -> None: + observed = ObservabilityRuntime.disabled() + context = RequestObservabilityContext( + config=observed.config, + request_id="request-1", + method="GET", + route="/disabled", + path="/disabled", + endpoint="disabled", + query_keys=[], + content_length_bytes=None, + inbound={}, + ) + + handle = observed.start_operation("disabled") + observed.end_operation(handle) + observed.begin_request(context) + observed.complete_request(200) + observed.update_request_route(route="/other") + observed.teardown_request(None) + observed.set_attribute("key", "value") + observed.record_error(RuntimeError("ignored"), handled=True) + observed.record_event("ignored") + + with observed.segment(SegmentName.LOAD) as span: + assert span is None + + async def run() -> None: + async with observed.asegment(SegmentName.LOAD) as async_span: + assert async_span is None + + asyncio.run(run()) + assert observed.prepare_response(200) == {} + assert observed.current_context() is None + assert observed.current_operation() is None + + +def test_span_attribute_failure_does_not_drop_span_lifecycle() -> None: + observed = runtime() + span = AttributeFailingSpan() + observed.tracer = RecordingTracer(span=span) + failures = [] + observed.log_observability_failure = lambda operation, exc, **fields: ( + failures.append(operation) + ) + + with observed.segment(SegmentName.LOAD, tool="loader"): + pass + + assert "otel.span_attributes" in failures + assert observed.tracer.last_context_manager.exited + + +def test_collect_timings_records_block_exception_on_scope_span() -> None: + observed = runtime() + span = RecordingSpan() + observed.tracer = RecordingTracer(span=span) + + with pytest.raises(RuntimeError, match="scope failed"): + with observed.collect_timings("turn"): + raise RuntimeError("scope failed") + + assert len(span.exceptions) == 1 + assert isinstance(span.exceptions[0], RuntimeError) diff --git a/tests/test_runtime_tracing.py b/tests/test_runtime_tracing.py new file mode 100644 index 0000000..cef5cf3 --- /dev/null +++ b/tests/test_runtime_tracing.py @@ -0,0 +1,420 @@ +from __future__ import annotations + +import builtins +import time +from types import SimpleNamespace + +import pytest +from runtime_helpers import ( + AttributeFailingSpan, + ExceptionFailingSpan, + RecordingMeter, + RecordingPropagator, + RecordingTracer, + ValidContextSpan, + runtime, +) + +from policyengine_observability import ( + RequestObservabilityContext, +) +from policyengine_observability import _state as state_module +from policyengine_observability import runtime as runtime_module + + +def test_shutdown_logs_provider_failures_and_timeout() -> None: + class FailingProvider: + def shutdown(self) -> None: + raise RuntimeError("shutdown failed") + + class SlowProvider: + def shutdown(self) -> None: + time.sleep(0.05) + + observed = runtime(shutdown_timeout_seconds=0.001) + observed.tracer_provider = FailingProvider() + observed.meter_provider = SlowProvider() + failures = [] + observed.log_observability_failure = lambda operation, exc, **fields: ( + failures.append(operation) + ) + + observed.shutdown() + + assert "otel.trace_shutdown" in failures + assert "otel.shutdown_timeout" in failures + + +def test_shutdown_closes_destinations_with_full_budget_and_no_thread( + monkeypatch, +) -> None: + observed = runtime(shutdown_timeout_seconds=2.0) + close_calls = [] + monkeypatch.setattr( + observed.log_destination_manager, + "close", + lambda deadline=None: close_calls.append(deadline), + ) + + def fail_thread(*args, **kwargs): + raise AssertionError( + "no watchdog thread should exist without providers" + ) + + monkeypatch.setattr(runtime_module.threading, "Thread", fail_thread) + + observed.shutdown() + + assert close_calls == [2.0] + + +def test_shutdown_destination_deadline_fits_inside_provider_budget( + monkeypatch, +) -> None: + """The prior design gave the log flush a deadline larger than the + join bounding it, starving provider shutdown; the deadline must be + derived from (and smaller than) the shutdown budget.""" + + class Provider: + def __init__(self) -> None: + self.shutdown_called = False + + def shutdown(self) -> None: + self.shutdown_called = True + + observed = runtime(shutdown_timeout_seconds=2.0) + provider = Provider() + observed.tracer_provider = provider + close_calls = [] + monkeypatch.setattr( + observed.log_destination_manager, + "close", + lambda deadline=None: close_calls.append(deadline), + ) + + observed.shutdown() + + assert close_calls == [1.0] + assert provider.shutdown_called + + +def test_shutdown_slow_destination_close_still_runs_providers() -> None: + class Provider: + def __init__(self) -> None: + self.shutdown_called = False + + def shutdown(self) -> None: + self.shutdown_called = True + + observed = runtime(shutdown_timeout_seconds=0.2) + provider = Provider() + observed.tracer_provider = provider + observed.log_destination_manager.close = lambda deadline=None: time.sleep( + 0.05 + ) + + observed.shutdown() + + assert provider.shutdown_called + + +def test_shutdown_clamps_pathological_budget(monkeypatch) -> None: + observed = runtime(shutdown_timeout_seconds=float("inf")) + close_calls = [] + monkeypatch.setattr( + observed.log_destination_manager, + "close", + lambda deadline=None: close_calls.append(deadline), + ) + + observed.shutdown() + + assert close_calls == [3.0] + + +def test_restart_log_destinations_rebuilds_from_config() -> None: + observed = runtime(otel_enabled=False) + observed.configure() + first = observed.log_destination_manager.destinations[0] + + observed.restart_log_destinations() + + rebuilt = observed.log_destination_manager.destinations + assert len(rebuilt) == 1 + assert rebuilt[0] is not first + assert observed.log_destination_manager.configured is True + + +def test_restart_log_destinations_noops_when_disabled(monkeypatch) -> None: + """The kill switch must hold across forks and snapshot restores: + a disabled runtime's restart must not build destinations.""" + observed = runtime(enabled=False) + configure_calls = [] + monkeypatch.setattr( + observed.log_destination_manager, + "configure", + lambda: configure_calls.append(True), + ) + + observed.restart_log_destinations() + + assert configure_calls == [] + + +def test_shutdown_survives_destination_close_failure(monkeypatch) -> None: + class Provider: + def __init__(self) -> None: + self.shutdown_called = False + + def shutdown(self) -> None: + self.shutdown_called = True + + observed = runtime(shutdown_timeout_seconds=1.0) + provider = Provider() + observed.tracer_provider = provider + failures = [] + observed.log_observability_failure = lambda operation, exc, **fields: ( + failures.append(operation) + ) + + def broken_close(deadline=None): + raise RuntimeError("close exploded") + + monkeypatch.setattr( + observed.log_destination_manager, "close", broken_close + ) + + observed.shutdown() + + assert provider.shutdown_called + assert "logging.destination_close" in failures + + +def test_configure_otel_creates_real_providers_and_instruments() -> None: + observed = runtime(otel_enabled=True) + + observed.configure() + + assert observed.tracer is not None + assert observed.meter is not None + assert observed.trace is not None + assert observed.propagate is not None + + +def test_configure_otel_with_exporters_does_not_throw() -> None: + observed = runtime( + otel_enabled=True, + otlp_endpoint="http://localhost:4318", + otlp_protocol="http/protobuf", + ) + failures = [] + observed.log_observability_failure = lambda operation, exc, **fields: ( + failures.append(operation) + ) + + observed.configure() + + assert observed.tracer_provider is not None + assert observed.meter_provider is not None + + +def test_configure_otel_import_failure_is_logged(monkeypatch) -> None: + observed = runtime(otel_enabled=True) + failures = [] + original_import = builtins.__import__ + + def failing_import(name, *args, **kwargs): + if name == "opentelemetry": + raise RuntimeError("otel missing") + return original_import(name, *args, **kwargs) + + monkeypatch.setattr(builtins, "__import__", failing_import) + observed.log_observability_failure = lambda operation, exc, **fields: ( + failures.append(operation) + ) + + observed.configure() + + assert failures == ["otel.configure_imports"] + + +def test_configure_instruments_and_instrument_failures() -> None: + observed = runtime() + meter = RecordingMeter() + observed.meter = meter + + observed._configure_instruments() + + assert len(meter.created) == 11 + + failures = [] + observed.log_observability_failure = lambda operation, exc, **fields: ( + failures.append((operation, fields.get("instrument"))) + ) + noop = observed._instrument( + lambda *_args, **_kwargs: (_ for _ in ()).throw( + RuntimeError("factory failed") + ), + "broken", + ) + + noop.add(1) + noop.record(1) + assert failures == [("metrics.create_instrument", "broken")] + + +def test_request_span_lifecycle_records_enter_and_exit_failures() -> None: + observed = runtime() + observed.tracer = RecordingTracer(fail_enter=True) + failures = [] + observed.log_observability_failure = lambda operation, exc, **fields: ( + failures.append(operation) + ) + context = RequestObservabilityContext( + config=observed.config, + request_id="request-1", + method="GET", + route="/span", + path="/span", + endpoint="span", + query_keys=[], + content_length_bytes=None, + inbound={}, + ) + + observed._start_request_span(context) + assert context.server_span is None + + observed.tracer = RecordingTracer(fail_exit=True) + observed._start_request_span(context) + observed._close_request_span(context, RuntimeError("failed")) + observed._close_request_span(context, None) + + assert failures == ["otel.request_span_enter", "otel.request_span_exit"] + + +def test_safe_span_records_exception_and_preserves_user_error() -> None: + observed = runtime() + observed.tracer = RecordingTracer() + + with pytest.raises(RuntimeError, match="business failed"): + with observed._safe_span("safe", {}): + raise RuntimeError("business failed") + + assert isinstance(observed.tracer.span.exceptions[0], RuntimeError) + + +def test_span_and_segment_failure_helpers_are_logged(monkeypatch) -> None: + observed = runtime() + failures = [] + observed.log_observability_failure = lambda operation, exc, **fields: ( + failures.append(operation) + ) + observed.trace = SimpleNamespace( + get_current_span=lambda: AttributeFailingSpan() + ) + observed._set_current_span_attributes({"key": "value"}) + + observed.trace = SimpleNamespace( + get_current_span=lambda: ExceptionFailingSpan() + ) + observed._record_exception_on_span( + ExceptionFailingSpan(), + RuntimeError("failed"), + handled=False, + status_code=500, + ) + observed._add_span_event("event", {"safe": "yes", "unsafe": object()}) + observed._record_segment_safely("missing_start", None, {}) + monkeypatch.setattr( + observed, + "_safe_perf_counter", + lambda _operation: None, + ) + observed._record_segment_safely("missing_end", 1.0, {}) + + assert "otel.set_span_attributes" in failures + assert "otel.record_exception" in failures + + +def test_segment_helpers_cover_operation_attrs_and_span_prefix() -> None: + observed = runtime(span_prefix="svc") + with observed.operation("job", flavor="cli"): + attrs = observed._segment_span_attributes({"tool": "loader"}) + + assert attrs["policyengine.operation"] == "job" + assert attrs["tool"] == "loader" + assert observed._span_name("load") == "svc.load" + + +def test_contextvar_failure_paths_are_logged(monkeypatch) -> None: + class BrokenVar: + def get(self): + raise RuntimeError("get failed") + + observed = runtime() + failures = [] + observed.log_observability_failure = lambda operation, exc, **fields: ( + failures.append(operation) + ) + monkeypatch.setattr(state_module, "_REQUEST_CONTEXT", BrokenVar()) + monkeypatch.setattr(state_module, "_OPERATION_CONTEXT", BrokenVar()) + + assert observed.current_context() is None + assert observed.current_operation() is None + assert failures == ["context.current", "operation.current"] + + +def test_runtime_owned_httpx_instrumentation_failure_does_not_throw( + monkeypatch, +) -> None: + observed = runtime(otel_enabled=True) + failures = [] + original_import = builtins.__import__ + + def failing_import(name, *args, **kwargs): + if name == "opentelemetry.instrumentation.httpx": + raise RuntimeError("instrumentation failed") + return original_import(name, *args, **kwargs) + + monkeypatch.setattr(builtins, "__import__", failing_import) + monkeypatch.setattr( + observed, + "log_observability_failure", + lambda operation, exc, **fields: failures.append(operation), + ) + + observed.instrument_httpx() + + assert failures == ["httpx.auto_instrument"] + + +def test_runtime_owned_httpx_instrumentation_success_and_wrapper() -> None: + from policyengine_observability.integrations.httpx import ( + instrument_httpx, + ) + + observed = runtime(otel_enabled=True) + + instrument_httpx(observed) + instrument_httpx(observed) + + assert observed._httpx_instrumented is True + + +def test_traceparent_capture_and_valid_trace_ids() -> None: + observed = runtime() + propagator = RecordingPropagator() + observed.propagate = propagator + span = ValidContextSpan() + observed.trace = SimpleNamespace(get_current_span=lambda: span) + + trace_id, span_id = observed._trace_ids() + + assert observed.traceparent_header().startswith("00-4bf92f") + assert observed._extract_context({"traceparent": "parent"}) == { + "parent": {"traceparent": "parent"} + } + assert propagator.extracted == {"traceparent": "parent"} + assert trace_id == "4bf92f3577b34da6a3ce929d0e0e4736" + assert span_id == "00f067aa0ba902b7" diff --git a/tests/test_destinations.py b/tests/test_stdout_destinations.py similarity index 55% rename from tests/test_destinations.py rename to tests/test_stdout_destinations.py index ae51fb5..aef9c4c 100644 --- a/tests/test_destinations.py +++ b/tests/test_stdout_destinations.py @@ -1,9 +1,7 @@ from __future__ import annotations import json -import math -import pytest from fakes import ( ClosableRecordingDestination, FailingDestination, @@ -13,338 +11,12 @@ ) from policyengine_observability.config import ObservabilityConfig -from policyengine_observability.destinations import ( - GoogleCloudLoggingDestination, - google_cloud_logging, - normalize_payload, -) -from policyengine_observability.destinations.base import ( - accepts_keyword, - clamped, -) from policyengine_observability.destinations.stdout import ( StdoutJsonDestination, resolve_stdout_formatter, ) -class Unprintable: - def __str__(self) -> str: - raise RuntimeError("cannot stringify") - - -class FakeLogger: - def __init__(self) -> None: - self.calls = [] - - def log_struct(self, payload, **kwargs) -> None: - self.calls.append((payload, kwargs)) - - -class FakeGapicApi: - def __init__(self) -> None: - self.calls = [] - - def write_log_entries(self, *args, **kwargs) -> None: - self.calls.append((args, kwargs)) - - -class FakeLoggingApi: - def __init__(self) -> None: - self._gapic_api = FakeGapicApi() - - -class FakeClient: - def __init__(self, *, gapic: bool = False) -> None: - self.project = "resolved-project" - self.fake_logger = FakeLogger() - self.log_names = [] - if gapic: - self.logging_api = FakeLoggingApi() - - def logger(self, log_name: str) -> FakeLogger: - self.log_names.append(log_name) - return self.fake_logger - - -def _google_destination(monkeypatch, client, **kwargs): - monkeypatch.setattr( - google_cloud_logging, - "load_google_credentials", - lambda *, prefer_workload_identity: None, - ) - monkeypatch.setattr( - google_cloud_logging, - "configure_google_application_credentials", - lambda: None, - ) - return GoogleCloudLoggingDestination( - project=None, - log_name="policyengine-observability", - client_factory=lambda _project, _credentials: client, - **kwargs, - ) - - -@pytest.mark.parametrize( - ("value", "expected"), - [ - (5.0, 5.0), - ("5", 5.0), - (0, 0.5), - (-3, 0.5), - (1000, 60.0), - (float("inf"), 10.0), - (float("nan"), 10.0), - (None, 10.0), - ("garbage", 10.0), - ], -) -def test_clamped_bounds_and_rejects_non_finite(value, expected) -> None: - result = clamped(value, low=0.5, high=60.0, default=10.0) - - assert result == expected - assert math.isfinite(result) - - -def test_accepts_keyword_covers_named_var_keyword_and_uninspectable() -> None: - def named(payload, *, timestamp=None): - pass - - def var_keyword(payload, **kwargs): - pass - - def blind(payload): - pass - - assert accepts_keyword(named, "timestamp") is True - assert accepts_keyword(var_keyword, "timestamp") is True - assert accepts_keyword(blind, "timestamp") is False - # Builtins without introspectable signatures degrade to False - # instead of raising at construction time. - assert accepts_keyword(min, "timestamp") is False - - -def test_normalize_payload_recursively_stringifies_unsafe_values() -> None: - normalized = normalize_payload( - { - "keep": "value", - "drop_none": None, - "bytes": b"value", - "list": [1, Unprintable()], - "nested": {"object": object()}, - } - ) - - assert normalized["keep"] == "value" - assert normalized["drop_none"] is None - assert normalized["bytes"] == "value" - assert normalized["list"] == [1, ""] - assert normalized["nested"]["object"].startswith(" None: - monkeypatch.setattr( - google_cloud_logging, - "load_google_credentials", - lambda *, prefer_workload_identity: None, - ) - monkeypatch.setattr( - google_cloud_logging, - "configure_google_application_credentials", - lambda: None, - ) - client = FakeClient() - destination = GoogleCloudLoggingDestination( - project=None, - log_name="policyengine-observability", - client_factory=lambda _project, _credentials: client, - ) - - destination.emit( - { - "schema_version": "policyengine.observability.request.v1", - "service_name": "svc", - "service_role": "api", - "environment": "production", - "request_id": "request-1", - "trace_id": "abc123", - "span_id": "def456", - "path": "/calculate", - "object": object(), - }, - log_type="request", - severity="ERROR", - ) - - payload, kwargs = client.fake_logger.calls[0] - assert client.log_names == ["policyengine-observability"] - assert payload["object"].startswith(" None: - client = FakeClient(gapic=True) - - destination = _google_destination( - monkeypatch, client, write_timeout_seconds=5.0 - ) - # The write path under log_struct funnels through this method; the - # rebinding must inject the bounded retry and per-call timeout. - client.logging_api._gapic_api.write_log_entries(request="sentinel") - - assert destination.write_timeout_seconds == 5.0 - ((args, kwargs),) = client.logging_api._gapic_api.calls - assert kwargs["request"] == "sentinel" - assert kwargs["timeout"] == 5.0 - assert kwargs["retry"].timeout == 5.0 - - -def test_google_destination_clamps_write_timeout(monkeypatch) -> None: - destination = _google_destination( - monkeypatch, FakeClient(gapic=True), write_timeout_seconds=0.0 - ) - - assert destination.write_timeout_seconds == 0.5 - - -def test_google_destination_without_gapic_transport_still_works( - monkeypatch, -) -> None: - client = FakeClient() - - destination = _google_destination(monkeypatch, client) - destination.emit({"event": "x"}, log_type="event", severity="INFO") - - assert len(client.fake_logger.calls) == 1 - - -def test_google_destination_forwards_enqueue_timestamp(monkeypatch) -> None: - from datetime import UTC, datetime - - client = FakeClient() - destination = _google_destination(monkeypatch, client) - stamp = datetime(2026, 7, 8, 12, 0, 0, tzinfo=UTC) - - destination.emit( - {"event": "x"}, log_type="event", severity="INFO", timestamp=stamp - ) - destination.emit({"event": "y"}, log_type="event", severity="INFO") - - (_, stamped_kwargs), (_, plain_kwargs) = client.fake_logger.calls - assert stamped_kwargs["timestamp"] is stamp - assert "timestamp" not in plain_kwargs - - -def test_google_destination_close_closes_client(monkeypatch) -> None: - class ClosableFakeClient(FakeClient): - def __init__(self) -> None: - super().__init__() - self.closed = 0 - - def close(self) -> None: - self.closed += 1 - - client = ClosableFakeClient() - destination = _google_destination(monkeypatch, client) - - destination.close() - - assert client.closed == 1 - - -def test_google_destination_close_tolerates_closeless_client( - monkeypatch, -) -> None: - destination = _google_destination(monkeypatch, FakeClient()) - - destination.close() # FakeClient has no close; must be a no-op - - -def test_google_destination_suppresses_instrumentation_entry( - monkeypatch, -) -> None: - logging_v2 = pytest.importorskip("google.cloud.logging_v2") - monkeypatch.setattr( - logging_v2, "_instrumentation_emitted", False, raising=False - ) - - _google_destination(monkeypatch, FakeClient()) - - assert logging_v2._instrumentation_emitted is True - - -def test_google_factory_reads_write_timeout_env(monkeypatch) -> None: - captured = {} - - class StubDestination: - def __init__(self, **kwargs) -> None: - captured.update(kwargs) - - monkeypatch.setattr( - google_cloud_logging, "GoogleCloudLoggingDestination", StubDestination - ) - monkeypatch.setenv("OBSERVABILITY_GOOGLE_WRITE_TIMEOUT_SECONDS", "2.5") - - from policyengine_observability.destinations.registry import ( - destination_strategy, - ) - - destination_strategy("google_cloud_logging").factory( - config=ObservabilityConfig(google_cloud_project="proj"), - loggers={}, - serializer=json.dumps, - ) - - assert captured["project"] == "proj" - assert captured["write_timeout_seconds"] == 2.5 - - -def test_google_factory_write_timeout_defaults_without_env( - monkeypatch, -) -> None: - captured = {} - - class StubDestination: - def __init__(self, **kwargs) -> None: - captured.update(kwargs) - - monkeypatch.setattr( - google_cloud_logging, "GoogleCloudLoggingDestination", StubDestination - ) - monkeypatch.delenv( - "OBSERVABILITY_GOOGLE_WRITE_TIMEOUT_SECONDS", raising=False - ) - - from policyengine_observability.destinations.registry import ( - destination_strategy, - ) - - destination_strategy("google_cloud_logging").factory( - config=ObservabilityConfig(google_cloud_project="proj"), - loggers={}, - serializer=json.dumps, - ) - - assert captured["write_timeout_seconds"] == 10.0 - - -# ── Stdout formatters ──────────────────────────────────────────────────── - - def _stdout_destination(config=None, formatter=None): logger = RecordingLogger() if formatter is None and config is not None: