mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-15 23:31:29 +00:00
fix(langfuse): preserve v4 observation hierarchy
This commit is contained in:
parent
620c1b4d37
commit
2c1dfdb4f0
5 changed files with 82 additions and 55 deletions
|
|
@ -18,6 +18,7 @@ from typing import (
|
|||
)
|
||||
|
||||
from packaging.version import Version
|
||||
from opentelemetry import trace as otel_trace
|
||||
|
||||
import litellm
|
||||
from litellm._logging import verbose_logger
|
||||
|
|
@ -713,6 +714,8 @@ class LangFuseLogger:
|
|||
)
|
||||
trace_version = trace_params.get("version")
|
||||
propagated_version = trace_version if isinstance(trace_version, str) else None
|
||||
trace_release = trace_params.get("release")
|
||||
propagated_release = trace_release if isinstance(trace_release, str) else None
|
||||
propagated_tags = tags if tags else None
|
||||
trace_name_value = trace_params.get("name")
|
||||
propagated_trace_name = (
|
||||
|
|
@ -726,24 +729,21 @@ class LangFuseLogger:
|
|||
tags=propagated_tags,
|
||||
trace_name=propagated_trace_name,
|
||||
):
|
||||
trace = self.Langfuse.start_observation(
|
||||
with self.Langfuse.start_as_current_observation(
|
||||
trace_context=trace_context,
|
||||
name=propagated_trace_name or generation_name,
|
||||
input=trace_params.get("input"),
|
||||
output=trace_params.get("output"),
|
||||
metadata=propagated_metadata,
|
||||
version=propagated_version,
|
||||
level=level,
|
||||
status_message=trace_params.get("status_message"),
|
||||
)
|
||||
log_provider_specific_information_as_span(trace, clean_metadata)
|
||||
self._log_guardrail_information_as_span(
|
||||
trace=trace,
|
||||
standard_logging_object=standard_logging_object,
|
||||
)
|
||||
generation_client = trace.start_observation(**generation_params)
|
||||
generation_client.end()
|
||||
trace.end()
|
||||
**generation_params,
|
||||
) as generation_client:
|
||||
_set_langfuse_release(propagated_release)
|
||||
log_provider_specific_information_as_span(
|
||||
generation_client,
|
||||
clean_metadata,
|
||||
propagated_release,
|
||||
)
|
||||
self._log_guardrail_information_as_span(
|
||||
trace=generation_client,
|
||||
standard_logging_object=standard_logging_object,
|
||||
release=propagated_release,
|
||||
)
|
||||
|
||||
return trace_context["trace_id"], generation_client.id
|
||||
except Exception:
|
||||
|
|
@ -920,6 +920,7 @@ class LangFuseLogger:
|
|||
self,
|
||||
trace: "LangfuseSpan",
|
||||
standard_logging_object: Optional[StandardLoggingPayload],
|
||||
release: str | None = None,
|
||||
):
|
||||
"""
|
||||
Log guardrail information as a span
|
||||
|
|
@ -948,7 +949,7 @@ class LangFuseLogger:
|
|||
)
|
||||
continue
|
||||
|
||||
span = trace.start_observation(
|
||||
with trace.start_as_current_observation(
|
||||
name="guardrail",
|
||||
as_type="guardrail",
|
||||
input=guardrail_entry.get("guardrail_request", None),
|
||||
|
|
@ -958,10 +959,9 @@ class LangFuseLogger:
|
|||
"guardrail_mode": guardrail_entry.get("guardrail_mode", None),
|
||||
"guardrail_masked_entity_count": guardrail_entry.get("masked_entity_count", None),
|
||||
},
|
||||
)
|
||||
|
||||
verbose_logger.debug(f"Logged guardrail information as span: {span}")
|
||||
span.end()
|
||||
) as span:
|
||||
_set_langfuse_release(release)
|
||||
verbose_logger.debug(f"Logged guardrail information as span: {span}")
|
||||
|
||||
|
||||
def _add_prompt_to_generation_params(
|
||||
|
|
@ -1041,6 +1041,7 @@ def _add_prompt_to_generation_params(
|
|||
def log_provider_specific_information_as_span(
|
||||
trace: "LangfuseSpan",
|
||||
clean_metadata,
|
||||
release: str | None = None,
|
||||
):
|
||||
"""
|
||||
Logs provider-specific information as spans.
|
||||
|
|
@ -1064,23 +1065,28 @@ def log_provider_specific_information_as_span(
|
|||
for elem in vertex_ai_grounding_metadata:
|
||||
if isinstance(elem, dict):
|
||||
for key, value in elem.items():
|
||||
span = trace.start_observation(
|
||||
with trace.start_as_current_observation(
|
||||
name=key,
|
||||
input=value,
|
||||
)
|
||||
span.end()
|
||||
):
|
||||
_set_langfuse_release(release)
|
||||
else:
|
||||
span = trace.start_observation(
|
||||
with trace.start_as_current_observation(
|
||||
name="vertex_ai_grounding_metadata",
|
||||
input=elem,
|
||||
)
|
||||
span.end()
|
||||
):
|
||||
_set_langfuse_release(release)
|
||||
else:
|
||||
span = trace.start_observation(
|
||||
with trace.start_as_current_observation(
|
||||
name="vertex_ai_grounding_metadata",
|
||||
input=vertex_ai_grounding_metadata,
|
||||
)
|
||||
span.end()
|
||||
):
|
||||
_set_langfuse_release(release)
|
||||
|
||||
|
||||
def _set_langfuse_release(release: str | None) -> None:
|
||||
if release:
|
||||
otel_trace.get_current_span().set_attribute("langfuse.release", release)
|
||||
|
||||
|
||||
def log_requester_metadata(clean_metadata: dict):
|
||||
|
|
|
|||
|
|
@ -105,7 +105,6 @@ class LangfuseOtelLogger(OpenTelemetry):
|
|||
"generation_name": LangfuseSpanAttributes.GENERATION_NAME,
|
||||
"generation_id": LangfuseSpanAttributes.GENERATION_ID,
|
||||
"parent_observation_id": LangfuseSpanAttributes.PARENT_OBSERVATION_ID,
|
||||
"version": LangfuseSpanAttributes.GENERATION_VERSION,
|
||||
"mask_input": LangfuseSpanAttributes.MASK_INPUT,
|
||||
"mask_output": LangfuseSpanAttributes.MASK_OUTPUT,
|
||||
"user_id": LangfuseSpanAttributes.TRACE_USER_ID,
|
||||
|
|
@ -114,7 +113,6 @@ class LangfuseOtelLogger(OpenTelemetry):
|
|||
"tags": LangfuseSpanAttributes.TAGS,
|
||||
"trace_name": LangfuseSpanAttributes.TRACE_NAME,
|
||||
"trace_id": LangfuseSpanAttributes.TRACE_ID,
|
||||
"trace_version": LangfuseSpanAttributes.TRACE_VERSION,
|
||||
"trace_release": LangfuseSpanAttributes.TRACE_RELEASE,
|
||||
"existing_trace_id": LangfuseSpanAttributes.EXISTING_TRACE_ID,
|
||||
"update_trace_keys": LangfuseSpanAttributes.UPDATE_TRACE_KEYS,
|
||||
|
|
@ -137,6 +135,14 @@ class LangfuseOtelLogger(OpenTelemetry):
|
|||
trace_metadata,
|
||||
)
|
||||
|
||||
version = metadata.get("version") if metadata.get("version") is not None else metadata.get("trace_version")
|
||||
if version is not None:
|
||||
safe_set_attribute(
|
||||
span,
|
||||
LangfuseSpanAttributes.GENERATION_VERSION.value,
|
||||
version,
|
||||
)
|
||||
|
||||
for key, enum_attr in mapping.items():
|
||||
if key in metadata and metadata[key] is not None:
|
||||
value = metadata[key]
|
||||
|
|
|
|||
|
|
@ -306,9 +306,16 @@ def test_langfuse_v4_observations_propagate_trace_attributes():
|
|||
metadata={
|
||||
"session_id": "session-123",
|
||||
"trace_id": "a" * 32,
|
||||
"parent_observation_id": "b" * 16,
|
||||
"trace_name": "completion-trace",
|
||||
"trace_version": "v4",
|
||||
"trace_release": "release-123",
|
||||
"trace_metadata": {"request_type": "completion"},
|
||||
"hidden_params": {
|
||||
"vertex_ai_grounding_metadata": {
|
||||
"grounding": "enabled",
|
||||
}
|
||||
},
|
||||
},
|
||||
litellm_params={"metadata": {}},
|
||||
output={"role": "assistant", "content": "Hello"},
|
||||
|
|
@ -332,11 +339,16 @@ def test_langfuse_v4_observations_propagate_trace_attributes():
|
|||
assert trace_id == "a" * 32
|
||||
assert generation_id is not None
|
||||
assert len(spans) == 2
|
||||
generation_span = next(span for span in spans if span.attributes["langfuse.observation.type"] == "generation")
|
||||
provider_span = next(span for span in spans if span is not generation_span)
|
||||
assert generation_span.parent.span_id == int("b" * 16, 16)
|
||||
assert provider_span.parent.span_id == generation_span.context.span_id
|
||||
for span in spans:
|
||||
assert span.attributes["user.id"] == "user-123"
|
||||
assert span.attributes["session.id"] == "session-123"
|
||||
assert span.attributes["langfuse.trace.name"] == "completion-trace"
|
||||
assert span.attributes["langfuse.version"] == "v4"
|
||||
assert span.attributes["langfuse.release"] == "release-123"
|
||||
assert span.attributes["langfuse.trace.metadata.request_type"] == "completion"
|
||||
|
||||
|
||||
|
|
@ -375,7 +387,7 @@ def test_langfuse_v4_observations_do_not_use_historical_end_times():
|
|||
|
||||
langfuse_client.flush()
|
||||
spans = span_exporter.get_finished_spans()
|
||||
assert len(spans) == 2
|
||||
assert len(spans) == 1
|
||||
assert all(span.end_time >= span.start_time for span in spans)
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -40,30 +40,25 @@ class TestLangfuseUsageDetails(unittest.TestCase):
|
|||
self.mock_langfuse_client = MagicMock()
|
||||
# Mock the client attribute to prevent errors during logger initialization
|
||||
self.mock_langfuse_client.client = MagicMock()
|
||||
self.mock_langfuse_trace = MagicMock()
|
||||
self.mock_langfuse_generation = MagicMock()
|
||||
self.mock_langfuse_generation.trace_id = "test-trace-id"
|
||||
self.mock_langfuse_generation.id = "test-generation-id"
|
||||
self.mock_langfuse_trace.trace_id = "test-trace-id"
|
||||
|
||||
# Mock span method for trace (used by log_provider_specific_information_as_span and _log_guardrail_information_as_span)
|
||||
self.mock_langfuse_span = MagicMock()
|
||||
self.mock_langfuse_span.end = MagicMock()
|
||||
self.mock_langfuse_trace.start_observation.return_value = self.mock_langfuse_generation
|
||||
self.mock_generation_context = MagicMock()
|
||||
self.mock_generation_context.__enter__.return_value = self.mock_langfuse_generation
|
||||
|
||||
# Setup the trace and generation chain
|
||||
self.last_trace_kwargs = {}
|
||||
|
||||
def _trace_side_effect(*args, **kwargs):
|
||||
def _observation_side_effect(*args, **kwargs):
|
||||
propagated = self.mock_langfuse.propagate_attributes.call_args.kwargs
|
||||
self.last_trace_kwargs = {
|
||||
**kwargs,
|
||||
"id": kwargs["trace_context"]["trace_id"],
|
||||
"session_id": propagated.get("session_id"),
|
||||
}
|
||||
return self.mock_langfuse_trace
|
||||
return self.mock_generation_context
|
||||
|
||||
self.mock_langfuse_client.start_observation.side_effect = _trace_side_effect
|
||||
self.mock_langfuse_client.start_as_current_observation.side_effect = _observation_side_effect
|
||||
self.mock_langfuse_client.create_trace_id.side_effect = lambda seed=None: (
|
||||
seed or "00000000000000000000000000000001"
|
||||
)
|
||||
|
|
@ -300,17 +295,14 @@ class TestLangfuseUsageDetails(unittest.TestCase):
|
|||
"""
|
||||
# Reset the mock to ensure clean state; clear side_effect so return_value takes effect
|
||||
self.mock_langfuse_client.reset_mock(side_effect=True)
|
||||
self.mock_langfuse_trace.reset_mock(side_effect=True)
|
||||
self.mock_langfuse_generation.reset_mock(side_effect=True)
|
||||
|
||||
# Re-setup the trace and generation chain with clean state
|
||||
self.mock_langfuse_generation.id = "test-generation-id"
|
||||
self.mock_langfuse_trace.trace_id = "test-trace-id"
|
||||
mock_span = MagicMock()
|
||||
mock_span.end = MagicMock()
|
||||
self.mock_langfuse_trace.start_observation.return_value = self.mock_langfuse_generation
|
||||
self.mock_generation_context = MagicMock()
|
||||
self.mock_generation_context.__enter__.return_value = self.mock_langfuse_generation
|
||||
|
||||
self.mock_langfuse_client.start_observation.return_value = self.mock_langfuse_trace
|
||||
self.mock_langfuse_client.start_as_current_observation.return_value = self.mock_generation_context
|
||||
self.mock_langfuse_client.create_trace_id.side_effect = lambda seed=None: (
|
||||
seed or "00000000000000000000000000000001"
|
||||
)
|
||||
|
|
@ -372,10 +364,8 @@ class TestLangfuseUsageDetails(unittest.TestCase):
|
|||
except Exception as e:
|
||||
self.fail(f"_log_langfuse_v2 raised an exception: {e}")
|
||||
|
||||
self.mock_langfuse_client.start_observation.assert_called()
|
||||
|
||||
self.mock_langfuse_trace.start_observation.assert_called_once()
|
||||
call_args, call_kwargs = self.mock_langfuse_trace.start_observation.call_args
|
||||
self.mock_langfuse_client.start_as_current_observation.assert_called()
|
||||
call_args, call_kwargs = self.mock_langfuse_client.start_as_current_observation.call_args
|
||||
|
||||
usage_details_arg = call_kwargs.get("usage_details")
|
||||
|
||||
|
|
|
|||
|
|
@ -210,7 +210,6 @@ class TestLangfuseOtelIntegration:
|
|||
LangfuseSpanAttributes.TRACE_NAME.value: "trace-name",
|
||||
LangfuseSpanAttributes.TRACE_ID.value: "traceid", # stripped dashes
|
||||
f"{LangfuseSpanAttributes.TRACE_METADATA.value}.k": "v",
|
||||
LangfuseSpanAttributes.TRACE_VERSION.value: "t-ver",
|
||||
LangfuseSpanAttributes.TRACE_RELEASE.value: "rel-1",
|
||||
LangfuseSpanAttributes.EXISTING_TRACE_ID.value: "existing-id",
|
||||
LangfuseSpanAttributes.UPDATE_TRACE_KEYS.value: json.dumps(["key1", "key2"]),
|
||||
|
|
@ -225,6 +224,20 @@ class TestLangfuseOtelIntegration:
|
|||
|
||||
assert actual == expected, "Mismatch between expected and actual OTEL attribute mapping."
|
||||
|
||||
def test_trace_version_is_used_when_generation_version_is_missing(self):
|
||||
from litellm.types.integrations.langfuse_otel import LangfuseSpanAttributes
|
||||
|
||||
span = MagicMock()
|
||||
|
||||
with patch("litellm.integrations.arize._utils.safe_set_attribute") as mock_safe_set_attribute:
|
||||
LangfuseOtelLogger._set_metadata_attributes(span, {"trace_version": "trace-v"})
|
||||
|
||||
mock_safe_set_attribute.assert_called_once_with(
|
||||
span,
|
||||
LangfuseSpanAttributes.GENERATION_VERSION.value,
|
||||
"trace-v",
|
||||
)
|
||||
|
||||
@pytest.mark.parametrize("metadata_key", ["user_id", "trace_user_id"])
|
||||
def test_set_langfuse_user_id_attribute(self, metadata_key):
|
||||
from litellm.types.integrations.langfuse_otel import LangfuseSpanAttributes
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue