mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-07 02:59:05 +00:00
fix(proxy): generated session uses promoted caller trace id
Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
parent
5baeb9e80b
commit
d656a4f18b
3 changed files with 109 additions and 14 deletions
|
|
@ -150,17 +150,17 @@ def _session_id_from_baggage(baggage: str) -> str | None:
|
|||
return None
|
||||
|
||||
|
||||
def _caller_set_trace_field(data: Mapping[str, object], metadata_variable_name: str, field: str) -> bool:
|
||||
def _caller_trace_field(data: Mapping[str, object], metadata_variable_name: str, field: str) -> object | None:
|
||||
active: Final = data.get(metadata_variable_name)
|
||||
if isinstance(active, Mapping) and field in active:
|
||||
active_map: Final = cast(Mapping[str, object], active) # cast-ok: isinstance above, free-form JSON values
|
||||
return bool(active_map[field])
|
||||
return active_map[field] or None
|
||||
promoted: Final = metadata_variable_name == "litellm_metadata" and field in LITELLM_TRACE_CONTROL_METADATA_FIELDS
|
||||
requester: Final = data.get("metadata")
|
||||
if not promoted or not isinstance(requester, Mapping):
|
||||
return False
|
||||
return None
|
||||
requester_map: Final = cast(Mapping[str, object], requester) # cast-ok: isinstance above, free-form JSON values
|
||||
return bool(requester_map.get(field))
|
||||
return requester_map.get(field) or None
|
||||
|
||||
|
||||
def _stampable_key_hash(user_api_key_dict: UserAPIKeyAuth) -> str | None:
|
||||
|
|
@ -841,11 +841,15 @@ def apply_missing_session_id_policy(
|
|||
):
|
||||
metadata["session_id"] = body_session_id
|
||||
return
|
||||
if data.get("litellm_session_id") or _caller_set_trace_field(data, _metadata_variable_name, "session_id"):
|
||||
if data.get("litellm_session_id") or _caller_trace_field(data, _metadata_variable_name, "session_id") is not None:
|
||||
return
|
||||
match policy:
|
||||
case "generate":
|
||||
session_id: Final = str(data.get("litellm_trace_id") or metadata.get("trace_id") or uuid.uuid4())
|
||||
session_id: Final = str(
|
||||
data.get("litellm_trace_id")
|
||||
or _caller_trace_field(data, _metadata_variable_name, "trace_id")
|
||||
or uuid.uuid4()
|
||||
)
|
||||
data["litellm_session_id"] = session_id # rebind-ok: data is an out-param
|
||||
data.setdefault("litellm_trace_id", session_id)
|
||||
metadata["session_id"] = session_id
|
||||
|
|
@ -1598,10 +1602,14 @@ class LiteLLMProxyRequestSetup:
|
|||
# correlate with litellm's own logs instead of generating an unrelated
|
||||
# trace_id.
|
||||
normalized_headers: Final = MappingProxyType({k.lower(): v for k, v in headers.items() if isinstance(k, str)})
|
||||
if "litellm_trace_id" not in data and not _caller_set_trace_field(
|
||||
cast(Mapping[str, object], data), # cast-ok: request body is a str-keyed JSON object
|
||||
_metadata_variable_name,
|
||||
"trace_id",
|
||||
if (
|
||||
"litellm_trace_id" not in data
|
||||
and _caller_trace_field(
|
||||
cast(Mapping[str, object], data), # cast-ok: request body is a str-keyed JSON object
|
||||
_metadata_variable_name,
|
||||
"trace_id",
|
||||
)
|
||||
is None
|
||||
):
|
||||
traceparent: Final = normalized_headers.get("traceparent")
|
||||
if isinstance(traceparent, str):
|
||||
|
|
@ -1612,10 +1620,14 @@ class LiteLLMProxyRequestSetup:
|
|||
verbose_proxy_logger.debug(
|
||||
"Extracted trace_id from W3C traceparent header: %s", trace_id_from_traceparent
|
||||
)
|
||||
if "litellm_session_id" not in data and not _caller_set_trace_field(
|
||||
cast(Mapping[str, object], data), # cast-ok: request body is a str-keyed JSON object
|
||||
_metadata_variable_name,
|
||||
"session_id",
|
||||
if (
|
||||
"litellm_session_id" not in data
|
||||
and _caller_trace_field(
|
||||
cast(Mapping[str, object], data), # cast-ok: request body is a str-keyed JSON object
|
||||
_metadata_variable_name,
|
||||
"session_id",
|
||||
)
|
||||
is None
|
||||
):
|
||||
baggage: Final = normalized_headers.get("baggage")
|
||||
if isinstance(baggage, str):
|
||||
|
|
|
|||
|
|
@ -602,3 +602,56 @@ def test_missing_session_id_reject_accepts_caller_metadata_and_baggage_fallback(
|
|||
)
|
||||
assert third.status_code == 400, third.text
|
||||
assert upstream_targets == ["/v1/responses", "/v1/chat/completions"], upstream_targets
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
("endpoint", "kind", "expected_target"),
|
||||
(
|
||||
pytest.param("/v1/responses", "responses", "/v1/responses", id="responses_caller_trace"),
|
||||
pytest.param("/v1/messages", "messages", "/v1/responses", id="messages_caller_trace"),
|
||||
),
|
||||
)
|
||||
def test_missing_session_id_generate_derives_session_from_caller_trace(
|
||||
gateway: Gateway, tmp_path: Path, endpoint: str, kind: str, expected_target: str
|
||||
) -> None:
|
||||
marker: Final = "gen" + uuid.uuid4().hex
|
||||
provider_secret: Final = "synthetic-provider-secret-" + marker
|
||||
caller_trace: Final = uuid.uuid4().hex
|
||||
upstream_targets: Final[list[str]] = [] # mutable-ok: records which upstream endpoint each call hit
|
||||
|
||||
def upstream(request: Request) -> Reply:
|
||||
assert request.headers["authorization"] == f"Bearer {provider_secret}"
|
||||
upstream_targets.append(request.target)
|
||||
if request.target == "/v1/responses":
|
||||
return _responses_result("resp-" + marker)
|
||||
assert request.target == "/v1/chat/completions", request.target
|
||||
return _completion(marker + "-answer")
|
||||
|
||||
def langfuse(request: Request) -> Reply:
|
||||
if request.method == "GET" and request.target.startswith(PROJECTS_PATH):
|
||||
return _projects()
|
||||
return Reply(body=b"", content_type="application/x-protobuf")
|
||||
|
||||
with (
|
||||
wire_server(upstream) as provider,
|
||||
wire_server(langfuse) as destination,
|
||||
owned_proxy(
|
||||
gateway,
|
||||
tmp_path,
|
||||
_langfuse_environment(destination),
|
||||
config=_langfuse_config(tmp_path, {"missing_session_id": "generate"}),
|
||||
) as candidate,
|
||||
candidate.scenario() as scenario,
|
||||
):
|
||||
model: Final = scenario.model(api_base=provider.url + "/v1", api_key=provider_secret)
|
||||
response: Final = candidate.request(
|
||||
"POST",
|
||||
endpoint,
|
||||
_trace_body(kind, model, marker, {"trace_id": caller_trace}),
|
||||
headers=_w3c_headers(uuid.uuid4().hex, None),
|
||||
)
|
||||
assert response.status_code == 200, response.text
|
||||
assert upstream_targets == [expected_target], upstream_targets
|
||||
received: Final[list[Request]] = [] # mutable-ok: drain() consumes the queue, later polls keep earlier ones
|
||||
span: Final = _await_span(received, destination, response.headers["x-litellm-call-id"])
|
||||
assert _attribute(span.attributes, "session.id") == caller_trace
|
||||
|
|
|
|||
|
|
@ -8423,6 +8423,36 @@ async def test_missing_session_id_generate_reuses_traceparent_trace_id():
|
|||
assert _spend_log_session_id(updated) == "4bf92f3577b34da6a3ce929d0e0e4736"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("path", ["/v1/responses", "/v1/messages"])
|
||||
@pytest.mark.parametrize(
|
||||
"headers",
|
||||
[
|
||||
{"traceparent": "00-4bf92f3577b34da6a3ce929d0e0e4736-00f067aa0ba902b7-01"},
|
||||
{},
|
||||
],
|
||||
ids=["with_traceparent", "no_traceparent"],
|
||||
)
|
||||
async def test_missing_session_id_generate_reuses_promoted_caller_trace_id(path: str, headers: dict[str, str]):
|
||||
"""On litellm_metadata routes the caller's metadata.trace_id still counts as caller set even though
|
||||
it has not been promoted yet, so the generated session id derives from it instead of a fresh uuid."""
|
||||
request = _request_for(path)
|
||||
request.headers = headers
|
||||
|
||||
updated = await add_litellm_data_to_request(
|
||||
data={"model": "gpt-4o", "input": "hi", "metadata": {"trace_id": "caller-trace"}},
|
||||
request=request,
|
||||
user_api_key_dict=UserAPIKeyAuth(api_key="hashed-key"),
|
||||
proxy_config=MagicMock(),
|
||||
general_settings={"missing_session_id": "generate"},
|
||||
)
|
||||
|
||||
assert updated["litellm_session_id"] == "caller-trace"
|
||||
assert updated["litellm_metadata"]["session_id"] == "caller-trace"
|
||||
assert updated["litellm_metadata"][SESSION_ID_GENERATED_METADATA_KEY] is True
|
||||
assert _spend_log_session_id(updated, "litellm_metadata") == "caller-trace"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("policy", ["generate", "reject"])
|
||||
async def test_missing_session_id_policy_keeps_client_supplied_session_id(policy: str):
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue