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
This commit is contained in:
moe-berri 2026-10-06 14:00:40 -07:00 • committed by GitHub
parent fe29c55d11
commit a3e74a5223
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
5 changed files with 69 additions and 25 deletions

View file

@ -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<SeededDatabase>,
#[case] id: &str,
#[case] expected: Option<f64>,
) -> 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(())
}

View file

@ -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]

View file

@ -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,

View file

@ -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,

View file

@ -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}"