File size: 18,474 Bytes
2e818da
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
252
253
254
255
256
257
258
259
260
261
262
263
264
265
266
267
268
269
270
271
272
273
274
275
276
277
278
279
280
281
282
283
284
285
286
287
288
289
290
291
292
293
294
295
296
297
298
299
300
301
302
303
304
305
306
307
308
309
310
311
312
313
314
315
316
317
318
319
320
321
322
323
324
325
326
327
328
329
330
331
332
333
334
335
336
337
338
339
340
341
342
343
344
345
346
347
348
349
350
351
352
353
354
355
356
357
358
359
360
361
362
363
364
365
366
367
368
369
370
371
372
373
374
375
376
377
378
379
380
381
382
383
384
385
386
387
388
389
390
391
392
393
394
395
396
397
398
399
400
401
402
403
404
405
406
407
408
409
410
411
412
413
414
415
416
417
418
419
420
421
422
423
424
425
426
427
428
429
430
431
432
433
434
435
436
437
438
439
440
441
442
443
444
445
446
447
448
449
450
451
452
453
454
455
456
457
458
459
460
461
462
463
464
465
466
467
468
469
470
471
472
473
474
475
476
477
478
479
480
481
482
"""The single operation boundary: one span, bounded metrics, one log, one row.

Every instrumented call site in later tasks wraps its work in
:func:`observe_operation`. The context manager reads the clock **once** and
feeds that single duration to the span, the duration histogram, and the local
observation -- there is no second timer anywhere downstream. On completion it:

1. records one monotonic start time;
2. creates an OTel span when mode is ``local`` or ``full``;
3. attaches sanitized trace attributes;
4. exposes ``set()`` / ``event()`` / ``mark_error()`` / ``mark_terminal()``;
5. records the duration plus bounded RAG stage/measurement histograms;
6. emits a structured correlated log for errors/fallbacks;
7. writes the same completed :class:`OperationObservation` to the local store;
8. suppresses and counts every exporter/store/SDK failure so a telemetry
   problem can never change or block a product result.

``store`` and ``exporter`` are injectable for tests and specialised callers;
when omitted they resolve from the initialized global state. An explicitly
passed ``store`` is always honoured, even before/without global init.
"""

from __future__ import annotations

import logging
import time
from contextlib import contextmanager
from typing import Any, Iterator, Protocol

from opentelemetry import trace
from opentelemetry.trace import Span, Status, StatusCode

from app.observability import bootstrap
from app.observability.contracts import (
    METRIC_ATTRIBUTE_KEYS,
    ExperimentIdentity,
    OperationObservation,
    RetrievalOutcome,
    metric_attributes,
)
from app.observability.logging import emit_correlated_log, trace_context_ids
from app.observability.rag_signals import normalize_measure, normalize_stage
from app.observability.sanitize import (
    REDACTED,
    is_forbidden_key,
    sanitize_attribute_value,
    sanitize_attributes,
    sanitize_exception,
)

logger = logging.getLogger(__name__)

# Statuses that count as a failure/fallback for logging purposes.
_ERROR_STATUSES = frozenset({"error", "error_fallback", "degraded"})


class _Exporter(Protocol):
    def export(self, observation: OperationObservation) -> Any: ...


def _guard(fn, *, count: bool = True) -> None:
    """Run a telemetry side-effect, swallowing (and optionally counting) failure."""
    try:
        fn()
    except Exception:
        if count:
            bootstrap.note_suppressed_failure()
        logger.debug("observability side-effect suppressed", exc_info=True)


class Operation:
    """Handle yielded by :func:`observe_operation`.

    Product code calls ``set``/``event``/``add_stage``/``add_count`` and the
    ``set_experiment``/``set_retrieval``/``mark_*`` helpers; everything is
    buffered and sanitized, then flushed once at ``__exit__``.
    """

    def __init__(
        self,
        name: str,
        *,
        subsystem: str | None,
        consumer: str | None,
        evaluation_run: bool,
    ) -> None:
        self.name = name
        self.subsystem = subsystem
        self.consumer = consumer
        self.evaluation_run = evaluation_run
        self._span: Span | None = None
        self._attributes: dict[str, Any] = {}
        self._stages: dict[str, float] = {}
        self._counts: dict[str, int | float] = {}
        self._exportable_stages: dict[str, float] = {}
        self._exportable_measures: dict[str, float] = {}
        self._experiment: ExperimentIdentity | None = None
        self._retrieval: RetrievalOutcome | None = None
        self._status: str | None = None
        self._errored = False
        self._exception_recorded = False
        self._recorded_exception: BaseException | None = None

    # -- span binding -------------------------------------------------

    def _bind_span(self, span: Span | None) -> None:
        self._span = span
        if span is None:
            return
        self._span_set("operation", self.name)
        if self.subsystem:
            self._span_set("subsystem", self.subsystem)
        if self.consumer:
            self._span_set("consumer", self.consumer)
        if self.evaluation_run:
            self._span_set("evaluation_run", True)

    def _span_set(self, key: str, value: Any) -> None:
        if self._span is None:
            return
        _guard(lambda: self._span.set_attribute(key, value), count=False)

    # -- public API -------------------------------------------------

    def set(self, key: str, value: Any) -> None:
        """Attach one sanitized attribute to the observation and span."""
        skey = str(key)
        if is_forbidden_key(skey):
            sanitized: Any = REDACTED
        else:
            sanitized = sanitize_attribute_value(value)
            if sanitized is None:
                return
        self._attributes[skey] = sanitized
        self._span_set(skey, sanitized)

    def event(self, name: str, attributes: dict[str, Any] | None = None) -> None:
        """Record a timestamped span event (a discrete in-operation moment)."""
        clean = sanitize_attributes(attributes or {})
        if self._attributes.get("diagnostic.smoke") is True:
            clean["diagnostic.smoke"] = True
        if self._span is not None:
            _guard(lambda: self._span.add_event(name, attributes=clean), count=False)

    def add_stage(self, name: str, duration_ms: float) -> None:
        stage = normalize_stage(name)
        self._stages[str(name)] = float(duration_ms)
        if stage:
            self._exportable_stages[stage] = float(duration_ms)
            self._span_set(f"rag.stage.{stage}.ms", float(duration_ms))

    @property
    def stage_durations_ms(self) -> dict[str, float]:
        """Return a read-only snapshot of recorded stage durations."""

        return dict(self._stages)

    def add_count(self, name: str, value: int | float) -> None:
        measure = normalize_measure(name)
        self._counts[str(name)] = value
        if measure:
            self._exportable_measures[measure] = float(value)
            self._span_set(f"rag.measure.{measure}", float(value))

    def set_experiment(self, experiment: ExperimentIdentity) -> None:
        self._experiment = experiment
        for key, value in {
            "experiment.id": experiment.experiment_id,
            "experiment.run_id": experiment.run_id,
            "pipeline.version": experiment.pipeline_version,
            "query.category": experiment.query_category,
            "cache.state": experiment.cache_state,
            "telemetry.mode": experiment.telemetry_mode,
            "cold_start": experiment.cold_start,
        }.items():
            if value not in (None, ""):
                self._span_set(key, value)

    def set_retrieval(self, retrieval: RetrievalOutcome) -> None:
        self._retrieval = retrieval

    def mark_error(self, error_type: str | None = None) -> None:
        """Flag this operation as failed without raising an exception."""
        self._errored = True
        self._status = "error"
        if error_type:
            self._attributes["error.type"] = sanitize_attribute_value(str(error_type))

    def mark_terminal(self, status: str) -> None:
        """Set an explicit terminal status (e.g. ``success_empty``)."""
        self._status = str(status)
        if status in _ERROR_STATUSES:
            self._errored = True

    def record_exception(
        self,
        exc: BaseException,
        *,
        escaped: bool,
        stage: str | None = None,
        category: str | None = None,
    ) -> None:
        """Add one sanitized semantic exception event to this operation.

        This intentionally does not call OpenTelemetry's default exception
        recorder: its traceback payload is not safe for the research path.
        """
        if escaped:
            self._errored = True
            self._status = "error"
        try:
            exception_type, message = sanitize_exception(exc)
            normalized_stage = normalize_stage(stage) if stage is not None else None
            event: dict[str, Any] = {
                "exception.type": exception_type,
                "exception.message": message,
                "exception.escaped": bool(escaped),
            }
            if normalized_stage:
                event["rag.stage"] = normalized_stage
            if category:
                event["error.category"] = sanitize_attribute_value(category)
            self.event("exception", event)
            if escaped:
                self._attributes.setdefault("error.type", exception_type)
                self._attributes.setdefault("error.message", message)
                self._span_set("error.type", self._attributes["error.type"])
                self._span_set("error.message", self._attributes["error.message"])
        except BaseException:
            # Observability must never replace the product exception, including
            # if a third-party span implementation is itself defective.
            logger.debug("exception observability suppressed")
        finally:
            self._exception_recorded = True
            self._recorded_exception = exc

    # -- finalisation -------------------------------------------------

    def _resolve_status(self) -> str:
        if self._status is not None:
            return self._status
        return "error" if self._errored else "success"

    def _metric_dimensions(self, status: str) -> dict[str, Any]:
        dims: dict[str, Any] = {"operation": self.name, "status": status}
        if self.subsystem:
            dims["subsystem"] = self.subsystem
        if self.consumer:
            dims["consumer"] = self.consumer
        dims["evaluation_run"] = self.evaluation_run
        if self._experiment is not None:
            if self._experiment.pipeline_version:
                dims["pipeline.version"] = self._experiment.pipeline_version
            if self._experiment.query_category:
                dims["query.category"] = self._experiment.query_category
            if self._experiment.cache_state:
                dims["cache.state"] = self._experiment.cache_state
        for key in METRIC_ATTRIBUTE_KEYS:
            if key in self._attributes:
                dims.setdefault(key, self._attributes[key])
        return metric_attributes(dims)

    def _build_observation(self, status: str, duration_ms: float | None) -> OperationObservation:
        trace_id, _span_id = trace_context_ids(self._span)
        return OperationObservation(
            operation=self.name,
            status=status,
            subsystem=self.subsystem,
            consumer=self.consumer,
            duration_ms=duration_ms,
            stage_durations_ms=dict(self._stages),
            counts=dict(self._counts),
            trace_id=trace_id,
            experiment=self._experiment,
            retrieval=self._retrieval,
            evaluation_run=self.evaluation_run,
            attributes=dict(self._attributes),
        )

    def _finalize_span(self, status: str) -> None:
        if self._span is None:
            return
        self._span_set("status", status)
        if self._errored or status in _ERROR_STATUSES:
            _guard(lambda: self._span.set_status(Status(StatusCode.ERROR)), count=False)
        else:
            _guard(lambda: self._span.set_status(Status(StatusCode.OK)), count=False)


@contextmanager
def observe_operation(
    name: str,
    attributes: dict[str, Any] | None = None,
    *,
    store: Any = None,
    exporter: _Exporter | None = None,
    subsystem: str | None = None,
    consumer: str | None = None,
    evaluation_run: bool = False,
    record_escaped_exceptions: bool = True,
) -> Iterator[Operation]:
    """Wrap one product operation in the canonical observability boundary.

    Always yields a valid :class:`Operation` so call sites never need
    mode-conditional code. In ``disabled`` mode with no explicitly injected
    store this is a plain timing no-op: no span, no metric, no store write.
    """
    state = bootstrap.get_state()
    resolved_store = store if store is not None else state.local_store
    tracer = state.tracer
    histogram = state.duration_histogram
    stage_histogram = state.stage_duration_histogram
    measurement_histogram = state.rag_measurement_histogram

    op = Operation(
        name,
        subsystem=subsystem,
        consumer=consumer,
        evaluation_run=evaluation_run,
    )

    # (1) one monotonic clock read feeds span, metric, and observation.
    start = time.monotonic()

    # Span setup is guarded: a tracer that fails to start (or enter) a span
    # must degrade to a no-span operation, never propagate into the caller's
    # `with` block. span_cm is only bound once __enter__ has succeeded, so the
    # finally-block never tries to __exit__ a half-started context manager.
    span_cm: Any = None
    span: Span | None = None

    if tracer is not None:
        def _start_span() -> None:
            nonlocal span_cm, span
            cm = tracer.start_as_current_span(name)
            span = cm.__enter__()
            span_cm = cm

        _guard(_start_span, count=False)

    op._bind_span(span)

    if attributes:
        for key, value in attributes.items():
            op.set(key, value)

    exc_to_raise: BaseException | None = None
    try:
        yield op
    except BaseException as exc:  # noqa: BLE001 - re-raised below after recording
        if record_escaped_exceptions:
            if not op._exception_recorded or op._recorded_exception is not exc:
                op.record_exception(exc, escaped=True)
        else:
            # A nested technical boundary (currently a RAG stage) does not
            # own the product fallback decision. Mark its span failed, but
            # defer the single semantic handled/escaped exception event to
            # the parent that either applies a fallback or lets the request
            # fail. This prevents a handled child failure being counted as an
            # escaped request exception.
            op.mark_error(type(exc).__name__)
        exc_to_raise = exc
    finally:
        duration_ms = (time.monotonic() - start) * 1000.0
        status = op._resolve_status()

        op._finalize_span(status)

        dims = op._metric_dimensions(status)
        if histogram is not None:
            _guard(lambda: histogram.record(duration_ms, dims))

        if stage_histogram is not None:
            for stage, stage_duration_ms in op._exportable_stages.items():
                stage_dims = {**dims, "stage": stage}
                _guard(
                    lambda stage_duration_ms=stage_duration_ms, stage_dims=stage_dims: (
                        stage_histogram.record(stage_duration_ms, stage_dims)
                    )
                )

        if measurement_histogram is not None:
            for measure, value in op._exportable_measures.items():
                measure_dims = {**dims, "measure": measure}
                _guard(
                    lambda value=value, measure_dims=measure_dims: (
                        measurement_histogram.record(value, measure_dims)
                    )
                )

        if status in _ERROR_STATUSES or (op._retrieval is not None and op._retrieval.retrieval_error):
            trace_id, span_id = trace_context_ids(span)
            pipeline_version = op._experiment.pipeline_version if op._experiment else None
            error_type = op._attributes.get("error.type")
            diagnostic_smoke = (
                True if op._attributes.get("diagnostic.smoke") is True else None
            )
            _guard(
                lambda: emit_correlated_log(
                    f"operation {op.name} finished with status={status}",
                    operation=op.name,
                    status=status,
                    trace_id=trace_id,
                    span_id=span_id,
                    pipeline_version=pipeline_version,
                    consumer=op.consumer,
                    subsystem=op.subsystem,
                    error_type=error_type,
                    diagnostic_smoke=diagnostic_smoke,
                ),
                count=False,
            )

        observation = op._build_observation(status, round(duration_ms, 3))

        if resolved_store is not None:
            _guard(lambda: resolved_store.record(observation))

        if exporter is not None:
            _guard(lambda: exporter.export(observation))

        if span_cm is not None:
            _guard(lambda: span_cm.__exit__(None, None, None), count=False)

    if exc_to_raise is not None:
        raise exc_to_raise


@contextmanager
def observe_stage(
    parent: Operation,
    stage: str,
    *,
    subsystem: str,
    consumer: str | None = None,
    attributes: dict[str, Any] | None = None,
) -> Iterator[Operation]:
    """Observe one registered RAG stage without risking product behaviour.

    Unknown stage names remain useful in local operation observations, but are
    deliberately not emitted as child spans or bounded stage dimensions.
    """
    normalized = normalize_stage(stage)
    if normalized is None:
        started = time.monotonic()
        try:
            yield parent
        finally:
            parent.add_stage(str(stage), (time.monotonic() - started) * 1000.0)
        return

    started = time.monotonic()
    stage_attributes = {**(attributes or {}), "rag.stage": normalized}
    with observe_operation(
        f"rag.stage.{normalized}",
        subsystem=subsystem,
        consumer=consumer,
        attributes=stage_attributes,
        record_escaped_exceptions=False,
    ) as child:
        try:
            yield child
        except BaseException:
            raise
        finally:
            parent.add_stage(normalized, (time.monotonic() - started) * 1000.0)


def record_event(name: str, attributes: dict[str, Any] | None = None) -> None:
    """Record a standalone event on the currently active span, if any.

    A convenience for call sites that want to mark a discrete moment (a cache
    hit, a fallback taken) without holding the :class:`Operation` handle.
    Sanitizes attributes and is a no-op when no recording span is active.
    """
    span = trace.get_current_span()
    ctx = span.get_span_context() if span is not None else None
    if ctx is None or not ctx.is_valid:
        return
    clean = sanitize_attributes(attributes or {})
    _guard(lambda: span.add_event(name, attributes=clean), count=False)