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:
yucheng-berri 2026-08-13 00:44:26 -07:00 committed by GitHub
parent a7397b2459
commit 09889e1986
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
2 changed files with 153 additions and 4 deletions

View file

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

View file

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