diff --git a/litellm/tracing/store.py b/litellm/tracing/store.py new file mode 100644 index 00000000000..b3cefa8eac0 --- /dev/null +++ b/litellm/tracing/store.py @@ -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"]), + )