mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-08 03:08:45 +00:00
fix(langfuse): normalize non-string observation ids
This commit is contained in:
parent
b08598534c
commit
36f1af83cc
2 changed files with 35 additions and 15 deletions
|
|
@ -54,25 +54,20 @@ def to_unix_nanos(value: datetime | float | None) -> int | None:
|
|||
return int(seconds * 1_000_000_000)
|
||||
|
||||
|
||||
def resolve_trace_id(trace_id: str | None) -> str:
|
||||
"""Map litellm's trace id onto the 32 lowercase hex characters v4 requires.
|
||||
|
||||
Anything else raises inside the SDK rather than being ignored, so a plain
|
||||
uuid is dash-stripped and any other identifier is hashed deterministically,
|
||||
which keeps repeat calls with the same id on the same trace.
|
||||
"""
|
||||
normalized: Final = trace_id.lower().replace("-", "") if trace_id else ""
|
||||
if _TRACE_ID_PATTERN.match(normalized):
|
||||
def resolve_trace_id(trace_id: object | None) -> str:
|
||||
"""Map a caller's trace id onto the 32 lowercase hex characters v4 requires."""
|
||||
normalized: Final = "" if trace_id is None else str(trace_id).lower().replace("-", "")
|
||||
if _TRACE_ID_PATTERN.fullmatch(normalized):
|
||||
return normalized
|
||||
return Langfuse.create_trace_id(seed=trace_id) if trace_id else Langfuse.create_trace_id()
|
||||
return Langfuse.create_trace_id(seed=str(trace_id)) if normalized else Langfuse.create_trace_id()
|
||||
|
||||
|
||||
def resolve_observation_id(observation_id: str | None) -> str | None:
|
||||
"""Same for a caller-supplied parent, which v4 requires to be 16 hex characters."""
|
||||
normalized: Final = observation_id.lower().replace("-", "") if observation_id else ""
|
||||
def resolve_observation_id(observation_id: object | None) -> str | None:
|
||||
"""Map a caller's parent observation id onto v4's 16 lowercase hex characters."""
|
||||
normalized: Final = "" if observation_id is None else str(observation_id).lower().replace("-", "")
|
||||
if not normalized:
|
||||
return None
|
||||
if _OBSERVATION_ID_PATTERN.match(normalized):
|
||||
if _OBSERVATION_ID_PATTERN.fullmatch(normalized):
|
||||
return normalized
|
||||
return sha256(normalized.encode("utf-8")).digest()[:8].hex()
|
||||
|
||||
|
|
|
|||
|
|
@ -18,8 +18,8 @@ from opentelemetry.sdk.trace.export.in_memory_span_exporter import InMemorySpanE
|
|||
|
||||
from litellm.integrations.langfuse.langfuse import (
|
||||
MINIMUM_LANGFUSE_VERSION,
|
||||
raise_if_unsupported_langfuse_version,
|
||||
installed_langfuse_version,
|
||||
raise_if_unsupported_langfuse_version,
|
||||
)
|
||||
from litellm.integrations.langfuse.langfuse_sdk import (
|
||||
AS_ROOT_ATTRIBUTE,
|
||||
|
|
@ -251,6 +251,31 @@ def test_arbitrary_trace_id_is_hashed_deterministically():
|
|||
assert first != resolve_trace_id("order-4472")
|
||||
|
||||
|
||||
@pytest.mark.parametrize("supplied", [12345, 12.5, True, False], ids=["int", "float", "true", "false"])
|
||||
def test_non_string_trace_id_is_normalized(supplied):
|
||||
resolved = resolve_trace_id(supplied)
|
||||
|
||||
assert len(resolved) == 32
|
||||
assert resolved == resolve_trace_id(supplied)
|
||||
|
||||
|
||||
@pytest.mark.parametrize("supplied", [12345, 12.5, True, False], ids=["int", "float", "true", "false"])
|
||||
def test_non_string_observation_id_is_normalized(supplied):
|
||||
resolved = resolve_observation_id(supplied)
|
||||
|
||||
assert len(resolved) == 16
|
||||
assert resolved == resolve_observation_id(supplied)
|
||||
|
||||
|
||||
def test_trace_id_with_trailing_newline_is_hashed():
|
||||
supplied = "a" * 32 + "\n"
|
||||
|
||||
resolved = resolve_trace_id(supplied)
|
||||
|
||||
assert resolved != supplied
|
||||
assert len(resolved) == 32
|
||||
|
||||
|
||||
def test_missing_trace_id_still_yields_a_valid_trace_id():
|
||||
generated = resolve_trace_id(None)
|
||||
assert len(generated) == 32
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue