mirror of
https://github.com/tiennm99/DocsGPT.git
synced 2026-10-11 12:11:45 +00:00
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:
1 parent
b5e9257659
commit
e8166e6c29
4 files changed
+385
-4
No files matched your search
@@ -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
|
||||
|
||||
@@ -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
@@ -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:
|
||||
|
||||
@@ -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
|
||||
Reference in new issue
Block a user