diff --git a/litellm-rust/crates/python-bridge/src/routes/traces.rs b/litellm-rust/crates/python-bridge/src/routes/traces.rs index 3f2589e4b0e..6a18273ed4c 100644 --- a/litellm-rust/crates/python-bridge/src/routes/traces.rs +++ b/litellm-rust/crates/python-bridge/src/routes/traces.rs @@ -129,6 +129,9 @@ pub fn trace_decode_otlp<'py>( max_decompressed_bytes, ) }) - .map_err(|error| PyValueError::new_err(error.to_string()))?; + .map_err(|error| match error { + litellm_traces::DecodeError::TooLarge => PyOverflowError::new_err(error.to_string()), + _ => PyValueError::new_err(error.to_string()), + })?; litellm_host_python::Pythonized(spans).into_pyobject(py) } diff --git a/litellm-rust/crates/traces/migrations/0001_otel_traces.sql b/litellm-rust/crates/traces/migrations/0001_otel_traces.sql index aed869e6ee0..d8e0184b5a3 100644 --- a/litellm-rust/crates/traces/migrations/0001_otel_traces.sql +++ b/litellm-rust/crates/traces/migrations/0001_otel_traces.sql @@ -44,4 +44,4 @@ CREATE TABLE IF NOT EXISTS {database}.otel_traces ENGINE = MergeTree PARTITION BY toDate(Timestamp) ORDER BY (TeamId, ServiceName, toDateTime(Timestamp), TraceId) -SETTINGS ttl_only_drop_parts = 1 +SETTINGS ttl_only_drop_parts = 1, non_replicated_deduplication_window = 1000 diff --git a/litellm-rust/crates/traces/migrations/0002_agent_traces.sql b/litellm-rust/crates/traces/migrations/0002_agent_traces.sql index bb5cf67794d..0c3547872bb 100644 --- a/litellm-rust/crates/traces/migrations/0002_agent_traces.sql +++ b/litellm-rust/crates/traces/migrations/0002_agent_traces.sql @@ -22,3 +22,4 @@ CREATE TABLE IF NOT EXISTS {database}.agent_traces_by_key ) ENGINE = AggregatingMergeTree ORDER BY (TeamId, ApiKeyHash, TraceId) +SETTINGS non_replicated_deduplication_window = 1000 diff --git a/litellm-rust/crates/traces/src/insert.rs b/litellm-rust/crates/traces/src/insert.rs index e6b58903964..6f2d6acb023 100644 --- a/litellm-rust/crates/traces/src/insert.rs +++ b/litellm-rust/crates/traces/src/insert.rs @@ -58,6 +58,7 @@ pub async fn insert_rows( ), ) .append_pair("async_insert", "1") + .append_pair("async_insert_deduplicate", "1") .append_pair("wait_for_async_insert", "1") .append_pair("date_time_input_format", "best_effort"); let response = client diff --git a/litellm-rust/crates/traces/tests/migrations.rs b/litellm-rust/crates/traces/tests/migrations.rs index bf14710cf48..3fbd2207e8f 100644 --- a/litellm-rust/crates/traces/tests/migrations.rs +++ b/litellm-rust/crates/traces/tests/migrations.rs @@ -2,7 +2,7 @@ use std::{collections::BTreeMap, time::Duration}; use litellm_http::Client; use litellm_traces::{ - Connection, Error, encode_rows, ensure_schema, execute_read, schema_statements, + Connection, Error, InsertTable, encode_rows, ensure_schema, execute_read, schema_statements, }; use rstest::{fixture, rstest}; use testcontainers_modules::{ @@ -154,6 +154,41 @@ async fn schema_supports_span_rollups_and_spend_joins( Ok(()) } +#[rstest] +#[tokio::test] +async fn retried_trace_insert_does_not_inflate_rollup( + #[future(awt)] database: TestResult, +) -> TestResult { + let database = database?; + let writer = Connection::writer(&database.url)?; + ensure_schema(&database.client, &writer, "trace_test", 7, 14).await?; + let row = serde_json::from_value(serde_json::json!({ + "Timestamp": time::OffsetDateTime::now_utc().unix_timestamp_nanos() as i64, + "TraceId": "retried-trace", "SpanId": "span-1", "ParentSpanId": "", + "TeamId": "team-1", "ApiKeyHash": "key-1", "SpanName": "root", "InputTokens": 7 + }))?; + for _ in 0..2 { + litellm_traces::insert_rows( + &database.client, + &writer, + "trace_test", + InsertTable::OtelTraces, + vec![row.clone()], + ) + .await?; + } + let counts = read_json( + &database, + "SELECT toUInt32(sum(SpanCount)) AS spans, toUInt32(sum(InputTokens)) AS tokens \ + FROM trace_test.agent_traces_by_key WHERE TraceId = 'retried-trace'", + ) + .await?; + assert_eq!(table_rows(&database, "otel_traces").await?, 1); + assert_eq!(counts["data"][0]["spans"], 1); + assert_eq!(counts["data"][0]["tokens"], 7); + Ok(()) +} + #[rstest] #[tokio::test] async fn keyed_rollup_keeps_same_trace_ids_separate_by_api_key( @@ -370,7 +405,7 @@ async fn schema_statement_timeout_maps_to_transport_error() -> TestResult { let url = format!("http://{address}"); let writer = Connection::writer(&url)?; let result = tokio::time::timeout( - Duration::from_secs(12), + Duration::from_secs(35), ensure_schema(&client, &writer, "trace_test", 7, 14), ) .await; diff --git a/litellm/proxy/common_utils/http_parsing_utils.py b/litellm/proxy/common_utils/http_parsing_utils.py index 13a9b7ebf2b..e87279d02f7 100644 --- a/litellm/proxy/common_utils/http_parsing_utils.py +++ b/litellm/proxy/common_utils/http_parsing_utils.py @@ -155,7 +155,11 @@ async def _read_request_body(request: Request | None) -> dict: _request_headers: Final[dict] = _safe_get_request_headers(request=request) content_type: Final = _request_headers.get("content-type", "") - if _normalize_media_type(content_type) in _BINARY_CONTENT_TYPES: + if _normalize_media_type(content_type) in _BINARY_CONTENT_TYPES or ( + request.scope.get("path") == "/v1/traces" + and request.scope.get("method") == "POST" + and _request_headers.get("content-encoding", "").lower() == "gzip" + ): parsed_body = _parse_binary_body(await request.body()) elif _is_form_content_type(content_type): try: diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index 7f2742fcffc..4001f00a124 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -1523,7 +1523,7 @@ async def proxy_startup_event(app: FastAPI) -> AsyncGenerator[None, None]: asyncio.create_task(_adaptive_router_flusher_loop()) ## [Optional] Initialize agent tracing - await ProxyStartupEvent._init_tracing(general_settings) + asyncio.create_task(ProxyStartupEvent._init_tracing(general_settings)) ## [Optional] Initialize dd tracer ProxyStartupEvent._init_dd_tracer() @@ -11326,8 +11326,13 @@ class ProxyStartupEvent: settings = general_settings.get("tracing") or {} if settings.get("store") != "clickhouse": return - tracing = TraceReceiver.from_env() - await tracing.start() + try: + tracing: Final = TraceReceiver.from_env() + await tracing.start() + except (KeyError, OSError, RuntimeError, ValueError) as error: + tracing_endpoints.receiver = None + verbose_proxy_logger.warning("Agent tracing unavailable: %s", error) + return tracing_endpoints.receiver = tracing verbose_proxy_logger.info("Agent tracing enabled (store=clickhouse)") diff --git a/litellm/proxy/tracing_endpoints.py b/litellm/proxy/tracing_endpoints.py index 1486c218683..77a4e72ff57 100644 --- a/litellm/proxy/tracing_endpoints.py +++ b/litellm/proxy/tracing_endpoints.py @@ -20,7 +20,7 @@ from litellm.tracing import ( TraceReceiver, TracingPayloadTooLargeError, ) -from litellm.tracing.decode import encode_otlp_response +from litellm.tracing.decode import InvalidOTLPPayloadError, encode_otlp_response from litellm.tracing.types import SpanDetail, Trace, TracePage, TraceScope router = APIRouter(tags=["agent tracing"]) @@ -84,10 +84,12 @@ async def ingest_otlp_traces( content_encoding=request.headers.get("content-encoding"), tenant=tenant_for(user_api_key_dict), ) - except RuntimeError: - raise HTTPException(status_code=503, headers={"Retry-After": str(OTLP_RETRY_AFTER_SECONDS)}) except TracingPayloadTooLargeError as e: raise HTTPException(status_code=413, detail=str(e)) + except InvalidOTLPPayloadError as error: + raise HTTPException(status_code=400, detail=str(error)) from error + except RuntimeError: + raise HTTPException(status_code=503, headers={"Retry-After": str(OTLP_RETRY_AFTER_SECONDS)}) body, media_type = encode_otlp_response(content_type) return Response(content=body, media_type=media_type) @@ -100,20 +102,24 @@ async def list_agent_traces( cursor: Annotated[str | None, Query()] = None, ) -> TracePage: now_ms: Final = int(time.time() * 1000) - return await get_receiver().list_traces( - scope=scope_for(user_api_key_dict), - start_ms=start_ms if start_ms is not None else now_ms - MS_PER_DAY, - end_ms=end_ms if end_ms is not None else now_ms, - cursor=cursor, - ) + try: + return await get_receiver().list_traces( + scope=scope_for(user_api_key_dict), + start_ms=start_ms if start_ms is not None else now_ms - MS_PER_DAY, + end_ms=end_ms if end_ms is not None else now_ms, + cursor=cursor, + ) + except ValueError as error: + raise HTTPException(status_code=400, detail=str(error)) from error @router.get("/v1/traces/{trace_id}", response_model=None) async def get_agent_trace( trace_id: str, user_api_key_dict: Annotated[UserAPIKeyAuth, Depends(user_api_key_auth)], + trace_ref: Annotated[str, Query()] = "", ) -> Trace: - trace: Final = await get_receiver().get_trace(trace_id, scope_for(user_api_key_dict)) + trace: Final = await get_receiver().get_trace(trace_id, scope_for(user_api_key_dict), trace_ref) if trace is None: raise HTTPException(status_code=404, detail=f"Trace {trace_id} not found") return trace @@ -124,8 +130,9 @@ async def get_agent_trace_span( trace_id: str, span_id: str, user_api_key_dict: Annotated[UserAPIKeyAuth, Depends(user_api_key_auth)], + trace_ref: Annotated[str, Query()] = "", ) -> SpanDetail: - span: Final = await get_receiver().get_span(trace_id, span_id, scope_for(user_api_key_dict)) + span: Final = await get_receiver().get_span(trace_id, span_id, scope_for(user_api_key_dict), trace_ref) if span is None: raise HTTPException(status_code=404, detail=f"Span {span_id} not found") return span diff --git a/litellm/tracing/decode.py b/litellm/tracing/decode.py index ba13285e8bc..f4534645b55 100644 --- a/litellm/tracing/decode.py +++ b/litellm/tracing/decode.py @@ -43,6 +43,14 @@ _LC_ROLES: Final = {"human": "user", "ai": "assistant", "system": "system", "too _OPENINFERENCE_TYPES: Final[dict[str, SpanType]] = {"AGENT": "agent", "LLM": "llm", "TOOL": "tool"} +class InvalidOTLPPayloadError(ValueError): + pass + + +class OTLPPayloadTooLargeError(OverflowError): + pass + + # ---------------------------------------------------------------- decode @@ -56,7 +64,13 @@ def _truncate(value: str) -> str: def decode_otlp(body: bytes, content_type: str | None = None, content_encoding: str | None = None) -> list[SpanRow]: """Decode an OTLP trace export and normalize every span.""" - return [_span_row(span) for span in native_decode_otlp(body, content_type, content_encoding, OTLP_MAX_BODY_BYTES)] + try: + spans: Final = native_decode_otlp(body, content_type, content_encoding, OTLP_MAX_BODY_BYTES) + except OverflowError as error: + raise OTLPPayloadTooLargeError(str(error)) from error + except ValueError as error: + raise InvalidOTLPPayloadError(str(error)) from error + return [_span_row(span) for span in spans] def _exception_message(span: DecodedSpan) -> str: @@ -146,10 +160,18 @@ def _langsmith_io(row: SpanRow, attributes: Mapping[str, str]) -> None: if row["ObservationType"] == "llm" and isinstance(completion, dict): messages = prompt_payload.get("messages") or [[]] batch = messages[0] if messages and isinstance(messages[0], list) else messages - row["Input"] = json.dumps([_lc_message(m) for m in batch]) - generation = completion["generations"][0][0]["message"]["kwargs"] - row["Output"] = json.dumps(_lc_message({"kwargs": generation})) - row["LiteLLMRequestId"] = (generation.get("response_metadata") or {}).get("id") or "" + row["Input"] = json.dumps([_lc_message(m) for m in batch if isinstance(m, dict)]) if isinstance(batch, list) else "" + generations: Final = completion.get("generations") + first: Final = generations[0] if isinstance(generations, list) and generations else None + item: Final = first[0] if isinstance(first, list) and first else None + message: Final = item.get("message") if isinstance(item, dict) else None + generation: Final = message.get("kwargs") if isinstance(message, dict) else None + if isinstance(generation, dict): + row["Output"] = json.dumps(_lc_message({"kwargs": generation})) + metadata: Final = generation.get("response_metadata") + row["LiteLLMRequestId"] = metadata.get("id", "") if isinstance(metadata, dict) else "" + else: + row["Output"] = attributes.get("gen_ai.completion", "") return if row["ObservationType"] == "tool": output = (completion or {}).get("output", completion) if isinstance(completion, dict) else completion diff --git a/litellm/tracing/receiver.py b/litellm/tracing/receiver.py index 1efe52e7c73..cf664d5caef 100644 --- a/litellm/tracing/receiver.py +++ b/litellm/tracing/receiver.py @@ -24,7 +24,7 @@ from litellm.constants import ( ) from litellm.integrations.clickhouse.schema import ensure_schema from litellm.rust_bridge.traces import TraceStorage -from litellm.tracing.decode import decode_otlp +from litellm.tracing.decode import InvalidOTLPPayloadError, OTLPPayloadTooLargeError, decode_otlp from litellm.tracing.store import ClickHouseTraceStore from litellm.tracing.types import ( SpanDetail, @@ -94,11 +94,14 @@ class TraceReceiver: """Decode an OTLP trace export and store its authenticated spans.""" if len(body) > OTLP_MAX_BODY_BYTES: raise TracingPayloadTooLargeError(f"OTLP body exceeds {OTLP_MAX_BODY_BYTES} bytes") - rows: Final = ( - await asyncio.to_thread(decode_otlp, body, content_type, content_encoding) - if len(body) > OTLP_OFFLOAD_DECODE_BYTES - else decode_otlp(body, content_type, content_encoding) - ) + try: + rows: Final = ( + await asyncio.to_thread(decode_otlp, body, content_type, content_encoding) + if len(body) > OTLP_OFFLOAD_DECODE_BYTES + else decode_otlp(body, content_type, content_encoding) + ) + except OTLPPayloadTooLargeError as error: + raise TracingPayloadTooLargeError(str(error)) from error try: await self.store.insert_spans([tenant.stamp(r) for r in rows]) except OverflowError as error: @@ -110,8 +113,8 @@ class TraceReceiver: async def list_traces(self, scope: TraceScope, start_ms: int, end_ms: int, cursor: str | None = None) -> TracePage: return await self.store.list_traces(scope, start_ms, end_ms, cursor) - async def get_trace(self, trace_id: str, scope: TraceScope) -> Trace | None: - return await self.store.get_trace(trace_id, scope) + async def get_trace(self, trace_id: str, scope: TraceScope, trace_ref: str = "") -> Trace | None: + return await self.store.get_trace(trace_id, scope, trace_ref) - async def get_span(self, trace_id: str, span_id: str, scope: TraceScope) -> SpanDetail | None: - return await self.store.get_span(trace_id, span_id, scope) + async def get_span(self, trace_id: str, span_id: str, scope: TraceScope, trace_ref: str = "") -> SpanDetail | None: + return await self.store.get_span(trace_id, span_id, scope, trace_ref) diff --git a/litellm/tracing/store.py b/litellm/tracing/store.py index 404aa3acd3b..732d148a29b 100644 --- a/litellm/tracing/store.py +++ b/litellm/tracing/store.py @@ -1,6 +1,7 @@ """ClickHouse-backed trace store: batched span writes and scoped reads.""" import base64 +import binascii import json from datetime import datetime, timezone from typing import Any, Final @@ -30,38 +31,28 @@ _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})" ) +_TRACE_REF_SQL: Final = "hex(SHA256(concat(TeamId, char(0), ApiKeyHash, char(0), TraceId)))" LIST_TRACES_SQL: Final = f""" -SELECT t.TraceId AS trace_id, ifNull(any(t.RootName), '') AS name, any(t.ServiceName) AS service, - ifNull(any(t.RootInput), '') AS input_preview, ifNull(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 -FROM ( - SELECT TeamId, TraceId, min(StartTs) AS StartTs, max(EndTs) AS EndTs, - any(ServiceName) AS ServiceName, anyLast(a.RootName) AS RootName, - anyLast(a.RootInput) AS RootInput, - anyLast(a.RootStatus) 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 - FROM {AGENT_TRACES_BY_KEY_TABLE} AS a - WHERE (empty({{team_ids:Array(String)}}) OR TeamId IN {{team_ids:Array(String)}}) - AND ({{api_key_hash:String}} = '' OR ApiKeyHash = {{api_key_hash:String}}) - 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 -GROUP BY t.TraceId -ORDER BY start_ms DESC, t.TraceId DESC +SELECT TraceId AS trace_id, {_TRACE_REF_SQL} AS trace_ref, + ifNull(any(RootName), '') AS name, any(ServiceName) AS service, + ifNull(any(RootInput), '') AS input_preview, ifNull(any(RootStatus), '') AS status, + toUnixTimestamp64Milli(min(StartTs)) AS start_ms, + dateDiff('millisecond', min(StartTs), max(EndTs)) AS duration_ms, + sum(SpanCount) AS span_count, length(groupUniqArrayArray(AgentNames)) AS agent_count, + sum(AgentCount) AS agent_invocations, + sum(LlmCount) AS llm_calls, sum(ToolCount) AS tool_calls, + sum(InputTokens) AS input_tokens, sum(OutputTokens) AS output_tokens, + groupUniqArrayArray(Models) AS models, sum(ErrorCount) AS error_count +FROM {AGENT_TRACES_BY_KEY_TABLE} +WHERE (empty({{team_ids:Array(String)}}) OR TeamId IN {{team_ids:Array(String)}}) + AND ({{api_key_hash:String}} = '' OR ApiKeyHash = {{api_key_hash:String}}) +GROUP BY TeamId, ApiKeyHash, TraceId +HAVING min(StartTs) >= fromUnixTimestamp64Milli({{start_ms:Int64}}) + AND min(StartTs) < fromUnixTimestamp64Milli({{end_ms:Int64}}) + AND ({{cursor_ms:Int64}} = 0 OR (toUnixTimestamp64Milli(min(StartTs)), trace_ref) + < ({{cursor_ms:Int64}}, {{cursor_trace_id:String}})) +ORDER BY start_ms DESC, trace_ref DESC +LIMIT {{limit:UInt32}} """ TRACE_SPANS_SQL: Final = f""" @@ -74,6 +65,7 @@ SELECT o.SpanId AS span_id, o.ParentSpanId AS parent_span_id, o.SpanName AS name o.LiteLLMRequestId AS litellm_request_id FROM {OTEL_TRACES_TABLE} AS o WHERE o.TraceId = {{trace_id:String}} AND {_SCOPE_OTEL} + AND ({{trace_ref:String}} = '' OR {_TRACE_REF_SQL} = {{trace_ref:String}}) ORDER BY o.Timestamp LIMIT 1 BY o.SpanId """ @@ -82,6 +74,7 @@ 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} + AND ({{trace_ref:String}} = '' OR {_TRACE_REF_SQL} = {{trace_ref:String}}) LIMIT 1 """ @@ -93,8 +86,21 @@ def encode_cursor(start_ms: int, trace_id: str) -> str: 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) + try: + value: Final = json.loads(base64.b64decode(cursor, altchars=b"-_", validate=True)) + if ( + not isinstance(value, list) + or len(value) != 2 + or not isinstance(value[0], int) + or isinstance(value[0], bool) + or value[0] <= 0 + or not isinstance(value[1], str) + or not value[1] + ): + raise ValueError("Invalid trace cursor") + return value[0], value[1] + except (ValueError, UnicodeError, binascii.Error) as error: + raise ValueError("Invalid trace cursor") from error def _iso(ms: int) -> str: @@ -108,6 +114,7 @@ def _status(code: str) -> SpanStatus: def trace_summary_from_row(row: dict[str, Any]) -> TraceSummary: return TraceSummary( trace_id=row["trace_id"], + trace_ref=row.get("trace_ref", ""), name=row["name"], service=row["service"], input_preview=row["input_preview"], @@ -147,7 +154,9 @@ def span_from_row(row: dict[str, Any], trace_start_ns: int) -> Span: 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: + visited: Final = {span["span_id"]} + while parent_id is not None and parent_id in by_id and parent_id not in visited: + visited.add(parent_id) parent = by_id[parent_id] if parent["type"] == "agent" and parent["name"] != span["name"]: return parent["name"] @@ -186,7 +195,7 @@ def agent_nodes(spans: list[Span]) -> list[AgentNode]: return list(agents.values()) -def trace_from_rows(trace_id: str, rows: list[dict[str, Any]]) -> Trace | None: +def trace_from_rows(trace_id: str, rows: list[dict[str, Any]], trace_ref: str = "") -> Trace | None: if not rows: return None trace_start_ns: Final = min(int(r["start_ns"]) for r in rows) @@ -198,6 +207,7 @@ def trace_from_rows(trace_id: str, rows: list[dict[str, Any]]) -> Trace | None: return Trace( summary=TraceSummary( trace_id=trace_id, + trace_ref=trace_ref, name=root["name"], service=rows[0]["service"], input_preview=root["input_preview"], @@ -248,15 +258,17 @@ class ClickHouseTraceStore: "limit": limit, }, ) - next_cursor = encode_cursor(int(rows[-1]["start_ms"]), rows[-1]["trace_id"]) if len(rows) == limit else None + next_cursor = encode_cursor(int(rows[-1]["start_ms"]), rows[-1]["trace_ref"]) 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.storage.query(TRACE_SPANS_SQL, {**scope, "trace_id": trace_id}) - return trace_from_rows(trace_id, rows) + async def get_trace(self, trace_id: str, scope: TraceScope, trace_ref: str = "") -> Trace | None: + rows = await self.storage.query(TRACE_SPANS_SQL, {**scope, "trace_id": trace_id, "trace_ref": trace_ref}) + return trace_from_rows(trace_id, rows, trace_ref) - async def get_span(self, trace_id: str, span_id: str, scope: TraceScope) -> SpanDetail | None: - rows = await self.storage.query(SPAN_DETAIL_SQL, {**scope, "trace_id": trace_id, "span_id": span_id}) + async def get_span(self, trace_id: str, span_id: str, scope: TraceScope, trace_ref: str = "") -> SpanDetail | None: + rows = await self.storage.query( + SPAN_DETAIL_SQL, {**scope, "trace_id": trace_id, "span_id": span_id, "trace_ref": trace_ref} + ) if not rows: return None return SpanDetail( diff --git a/litellm/tracing/types.py b/litellm/tracing/types.py index bf0e5b3a722..eb6acc541c3 100644 --- a/litellm/tracing/types.py +++ b/litellm/tracing/types.py @@ -11,7 +11,7 @@ A trace is one agent run. It's made of spans (agent / llm / tool / chain / frame from typing import Literal -from typing_extensions import TypedDict +from typing_extensions import NotRequired, ReadOnly, TypedDict SpanType = Literal["agent", "llm", "tool", "chain", "framework"] SpanStatus = Literal["ok", "error", "unset"] @@ -47,6 +47,7 @@ class AgentNode(TypedDict): class TraceSummary(TypedDict): trace_id: str + trace_ref: ReadOnly[NotRequired[str]] name: str service: str input_preview: str diff --git a/tests/test_litellm/proxy/common_utils/test_http_parsing_utils.py b/tests/test_litellm/proxy/common_utils/test_http_parsing_utils.py index d4c2fb3b9a1..bb1bca45c05 100644 --- a/tests/test_litellm/proxy/common_utils/test_http_parsing_utils.py +++ b/tests/test_litellm/proxy/common_utils/test_http_parsing_utils.py @@ -1,4 +1,5 @@ import io +import gzip import json from typing import get_type_hints from unittest.mock import AsyncMock, MagicMock, patch @@ -30,12 +31,12 @@ from litellm.proxy.common_utils.http_parsing_utils import ( ) -def _starlette_request(body: bytes, content_type: str) -> Request: +def _starlette_request(body: bytes, content_type: str, path: str = "/v1/messages", content_encoding: str = "") -> Request: scope = { "type": "http", "method": "POST", - "path": "/v1/messages", - "headers": [(b"content-type", content_type.encode())], + "path": path, + "headers": [(b"content-type", content_type.encode()), (b"content-encoding", content_encoding.encode())], "query_string": b"", } chunks = iter((body,)) @@ -83,6 +84,14 @@ async def test_protobuf_body_is_not_parsed_as_json(content_type): assert await request.body() == body # body is still readable by the endpoint +@pytest.mark.asyncio +async def test_gzipped_json_trace_body_survives_auth_pre_read(): + body = gzip.compress(b'{"resourceSpans": []}') + request = _starlette_request(body, "application/json", "/v1/traces", "gzip") + assert await _read_request_body(request) == {} + assert await request.body() == body + + @pytest.mark.asyncio async def test_read_raw_json_body_is_none_for_a_request_that_only_mocks_the_parsed_body_path(): mock_request = MagicMock() diff --git a/tests/test_litellm/proxy/test_tracing_endpoints.py b/tests/test_litellm/proxy/test_tracing_endpoints.py index 02b2056a63b..4a7fdfdb33a 100644 --- a/tests/test_litellm/proxy/test_tracing_endpoints.py +++ b/tests/test_litellm/proxy/test_tracing_endpoints.py @@ -140,7 +140,7 @@ def test_get_trace_404_and_200(client, receiver): response = client.get("/v1/traces/t1") assert response.status_code == 200 assert response.json() == trace - receiver.get_trace.assert_awaited_with("t1", {"team_ids": ["team-research"], "api_key_hash": ""}) + receiver.get_trace.assert_awaited_with("t1", {"team_ids": ["team-research"], "api_key_hash": ""}, "") def test_get_span_404_and_200(client, receiver): @@ -149,7 +149,22 @@ def test_get_span_404_and_200(client, receiver): response = client.get("/v1/traces/t1/spans/s1") assert response.status_code == 200 assert response.json()["span_id"] == "s1" - receiver.get_span.assert_awaited_with("t1", "s1", {"team_ids": ["team-research"], "api_key_hash": ""}) + receiver.get_span.assert_awaited_with("t1", "s1", {"team_ids": ["team-research"], "api_key_hash": ""}, "") + + +def test_trace_detail_passes_scoped_reference(client, receiver): + receiver.get_trace.return_value = {"summary": {"trace_id": "t1"}, "agents": [], "spans": []} + assert client.get("/v1/traces/t1?trace_ref=run-one").status_code == 200 + receiver.get_trace.assert_awaited_with("t1", {"team_ids": ["team-research"], "api_key_hash": ""}, "run-one") + + +def test_invalid_export_and_cursor_are_client_errors(client, receiver): + from litellm.tracing.decode import InvalidOTLPPayloadError + + receiver.ingest.side_effect = InvalidOTLPPayloadError("invalid OTLP trace payload") + assert client.post("/v1/traces", content=b"broken").status_code == 400 + receiver.list_traces.side_effect = ValueError("Invalid trace cursor") + assert client.get("/v1/traces?cursor=broken").status_code == 400 def test_teamless_key_without_token_gets_403_on_reads(client, receiver): diff --git a/tests/test_litellm/tracing/test_decode.py b/tests/test_litellm/tracing/test_decode.py index ea252f21e16..b213d7fb3df 100644 --- a/tests/test_litellm/tracing/test_decode.py +++ b/tests/test_litellm/tracing/test_decode.py @@ -113,6 +113,22 @@ def test_llm_input_output_are_normalized_messages(rows_by_name): assert output["tool_calls"][0]["name"] +@pytest.mark.parametrize("completion", ["{}", '{"generations": []}', '{"generations": [[{}]]}']) +def test_incomplete_langsmith_completion_preserves_the_export(completion): + span = _span( + "ChatOpenAI", + b"\x03" * 8, + b"\x02" * 8, + langsmith__span__kind="llm", + gen_ai__prompt='{"messages": [[{"kwargs": {"type": "human", "content": "hi"}}]]}', + gen_ai__completion=completion, + ) + rows = decode_otlp(_export(span, scope="langsmith"), "application/x-protobuf") + assert len(rows) == 1 + assert json.loads(rows[0]["Input"])[0]["content"] == "hi" + assert rows[0]["Output"] == completion + + def test_task_tool_output_is_subagent_final_message_text(rows_by_name): task = rows_by_name["task"] assert json.loads(task["Input"])["subagent_type"] == "researcher" diff --git a/tests/test_litellm/tracing/test_receiver.py b/tests/test_litellm/tracing/test_receiver.py index 525937bf552..3cf244e837c 100644 --- a/tests/test_litellm/tracing/test_receiver.py +++ b/tests/test_litellm/tracing/test_receiver.py @@ -118,4 +118,4 @@ async def test_reads_delegate_to_store(): tracing = TraceReceiver(store) scope: TraceScope = {"team_ids": ["team-research"], "api_key_hash": ""} assert await tracing.get_trace("t1", scope) is None - store.get_trace.assert_awaited_once_with("t1", scope) + store.get_trace.assert_awaited_once_with("t1", scope, "") diff --git a/tests/test_litellm/tracing/test_store.py b/tests/test_litellm/tracing/test_store.py index e5e5ee0f942..04dd84241f5 100644 --- a/tests/test_litellm/tracing/test_store.py +++ b/tests/test_litellm/tracing/test_store.py @@ -206,6 +206,16 @@ def test_parent_agent_skips_same_name_ancestors(): assert nodes["researcher"]["invocations"] == 2 +def test_parent_agent_stops_at_cyclic_parents(): + rows = [ + _row("self", "self", "researcher", "agent", "researcher"), + _row("first", "second", "researcher", "agent", "researcher"), + _row("second", "first", "researcher", "agent", "researcher"), + ] + spans = [span_from_row(row, T0) for row in rows] + assert agent_nodes(spans)[0]["parent_agent"] is None + + def test_agent_nodes_ignores_spans_of_unknown_agents(): spans = [span_from_row(_row("t", "", "tool", "tool", "ghost"), T0)] assert agent_nodes(spans) == [] @@ -221,6 +231,12 @@ def test_cursor_round_trip(): assert decode_cursor("") == (0, "") +@pytest.mark.parametrize("cursor", ["abc", "bm90LWpzb24=", "WzEsIDJd", "WzAsICJ0Il0="]) +def test_invalid_cursor_is_rejected(cursor): + with pytest.raises(ValueError, match="Invalid trace cursor"): + decode_cursor(cursor) + + def test_trace_summary_from_row(): summary = trace_summary_from_row( { @@ -251,6 +267,7 @@ async def test_list_traces_sets_next_cursor_on_full_page(): client = MagicMock() row = { "trace_id": "t2", + "trace_ref": "ref2", "name": "a", "service": "s", "input_preview": "", @@ -266,20 +283,20 @@ async def test_list_traces_sets_next_cursor_on_full_page(): "output_tokens": 0, "models": [], } - client.query = AsyncMock(return_value=[row, {**row, "trace_id": "t1", "start_ms": 900}]) + client.query = AsyncMock(return_value=[row, {**row, "trace_id": "t1", "trace_ref": "ref1", "start_ms": 900}]) store = ClickHouseTraceStore(client) scope: TraceScope = {"team_ids": ["team-a"], "api_key_hash": ""} page = await store.list_traces(scope, 0, 2000, limit=2) assert [t["trace_id"] for t in page["data"]] == ["t2", "t1"] assert page["next_cursor"] is not None - assert decode_cursor(page["next_cursor"]) == (900, "t1") + assert decode_cursor(page["next_cursor"]) == (900, "ref1") params = client.query.call_args.args[1] assert params["team_ids"] == ["team-a"] and params["limit"] == 2 and params["cursor_ms"] == 0 page = await store.list_traces(scope, 0, 2000, cursor=page["next_cursor"], limit=3) assert page["next_cursor"] is None - assert client.query.call_args.args[1]["cursor_trace_id"] == "t1" + assert client.query.call_args.args[1]["cursor_trace_id"] == "ref1" @pytest.mark.asyncio diff --git a/ui/litellm-dashboard/src/components/networking.tsx b/ui/litellm-dashboard/src/components/networking.tsx index d0a5349364f..7ffe0e0ec85 100644 --- a/ui/litellm-dashboard/src/components/networking.tsx +++ b/ui/litellm-dashboard/src/components/networking.tsx @@ -2121,12 +2121,21 @@ export const agentTraceListCall = async ({ return apiClient.get(`/v1/traces`, { accessToken, query }); }; -export const agentTraceCall = async (accessToken: string, traceId: string): Promise => - apiClient.get(`/v1/traces/${encodeURIComponent(traceId)}`, { accessToken }); +export const agentTraceCall = async (accessToken: string, traceId: string, traceRef?: string): Promise => + apiClient.get(`/v1/traces/${encodeURIComponent(traceId)}`, { + accessToken, + query: { trace_ref: traceRef || undefined }, + }); -export const agentTraceSpanCall = async (accessToken: string, traceId: string, spanId: string): Promise => +export const agentTraceSpanCall = async ( + accessToken: string, + traceId: string, + spanId: string, + traceRef?: string, +): Promise => apiClient.get(`/v1/traces/${encodeURIComponent(traceId)}/spans/${encodeURIComponent(spanId)}`, { accessToken, + query: { trace_ref: traceRef || undefined }, }); export const adminSpendLogsCall = async (accessToken: string) => { diff --git a/ui/litellm-dashboard/src/components/view_logs/TraceView/AgentTracesSection.tsx b/ui/litellm-dashboard/src/components/view_logs/TraceView/AgentTracesSection.tsx index 2bfa3c320a7..137462b38ec 100644 --- a/ui/litellm-dashboard/src/components/view_logs/TraceView/AgentTracesSection.tsx +++ b/ui/litellm-dashboard/src/components/view_logs/TraceView/AgentTracesSection.tsx @@ -67,7 +67,7 @@ export function AgentTracesSection({ timeControls, onRunOpenChange, }: AgentTracesSectionProps) { - const [openTraceId, setOpenTraceId] = useState(null); + const [openTrace, setOpenTrace] = useState(null); const [query, setQuery] = useState(""); const [service, setService] = useState(ALL_SERVICES); const [status, setStatus] = useState("all"); @@ -96,9 +96,9 @@ export function AgentTracesSection({ apply(hours); }; - const openRun = (traceId: string | null) => { - setOpenTraceId(traceId); - onRunOpenChange?.(traceId !== null); + const openRun = (trace: TraceSummary | null) => { + setOpenTrace(trace); + onRunOpenChange?.(trace !== null); }; if (traces.notEnabledDetail !== null) return ; @@ -120,8 +120,8 @@ export function AgentTracesSection({ ); } - if (openTraceId !== null) { - return openRun(null)} />; + if (openTrace !== null) { + return openRun(null)} />; } return ( diff --git a/ui/litellm-dashboard/src/components/view_logs/TraceView/AgentTracesTable.tsx b/ui/litellm-dashboard/src/components/view_logs/TraceView/AgentTracesTable.tsx index e4faff1621a..d3912b7be5c 100644 --- a/ui/litellm-dashboard/src/components/view_logs/TraceView/AgentTracesTable.tsx +++ b/ui/litellm-dashboard/src/components/view_logs/TraceView/AgentTracesTable.tsx @@ -14,7 +14,7 @@ interface AgentTracesTableProps { error: Error | null; hasMore: boolean; onLoadMore: () => void; - onOpenTrace: (traceId: string) => void; + onOpenTrace: (trace: TraceSummary) => void; } /** Spend is only on summaries once the spend-enrichment PR lands; show Cost when it's there. */ @@ -82,9 +82,9 @@ export function AgentTracesTable({ {traces.map((run) => ( onOpenTrace(run.trace_id)} + onClick={() => onOpenTrace(run)} className="h-9 cursor-pointer border-b border-border/60 text-[12px] hover:bg-accent/50" > const errorReason = (headline: string): string => /^([A-Za-z_][\w.]*)\(/.exec(headline)?.[1] ?? "error"; /** Shared lazy fetch of one span's full input / output / attributes. */ -export function useSpanDetail(accessToken: string, traceId: string, spanId: string | null) { +export function useSpanDetail(accessToken: string, traceId: string, spanId: string | null, traceRef?: string) { const queryOptions: UseQueryOptions = { - queryKey: ["agentTraceSpan", traceId, spanId, accessToken], - queryFn: () => agentTraceSpanCall(accessToken, traceId, spanId as string), + queryKey: ["agentTraceSpan", traceId, traceRef, spanId, accessToken], + queryFn: () => agentTraceSpanCall(accessToken, traceId, spanId as string, traceRef), enabled: spanId !== null, staleTime: Infinity, }; @@ -121,12 +121,13 @@ function Payload({ label, value, mono }: { label: string; value: string; mono: b interface DetailContentProps { accessToken: string; traceId: string; + traceRef?: string; span: Span; } /** Content tab: the error first (if any), then what went in and what came out. */ -export function DetailContent({ accessToken, traceId, span }: DetailContentProps) { - const detailQuery = useSpanDetail(accessToken, traceId, span.span_id); +export function DetailContent({ accessToken, traceId, traceRef, span }: DetailContentProps) { + const detailQuery = useSpanDetail(accessToken, traceId, span.span_id, traceRef); const detail = detailQuery.data; const isTool = span.type === "tool"; const empty = detail && !detail.input && !detail.output; diff --git a/ui/litellm-dashboard/src/components/view_logs/TraceView/DetailPane.test.tsx b/ui/litellm-dashboard/src/components/view_logs/TraceView/DetailPane.test.tsx index a15584e6d41..9a541d26966 100644 --- a/ui/litellm-dashboard/src/components/view_logs/TraceView/DetailPane.test.tsx +++ b/ui/litellm-dashboard/src/components/view_logs/TraceView/DetailPane.test.tsx @@ -125,7 +125,7 @@ describe("DetailPane", () => { expect(screen.getByRole("tab", { name: "Attributes" })).toBeInTheDocument(); expect(await screen.findByText("You are a LiteLLM support agent.")).toBeInTheDocument(); expect(screen.getByText("get_customer_plan")).toBeInTheDocument(); - expect(vi.mocked(agentTraceSpanCall)).toHaveBeenCalledWith("sk-test", "t1", "llm1"); + expect(vi.mocked(agentTraceSpanCall)).toHaveBeenCalledWith("sk-test", "t1", "llm1", undefined); }); it("shows a tool failure as 'Tool ยท ' with the exception line and no traceback", async () => { diff --git a/ui/litellm-dashboard/src/components/view_logs/TraceView/DetailPane.tsx b/ui/litellm-dashboard/src/components/view_logs/TraceView/DetailPane.tsx index 0a4dd3c5029..45117618faa 100644 --- a/ui/litellm-dashboard/src/components/view_logs/TraceView/DetailPane.tsx +++ b/ui/litellm-dashboard/src/components/view_logs/TraceView/DetailPane.tsx @@ -67,7 +67,7 @@ function SpanPane({ }) { const [tab, setTab] = useState("content"); const traceId = trace.summary.trace_id; - const detailQuery = useSpanDetail(accessToken, traceId, tab === "attributes" ? span.span_id : null); + const detailQuery = useSpanDetail(accessToken, traceId, tab === "attributes" ? span.span_id : null, trace.summary.trace_ref); const tokens = span.input_tokens + span.output_tokens; return (