mirror of
https://github.com/BerriAI/litellm.git
synced 2026-08-28 05:25:59 +00:00
fix(langfuse): coerce header-sourced mask and trace-update steering values (#36740)
langfuse_* request headers land in metadata as strings, but the trace path reads mask_input/mask_output with a bare truthiness check and iterates update_trace_keys directly. A header saying mask_input: false redacted the payload it was asked to keep, and update_trace_keys was walked one character at a time so every requested key silently failed to match
This commit is contained in:
parent
a7397b2459
commit
09889e1986
2 changed files with 153 additions and 4 deletions
|
|
@ -2,7 +2,7 @@
|
|||
# On success, logs events to Langfuse
|
||||
import os
|
||||
import traceback
|
||||
from collections.abc import Callable
|
||||
from collections.abc import Callable, Iterable
|
||||
from datetime import datetime
|
||||
from typing import TYPE_CHECKING, Any, Final, cast
|
||||
|
||||
|
|
@ -75,6 +75,22 @@ def _extract_cache_read_input_tokens(usage_obj) -> int:
|
|||
return cache_read_input_tokens
|
||||
|
||||
|
||||
def _as_steering_flag(value: object) -> bool:
|
||||
"""A string ``str_to_bool`` does not recognise falls back to its truthiness."""
|
||||
if isinstance(value, str):
|
||||
parsed: Final = str_to_bool(value)
|
||||
return bool(value) if parsed is None else parsed
|
||||
return bool(value)
|
||||
|
||||
|
||||
def _as_steering_key_sequence(value: object) -> tuple[str, ...]:
|
||||
if isinstance(value, str):
|
||||
return tuple(key.strip() for key in value.split(",") if key.strip())
|
||||
if isinstance(value, Iterable):
|
||||
return tuple(str(key) for key in value)
|
||||
return ()
|
||||
|
||||
|
||||
def resolve_langfuse_credentials(
|
||||
langfuse_public_key=None,
|
||||
langfuse_secret=None,
|
||||
|
|
@ -552,10 +568,10 @@ class LangFuseLogger:
|
|||
# This allows continuing an existing trace while still returning the correct trace_id
|
||||
if existing_trace_id is not None:
|
||||
trace_id = existing_trace_id
|
||||
update_trace_keys: Final = cast(list, clean_metadata.pop("update_trace_keys", []))
|
||||
update_trace_keys: Final = _as_steering_key_sequence(clean_metadata.pop("update_trace_keys", ()))
|
||||
debug: Final = clean_metadata.pop("debug_langfuse", None)
|
||||
mask_input: Final = clean_metadata.pop("mask_input", False)
|
||||
mask_output: Final = clean_metadata.pop("mask_output", False)
|
||||
mask_input: Final = _as_steering_flag(clean_metadata.pop("mask_input", False))
|
||||
mask_output: Final = _as_steering_flag(clean_metadata.pop("mask_output", False))
|
||||
# Look for masking function in the dedicated location first (set by scrub_sensitive_keys_in_metadata)
|
||||
# Fall back to metadata for backwards compatibility
|
||||
masking_function: Final = litellm_params.get("_langfuse_masking_function") or clean_metadata.pop(
|
||||
|
|
|
|||
|
|
@ -994,3 +994,136 @@ def test_langfuse_logger_reuses_the_shared_cached_client(monkeypatch):
|
|||
gc.collect()
|
||||
|
||||
assert not first.langfuse_client.is_closed
|
||||
|
||||
|
||||
_LANGFUSE_REDACTED = "redacted-by-litellm"
|
||||
|
||||
|
||||
def _steering_logger() -> LangFuseLogger:
|
||||
"""``__new__`` skips the SDK and network setup in ``__init__``."""
|
||||
logger = LangFuseLogger.__new__(LangFuseLogger)
|
||||
logger.Langfuse = MagicMock()
|
||||
logger.langfuse_sdk_version = "2.60.0"
|
||||
return logger
|
||||
|
||||
|
||||
def _emit(logger: LangFuseLogger, *, metadata=None, headers=None):
|
||||
"""``log_event_on_langfuse`` is the entry point that folds ``langfuse_*`` headers into metadata."""
|
||||
now = datetime.datetime.now()
|
||||
response_obj = litellm.ModelResponse(
|
||||
choices=[{"message": {"role": "assistant", "content": "the-output"}}]
|
||||
)
|
||||
logger.log_event_on_langfuse(
|
||||
kwargs={
|
||||
"call_type": "completion",
|
||||
"litellm_params": {
|
||||
"metadata": dict(metadata or {}),
|
||||
"proxy_server_request": {"headers": dict(headers or {})},
|
||||
},
|
||||
"messages": [{"role": "user", "content": "the-input"}],
|
||||
"optional_params": {},
|
||||
},
|
||||
response_obj=response_obj,
|
||||
start_time=now,
|
||||
end_time=now,
|
||||
)
|
||||
return (
|
||||
logger.Langfuse.trace.call_args.kwargs,
|
||||
logger.Langfuse.trace.return_value.generation.call_args.kwargs,
|
||||
)
|
||||
|
||||
|
||||
def test_mask_input_header_false_keeps_the_prompt():
|
||||
logger = _steering_logger()
|
||||
|
||||
trace_params, generation_params = _emit(logger, headers={"langfuse_mask_input": "false"})
|
||||
|
||||
assert trace_params["input"] == {"messages": [{"role": "user", "content": "the-input"}]}
|
||||
assert generation_params["input"] == {"messages": [{"role": "user", "content": "the-input"}]}
|
||||
|
||||
|
||||
def test_mask_input_header_true_redacts_the_prompt():
|
||||
logger = _steering_logger()
|
||||
|
||||
trace_params, generation_params = _emit(logger, headers={"langfuse_mask_input": "true"})
|
||||
|
||||
assert trace_params["input"] == _LANGFUSE_REDACTED
|
||||
assert generation_params["input"] == _LANGFUSE_REDACTED
|
||||
|
||||
|
||||
def test_mask_output_header_false_keeps_the_completion():
|
||||
logger = _steering_logger()
|
||||
|
||||
trace_params, generation_params = _emit(logger, headers={"langfuse_mask_output": "false"})
|
||||
|
||||
assert trace_params["output"] != _LANGFUSE_REDACTED
|
||||
assert generation_params["output"] != _LANGFUSE_REDACTED
|
||||
|
||||
|
||||
def test_mask_output_header_true_redacts_the_completion():
|
||||
logger = _steering_logger()
|
||||
|
||||
trace_params, generation_params = _emit(logger, headers={"langfuse_mask_output": "true"})
|
||||
|
||||
assert trace_params["output"] == _LANGFUSE_REDACTED
|
||||
assert generation_params["output"] == _LANGFUSE_REDACTED
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"mask_input, expect_redacted",
|
||||
[
|
||||
(False, False),
|
||||
(True, True),
|
||||
# An unrecognised string keeps its truthiness, so existing behaviour is unchanged
|
||||
("yes", True),
|
||||
],
|
||||
)
|
||||
def test_mask_input_from_the_request_body_is_unchanged(mask_input, expect_redacted):
|
||||
logger = _steering_logger()
|
||||
|
||||
trace_params, _ = _emit(logger, metadata={"mask_input": mask_input})
|
||||
|
||||
assert (trace_params["input"] == _LANGFUSE_REDACTED) is expect_redacted
|
||||
|
||||
|
||||
def test_update_trace_keys_header_applies_every_key():
|
||||
logger = _steering_logger()
|
||||
|
||||
trace_params, _ = _emit(
|
||||
logger,
|
||||
headers={
|
||||
"langfuse_existing_trace_id": "trace-1",
|
||||
"langfuse_update_trace_keys": "trace_release, trace_tail",
|
||||
"langfuse_trace_release": "v1.2.3",
|
||||
"langfuse_trace_tail": "last",
|
||||
},
|
||||
)
|
||||
|
||||
assert trace_params["release"] == "v1.2.3"
|
||||
assert trace_params["tail"] == "last"
|
||||
|
||||
|
||||
def test_update_trace_keys_from_the_request_body_list_is_unchanged():
|
||||
logger = _steering_logger()
|
||||
|
||||
trace_params, _ = _emit(
|
||||
logger,
|
||||
metadata={
|
||||
"existing_trace_id": "trace-1",
|
||||
"update_trace_keys": ["trace_release"],
|
||||
"trace_release": "v1.2.3",
|
||||
},
|
||||
)
|
||||
|
||||
assert trace_params["release"] == "v1.2.3"
|
||||
|
||||
|
||||
def test_update_trace_keys_matches_whole_keys_not_substrings():
|
||||
logger = _steering_logger()
|
||||
|
||||
trace_params, _ = _emit(
|
||||
logger,
|
||||
headers={"langfuse_existing_trace_id": "trace-1", "langfuse_update_trace_keys": "my_input"},
|
||||
)
|
||||
|
||||
assert "input" not in trace_params
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue