From 3f45d019b05967fec905fec07e4f8893fb0452ff Mon Sep 17 00:00:00 2001 From: arc53-machine <232052973+arc53-machine@users.noreply.github.com> Date: Wed, 23 Sep 2026 17:19:54 +0100 Subject: [PATCH] Add the per-request trace recorder docsgpt.tracing records agent, LLM, tool, retrieval and embedding spans in memory and writes the finished trace once. Container spans nest by a per-thread stack so suspended generators cannot corrupt the tree; open spans are cancelled at flush and previews are bounded and redacted. --- docsgpt/tracing/__init__.py | 87 ++++++ docsgpt/tracing/core.py | 585 ++++++++++++++++++++++++++++++++++++ docsgpt/tracing/preview.py | 58 ++++ docsgpt/tracing/sink.py | 53 ++++ tests/tracing/__init__.py | 0 tests/tracing/test_trace.py | 341 +++++++++++++++++++++ 6 files changed, 1124 insertions(+) create mode 100644 docsgpt/tracing/__init__.py create mode 100644 docsgpt/tracing/core.py create mode 100644 docsgpt/tracing/preview.py create mode 100644 docsgpt/tracing/sink.py create mode 100644 tests/tracing/__init__.py create mode 100644 tests/tracing/test_trace.py diff --git a/docsgpt/tracing/__init__.py b/docsgpt/tracing/__init__.py new file mode 100644 index 00000000..0cf018d2 --- /dev/null +++ b/docsgpt/tracing/__init__.py @@ -0,0 +1,87 @@ +"""Per-request execution traces. + +Records agent runs, LLM calls, tool calls, retrieval and embeddings as a tree +of timed spans. A finished trace is stored in ``request_traces`` (rendered as +a waterfall in the Logs UI) and replayed as OpenTelemetry GenAI spans when an +OTel SDK is configured. See ``docs/content/Deploying/Observability.mdx``. + +Typical use:: + + trace = tracing.start_trace(source="stream", request_id=request_id) + with tracing.activate(trace): + with tracing.span(tracing.KIND_TOOL, "execute_tool search") as s: + s.preview("arguments", args) + ... + tracing.flush(trace) + +Every call is a no-op when no trace is active. +""" + +from docsgpt.tracing.core import ( + BINDABLE_IDS, + CONTAINER_KINDS, + KIND_AGENT, + KIND_EMBEDDING, + KIND_GUARDRAIL, + KIND_LLM, + KIND_RERANK, + KIND_RETRIEVAL, + KIND_SEARCH, + KIND_STEP, + KIND_TOOL, + NOOP_SPAN, + STATUS_CANCELLED, + STATUS_DENIED, + STATUS_ERROR, + STATUS_OK, + STATUS_PAUSED, + STATUS_PENDING, + STATUS_SKIPPED, + Span, + Trace, + activate, + bind, + bind_if_unset, + current_trace, + mark_content_blocked, + span, + start_span, + start_trace, + wrap, +) +from docsgpt.tracing.sink import discard, flush + +__all__ = [ + "BINDABLE_IDS", + "CONTAINER_KINDS", + "KIND_AGENT", + "KIND_EMBEDDING", + "KIND_GUARDRAIL", + "KIND_LLM", + "KIND_RERANK", + "KIND_RETRIEVAL", + "KIND_SEARCH", + "KIND_STEP", + "KIND_TOOL", + "NOOP_SPAN", + "STATUS_CANCELLED", + "STATUS_DENIED", + "STATUS_ERROR", + "STATUS_OK", + "STATUS_PAUSED", + "STATUS_PENDING", + "STATUS_SKIPPED", + "Span", + "Trace", + "activate", + "bind", + "bind_if_unset", + "current_trace", + "discard", + "flush", + "mark_content_blocked", + "span", + "start_span", + "start_trace", + "wrap", +] diff --git a/docsgpt/tracing/core.py b/docsgpt/tracing/core.py new file mode 100644 index 00000000..b0289152 --- /dev/null +++ b/docsgpt/tracing/core.py @@ -0,0 +1,585 @@ +"""In-memory recording of one execution trace. + +A :class:`Trace` collects spans for a single request (a chat turn, a +scheduled run, a search). Spans are recorded while the request runs and the +whole trace is written once, when its owner calls :func:`flush`. + +Nesting is tracked with one span stack per thread. Only *container* spans +(agent, tool, retrieval, step) are pushed, so a leaf such as an LLM call can +never become the parent of a sibling that starts while it is still open -- +which matters because DocsGPT's agent loop is a chain of suspended +generators. Ending a span pops anything left above it (marked +``cancelled``), and :meth:`Trace.finish` closes whatever is still open, so a +generator finalized late can never corrupt the tree or the stored record. +""" + +from __future__ import annotations + +import contextlib +import functools +import logging +import threading +import time +import uuid +from contextvars import ContextVar +from typing import Any, Callable, Dict, Iterator, List, Optional + +from docsgpt.core.settings import settings +from docsgpt.tracing.preview import make_preview + +logger = logging.getLogger(__name__) + +KIND_AGENT = "agent" +KIND_LLM = "llm" +KIND_TOOL = "tool" +KIND_RETRIEVAL = "retrieval" +KIND_SEARCH = "search" +KIND_EMBEDDING = "embedding" +KIND_RERANK = "rerank" +KIND_GUARDRAIL = "guardrail" +KIND_STEP = "step" + +#: Kinds that become the implicit parent of spans started while they are open. +CONTAINER_KINDS = frozenset({KIND_AGENT, KIND_TOOL, KIND_RETRIEVAL, KIND_STEP}) + +STATUS_OK = "ok" +STATUS_ERROR = "error" +STATUS_CANCELLED = "cancelled" +STATUS_PAUSED = "paused" +STATUS_PENDING = "pending" +STATUS_DENIED = "denied" +STATUS_SKIPPED = "skipped" + +#: Trace ids ``bind`` may set; anything else is ignored. +BINDABLE_IDS = frozenset( + { + "request_id", + "message_id", + "conversation_id", + "activity_id", + "workflow_run_id", + "user_id", + "agent_id", + "name", + } +) + +_current: ContextVar[Optional["Trace"]] = ContextVar("docsgpt_trace", default=None) + + +def _new_id() -> str: + return uuid.uuid4().hex[:16] + + +class Span: + """One timed step inside a trace. + + Attribute keys follow the OTel GenAI conventions (``gen_ai.*``) where one + exists and ``docsgpt.*`` otherwise, so the stored trace and the exported + spans share a vocabulary. + """ + + __slots__ = ( + "trace", + "id", + "parent_id", + "kind", + "name", + "attributes", + "previews", + "status", + "error", + "start_perf_ns", + "end_perf_ns", + "_thread", + "_pushed", + ) + + def __init__( + self, + trace: "Trace", + kind: str, + name: str, + parent_id: Optional[str], + attributes: Optional[Dict[str, Any]] = None, + ) -> None: + self.trace = trace + self.id = _new_id() + self.parent_id = parent_id + self.kind = kind + self.name = name + self.attributes: Dict[str, Any] = dict(attributes or {}) + self.previews: Dict[str, Any] = {} + self.status: Optional[str] = None + self.error: Optional[str] = None + self.start_perf_ns = time.perf_counter_ns() + self.end_perf_ns: Optional[int] = None + self._thread: Optional[int] = None + self._pushed = False + + @property + def ended(self) -> bool: + return self.end_perf_ns is not None + + @property + def duration_ms(self) -> Optional[float]: + if self.end_perf_ns is None: + return None + return (self.end_perf_ns - self.start_perf_ns) / 1e6 + + def set(self, **attributes: Any) -> "Span": + """Merge attributes; ``None`` values are ignored. Keys may contain dots via ``**{...}``.""" + if not self.ended: + self.attributes.update({k: v for k, v in attributes.items() if v is not None}) + return self + + def preview(self, key: str, value: Any) -> "Span": + """Attach a bounded, redacted content preview (skipped when capture is off).""" + if self.ended or value is None or not settings.TRACES_CAPTURE_CONTENT: + return self + try: + self.previews[key] = make_preview(value) + except Exception: # noqa: BLE001 - a preview must never break a request + logger.debug("trace preview failed for %s", key, exc_info=True) + return self + + def fail(self, exc: BaseException) -> "Span": + """Record ``exc`` on the span without ending it.""" + self.status = STATUS_ERROR + self.error = str(exc)[:500] or type(exc).__name__ + self.attributes["error.type"] = type(exc).__name__ + return self + + def end( + self, + status: Optional[str] = None, + *, + error: Optional[BaseException] = None, + attributes: Optional[Dict[str, Any]] = None, + ) -> None: + """Close the span. A second end, or an end after the trace finished, is ignored.""" + if self.ended: + return + if attributes: + self.set(**attributes) + if error is not None: + self.fail(error) + if status is not None: + self.status = status + elif self.status is None: + self.status = STATUS_OK + self.end_perf_ns = time.perf_counter_ns() + self.trace._on_end(self) + + def __enter__(self) -> "Span": + return self + + def __exit__(self, exc_type, exc, tb) -> bool: + if exc is None: + self.end() + elif isinstance(exc, GeneratorExit): + self.end(STATUS_CANCELLED) + else: + self.end(error=exc) + return False + + +class _NoopSpan: + """Stand-in returned when no trace is active or the span cap is reached.""" + + id = None + parent_id = None + kind = None + name = "" + status = None + ended = True + duration_ms = None + + @property + def attributes(self) -> Dict[str, Any]: + return {} + + @property + def previews(self) -> Dict[str, Any]: + return {} + + def set(self, **attributes: Any) -> "_NoopSpan": + return self + + def preview(self, key: str, value: Any) -> "_NoopSpan": + return self + + def fail(self, exc: BaseException) -> "_NoopSpan": + return self + + def end(self, status=None, *, error=None, attributes=None) -> None: + return None + + def __enter__(self) -> "_NoopSpan": + return self + + def __exit__(self, exc_type, exc, tb) -> bool: + return False + + def __bool__(self) -> bool: + return False + + +NOOP_SPAN = _NoopSpan() + + +class Trace: + """All spans recorded for one execution, plus the ids that link it to logs.""" + + def __init__( + self, + *, + source: str, + name: Optional[str] = None, + request_id: Optional[str] = None, + message_id: Optional[str] = None, + conversation_id: Optional[str] = None, + activity_id: Optional[str] = None, + workflow_run_id: Optional[str] = None, + user_id: Optional[str] = None, + agent_id: Optional[str] = None, + otel_context: Any = None, + ) -> None: + self.id = str(uuid.uuid4()) + self.source = source + self.name = name or source + self.request_id = request_id + self.message_id = message_id + self.conversation_id = conversation_id + self.activity_id = activity_id + self.workflow_run_id = workflow_run_id + self.user_id = user_id + self.agent_id = agent_id + self.otel_context = otel_context + self.otel_trace_id: Optional[str] = None + self.start_ns = time.time_ns() + self.start_perf_ns = time.perf_counter_ns() + self.end_perf_ns: Optional[int] = None + self.status: Optional[str] = None + self.spans: List[Span] = [] + self.dropped_spans = 0 + self.content_blocked = False + self.attributes: Dict[str, Any] = {} + self.finished = False + self.flushed = False + self._lock = threading.Lock() + self._stacks: Dict[int, List[Any]] = {} + + # -- span lifecycle ------------------------------------------------- + + def _parent_for_current_thread(self) -> Optional[str]: + stack = self._stacks.get(threading.get_ident()) + return stack[-1].id if stack else None + + def start_span( + self, + kind: str, + name: str, + *, + parent: Any = None, + attributes: Optional[Dict[str, Any]] = None, + ): + """Record a new span; returns :data:`NOOP_SPAN` once finished or over the cap.""" + with self._lock: + if self.finished: + return NOOP_SPAN + if len(self.spans) >= settings.TRACES_MAX_SPANS: + self.dropped_spans += 1 + return NOOP_SPAN + if parent is not None: + parent_id = getattr(parent, "id", None) + else: + parent_id = self._parent_for_current_thread() + span = Span(self, kind, name, parent_id, attributes) + self.spans.append(span) + if kind in CONTAINER_KINDS: + ident = threading.get_ident() + self._stacks.setdefault(ident, []).append(span) + span._thread = ident + span._pushed = True + return span + + def _on_end(self, span: Span) -> None: + if not span._pushed: + return + with self._lock: + stack = self._stacks.get(span._thread) + if not stack or span not in stack: + return + index = stack.index(span) + abandoned = stack[index + 1:] + del stack[index:] + if not stack: + self._stacks.pop(span._thread, None) + for child in reversed(abandoned): + if isinstance(child, Span) and not child.ended: + child.end(STATUS_CANCELLED) + + def _seed_thread(self, parent: Optional[Span]) -> Callable[[], None]: + """Make ``parent`` the implicit parent in this thread; returns an undo callable.""" + if parent is None: + return lambda: None + ident = threading.get_ident() + marker = _Seed(parent.id) + with self._lock: + self._stacks.setdefault(ident, []).append(marker) + + def _undo() -> None: + with self._lock: + stack = self._stacks.get(ident) + if stack and marker in stack: + del stack[stack.index(marker):] + if not stack: + self._stacks.pop(ident, None) + + return _undo + + def current_parent(self) -> Optional["_Seed | Span"]: + stack = self._stacks.get(threading.get_ident()) + return stack[-1] if stack else None + + # -- ids -------------------------------------------------------------- + + def bind(self, *, only_if_unset: bool = False, **ids: Any) -> None: + for key, value in ids.items(): + if key not in BINDABLE_IDS or value is None: + continue + if only_if_unset and getattr(self, key, None): + continue + setattr(self, key, str(value)) + + # -- completion --------------------------------------------------------- + + def finish(self, status: Optional[str] = None) -> None: + """Freeze the trace: close open spans as ``cancelled`` and set the status.""" + with self._lock: + if self.finished: + return + self.finished = True + open_spans = [s for s in self.spans if not s.ended] + self._stacks.clear() + now = time.perf_counter_ns() + for span in open_spans: + span.status = STATUS_CANCELLED + span.end_perf_ns = now + self.end_perf_ns = now + if status is not None: + self.status = status + else: + failed = any( + s.parent_id is None and s.status == STATUS_ERROR for s in self.spans + ) + self.status = STATUS_ERROR if failed else STATUS_OK + + @property + def duration_ms(self) -> Optional[float]: + if self.end_perf_ns is None: + return None + return (self.end_perf_ns - self.start_perf_ns) / 1e6 + + def span_start_ns(self, span: Span) -> int: + """Wall-clock start of ``span`` in ns, derived from the monotonic offset.""" + return self.start_ns + (span.start_perf_ns - self.start_perf_ns) + + def span_end_ns(self, span: Span) -> int: + end = span.end_perf_ns if span.end_perf_ns is not None else self.end_perf_ns + return self.start_ns + ((end or span.start_perf_ns) - self.start_perf_ns) + + def summary(self) -> Dict[str, Any]: + """Aggregate counts shown as chips in the Logs UI.""" + by_id = {s.id: s for s in self.spans} + llm = [s for s in self.spans if s.kind == KIND_LLM] + # Outermost retrieval spans only, so nested dispatcher/retriever + # spans are not double-counted. + retrieval = [ + s + for s in self.spans + if s.kind == KIND_RETRIEVAL + and not (s.parent_id in by_id and by_id[s.parent_id].kind == KIND_RETRIEVAL) + ] + tools = [s for s in self.spans if s.kind == KIND_TOOL] + + def _tokens(key: str) -> int: + total = 0 + for s in llm: + value = s.attributes.get(key) + if isinstance(value, (int, float)): + total += int(value) + return total + + return { + "llm_calls": len(llm), + "tool_calls": len(tools), + "retrieval_calls": len(retrieval), + "retrieval_ms": round(sum(s.duration_ms or 0 for s in retrieval), 1), + "input_tokens": _tokens("gen_ai.usage.input_tokens"), + "output_tokens": _tokens("gen_ai.usage.output_tokens"), + "errors": sum(1 for s in self.spans if s.status == STATUS_ERROR), + } + + def to_record(self) -> Dict[str, Any]: + """The ``request_traces`` row for this (finished) trace.""" + spans = [] + for s in self.spans: + entry: Dict[str, Any] = { + "id": s.id, + "parent_id": s.parent_id, + "kind": s.kind, + "name": s.name, + "status": s.status or STATUS_CANCELLED, + "offset_ms": round((s.start_perf_ns - self.start_perf_ns) / 1e6, 2), + "duration_ms": round(s.duration_ms or 0.0, 2), + "attributes": s.attributes, + } + if s.error: + entry["error"] = s.error + if s.previews and not self.content_blocked: + entry["preview"] = s.previews + spans.append(entry) + return { + "id": self.id, + "request_id": self.request_id, + "message_id": self.message_id, + "conversation_id": self.conversation_id, + "activity_id": self.activity_id, + "workflow_run_id": self.workflow_run_id, + "user_id": self.user_id, + "agent_id": self.agent_id, + "source": self.source, + "name": self.name, + "status": self.status or STATUS_OK, + "started_at_ns": self.start_ns, + "duration_ms": int(round(self.duration_ms or 0)), + "span_count": len(spans), + "dropped_spans": self.dropped_spans, + "summary": self.summary(), + "spans": spans, + "otel_trace_id": self.otel_trace_id, + } + + +class _Seed: + """Placeholder stack entry naming a parent span owned by another thread.""" + + __slots__ = ("id",) + + def __init__(self, span_id: Optional[str]) -> None: + self.id = span_id + + +# -- module-level API -------------------------------------------------------- + + +def current_trace() -> Optional[Trace]: + """The trace active in this context, or ``None``.""" + return _current.get() + + +def start_trace(*, source: str, capture_otel_context: bool = True, **ids: Any) -> Optional[Trace]: + """Create a trace, or return ``None`` when ``TRACES_ENABLED`` is off. + + The OTel context current at this point (normally the HTTP server span) is + captured so the exported GenAI spans hang off the request's own trace. + """ + if not settings.TRACES_ENABLED: + return None + otel_context = None + if capture_otel_context: + try: + from opentelemetry import context as otel_ctx + + otel_context = otel_ctx.get_current() + except Exception: # noqa: BLE001 + otel_context = None + known = {k: v for k, v in ids.items() if k in BINDABLE_IDS} + known = {k: (str(v) if v is not None else None) for k, v in known.items()} + return Trace(source=source, otel_context=otel_context, **known) + + +@contextlib.contextmanager +def activate(trace: Optional[Trace]) -> Iterator[Optional[Trace]]: + """Make ``trace`` current for the enclosed block (a no-op for ``None``).""" + if trace is None: + yield None + return + token = _current.set(trace) + try: + yield trace + finally: + try: + _current.reset(token) + except ValueError: + # Reset from a different context (a generator finalized + # elsewhere); clearing is the closest safe equivalent. + _current.set(None) + + +def start_span(kind: str, name: str, *, parent: Any = None, attributes: Optional[Dict[str, Any]] = None): + """Start a span in the current trace; returns :data:`NOOP_SPAN` without one.""" + trace = _current.get() + if trace is None: + return NOOP_SPAN + try: + return trace.start_span(kind, name, parent=parent, attributes=attributes) + except Exception: # noqa: BLE001 - tracing must never break a request + logger.debug("trace start_span failed", exc_info=True) + return NOOP_SPAN + + +def span(kind: str, name: str, *, parent: Any = None, attributes: Optional[Dict[str, Any]] = None): + """Context-manager form of :func:`start_span` (errors mark the span and re-raise).""" + return start_span(kind, name, parent=parent, attributes=attributes) + + +def bind(**ids: Any) -> None: + """Set link ids (``message_id``, ``conversation_id``, ...) on the current trace.""" + trace = _current.get() + if trace is not None: + trace.bind(**ids) + + +def bind_if_unset(**ids: Any) -> None: + """Like :func:`bind` but keeps any value already set.""" + trace = _current.get() + if trace is not None: + trace.bind(only_if_unset=True, **ids) + + +def mark_content_blocked() -> None: + """A guardrail blocked or retracted content: drop every preview from the stored trace.""" + trace = _current.get() + if trace is not None: + trace.content_blocked = True + + +def wrap(fn: Callable[..., Any]) -> Callable[..., Any]: + """Bind the current trace and parent span into ``fn`` for another thread. + + Use around work handed to a thread pool or ``threading.Thread``; those do + not inherit context variables. Without an active trace ``fn`` is returned + unchanged. + """ + trace = _current.get() + if trace is None: + return fn + parent = trace.current_parent() + + @functools.wraps(fn) + def _run(*args: Any, **kwargs: Any) -> Any: + token = _current.set(trace) + undo = trace._seed_thread(parent) + try: + return fn(*args, **kwargs) + finally: + undo() + try: + _current.reset(token) + except ValueError: + _current.set(None) + + return _run diff --git a/docsgpt/tracing/preview.py b/docsgpt/tracing/preview.py new file mode 100644 index 00000000..892e00f2 --- /dev/null +++ b/docsgpt/tracing/preview.py @@ -0,0 +1,58 @@ +"""Bounded, secret-redacted content previews stored with a trace span.""" + +from __future__ import annotations + +import json +from typing import Any + +from docsgpt.core.settings import settings +from docsgpt.storage.db.redaction import redact_secrets +from docsgpt.utils import strip_null_bytes + +_ELLIPSIS = "…" +# A list longer than this keeps its head only; the preview is for reading, +# not for reconstructing the payload. +_MAX_LIST_ITEMS = 50 + + +def _bound(value: Any, limit: int) -> Any: + """Truncate strings and long lists inside ``value``; stringify unknown objects.""" + if isinstance(value, str): + return value if len(value) <= limit else value[:limit] + _ELLIPSIS + if isinstance(value, dict): + return {str(k): _bound(v, limit) for k, v in value.items()} + if isinstance(value, (list, tuple)): + items = [_bound(v, limit) for v in list(value)[:_MAX_LIST_ITEMS]] + if len(value) > _MAX_LIST_ITEMS: + items.append(f"{_ELLIPSIS} {len(value) - _MAX_LIST_ITEMS} more") + return items + if value is None or isinstance(value, (bool, int, float)): + return value + return _bound(str(value), limit) + + +def make_preview(value: Any, limit: int | None = None) -> Any: + """Return a storable preview of ``value``. + + Secret-keyed fields are redacted, NUL bytes stripped and every string + truncated to ``limit`` (``TRACES_PREVIEW_CHARS`` by default). A structure + whose JSON form is still far larger than ``limit`` collapses to a + truncated JSON string so one huge tool result cannot bloat the row. + + Args: + value: Any JSON-like value (tool arguments, a result, a query string). + limit: Maximum characters per string; defaults to the setting. + + Returns: + A JSON-serializable preview. + """ + limit = limit or settings.TRACES_PREVIEW_CHARS + bounded = strip_null_bytes(_bound(redact_secrets(value), limit)) + if isinstance(bounded, (dict, list)): + try: + encoded = json.dumps(bounded, ensure_ascii=False, default=str) + except (TypeError, ValueError): + encoded = str(bounded) + if len(encoded) > limit * 4: + return encoded[:limit] + _ELLIPSIS + return bounded diff --git a/docsgpt/tracing/sink.py b/docsgpt/tracing/sink.py new file mode 100644 index 00000000..c5f9ad6c --- /dev/null +++ b/docsgpt/tracing/sink.py @@ -0,0 +1,53 @@ +"""Write a finished trace to its two sinks: Postgres and OpenTelemetry.""" + +from __future__ import annotations + +import logging +from typing import Optional + +from docsgpt.tracing.core import Trace + +logger = logging.getLogger(__name__) + + +def flush(trace: Optional[Trace], status: Optional[str] = None) -> None: + """Finish ``trace`` and persist it once; later calls are no-ops. + + OTel replay runs first so the exported trace id can be stored with the + row. Both sinks swallow their own failures: a trace is diagnostic data + and must never fail the request that produced it. + + Args: + trace: The trace to write; ``None`` is accepted and ignored. + status: Final status; defaults to ``error`` when a top-level span + failed, else ``ok``. + """ + if trace is None or trace.flushed: + return + trace.flushed = True + trace.finish(status) + try: + from docsgpt.tracing.otel import export_trace + + trace.otel_trace_id = export_trace(trace) + except Exception: # noqa: BLE001 + logger.warning("Failed to export trace %s to OpenTelemetry", trace.id, exc_info=True) + if not trace.spans: + # Nothing ran worth showing (e.g. an early validation error). + return + try: + from docsgpt.storage.db.repositories.request_traces import RequestTracesRepository + from docsgpt.storage.db.session import db_session + + with db_session() as conn: + RequestTracesRepository(conn).insert(trace.to_record()) + except Exception: # noqa: BLE001 + logger.warning("Failed to store trace %s", trace.id, exc_info=True) + + +def discard(trace: Optional[Trace]) -> None: + """Drop ``trace`` without writing it (the request was rejected or superseded).""" + if trace is None or trace.flushed: + return + trace.flushed = True + trace.finish() diff --git a/tests/tracing/__init__.py b/tests/tracing/__init__.py new file mode 100644 index 00000000..e69de29b diff --git a/tests/tracing/test_trace.py b/tests/tracing/test_trace.py new file mode 100644 index 00000000..685a5595 --- /dev/null +++ b/tests/tracing/test_trace.py @@ -0,0 +1,341 @@ +"""Tests for the execution-trace recorder in ``docsgpt.tracing``.""" + +from __future__ import annotations + +import threading +from concurrent.futures import ThreadPoolExecutor + +import pytest + +from docsgpt import tracing +from docsgpt.core.settings import settings + + +@pytest.fixture(autouse=True) +def _tracing_on(monkeypatch): + monkeypatch.setattr(settings, "TRACES_ENABLED", True) + monkeypatch.setattr(settings, "TRACES_CAPTURE_CONTENT", True) + monkeypatch.setattr(settings, "TRACES_MAX_SPANS", 500) + monkeypatch.setattr(settings, "TRACES_PREVIEW_CHARS", 2000) + + +def _by_name(trace): + return {s.name: s for s in trace.spans} + + +class TestNoActiveTrace: + def test_span_calls_are_noops_without_a_trace(self): + assert tracing.current_trace() is None + with tracing.span(tracing.KIND_LLM, "chat gpt") as s: + s.set(foo=1) + s.preview("output", "hi") + handle = tracing.start_span(tracing.KIND_TOOL, "tool") + handle.end(status="error") + tracing.bind(message_id="m") + + def test_start_trace_returns_none_when_disabled(self, monkeypatch): + monkeypatch.setattr(settings, "TRACES_ENABLED", False) + assert tracing.start_trace(source="stream") is None + with tracing.activate(None): + assert tracing.current_trace() is None + + +class TestNesting: + def test_containers_push_and_leaves_record_parent(self): + trace = tracing.start_trace(source="stream") + with tracing.activate(trace): + with tracing.span(tracing.KIND_AGENT, "agent"): + with tracing.span(tracing.KIND_LLM, "llm-1"): + pass + with tracing.span(tracing.KIND_TOOL, "tool"): + with tracing.span(tracing.KIND_RETRIEVAL, "retrieval"): + with tracing.span(tracing.KIND_EMBEDDING, "embed"): + pass + with tracing.span(tracing.KIND_LLM, "llm-2"): + pass + spans = _by_name(trace) + assert spans["agent"].parent_id is None + assert spans["llm-1"].parent_id == spans["agent"].id + assert spans["tool"].parent_id == spans["agent"].id + assert spans["retrieval"].parent_id == spans["tool"].id + assert spans["embed"].parent_id == spans["retrieval"].id + assert spans["llm-2"].parent_id == spans["agent"].id + assert all(s.status == "ok" for s in trace.spans) + + def test_leaf_span_does_not_become_a_parent(self): + trace = tracing.start_trace(source="stream") + with tracing.activate(trace): + llm = tracing.start_span(tracing.KIND_LLM, "llm") + with tracing.span(tracing.KIND_TOOL, "tool"): + pass + llm.end() + spans = _by_name(trace) + assert spans["tool"].parent_id is None + + def test_explicit_parent_wins(self): + trace = tracing.start_trace(source="stream") + with tracing.activate(trace): + outer = tracing.start_span(tracing.KIND_RETRIEVAL, "outer") + with tracing.span(tracing.KIND_AGENT, "agent"): + child = tracing.start_span(tracing.KIND_SEARCH, "child", parent=outer) + child.end() + outer.end() + assert _by_name(trace)["child"].parent_id == _by_name(trace)["outer"].id + + def test_interleaved_generators_nest_by_start_order(self): + """An agent generator suspended at yield keeps its children nested.""" + trace = tracing.start_trace(source="stream") + + def llm_stream(): + handle = tracing.start_span(tracing.KIND_LLM, "llm") + try: + yield "a" + yield "b" + finally: + handle.end() + + def agent_gen(): + handle = tracing.start_span(tracing.KIND_AGENT, "agent") + try: + yield from llm_stream() + with tracing.span(tracing.KIND_TOOL, "tool"): + pass + yield "c" + finally: + handle.end() + + with tracing.activate(trace): + assert list(agent_gen()) == ["a", "b", "c"] + spans = _by_name(trace) + assert spans["llm"].parent_id == spans["agent"].id + assert spans["tool"].parent_id == spans["agent"].id + + def test_out_of_order_end_cancels_abandoned_children(self): + trace = tracing.start_trace(source="stream") + with tracing.activate(trace): + agent = tracing.start_span(tracing.KIND_AGENT, "agent") + tool = tracing.start_span(tracing.KIND_TOOL, "tool") + agent.end() + after = tracing.start_span(tracing.KIND_LLM, "after") + after.end() + tool.end() # late end is ignored + spans = _by_name(trace) + assert spans["tool"].status == "cancelled" + assert spans["agent"].status == "ok" + assert spans["after"].parent_id is None + + def test_exception_marks_span_error_and_propagates(self): + trace = tracing.start_trace(source="stream") + with tracing.activate(trace): + with pytest.raises(ValueError): + with tracing.span(tracing.KIND_TOOL, "tool"): + raise ValueError("boom") + span = trace.spans[0] + assert span.status == "error" + assert span.attributes["error.type"] == "ValueError" + assert span.error == "boom" + + def test_generator_exit_marks_span_cancelled(self): + trace = tracing.start_trace(source="stream") + + def gen(): + with tracing.span(tracing.KIND_AGENT, "agent"): + yield 1 + yield 2 + + with tracing.activate(trace): + g = gen() + next(g) + g.close() + assert trace.spans[0].status == "cancelled" + + +class TestFinish: + def test_finish_cancels_open_spans_and_ignores_late_ends(self): + trace = tracing.start_trace(source="stream") + with tracing.activate(trace): + open_span = tracing.start_span(tracing.KIND_AGENT, "agent") + trace.finish() + assert open_span.status == "cancelled" + duration = open_span.duration_ms + open_span.end(status="ok") + assert open_span.status == "cancelled" + assert open_span.duration_ms == duration + + def test_finish_status_defaults_to_error_when_a_top_level_span_failed(self): + trace = tracing.start_trace(source="stream") + with tracing.activate(trace): + tracing.start_span(tracing.KIND_AGENT, "agent").end(status="error") + trace.finish() + assert trace.status == "error" + + def test_explicit_finish_status(self): + trace = tracing.start_trace(source="stream") + trace.finish(status="paused") + assert trace.status == "paused" + + def test_spans_after_finish_are_dropped(self): + trace = tracing.start_trace(source="stream") + trace.finish() + with tracing.activate(trace): + tracing.start_span(tracing.KIND_LLM, "late").end() + assert trace.spans == [] + + def test_summary_rolls_up_llm_tool_retrieval(self): + trace = tracing.start_trace(source="stream") + with tracing.activate(trace): + with tracing.span(tracing.KIND_AGENT, "agent"): + with tracing.span( + tracing.KIND_LLM, + "chat", + attributes={ + "gen_ai.usage.input_tokens": 10, + "gen_ai.usage.output_tokens": 4, + }, + ): + pass + with tracing.span(tracing.KIND_RETRIEVAL, "r"): + with tracing.span(tracing.KIND_RETRIEVAL, "inner"): + pass + with tracing.span(tracing.KIND_TOOL, "t") as t: + t.end(status="error") + trace.finish() + summary = trace.summary() + assert summary["llm_calls"] == 1 + assert summary["tool_calls"] == 1 + assert summary["input_tokens"] == 10 + assert summary["output_tokens"] == 4 + assert summary["retrieval_calls"] == 1 + assert summary["errors"] == 1 + assert summary["retrieval_ms"] >= 0 + + +class TestCap: + def test_span_cap_counts_dropped(self, monkeypatch): + monkeypatch.setattr(settings, "TRACES_MAX_SPANS", 10) + trace = tracing.start_trace(source="graph_extraction") + with tracing.activate(trace): + for i in range(15): + with tracing.span(tracing.KIND_LLM, f"llm-{i}"): + pass + assert len(trace.spans) == 10 + assert trace.dropped_spans == 5 + + +class TestThreads: + def test_wrap_carries_trace_and_parent_into_pool_threads(self): + trace = tracing.start_trace(source="stream") + + def work(i): + with tracing.span(tracing.KIND_SEARCH, f"search-{i}"): + with tracing.span(tracing.KIND_EMBEDDING, f"embed-{i}"): + pass + return threading.get_ident() + + with tracing.activate(trace): + with tracing.span(tracing.KIND_RETRIEVAL, "retrieval"): + with ThreadPoolExecutor(max_workers=3) as pool: + list(pool.map(tracing.wrap(work), range(3))) + spans = _by_name(trace) + for i in range(3): + assert spans[f"search-{i}"].parent_id == spans["retrieval"].id + assert spans[f"embed-{i}"].parent_id == spans["retrieval"].id + + def test_pool_threads_without_wrap_see_no_trace(self): + trace = tracing.start_trace(source="stream") + seen = [] + with tracing.activate(trace): + t = threading.Thread(target=lambda: seen.append(tracing.current_trace())) + t.start() + t.join() + assert seen == [None] + + def test_wrap_without_trace_is_passthrough(self): + fn = tracing.wrap(lambda x: x + 1) + assert fn(1) == 2 + + +class TestBind: + def test_bind_sets_ids_on_active_trace(self): + trace = tracing.start_trace(source="stream", request_id="r1") + with tracing.activate(trace): + tracing.bind(message_id="m1", conversation_id="c1", unknown="x") + assert trace.request_id == "r1" + assert trace.message_id == "m1" + assert trace.conversation_id == "c1" + + def test_bind_if_unset_keeps_first_value(self): + trace = tracing.start_trace(source="stream") + with tracing.activate(trace): + tracing.bind_if_unset(activity_id="a1") + tracing.bind_if_unset(activity_id="a2") + assert trace.activity_id == "a1" + + +class TestPreviews: + def test_preview_redacts_secrets_and_bounds_strings(self, monkeypatch): + monkeypatch.setattr(settings, "TRACES_PREVIEW_CHARS", 100) + trace = tracing.start_trace(source="stream") + with tracing.activate(trace): + with tracing.span(tracing.KIND_TOOL, "tool") as s: + s.preview("arguments", {"api_key": "sk-123", "q": "x" * 500}) + s.preview("result", "y" * 500 + "\x00") + preview = trace.spans[0].previews + assert preview["arguments"]["api_key"] == "[REDACTED]" + assert len(preview["arguments"]["q"]) < 200 + assert preview["result"].endswith("…") + assert "\x00" not in preview["result"] + + def test_preview_skipped_when_capture_disabled(self, monkeypatch): + monkeypatch.setattr(settings, "TRACES_CAPTURE_CONTENT", False) + trace = tracing.start_trace(source="stream") + with tracing.activate(trace): + with tracing.span(tracing.KIND_TOOL, "tool") as s: + s.preview("result", "secret stuff") + assert trace.spans[0].previews == {} + + def test_guardrail_trigger_strips_all_previews(self): + trace = tracing.start_trace(source="stream") + with tracing.activate(trace): + with tracing.span(tracing.KIND_TOOL, "tool") as s: + s.preview("result", "leaked") + tracing.mark_content_blocked() + trace.finish() + record = trace.to_record() + assert all(not s.get("preview") for s in record["spans"]) + assert trace.content_blocked is True + + def test_huge_structures_collapse_to_a_string(self, monkeypatch): + monkeypatch.setattr(settings, "TRACES_PREVIEW_CHARS", 100) + trace = tracing.start_trace(source="stream") + with tracing.activate(trace): + with tracing.span(tracing.KIND_TOOL, "tool") as s: + s.preview("result", [{"k": "v" * 50} for _ in range(100)]) + value = trace.spans[0].previews["result"] + assert isinstance(value, str) + assert len(value) <= 101 + + +class TestRecord: + def test_to_record_shape(self): + trace = tracing.start_trace( + source="stream", request_id="r", user_id="u", agent_id="a" + ) + with tracing.activate(trace): + with tracing.span( + tracing.KIND_LLM, "chat m", attributes={"gen_ai.request.model": "m"} + ): + pass + trace.finish() + record = trace.to_record() + assert record["source"] == "stream" + assert record["request_id"] == "r" + assert record["status"] == "ok" + assert record["span_count"] == 1 + span = record["spans"][0] + assert set(span) >= { + "id", "parent_id", "kind", "name", "status", "offset_ms", + "duration_ms", "attributes", + } + assert span["attributes"]["gen_ai.request.model"] == "m" + assert span["offset_ms"] >= 0