mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-07 02:59:05 +00:00
feat(tracing): add ClickHouse trace store with batched writes and team-scoped reads
This commit is contained in:
parent
bab0811ba1
commit
aff5c28a33
1 changed files with 343 additions and 0 deletions
343
litellm/tracing/store.py
Normal file
343
litellm/tracing/store.py
Normal file
|
|
@ -0,0 +1,343 @@
|
|||
"""
|
||||
ClickHouse-backed trace store: batched span writes + scoped reads.
|
||||
|
||||
Reads join agent spans (otel_traces) to LiteLLM requests (spend_logs) on
|
||||
`otel_traces.LiteLLMRequestId = spend_logs.response_id`.
|
||||
"""
|
||||
|
||||
import base64
|
||||
import json
|
||||
from datetime import UTC, datetime
|
||||
from typing import Any, Final
|
||||
|
||||
from litellm.constants import AGENT_TRACING_LIST_PAGE_SIZE
|
||||
from litellm.integrations.clickhouse.clickhouse_batch_logger import ClickHouseBatchLogger
|
||||
from litellm.integrations.clickhouse.clickhouse_client import ClickHouseClient
|
||||
from litellm.integrations.clickhouse.schema import (
|
||||
AGENT_TRACES_TABLE,
|
||||
OTEL_TRACES_TABLE,
|
||||
SPEND_LOGS_TABLE,
|
||||
)
|
||||
from litellm.tracing.types import (
|
||||
AgentNode,
|
||||
LiteLLMRequest,
|
||||
Span,
|
||||
SpanDetail,
|
||||
SpanRow,
|
||||
SpanStatus,
|
||||
Trace,
|
||||
TracePage,
|
||||
TraceScope,
|
||||
TraceSummary,
|
||||
)
|
||||
|
||||
NANOS_PER_MS: Final = 1_000_000
|
||||
_STATUS: Final[dict[str, SpanStatus]] = {"STATUS_CODE_OK": "ok", "STATUS_CODE_ERROR": "error"}
|
||||
|
||||
_SCOPE_OTEL: Final = (
|
||||
"(empty({team_ids:Array(String)}) OR TeamId IN {team_ids:Array(String)})"
|
||||
" AND ({api_key_hash:String} = '' OR ApiKeyHash = {api_key_hash:String})"
|
||||
)
|
||||
_SCOPE_SPEND: Final = (
|
||||
"(empty({team_ids:Array(String)}) OR team_id IN {team_ids:Array(String)})"
|
||||
" AND ({api_key_hash:String} = '' OR api_key = {api_key_hash:String})"
|
||||
)
|
||||
|
||||
# Page of traces from the per-trace MV, then cost from spend logs via ARRAY JOIN on request ids.
|
||||
# The MV writes one partial row per insert, so root fields come from the partial that saw the root span.
|
||||
# agent_traces has no ApiKeyHash, so key-scoped (team-less) reads filter trace ids through otel_traces.
|
||||
LIST_TRACES_SQL: Final = f"""
|
||||
SELECT t.TraceId AS trace_id, any(t.RootName) AS name, any(t.ServiceName) AS service,
|
||||
any(t.RootInput) AS input_preview, any(t.RootStatus) AS status,
|
||||
toUnixTimestamp64Milli(any(t.StartTs)) AS start_ms,
|
||||
dateDiff('millisecond', any(t.StartTs), any(t.EndTs)) AS duration_ms,
|
||||
any(t.SpanCount) AS span_count, length(any(t.AgentNames)) AS agent_count,
|
||||
any(t.AgentCount) AS agent_invocations,
|
||||
any(t.LlmCount) AS llm_calls, any(t.ToolCount) AS tool_calls,
|
||||
any(t.InputTokens) AS input_tokens, any(t.OutputTokens) AS output_tokens,
|
||||
any(t.Models) AS models, any(t.ErrorCount) AS error_count, sum(s.spend) AS spend
|
||||
FROM (
|
||||
SELECT TeamId, TraceId, min(StartTs) AS StartTs, max(EndTs) AS EndTs,
|
||||
any(ServiceName) AS ServiceName, anyLastIf(a.RootName, a.RootName != '') AS RootName,
|
||||
anyLastIf(a.RootInput, a.RootName != '') AS RootInput,
|
||||
anyLastIf(a.RootStatus, a.RootName != '') AS RootStatus,
|
||||
sum(SpanCount) AS SpanCount, sum(AgentCount) AS AgentCount, sum(LlmCount) AS LlmCount,
|
||||
sum(ToolCount) AS ToolCount, sum(ErrorCount) AS ErrorCount, sum(InputTokens) AS InputTokens,
|
||||
sum(OutputTokens) AS OutputTokens, groupUniqArrayArray(Models) AS Models,
|
||||
groupUniqArrayArray(AgentNames) AS AgentNames, groupArrayArray(RequestIds) AS RequestIds
|
||||
FROM {AGENT_TRACES_TABLE} AS a
|
||||
WHERE (empty({{team_ids:Array(String)}}) OR TeamId IN {{team_ids:Array(String)}})
|
||||
AND ({{api_key_hash:String}} = '' OR TraceId IN (
|
||||
SELECT TraceId FROM {OTEL_TRACES_TABLE} WHERE {_SCOPE_OTEL}))
|
||||
GROUP BY TeamId, TraceId
|
||||
HAVING StartTs >= fromUnixTimestamp64Milli({{start_ms:Int64}})
|
||||
AND StartTs < fromUnixTimestamp64Milli({{end_ms:Int64}})
|
||||
AND (({{cursor_ms:Int64}} = 0) OR (toUnixTimestamp64Milli(StartTs), TraceId)
|
||||
< ({{cursor_ms:Int64}}, {{cursor_trace_id:String}}))
|
||||
ORDER BY StartTs DESC, TraceId DESC
|
||||
LIMIT {{limit:UInt32}}
|
||||
) AS t
|
||||
LEFT ARRAY JOIN t.RequestIds AS request_id
|
||||
LEFT JOIN (
|
||||
SELECT response_id, any(spend) AS spend FROM {SPEND_LOGS_TABLE} FINAL WHERE {_SCOPE_SPEND} GROUP BY response_id
|
||||
) AS s ON s.response_id = request_id
|
||||
GROUP BY t.TraceId
|
||||
ORDER BY start_ms DESC, t.TraceId DESC
|
||||
"""
|
||||
|
||||
# response_id is not unique (an upstream cache can replay the same provider id), so each span keeps
|
||||
# the spend-log row closest to it in time; LIMIT 1 BY also drops re-exported duplicate spans.
|
||||
TRACE_SPANS_SQL: Final = f"""
|
||||
SELECT o.SpanId AS span_id, o.ParentSpanId AS parent_span_id, o.SpanName AS name,
|
||||
o.ObservationType AS type, o.AgentName AS agent, o.StatusCode AS status,
|
||||
o.StatusMessage AS status_message,
|
||||
toUnixTimestamp64Nano(o.Timestamp) AS start_ns, o.Duration AS duration_ns,
|
||||
o.ServiceName AS service, o.InputPreview AS input_preview, o.Model AS model,
|
||||
o.InputTokens AS input_tokens, o.OutputTokens AS output_tokens,
|
||||
o.LiteLLMRequestId AS litellm_request_id,
|
||||
s.request_id AS s_request_id, s.model AS s_model, s.model_group AS s_model_group,
|
||||
s.custom_llm_provider AS s_provider, s.api_base AS s_api_base, s.key_alias AS s_key_alias,
|
||||
s.team_alias AS s_team_alias, s.spend AS s_spend, s.prompt_tokens AS s_prompt_tokens,
|
||||
s.completion_tokens AS s_completion_tokens, s.cache_read_tokens AS s_cache_read_tokens,
|
||||
s.cache_write_tokens AS s_cache_write_tokens,
|
||||
toUnixTimestamp64Milli(s.start_time) AS s_start_ms, toUnixTimestamp64Milli(s.end_time) AS s_end_ms,
|
||||
if(isNull(s.completion_start_time), 0, toUnixTimestamp64Milli(assumeNotNull(s.completion_start_time)))
|
||||
AS s_ttft_start_ms,
|
||||
s.status AS s_status
|
||||
FROM {OTEL_TRACES_TABLE} AS o
|
||||
LEFT JOIN (
|
||||
SELECT * FROM {SPEND_LOGS_TABLE} FINAL
|
||||
WHERE response_id IN (
|
||||
SELECT LiteLLMRequestId FROM {OTEL_TRACES_TABLE}
|
||||
WHERE TraceId = {{trace_id:String}} AND LiteLLMRequestId != '' AND {_SCOPE_OTEL})
|
||||
AND {_SCOPE_SPEND}
|
||||
) AS s ON s.response_id = o.LiteLLMRequestId
|
||||
WHERE o.TraceId = {{trace_id:String}} AND {_SCOPE_OTEL}
|
||||
ORDER BY o.Timestamp, abs(toUnixTimestamp64Milli(s.start_time) - toUnixTimestamp64Milli(o.Timestamp))
|
||||
LIMIT 1 BY o.SpanId
|
||||
"""
|
||||
|
||||
SPAN_DETAIL_SQL: Final = f"""
|
||||
SELECT SpanId AS span_id, Input AS input, Output AS output, SpanAttributes AS attributes
|
||||
FROM {OTEL_TRACES_TABLE}
|
||||
WHERE TraceId = {{trace_id:String}} AND SpanId = {{span_id:String}} AND {_SCOPE_OTEL}
|
||||
LIMIT 1
|
||||
"""
|
||||
|
||||
|
||||
def encode_cursor(start_ms: int, trace_id: str) -> str:
|
||||
return base64.urlsafe_b64encode(json.dumps([start_ms, trace_id]).encode()).decode()
|
||||
|
||||
|
||||
def decode_cursor(cursor: str | None) -> tuple[int, str]:
|
||||
if not cursor:
|
||||
return 0, ""
|
||||
start_ms, trace_id = json.loads(base64.urlsafe_b64decode(cursor.encode()))
|
||||
return int(start_ms), str(trace_id)
|
||||
|
||||
|
||||
def _iso(ms: int) -> str:
|
||||
return datetime.fromtimestamp(ms / 1000, tz=UTC).isoformat()
|
||||
|
||||
|
||||
def _status(code: str) -> SpanStatus:
|
||||
return _STATUS.get(code, "unset")
|
||||
|
||||
|
||||
def trace_summary_from_row(row: dict[str, Any]) -> TraceSummary:
|
||||
return TraceSummary(
|
||||
trace_id=row["trace_id"],
|
||||
name=row["name"],
|
||||
service=row["service"],
|
||||
input_preview=row["input_preview"],
|
||||
start_time=_iso(int(row["start_ms"])),
|
||||
duration_ms=float(row["duration_ms"]),
|
||||
status=_status(row["status"]),
|
||||
span_count=int(row["span_count"]),
|
||||
agent_count=int(row["agent_count"]),
|
||||
agent_invocations=int(row.get("agent_invocations") or row["agent_count"]),
|
||||
llm_calls=int(row["llm_calls"]),
|
||||
tool_calls=int(row["tool_calls"]),
|
||||
error_count=int(row.get("error_count") or 0),
|
||||
input_tokens=int(row["input_tokens"]),
|
||||
output_tokens=int(row["output_tokens"]),
|
||||
spend=float(row["spend"] or 0.0),
|
||||
models=list(row["models"]),
|
||||
)
|
||||
|
||||
|
||||
def _litellm_request(row: dict[str, Any]) -> LiteLLMRequest | None:
|
||||
if not row.get("s_request_id"):
|
||||
return None
|
||||
ttft_start: Final = int(row["s_ttft_start_ms"] or 0)
|
||||
return LiteLLMRequest(
|
||||
request_id=row["s_request_id"],
|
||||
model=row["s_model"],
|
||||
model_group=row["s_model_group"],
|
||||
provider=row["s_provider"],
|
||||
key_alias=row["s_key_alias"],
|
||||
team_alias=row["s_team_alias"],
|
||||
spend=float(row["s_spend"]),
|
||||
prompt_tokens=int(row["s_prompt_tokens"]),
|
||||
completion_tokens=int(row["s_completion_tokens"]),
|
||||
cache_read_tokens=int(row["s_cache_read_tokens"]),
|
||||
cache_write_tokens=int(row["s_cache_write_tokens"]),
|
||||
latency_ms=int(row["s_end_ms"]) - int(row["s_start_ms"]),
|
||||
ttft_ms=(ttft_start - int(row["s_start_ms"])) if ttft_start else None,
|
||||
status=row["s_status"],
|
||||
)
|
||||
|
||||
|
||||
def span_from_row(row: dict[str, Any], trace_start_ns: int) -> Span:
|
||||
return Span(
|
||||
span_id=row["span_id"],
|
||||
parent_span_id=row["parent_span_id"] or None,
|
||||
name=row["name"],
|
||||
type=row["type"],
|
||||
agent=row["agent"],
|
||||
start_offset_ms=(int(row["start_ns"]) - trace_start_ns) / NANOS_PER_MS,
|
||||
duration_ms=int(row["duration_ns"]) / NANOS_PER_MS,
|
||||
status=_status(row["status"]),
|
||||
error=row.get("status_message") or None,
|
||||
input_preview=row["input_preview"],
|
||||
model=row["model"] or None,
|
||||
input_tokens=int(row["input_tokens"]),
|
||||
output_tokens=int(row["output_tokens"]),
|
||||
litellm=_litellm_request(row),
|
||||
)
|
||||
|
||||
|
||||
def _parent_agent_of(span: Span, by_id: dict[str, Span]) -> str | None:
|
||||
parent_id = span["parent_span_id"]
|
||||
while parent_id is not None and parent_id in by_id:
|
||||
parent = by_id[parent_id]
|
||||
if parent["type"] == "agent" and parent["name"] != span["name"]:
|
||||
return parent["name"]
|
||||
parent_id = parent["parent_span_id"]
|
||||
return None
|
||||
|
||||
|
||||
def agent_nodes(spans: list[Span]) -> list[AgentNode]:
|
||||
"""One node per distinct agent name (200 `researcher` invocations = 1 node), with who invoked it."""
|
||||
by_id: Final = {s["span_id"]: s for s in spans}
|
||||
agents: dict[str, AgentNode] = {}
|
||||
for span in spans:
|
||||
if span["type"] != "agent":
|
||||
continue
|
||||
node = agents.setdefault(
|
||||
span["name"],
|
||||
AgentNode(
|
||||
name=span["name"],
|
||||
parent_agent=_parent_agent_of(span, by_id),
|
||||
invocations=0,
|
||||
llm_calls=0,
|
||||
tool_calls=0,
|
||||
spend=0.0,
|
||||
duration_ms=0.0,
|
||||
),
|
||||
)
|
||||
node["invocations"] += 1
|
||||
node["duration_ms"] += span["duration_ms"]
|
||||
for span in spans:
|
||||
owner = agents.get(span["agent"])
|
||||
if owner is None:
|
||||
continue
|
||||
if span["type"] == "llm":
|
||||
owner["llm_calls"] += 1
|
||||
owner["spend"] += span["litellm"]["spend"] if span["litellm"] else 0.0
|
||||
elif span["type"] == "tool":
|
||||
owner["tool_calls"] += 1
|
||||
return list(agents.values())
|
||||
|
||||
|
||||
def trace_from_rows(trace_id: str, rows: list[dict[str, Any]]) -> Trace | None:
|
||||
if not rows:
|
||||
return None
|
||||
trace_start_ns: Final = min(int(r["start_ns"]) for r in rows)
|
||||
trace_end_ns: Final = max(int(r["start_ns"]) + int(r["duration_ns"]) for r in rows)
|
||||
spans: Final = [span_from_row(r, trace_start_ns) for r in rows]
|
||||
root: Final = next((s for s in spans if s["parent_span_id"] is None), spans[0])
|
||||
agents: Final = agent_nodes(spans)
|
||||
llm_spans: Final = [s for s in spans if s["type"] == "llm"]
|
||||
return Trace(
|
||||
summary=TraceSummary(
|
||||
trace_id=trace_id,
|
||||
name=root["name"],
|
||||
service=rows[0]["service"],
|
||||
input_preview=root["input_preview"],
|
||||
start_time=_iso(trace_start_ns // NANOS_PER_MS),
|
||||
duration_ms=(trace_end_ns - trace_start_ns) / NANOS_PER_MS,
|
||||
status=root["status"],
|
||||
span_count=len(spans),
|
||||
agent_count=len(agents),
|
||||
agent_invocations=sum(a["invocations"] for a in agents),
|
||||
llm_calls=len(llm_spans),
|
||||
tool_calls=sum(1 for s in spans if s["type"] == "tool"),
|
||||
error_count=sum(1 for s in spans if s["status"] == "error"),
|
||||
input_tokens=sum(s["input_tokens"] for s in spans),
|
||||
output_tokens=sum(s["output_tokens"] for s in spans),
|
||||
spend=sum(s["litellm"]["spend"] for s in llm_spans if s["litellm"]),
|
||||
models=sorted({s["model"] for s in llm_spans if s["model"]}),
|
||||
),
|
||||
agents=agents,
|
||||
spans=spans,
|
||||
)
|
||||
|
||||
|
||||
class _SpanBatchWriter(ClickHouseBatchLogger):
|
||||
table = OTEL_TRACES_TABLE
|
||||
|
||||
|
||||
class ClickHouseTraceStore:
|
||||
"""Writes spans in batches (never blocks the caller) and runs scoped trace reads."""
|
||||
|
||||
def __init__(self, client: ClickHouseClient):
|
||||
self.client = client
|
||||
self.writer = _SpanBatchWriter(client=client)
|
||||
|
||||
def is_full(self) -> bool:
|
||||
return self.writer.is_full()
|
||||
|
||||
def write_spans(self, rows: list[SpanRow]) -> None:
|
||||
self.writer.enqueue([dict(r) for r in rows])
|
||||
|
||||
async def flush(self) -> None:
|
||||
await self.writer.flush_queue()
|
||||
|
||||
async def list_traces(
|
||||
self,
|
||||
scope: TraceScope,
|
||||
start_ms: int,
|
||||
end_ms: int,
|
||||
cursor: str | None = None,
|
||||
limit: int = AGENT_TRACING_LIST_PAGE_SIZE,
|
||||
) -> TracePage:
|
||||
cursor_ms, cursor_trace_id = decode_cursor(cursor)
|
||||
rows = await self.client.query(
|
||||
LIST_TRACES_SQL,
|
||||
{
|
||||
**scope,
|
||||
"start_ms": start_ms,
|
||||
"end_ms": end_ms,
|
||||
"cursor_ms": cursor_ms,
|
||||
"cursor_trace_id": cursor_trace_id,
|
||||
"limit": limit,
|
||||
},
|
||||
)
|
||||
next_cursor = encode_cursor(int(rows[-1]["start_ms"]), rows[-1]["trace_id"]) if len(rows) == limit else None
|
||||
return TracePage(data=[trace_summary_from_row(r) for r in rows], next_cursor=next_cursor)
|
||||
|
||||
async def get_trace(self, trace_id: str, scope: TraceScope) -> Trace | None:
|
||||
rows = await self.client.query(TRACE_SPANS_SQL, {**scope, "trace_id": trace_id})
|
||||
return trace_from_rows(trace_id, rows)
|
||||
|
||||
async def get_span(self, trace_id: str, span_id: str, scope: TraceScope) -> SpanDetail | None:
|
||||
rows = await self.client.query(SPAN_DETAIL_SQL, {**scope, "trace_id": trace_id, "span_id": span_id})
|
||||
if not rows:
|
||||
return None
|
||||
return SpanDetail(
|
||||
span_id=rows[0]["span_id"],
|
||||
input=rows[0]["input"],
|
||||
output=rows[0]["output"],
|
||||
attributes=dict(rows[0]["attributes"]),
|
||||
)
|
||||
Loading…
Add table
Reference in a new issue