mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-07 02:59:05 +00:00
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:
parent
fe29c55d11
commit
a3e74a5223
5 changed files with 69 additions and 25 deletions
|
|
@ -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(())
|
||||
}
|
||||
|
||||
|
|
|
|||
|
|
@ -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]
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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}"
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue