Record a chat span for every LLM call

The token-usage wrappers open one span per decorated invocation, so a
primary attempt, its retry and a fallback appear as siblings with the
provider that ran. Stream spans start on the first pull, carry tokens,
cost, time to first token and cache hits, and feed the GenAI metrics.
This commit is contained in:
arc53-machine committed 2026-09-23 17:23:06 +01:00
1 parent b5e9257659
commit e8166e6c29
4 files changed
+385 -4

No files matched your search

+8
View File
@@ -8,6 +8,7 @@ from threading import Lock
import redis
from docsgpt.core.settings import settings
from docsgpt.tracing.llm import CACHE_HIT_ATTR, record_cached_gen
from docsgpt.utils import get_hash
logger = logging.getLogger(__name__)
@@ -240,6 +241,7 @@ def gen_cache(func):
if cached_response:
decoded = cached_response.decode("utf-8")
if not _is_stream_payload(decoded):
record_cached_gen(self, model, decoded)
return decoded
except Exception as e:
logger.error(f"Error getting cached response: {e}", exc_info=True)
@@ -295,6 +297,12 @@ def stream_cache(func):
if cached_chunks is not None:
logger.info(f"Cache hit for stream key: {cache_key}")
# ``stream_token_usage`` wraps this cache and owns
# the call's span; flag it as served from cache.
try:
setattr(self, CACHE_HIT_ATTR, True)
except AttributeError:
pass
for chunk in cached_chunks:
yield chunk
time.sleep(0.03) # Simulate streaming delay
+119
View File
@@ -0,0 +1,119 @@
"""``chat`` spans and GenAI metrics for LLM calls.
Called from the token-usage wrappers in ``docsgpt/usage.py``: one span per
decorated invocation, so a primary attempt, its same-provider retry and a
fallback each get their own span with the provider that actually ran.
"""
from __future__ import annotations
from typing import Any, Dict, Iterable, Optional
from docsgpt.tracing import core
from docsgpt.tracing.otel import provider_name, record_llm_metrics
#: LLM attribute a cache wrapper sets when it served the call from Redis.
CACHE_HIT_ATTR = "_trace_cache_hit"
def start_llm_span(llm: Any, model: Optional[str], *, stream: bool, tools: Any = None):
"""Open a ``chat {model}`` span for a call on ``llm`` (no-op without a trace)."""
if core.current_trace() is None:
return core.NOOP_SPAN
attributes = {
"gen_ai.operation.name": "chat",
"gen_ai.provider.name": provider_name(getattr(llm, "provider_name", None)),
"gen_ai.request.model": str(model) if model else None,
"docsgpt.token_source": getattr(llm, "_token_usage_source", None) or "agent_stream",
"docsgpt.stream": bool(stream),
"docsgpt.tool_count": len(tools) if tools else None,
}
return core.start_span(
core.KIND_LLM,
f"chat {model}" if model else "chat",
attributes={k: v for k, v in attributes.items() if v is not None},
)
def output_text(chunks: Iterable[Any]) -> str:
"""Concatenate the text deltas of a streamed or returned response."""
if isinstance(chunks, str):
return chunks
return "".join(chunk for chunk in chunks if isinstance(chunk, str))
def finish_llm_call(
span: Any,
llm: Any,
model: Optional[str],
call_usage: Dict[str, Any],
*,
duration_ms: int,
error: Optional[BaseException],
completed: bool = True,
ttft_ms: Optional[int] = None,
cost_usd: Optional[float] = None,
estimated: bool = True,
output: Optional[str] = None,
) -> None:
"""Close the call's span and record the GenAI client metrics.
Args:
span: The span from :func:`start_llm_span`.
llm: The LLM instance that ran the call.
model: Model the call was made with.
call_usage: Final token counts (provider-reported when available).
duration_ms: Provider time for the call.
error: The exception the call raised, if any.
completed: False when a stream was abandoned before it finished.
ttft_ms: Time to first streamed chunk.
cost_usd: Cost priced for the call, when known.
estimated: True when token counts are local estimates.
output: Response text for the preview.
"""
cache_hit = bool(getattr(llm, CACHE_HIT_ATTR, False))
try:
setattr(llm, CACHE_HIT_ATTR, False)
except AttributeError:
pass
record_llm_metrics(
provider=getattr(llm, "provider_name", None),
model=str(model) if model else None,
input_tokens=call_usage.get("prompt_tokens", 0),
output_tokens=call_usage.get("generated_tokens", 0),
duration_s=max(duration_ms, 0) / 1000.0,
error_type=type(error).__name__ if error is not None else None,
)
if not span:
return
if output:
span.preview("output", output)
span.end(
None if completed or error is not None else core.STATUS_CANCELLED,
error=error,
attributes={
"gen_ai.usage.input_tokens": int(call_usage.get("prompt_tokens") or 0),
"gen_ai.usage.output_tokens": int(call_usage.get("generated_tokens") or 0),
"gen_ai.usage.cache_read.input_tokens": call_usage.get("cached_tokens"),
"gen_ai.usage.cache_creation.input_tokens": call_usage.get("cache_write_tokens"),
"docsgpt.usage_estimated": estimated,
"docsgpt.provider_ms": duration_ms,
"docsgpt.ttft_ms": ttft_ms,
"docsgpt.cost_usd": cost_usd if isinstance(cost_usd, (int, float)) else None,
"docsgpt.cache_hit": True if cache_hit else None,
},
)
def record_cached_gen(llm: Any, model: Optional[str], output: Optional[str]) -> None:
"""Record a non-streaming call answered from the response cache.
The gen cache wraps the usage wrapper, so a hit never reaches it; this
records the zero-cost call so the trace still shows it happened.
"""
span = start_llm_span(llm, model, stream=False)
if not span:
return
if output:
span.preview("output", output)
span.end(attributes={"docsgpt.cache_hit": True})
+42 -4
View File
@@ -3,6 +3,7 @@ import time
from typing import Any, Dict
from docsgpt.pricing import compute_cost_usd
from docsgpt.tracing.llm import finish_llm_call, output_text, start_llm_span
from docsgpt.storage.db.repositories.token_usage import TokenUsageRepository
from docsgpt.storage.db.session import db_session
from docsgpt.utils import num_tokens_from_object_or_list, num_tokens_from_string
@@ -108,9 +109,12 @@ def _persist_call_usage(llm, call_usage, *, duration_ms=None, ttft_ms=None):
duration_ms: Wall-clock for the call, measured by the wrapper.
ttft_ms: Time to the first streamed chunk; None for a non-streaming
call and for a stream that failed before yielding anything.
Returns:
The call's priced cost in USD, or None when no row was written.
"""
if call_usage["prompt_tokens"] == 0 and call_usage["generated_tokens"] == 0:
return
return None
decoded_token = getattr(llm, "decoded_token", None)
user_id = (
decoded_token.get("sub") if isinstance(decoded_token, dict) else None
@@ -126,7 +130,7 @@ def _persist_call_usage(llm, call_usage, *, duration_ms=None, ttft_ms=None):
"source": getattr(llm, "_token_usage_source", "agent_stream"),
},
)
return
return None
model_id = getattr(llm, "_canonical_model_id", None)
# Bring-your-own models run on the user's own provider key: recorded, never priced.
if getattr(llm, "_is_byom", False):
@@ -161,6 +165,7 @@ def _persist_call_usage(llm, call_usage, *, duration_ms=None, ttft_ms=None):
)
except Exception:
logger.exception("token_usage persist failed")
return cost
def _call_cost_usd(model_id, call_usage) -> float:
@@ -257,8 +262,10 @@ def gen_token_usage(func):
usage_attachments=usage_attachments,
**kwargs,
)
span = start_llm_span(self, model, stream=False, tools=tools)
started_at = time.monotonic()
error: BaseException | None = None
result = None
try:
result = func(self, model, messages, stream, tools, **kwargs)
call_usage["generated_tokens"] += _count_tokens(result)
@@ -268,11 +275,23 @@ def gen_token_usage(func):
raise
finally:
duration_ms = int((time.monotonic() - started_at) * 1000)
estimated_usage = call_usage
call_usage = _prefer_provider_usage(self, call_usage)
self.token_usage["prompt_tokens"] += call_usage["prompt_tokens"]
self.token_usage["generated_tokens"] += call_usage["generated_tokens"]
# A non-streaming call has no first-token moment; ttft stays NULL.
_persist_call_usage(self, call_usage, duration_ms=duration_ms)
cost = _persist_call_usage(self, call_usage, duration_ms=duration_ms)
finish_llm_call(
span,
self,
model,
call_usage,
duration_ms=duration_ms,
error=error,
cost_usd=cost,
estimated=call_usage is estimated_usage,
output=result if isinstance(result, str) else None,
)
emit = getattr(self, "_emit_gen_finished_log", None)
if callable(emit):
try:
@@ -314,6 +333,10 @@ def stream_token_usage(func):
# non-streaming durations under one p50.
provider_seconds = 0.0
error: BaseException | None = None
completed = False
# This body runs on the first ``next()``, not at ``gen_stream()``
# time, so the span starts when the provider call really does.
span = start_llm_span(self, model, stream=True, tools=tools)
try:
result = func(self, model, messages, stream, tools, **kwargs)
stream_iter = iter(result)
@@ -323,6 +346,7 @@ def stream_token_usage(func):
r = next(stream_iter)
except StopIteration:
provider_seconds += time.monotonic() - pull_started
completed = True
break
provider_seconds += time.monotonic() - pull_started
if first_chunk_at is None:
@@ -346,12 +370,26 @@ def stream_token_usage(func):
)
for line in batch:
call_usage["generated_tokens"] += _count_tokens(line)
estimated_usage = call_usage
call_usage = _prefer_provider_usage(self, call_usage)
self.token_usage["prompt_tokens"] += call_usage["prompt_tokens"]
self.token_usage["generated_tokens"] += call_usage["generated_tokens"]
_persist_call_usage(
cost = _persist_call_usage(
self, call_usage, duration_ms=duration_ms, ttft_ms=ttft_ms
)
finish_llm_call(
span,
self,
model,
call_usage,
duration_ms=duration_ms,
error=error,
completed=completed,
ttft_ms=ttft_ms,
cost_usd=cost,
estimated=call_usage is estimated_usage,
output=output_text(batch),
)
emit = getattr(self, "_emit_stream_finished_log", None)
if callable(emit):
try:
+216
View File
@@ -0,0 +1,216 @@
"""LLM calls become ``chat`` spans through the token-usage wrappers."""
from __future__ import annotations
from unittest.mock import patch
import pytest
from docsgpt import tracing
from docsgpt.cache import gen_cache, stream_cache
from docsgpt.core.settings import settings
from docsgpt.usage import gen_token_usage, stream_token_usage
class _LLM:
provider_name = "openai"
def __init__(self, source=None):
self.token_usage = {"prompt_tokens": 0, "generated_tokens": 0}
self.decoded_token = {"sub": "u1"}
self.user_api_key = None
self.agent_id = None
if source:
self._token_usage_source = source
@pytest.fixture(autouse=True)
def _env(monkeypatch):
monkeypatch.setattr(settings, "TRACES_ENABLED", True)
monkeypatch.setattr(settings, "TRACES_CAPTURE_CONTENT", True)
with patch("docsgpt.usage._persist_call_usage", return_value=0.0012):
yield
@pytest.fixture()
def trace():
t = tracing.start_trace(source="stream", capture_otel_context=False)
with tracing.activate(t):
yield t
@pytest.fixture()
def metrics():
with patch("docsgpt.tracing.llm.record_llm_metrics") as rec:
yield rec
class TestNonStreaming:
def test_span_with_usage_and_preview(self, trace, metrics):
@gen_token_usage
def _gen(self, model, messages, stream, tools, **kwargs):
return "the answer"
_gen(_LLM(), "gpt-4o", [{"role": "user", "content": "hi"}], False, None)
(span,) = trace.spans
assert span.kind == tracing.KIND_LLM
assert span.name == "chat gpt-4o"
assert span.status == "ok"
attrs = span.attributes
assert attrs["gen_ai.operation.name"] == "chat"
assert attrs["gen_ai.provider.name"] == "openai"
assert attrs["gen_ai.request.model"] == "gpt-4o"
assert attrs["gen_ai.usage.input_tokens"] > 0
assert attrs["gen_ai.usage.output_tokens"] > 0
assert attrs["docsgpt.token_source"] == "agent_stream"
assert attrs["docsgpt.cost_usd"] == 0.0012
assert span.previews["output"] == "the answer"
metrics.assert_called_once()
assert metrics.call_args.kwargs["error_type"] is None
def test_failure_marks_span_error(self, trace, metrics):
@gen_token_usage
def _gen(self, model, messages, stream, tools, **kwargs):
raise TimeoutError("slow")
with pytest.raises(TimeoutError):
_gen(_LLM(source="fallback"), "m", [], False, None)
(span,) = trace.spans
assert span.status == "error"
assert span.attributes["error.type"] == "TimeoutError"
assert span.attributes["docsgpt.token_source"] == "fallback"
assert metrics.call_args.kwargs["error_type"] == "TimeoutError"
def test_provider_reported_usage_is_flagged(self, trace, metrics):
llm = _LLM()
@gen_token_usage
def _gen(self, model, messages, stream, tools, **kwargs):
self._last_usage = {
"prompt_tokens": 100,
"completion_tokens": 7,
"prompt_tokens_details": {"cached_tokens": 40},
}
self._last_usage_claimed = False
return "x"
_gen(llm, "m", [], False, None)
attrs = trace.spans[0].attributes
assert attrs["gen_ai.usage.input_tokens"] == 100
assert attrs["gen_ai.usage.output_tokens"] == 7
assert attrs["gen_ai.usage.cache_read.input_tokens"] == 40
assert attrs["docsgpt.usage_estimated"] is False
def test_no_trace_no_span(self, metrics):
@gen_token_usage
def _gen(self, model, messages, stream, tools, **kwargs):
return "x"
assert _gen(_LLM(), "m", [], False, None) == "x"
metrics.assert_called_once()
class TestStreaming:
def test_span_starts_on_first_next_not_on_call(self, trace, metrics):
@stream_token_usage
def _stream(self, model, messages, stream, tools, **kwargs):
yield "a"
yield {"type": "thought", "thought": "hmm"}
yield "b"
gen = _stream(_LLM(), "m", [], True, None)
assert trace.spans == []
assert list(gen) == ["a", {"type": "thought", "thought": "hmm"}, "b"]
(span,) = trace.spans
assert span.status == "ok"
assert span.attributes["docsgpt.ttft_ms"] is not None
assert span.attributes["docsgpt.stream"] is True
assert span.previews["output"] == "ab"
def test_abandoned_stream_is_cancelled(self, trace, metrics):
@stream_token_usage
def _stream(self, model, messages, stream, tools, **kwargs):
yield "a"
yield "b"
gen = _stream(_LLM(), "m", [], True, None)
next(gen)
gen.close()
assert trace.spans[0].status == "cancelled"
def test_failed_stream(self, trace, metrics):
@stream_token_usage
def _stream(self, model, messages, stream, tools, **kwargs):
yield "a"
raise ConnectionError("reset")
with pytest.raises(ConnectionError):
list(_stream(_LLM(), "m", [], True, None))
assert trace.spans[0].status == "error"
def test_primary_and_fallback_are_siblings(self, trace, metrics):
@stream_token_usage
def _primary(self, model, messages, stream, tools, **kwargs):
raise ConnectionError("down")
yield # pragma: no cover
@stream_token_usage
def _fallback(self, model, messages, stream, tools, **kwargs):
yield "ok"
with tracing.span(tracing.KIND_AGENT, "agent"):
with pytest.raises(ConnectionError):
list(_primary(_LLM(), "m", [], True, None))
list(_fallback(_LLM(source="fallback"), "m2", [], True, None))
agent, primary, fallback = trace.spans
assert primary.parent_id == agent.id == fallback.parent_id
assert primary.status == "error"
assert fallback.attributes["docsgpt.token_source"] == "fallback"
class _FakeRedis:
def __init__(self):
self.store = {}
def get(self, key):
return self.store.get(key)
def set(self, key, value, ex=None):
self.store[key] = value.encode("utf-8") if isinstance(value, str) else value
def delete(self, key):
self.store.pop(key, None)
class TestCacheHits:
def test_gen_cache_hit_records_a_cached_span(self, trace, metrics):
redis = _FakeRedis()
@gen_cache
@gen_token_usage
def _gen(self, model, messages, stream, tools=None, **kwargs):
return "fresh"
with patch("docsgpt.cache.get_redis_instance", return_value=redis):
_gen(_LLM(), "m", [{"role": "user", "content": "q"}], False)
_gen(_LLM(), "m", [{"role": "user", "content": "q"}], False)
first, second = trace.spans
assert first.attributes.get("docsgpt.cache_hit") is None
assert second.attributes["docsgpt.cache_hit"] is True
assert second.status == "ok"
def test_stream_cache_hit_flags_the_open_span(self, trace, metrics, monkeypatch):
monkeypatch.setattr("docsgpt.cache.time.sleep", lambda _s: None)
redis = _FakeRedis()
@stream_token_usage
@stream_cache
def _stream(self, model, messages, stream, tools=None, **kwargs):
yield "fresh"
with patch("docsgpt.cache.get_redis_instance", return_value=redis):
list(_stream(_LLM(), "m", [{"role": "user", "content": "q"}], True, None))
list(_stream(_LLM(), "m", [{"role": "user", "content": "q"}], True, None))
first, second = trace.spans
assert first.attributes.get("docsgpt.cache_hit") is None
assert second.attributes["docsgpt.cache_hit"] is True