From a3e74a52233261fe4a66b1b5b0df806bb47adb2d Mon Sep 17 00:00:00 2001 From: moe-berri Date: Tue, 6 Oct 2026 14:00:40 -0700 Subject: [PATCH] fix(tracing): preserve optional provider evidence in fixture replay (#44930) * fix(tracing): preserve optional provider evidence in fixture replay * fix(tracing): keep copied provider identities consistent --- .../crates/traces-clickhouse/tests/reads.rs | 34 ++++++++++------- litellm/tracing/types.py | 4 +- scripts/seed_tracing_fixtures.py | 18 ++++++--- tests/unit/proxy/test_tracing_endpoints.py | 1 + tests/unit/test_seed_tracing_fixtures.py | 37 +++++++++++++++++-- 5 files changed, 69 insertions(+), 25 deletions(-) diff --git a/litellm-rust/crates/traces-clickhouse/tests/reads.rs b/litellm-rust/crates/traces-clickhouse/tests/reads.rs index b3d1b72bec1..c938b6f451b 100644 --- a/litellm-rust/crates/traces-clickhouse/tests/reads.rs +++ b/litellm-rust/crates/traces-clickhouse/tests/reads.rs @@ -600,9 +600,17 @@ async fn an_oversized_span_keeps_the_run_list_available_with_partial_totals( } #[rstest] +#[case::same_key("same-key", Some(0.25))] +#[case::other_key("other-key", None)] +#[case::foreign_team("foreign-team", None)] +#[case::call_id_other_key("call-id-other-key", None)] +#[case::call_id_foreign_team("call-id-foreign-team", None)] +#[case::transport_only("transport-only", None)] #[tokio::test] async fn assigned_call_ids_require_shared_ownership_through_detail_and_batch_reads( #[future(awt)] migrated_database: TestResult, + #[case] id: &str, + #[case] expected: Option, ) -> TestResult { let fixture = migrated_database?; let client = &fixture.database.client; @@ -705,20 +713,18 @@ async fn assigned_call_ids_require_shared_ownership_through_detail_and_batch_rea .list_traces(&store, &access, 0, 2_000_000_000_000, None, 50) .await?; assert_eq!(page.data.len(), cases.len()); - for (id, _, _, _, expected) in cases { - let summary = page - .data - .iter() - .find(|summary| summary.trace_id == id) - .ok_or("missing run")?; - let detail = reader - .get_trace(&store, &access, id, &summary.trace_ref) - .await? - .ok_or("missing trace")?; - assert_eq!(detail.summary.spend, expected, "{id}"); - assert_eq!(summary.spend, expected, "{id}"); - assert_eq!(summary.priced_calls, u64::from(expected.is_some()), "{id}"); - } + let summary = page + .data + .iter() + .find(|summary| summary.trace_id == id) + .ok_or("missing run")?; + let detail = reader + .get_trace(&store, &access, id, &summary.trace_ref) + .await? + .ok_or("missing trace")?; + assert_eq!(detail.summary.spend, expected, "{id}"); + assert_eq!(summary.spend, expected, "{id}"); + assert_eq!(summary.priced_calls, u64::from(expected.is_some()), "{id}"); Ok(()) } diff --git a/litellm/tracing/types.py b/litellm/tracing/types.py index 8bd119940c2..2e5a80b7647 100644 --- a/litellm/tracing/types.py +++ b/litellm/tracing/types.py @@ -1,6 +1,6 @@ from collections.abc import Sequence -from typing_extensions import ReadOnly, TypedDict +from typing_extensions import NotRequired, ReadOnly, TypedDict class SpendLogRecord(TypedDict): @@ -8,7 +8,7 @@ class SpendLogRecord(TypedDict): request_id: ReadOnly[str] response_id: ReadOnly[str] - provider_request_id: ReadOnly[str] + provider_request_id: NotRequired[ReadOnly[str]] litellm_call_id: ReadOnly[str] call_type: ReadOnly[str] api_key: ReadOnly[str] diff --git a/scripts/seed_tracing_fixtures.py b/scripts/seed_tracing_fixtures.py index 5b3cf30f173..bcd904caec5 100644 --- a/scripts/seed_tracing_fixtures.py +++ b/scripts/seed_tracing_fixtures.py @@ -108,7 +108,13 @@ def response_ids(rows: tuple[SpendLogRecord, ...]) -> Iterator[str]: def response_pattern(rows: tuple[SpendLogRecord, ...]) -> re.Pattern[str]: identities: Final = sorted( frozenset( - identity for identity in chain(response_ids(rows), (row["litellm_call_id"] for row in rows)) if identity + identity + for identity in chain( + response_ids(rows), + (row["litellm_call_id"] for row in rows), + (row.get("provider_request_id", "") for row in rows), + ) + if identity ), key=len, reverse=True, @@ -119,8 +125,11 @@ def response_pattern(rows: tuple[SpendLogRecord, ...]) -> re.Pattern[str]: def rebased_response(value: str, namespace: str, pattern: re.Pattern[str]) -> str: decoded: Final = managed_response(value) if decoded is None: + if value.startswith(("msg_", "req_")): + prefix, suffix = value.split("_", 1) + return f"{prefix}_seed-{namespace}-{suffix}" return f"seed-{namespace}-{value}" - payload: Final = pattern.sub(lambda match: f"seed-{namespace}-{match.group()}", decoded) + payload: Final = pattern.sub(lambda match: rebased_response(match.group(), namespace, pattern), decoded) return "resp_" + base64.b64encode(payload.encode()).decode() @@ -194,9 +203,7 @@ def rebase( if isinstance(value, dict): attribute_key: Final = value.get("key") attribute_value: Final = value.get("value") - session_id: Final = ( - attribute_value.get("stringValue") if isinstance(attribute_value, dict) else None - ) + session_id: Final = attribute_value.get("stringValue") if isinstance(attribute_value, dict) else None if ( isinstance(attribute_key, str) and attribute_key in TRACE_ID_ATTRIBUTES @@ -431,6 +438,7 @@ WHERE t.TraceId IN {{trace_ids:Array(String)}} SELECT s.* REPLACE ( {clickhouse_call_id("s.request_id")} AS request_id, {clickhouse_call_id("s.response_id")} AS response_id, + {clickhouse_call_id("s.provider_request_id")} AS provider_request_id, {clickhouse_call_id("s.litellm_call_id")} AS litellm_call_id, {clickhouse_hash("s.trace_id", trace_salt, 32)} AS trace_id, {clickhouse_hash("s.session_id", trace_salt, 32)} AS session_id, diff --git a/tests/unit/proxy/test_tracing_endpoints.py b/tests/unit/proxy/test_tracing_endpoints.py index 8c6cea57f7e..a3c8c8d9181 100644 --- a/tests/unit/proxy/test_tracing_endpoints.py +++ b/tests/unit/proxy/test_tracing_endpoints.py @@ -86,6 +86,7 @@ TRACE_RESPONSE: Final = { "agent_count": 0, "agent_invocations": 0, "llm_calls": 0, + "priced_calls": 0, "tool_calls": 0, "error_count": 0, "input_tokens": 0, diff --git a/tests/unit/test_seed_tracing_fixtures.py b/tests/unit/test_seed_tracing_fixtures.py index 626cb3543a9..5dca24487f1 100644 --- a/tests/unit/test_seed_tracing_fixtures.py +++ b/tests/unit/test_seed_tracing_fixtures.py @@ -1,3 +1,4 @@ +import base64 import hashlib import json import re @@ -15,13 +16,13 @@ from pydantic import InstanceOf, TypeAdapter from litellm.rust_bridge.trace.generated.responses import TraceSQLResponse from litellm.rust_bridge.trace.storage import Tenant, span_rows from litellm.tracing.types import SpendLogRecord -from tests._master_key import MASTER_KEY from scripts.seed_tracing_fixtures import ( JSON, TRACE_FIXTURES, bulk_span_rows, fixture_capture, fixture_replays, + managed_response, postgres_row, rebase, rebase_spend, @@ -33,6 +34,7 @@ from scripts.seed_tracing_fixtures import ( spend_fixtures, timestamps, ) +from tests._master_key import MASTER_KEY CALL_KEYS: Final = TypeAdapter(tuple[str, ...]) DATETIMES: Final = TypeAdapter(tuple[datetime, datetime]) @@ -77,9 +79,12 @@ def test_all_fixture_replays_are_recent_and_preserve_spans(path: Path) -> None: assert after_span_attributes["lens.original_trace_id"] == seed_id( before_original_trace_id, replay.namespace, 32 ) - assert after["TraceId"] == hashlib.sha256( - f"litellm.claude.session.v1\0{seed_id(before_session, replay.namespace, 32)}".encode() - ).hexdigest()[:32] + assert ( + after["TraceId"] + == hashlib.sha256( + f"litellm.claude.session.v1\0{seed_id(before_session, replay.namespace, 32)}".encode() + ).hexdigest()[:32] + ) assert after["TraceId"] != before["TraceId"] else: assert after["TraceId"] == seed_id(trace_id, replay.namespace, 32) @@ -255,3 +260,27 @@ def test_seed_cli_rejects_invalid_http_timeouts(timeout: str) -> None: with pytest.raises(SystemExit) as error: seed_arguments(["--timeout-seconds", timeout]) assert error.value.code == 2 + + +def test_replay_rebases_provider_request_evidence_with_the_spend_row() -> None: + spend: Final = {**dict(spend_fixtures())["deepagents_swarm"][0], "provider_request_id": "req_replay"} + pattern: Final = response_pattern((spend,)) + replayed: Final = rebase_spend((spend,), 0, "another-run", pattern)[0] + export: Final = rebase({"request_id": "req_replay"}, 0, "another-run", pattern) + + assert export == {"request_id": replayed["provider_request_id"]} + assert replayed["provider_request_id"] != spend["provider_request_id"] + assert replayed["provider_request_id"].startswith("req_") + + +@pytest.mark.parametrize("upstream", ("msg_response", "req_response", "chatcmpl-response")) +def test_replay_preserves_identity_between_managed_and_plain_response_ids(upstream: str) -> None: + payload: Final = f"model:example;response_id:{upstream}" + managed: Final = "resp_" + base64.b64encode(payload.encode()).decode() + spend: Final = {**dict(spend_fixtures())["deepagents_swarm"][0], "response_id": managed} + pattern: Final = response_pattern((spend,)) + replayed: Final = rebase_spend((spend,), 0, "managed-replay", pattern)[0] + plain: Final = rebase(upstream, 0, "managed-replay", pattern) + + assert plain != upstream + assert managed_response(replayed["response_id"]) == f"model:example;response_id:{plain}"