litellm/tests/unit/proxy/test_litellm_pre_call_utils.py
devin-ai-integration[bot] a76b59db9f
test(proxy): move middleware, spend_tracking, pass_through, common_utils and root proxy tests into tests/unit/proxy (#44015)
Co-authored-by: yuneng <yuneng@berri.ai>
Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
2026-10-01 18:23:31 +00:00

8592 lines
327 KiB
Python

import asyncio
import copy
import json
import os
import time
from datetime import datetime, timezone
from types import SimpleNamespace
from typing import Final
from unittest.mock import AsyncMock, MagicMock, patch
import pytest
from botocore.credentials import Credentials
from fastapi import Request
from opentelemetry.trace import INVALID_SPAN, NonRecordingSpan, SpanContext
from pydantic import ValidationError as PydanticValidationError
from starlette.datastructures import Headers
import litellm
from litellm.proxy._types import AddTeamCallback, ProxyException, TeamCallbackMetadata, UserAPIKeyAuth
from litellm.proxy.litellm_pre_call_utils import (
KeyAndTeamLoggingSettings,
LiteLLMProxyRequestSetup,
_apply_credential_overrides_from_model_config,
_extract_credential_from_entry,
_get_dynamic_logging_metadata,
_get_enforced_params,
_get_metadata_variable_name,
_match_and_track_policies,
_promoted_trace_control_fields,
_resolve_credential_from_model_config,
_resolve_provider_from_deployment,
_update_model_if_key_alias_exists,
add_guardrails_from_policy_engine,
add_litellm_data_to_request,
add_provider_specific_headers_to_request,
check_if_token_is_service_account,
clean_headers,
move_guardrails_to_metadata,
)
from litellm.litellm_core_utils.core_helpers import get_litellm_metadata_from_kwargs
from litellm.proxy.spend_tracking.spend_tracking_utils import get_logging_payload
from litellm.litellm_core_utils.internal_call_metadata import MODEL_ACCESS_GROUP_METADATA_KEY
from litellm.litellm_core_utils.redact_messages import _get_turn_off_message_logging_from_dynamic_params
from litellm.litellm_core_utils.get_provider_specific_headers import (
ProviderSpecificHeaderUtils,
)
from litellm.litellm_core_utils.initialize_dynamic_callback_params import (
TRUSTED_CALLBACK_VARS_FIELD,
)
from litellm.constants import (
ROUTER_USAGE_COUNTED_TOKENS_METADATA_KEY,
SESSION_ID_GENERATED_METADATA_KEY,
SESSION_ID_OMITTED_METADATA_KEY,
)
from litellm.llms.bedrock.base_aws_llm import BaseAWSLLM
from litellm.llms.fireworks_ai.common_utils import get_fireworks_session_id
from litellm.types.utils import CredentialItem
def test_check_if_token_is_service_account():
"""
Test that only keys with `service_account_id` in metadata are considered service accounts
"""
# Test case 1: Service account token
service_account_token = UserAPIKeyAuth(api_key="test-key", metadata={"service_account_id": "test-service-account"})
assert check_if_token_is_service_account(service_account_token) == True
# Test case 2: Regular user token
regular_token = UserAPIKeyAuth(api_key="test-key", metadata={})
assert check_if_token_is_service_account(regular_token) == False
# Test case 3: Token with other metadata
other_metadata_token = UserAPIKeyAuth(api_key="test-key", metadata={"user_id": "test-user"})
assert check_if_token_is_service_account(other_metadata_token) == False
class TestGetMetadataVariableName:
"""Tests for _get_metadata_variable_name()"""
def _make_request(self, path: str) -> MagicMock:
request = MagicMock(spec=Request)
request.url.path = path
return request
def test_returns_litellm_metadata_for_thread_routes(self):
request = self._make_request("/v1/threads/thread_123/messages")
assert _get_metadata_variable_name(request) == "litellm_metadata"
def test_returns_litellm_metadata_for_assistant_routes(self):
request = self._make_request("/v1/assistants/asst_123")
assert _get_metadata_variable_name(request) == "litellm_metadata"
def test_returns_litellm_metadata_for_batches_route(self):
request = self._make_request("/v1/batches")
assert _get_metadata_variable_name(request) == "litellm_metadata"
def test_returns_litellm_metadata_for_messages_route(self):
request = self._make_request("/v1/messages")
assert _get_metadata_variable_name(request) == "litellm_metadata"
def test_returns_litellm_metadata_for_files_route(self):
request = self._make_request("/v1/files")
assert _get_metadata_variable_name(request) == "litellm_metadata"
def test_returns_metadata_for_chat_completions(self):
request = self._make_request("/chat/completions")
assert _get_metadata_variable_name(request) == "metadata"
def test_returns_metadata_for_completions(self):
request = self._make_request("/v1/completions")
assert _get_metadata_variable_name(request) == "metadata"
def test_returns_metadata_for_embeddings(self):
request = self._make_request("/v1/embeddings")
assert _get_metadata_variable_name(request) == "metadata"
def test_returns_litellm_metadata_for_bedrock_invoke(self):
# GH#30629: bedrock passthrough must use litellm_metadata
# to prevent key-level tags from leaking into provider body
request = self._make_request("/bedrock/model/us.anthropic.claude-sonnet-4-6/invoke")
assert _get_metadata_variable_name(request) == "litellm_metadata"
def test_returns_litellm_metadata_for_bedrock_converse(self):
request = self._make_request("/bedrock/model/us.anthropic.claude-sonnet-4-6/converse")
assert _get_metadata_variable_name(request) == "litellm_metadata"
def test_get_enforced_params_for_service_account_settings():
"""
Test that service account enforced params are only added to service account keys
"""
service_account_token = UserAPIKeyAuth(api_key="test-key", metadata={"service_account_id": "test-service-account"})
general_settings_with_service_account_settings = {
"service_account_settings": {"enforced_params": ["metadata.service"]},
}
result = _get_enforced_params(
general_settings=general_settings_with_service_account_settings,
user_api_key_dict=service_account_token,
)
assert result == ["metadata.service"]
regular_token = UserAPIKeyAuth(api_key="test-key", metadata={"enforced_params": ["user"]})
result = _get_enforced_params(
general_settings=general_settings_with_service_account_settings,
user_api_key_dict=regular_token,
)
assert result == ["user"]
@pytest.mark.parametrize(
"general_settings, user_api_key_dict, expected_enforced_params",
[
(
{"enforced_params": ["param1", "param2"]},
UserAPIKeyAuth(api_key="test_api_key", user_id="test_user_id", org_id="test_org_id"),
["param1", "param2"],
),
(
{"service_account_settings": {"enforced_params": ["param1", "param2"]}},
UserAPIKeyAuth(
api_key="test_api_key",
user_id="test_user_id",
org_id="test_org_id",
metadata={"service_account_id": "test_service_account_id"},
),
["param1", "param2"],
),
(
{"service_account_settings": {"enforced_params": ["param1", "param2"]}},
UserAPIKeyAuth(
api_key="test_api_key",
metadata={
"enforced_params": ["param3", "param4"],
"service_account_id": "test_service_account_id",
},
),
["param1", "param2", "param3", "param4"],
),
],
)
def test_get_enforced_params(general_settings, user_api_key_dict, expected_enforced_params):
from litellm.proxy.litellm_pre_call_utils import _get_enforced_params
enforced_params = _get_enforced_params(general_settings, user_api_key_dict)
assert enforced_params == expected_enforced_params
@pytest.mark.asyncio
async def test_add_litellm_data_to_request_parses_string_metadata():
from litellm.proxy.litellm_pre_call_utils import add_litellm_data_to_request
# Setup
request_mock = MagicMock(spec=Request)
request_mock.url.path = "/v1/completions"
request_mock.url = MagicMock()
request_mock.url.__str__.return_value = "http://localhost/v1/completions"
request_mock.method = "POST"
request_mock.query_params = {}
request_mock.headers = {"Content-Type": "application/json"}
request_mock.client = MagicMock()
request_mock.client.host = "127.0.0.1"
# Simulate data with stringified metadata
fake_metadata = {"generation_name": "gen123"}
data = {"metadata": json.dumps(fake_metadata), "model": "gpt-3.5-turbo"}
user_api_key_dict = UserAPIKeyAuth(
api_key="hashed-key",
metadata={},
team_metadata={},
spend=0.0,
max_budget=100.0,
model_max_budget={}, # this one can be a dict
team_spend=0.0,
team_max_budget=200.0,
)
# Call
updated_data = await add_litellm_data_to_request(
data=data,
request=request_mock,
user_api_key_dict=user_api_key_dict,
proxy_config=MagicMock(),
general_settings={},
version="test-version",
)
# Assert
litellm_metadata = updated_data.get("metadata", {})
assert isinstance(litellm_metadata, dict)
assert updated_data["metadata"]["generation_name"] == "gen123"
@pytest.mark.asyncio
async def test_key_otel_service_name_outranks_team_metadata_merge():
from litellm.proxy.litellm_pre_call_utils import add_litellm_data_to_request
request_mock = MagicMock(spec=Request)
request_mock.url = MagicMock()
request_mock.url.path = "/v1/chat/completions"
request_mock.url.__str__.return_value = "http://localhost/v1/chat/completions"
request_mock.method = "POST"
request_mock.query_params = {}
request_mock.headers = {"Content-Type": "application/json"}
request_mock.client = MagicMock()
request_mock.client.host = "127.0.0.1"
user_api_key_dict = UserAPIKeyAuth(
api_key="hashed-key",
metadata={"otel_service_name": "key-svc"},
team_metadata={"otel_service_name": "team-svc", "other_setting": "team-val"},
)
updated_data = await add_litellm_data_to_request(
data={"model": "gpt-3.5-turbo"},
request=request_mock,
user_api_key_dict=user_api_key_dict,
proxy_config=MagicMock(),
general_settings={},
version="test-version",
)
auth_metadata = updated_data["metadata"]["user_api_key_auth_metadata"]
assert auth_metadata["otel_service_name"] == "key-svc"
assert auth_metadata["other_setting"] == "team-val"
@pytest.mark.asyncio
async def test_stamped_auth_object_reflects_header_derived_identity():
"""
Regression (LIT-5487): the stamped object is a copy taken partway through request setup,
so it only carries header-derived identity if the stamp still runs after those fields are
resolved. Moving the stamp earlier would silently misattribute spend.
"""
from litellm.proxy.litellm_pre_call_utils import add_litellm_data_to_request
request_mock = MagicMock(spec=Request)
request_mock.url = MagicMock()
request_mock.url.path = "/v1/chat/completions"
request_mock.url.__str__.return_value = "http://localhost/v1/chat/completions"
request_mock.method = "POST"
request_mock.query_params = {}
request_mock.headers = {"Content-Type": "application/json", "user": "end-user-from-header"}
request_mock.client = MagicMock()
request_mock.client.host = "127.0.0.1"
user_api_key_dict = UserAPIKeyAuth(api_key="hashed-key", metadata={}, team_metadata={})
updated_data = await add_litellm_data_to_request(
data={"model": "gpt-3.5-turbo"},
request=request_mock,
user_api_key_dict=user_api_key_dict,
proxy_config=MagicMock(),
general_settings={"user_header_name": "user"},
version="test-version",
)
# precondition: the header was actually resolved onto the live object
assert user_api_key_dict.end_user_id == "end-user-from-header"
stamped = updated_data["metadata"]["user_api_key_auth"]
assert stamped.end_user_id == "end-user-from-header"
@pytest.mark.asyncio
async def test_arrival_time_prefers_litellm_received_at_over_time_time():
"""LIT-6012: by the time this function runs, auth has already completed, so
time.time() here would silently exclude the whole auth phase from the
queue-time window. request.state.litellm_received_at (stamped at the top of
user_api_key_auth, before auth work) must win when present."""
from litellm.proxy.litellm_pre_call_utils import add_litellm_data_to_request
request_mock = MagicMock(spec=Request)
request_mock.url = MagicMock()
request_mock.url.path = "/v1/chat/completions"
request_mock.url.__str__.return_value = "http://localhost/v1/chat/completions"
request_mock.method = "POST"
request_mock.query_params = {}
request_mock.headers = {"Content-Type": "application/json"}
request_mock.client = MagicMock()
request_mock.client.host = "127.0.0.1"
received_at = datetime(2024, 1, 1, tzinfo=timezone.utc)
request_mock.state = SimpleNamespace(litellm_received_at=received_at)
user_api_key_dict = UserAPIKeyAuth(api_key="hashed-key", metadata={}, team_metadata={})
updated_data = await add_litellm_data_to_request(
data={"model": "gpt-3.5-turbo"},
request=request_mock,
user_api_key_dict=user_api_key_dict,
proxy_config=MagicMock(),
general_settings={},
version="test-version",
)
assert updated_data["proxy_server_request"]["arrival_time"] == received_at.timestamp()
@pytest.mark.asyncio
async def test_proxy_clears_client_supplied_timing_windows():
request_mock = MagicMock(spec=Request)
request_mock.url = MagicMock()
request_mock.url.path = "/v1/chat/completions"
request_mock.url.__str__.return_value = "http://localhost/v1/chat/completions"
request_mock.method = "POST"
request_mock.query_params = {}
request_mock.headers = {"Content-Type": "application/json"}
request_mock.client = MagicMock()
request_mock.client.host = "127.0.0.1"
request_mock.state = SimpleNamespace(litellm_received_at=datetime.now(timezone.utc))
user_api_key_dict = UserAPIKeyAuth(api_key="hashed-key", metadata={}, team_metadata={})
updated_data = await add_litellm_data_to_request(
data={
"model": "gpt-3.5-turbo",
"metadata": {"llm_api_timing_windows": ((0.0, 1.0),)},
},
request=request_mock,
user_api_key_dict=user_api_key_dict,
proxy_config=MagicMock(),
general_settings={},
version="test-version",
)
assert updated_data["metadata"]["llm_api_timing_windows"] == ()
@pytest.mark.asyncio
async def test_arrival_time_falls_back_to_time_time_without_litellm_received_at():
"""Callers that never went through user_api_key_auth (no stamp on request.state)
must still get a usable arrival_time instead of erroring."""
from litellm.proxy.litellm_pre_call_utils import add_litellm_data_to_request
request_mock = MagicMock(spec=Request)
request_mock.url = MagicMock()
request_mock.url.path = "/v1/chat/completions"
request_mock.url.__str__.return_value = "http://localhost/v1/chat/completions"
request_mock.method = "POST"
request_mock.query_params = {}
request_mock.headers = {"Content-Type": "application/json"}
request_mock.client = MagicMock()
request_mock.client.host = "127.0.0.1"
request_mock.state = SimpleNamespace() # no litellm_received_at attribute
user_api_key_dict = UserAPIKeyAuth(api_key="hashed-key", metadata={}, team_metadata={})
before = time.time()
updated_data = await add_litellm_data_to_request(
data={"model": "gpt-3.5-turbo"},
request=request_mock,
user_api_key_dict=user_api_key_dict,
proxy_config=MagicMock(),
general_settings={},
version="test-version",
)
after = time.time()
arrival_time = updated_data["proxy_server_request"]["arrival_time"]
assert isinstance(arrival_time, float)
assert before <= arrival_time <= after
@pytest.mark.asyncio
async def test_add_litellm_data_to_request_strips_admin_injection_slots():
"""User-supplied user_api_key_metadata / user_api_key_team_metadata /
_pipeline_managed_guardrails must be stripped from both metadata keys
before the proxy writes its own admin-populated values. Otherwise a
caller can shadow admin config via the non-`_metadata_variable_name`
metadata key (e.g. litellm_metadata while the proxy writes to metadata).
"""
from litellm.proxy.litellm_pre_call_utils import add_litellm_data_to_request
request_mock = MagicMock(spec=Request)
request_mock.url.path = "/v1/chat/completions"
request_mock.url = MagicMock()
request_mock.url.__str__.return_value = "http://localhost/v1/chat/completions"
request_mock.method = "POST"
request_mock.query_params = {}
request_mock.headers = {"Content-Type": "application/json"}
request_mock.client = MagicMock()
request_mock.client.host = "127.0.0.1"
# Caller tries to inject admin config into BOTH metadata keys
attacker_admin_payload = {"disable_global_guardrails": True}
data = {
"model": "gpt-3.5-turbo",
"metadata": {
"user_api_key_metadata": attacker_admin_payload,
"user_api_key_team_metadata": attacker_admin_payload,
"_pipeline_managed_guardrails": ["evaded"],
},
"litellm_metadata": {
"user_api_key_metadata": attacker_admin_payload,
"user_api_key_team_metadata": attacker_admin_payload,
"_pipeline_managed_guardrails": ["evaded"],
},
}
real_admin_metadata = {"admin_flag": "from_proxy"}
user_api_key_dict = UserAPIKeyAuth(
api_key="hashed-key",
metadata=real_admin_metadata,
team_metadata=real_admin_metadata,
spend=0.0,
max_budget=100.0,
model_max_budget={},
team_spend=0.0,
team_max_budget=200.0,
)
updated = await add_litellm_data_to_request(
data=data,
request=request_mock,
user_api_key_dict=user_api_key_dict,
proxy_config=MagicMock(),
general_settings={},
version="test-version",
)
# The key that matches `_metadata_variable_name` gets proxy-populated
# with the real admin payload; the OTHER key must not retain the
# attacker's injection.
populated = updated["metadata"]
assert populated["user_api_key_metadata"] == real_admin_metadata
assert populated["user_api_key_team_metadata"] == real_admin_metadata
assert "_pipeline_managed_guardrails" not in populated or populated["_pipeline_managed_guardrails"] != ["evaded"]
other = updated.get("litellm_metadata") or {}
assert other.get("user_api_key_metadata") in (None, {}, real_admin_metadata)
assert other.get("user_api_key_team_metadata") in (None, {}, real_admin_metadata)
assert "_pipeline_managed_guardrails" not in other
@pytest.mark.asyncio
async def test_add_litellm_data_to_request_strips_all_user_api_key_prefix_keys():
"""Strip must cover the full user_api_key_* family, not a hand-maintained
list of 2-3 names. Proxy writes a dozen such fields (user_id, alias,
spend, team_id, request_route, …) and an attacker populating any of them
in the non-authoritative metadata key would otherwise forge identity /
spend in audit logs and guardrails."""
from litellm.proxy.litellm_pre_call_utils import add_litellm_data_to_request
request_mock = MagicMock(spec=Request)
request_mock.url.path = "/v1/chat/completions"
request_mock.url = MagicMock()
request_mock.url.__str__.return_value = "http://localhost/v1/chat/completions"
request_mock.method = "POST"
request_mock.query_params = {}
request_mock.headers = {"Content-Type": "application/json"}
request_mock.client = MagicMock()
request_mock.client.host = "127.0.0.1"
attacker_injected = {
"user_api_key_user_id": "victim",
"user_api_key_alias": "admin-key",
"user_api_key_spend": 0.0,
"user_api_key_team_id": "victim-team",
"user_api_key_end_user_id": "victim-user",
"user_api_key_request_route": "/fake/route",
"user_api_key_hash": "fake-hash",
}
data = {
"model": "gpt-3.5-turbo",
"metadata": {**attacker_injected},
"litellm_metadata": {**attacker_injected},
}
user_api_key_dict = UserAPIKeyAuth(
api_key="hashed-key",
user_id="real-user",
metadata={},
team_metadata={},
spend=42.0,
max_budget=100.0,
model_max_budget={},
team_spend=0.0,
team_max_budget=200.0,
)
updated = await add_litellm_data_to_request(
data=data,
request=request_mock,
user_api_key_dict=user_api_key_dict,
proxy_config=MagicMock(),
general_settings={},
version="test-version",
)
# The non-authoritative metadata dict must not retain ANY attacker-injected
# user_api_key_* key.
other = updated.get("litellm_metadata") or {}
attacker_leaks = [k for k in other if k.startswith("user_api_key_")]
assert attacker_leaks == [], f"Unexpected leaked keys: {attacker_leaks}"
@pytest.mark.asyncio
async def test_add_litellm_data_to_request_string_metadata_does_not_crash():
"""Regression: pre-strip code that pre-populated data['metadata'][k]=v
before the string-to-dict parse would crash on JSON-string metadata.
The snapshot / strip / admin-population pipeline must survive metadata
arriving as a string."""
import json as _json
from litellm.proxy.litellm_pre_call_utils import add_litellm_data_to_request
request_mock = MagicMock(spec=Request)
request_mock.url.path = "/v1/chat/completions"
request_mock.url = MagicMock()
request_mock.url.__str__.return_value = "http://localhost/v1/chat/completions"
request_mock.method = "POST"
request_mock.query_params = {}
request_mock.headers = {"Content-Type": "multipart/form-data"}
request_mock.client = MagicMock()
request_mock.client.host = "127.0.0.1"
data = {
"model": "gpt-3.5-turbo",
"metadata": _json.dumps({"generation_name": "test"}),
}
user_api_key_dict = UserAPIKeyAuth(
api_key="hashed-key",
metadata={},
team_metadata={},
spend=0.0,
max_budget=100.0,
model_max_budget={},
team_spend=0.0,
team_max_budget=200.0,
)
# Must not raise TypeError / AttributeError.
updated = await add_litellm_data_to_request(
data=data,
request=request_mock,
user_api_key_dict=user_api_key_dict,
proxy_config=MagicMock(),
general_settings={},
version="test-version",
)
# The parsed metadata should be a dict and the proxy snapshot body
# should have been taken AFTER the strip (so no leaked user_api_key_*
# from a raw string snapshot).
assert isinstance(updated["metadata"], dict)
assert updated["metadata"].get("generation_name") == "test"
def _batches_request_mock() -> MagicMock:
request_mock = MagicMock(spec=Request)
request_mock.url = MagicMock()
request_mock.url.__str__.return_value = "http://localhost/v1/batches"
request_mock.url.path = "/v1/batches"
request_mock.method = "POST"
request_mock.query_params = {}
request_mock.headers = {"Content-Type": "application/json"}
request_mock.client = MagicMock()
request_mock.client.host = "127.0.0.1"
request_mock.state.parent_otel_span = None
return request_mock
@pytest.mark.asyncio
@pytest.mark.parametrize(
"field,value,received_type",
[
("metadata", "abc", "a string"),
("litellm_metadata", "abc", "a string"),
("metadata", 42, "an integer"),
("litellm_metadata", [1, 2], "an array"),
("metadata", True, "a boolean"),
],
)
async def test_add_litellm_data_to_request_rejects_non_object_metadata(field, value, received_type):
"""Regression for https://github.com/BerriAI/litellm/issues/37147: a
non-object metadata was silently dropped with a 200, and a non-object
litellm_metadata crashed later with a 500 ('str' object has no attribute
'update'). Both must be a 400 naming the field, like OpenAI returns."""
data = {"input_file_id": "file-abc", "endpoint": "/v1/chat/completions", field: value}
with pytest.raises(ProxyException) as exc_info:
await add_litellm_data_to_request(
data=data,
request=_batches_request_mock(),
user_api_key_dict=UserAPIKeyAuth(api_key="hashed-key"),
proxy_config=MagicMock(),
general_settings={},
version="test-version",
)
assert exc_info.value.code == "400"
assert exc_info.value.param == field
assert exc_info.value.message == f"Invalid type for '{field}': expected an object, but got {received_type} instead."
assert field not in data
@pytest.mark.asyncio
async def test_add_litellm_data_to_request_removes_every_invalid_metadata_field_before_raising():
"""When both fields are invalid, the raise for the first must not leave the
second invalid value in data, or failure-logging hooks that inspect the body
can crash on it and mask the 400 as a 500."""
data = {"input_file_id": "file-abc", "metadata": "abc", "litellm_metadata": "xyz"}
with pytest.raises(ProxyException) as exc_info:
await add_litellm_data_to_request(
data=data,
request=_batches_request_mock(),
user_api_key_dict=UserAPIKeyAuth(api_key="hashed-key"),
proxy_config=MagicMock(),
general_settings={},
version="test-version",
)
assert exc_info.value.param == "metadata"
assert "metadata" not in data
assert "litellm_metadata" not in data
@pytest.mark.asyncio
async def test_add_litellm_data_to_request_parses_json_object_string_litellm_metadata():
data = {"input_file_id": "file-abc", "litellm_metadata": json.dumps({"cost_centre": "research"})}
updated = await add_litellm_data_to_request(
data=data,
request=_batches_request_mock(),
user_api_key_dict=UserAPIKeyAuth(api_key="hashed-key"),
proxy_config=MagicMock(),
general_settings={},
version="test-version",
)
assert updated["litellm_metadata"]["cost_centre"] == "research"
@pytest.mark.asyncio
async def test_add_litellm_data_to_request_proxy_server_request_body_is_post_strip():
"""Regression: proxy_server_request['body'] used to be snapshotted before
the admin-slot strip, so standard_logging_object and spend-tracking
readers saw attacker-injected payload. Snapshot must now be post-strip."""
from litellm.proxy.litellm_pre_call_utils import add_litellm_data_to_request
request_mock = MagicMock(spec=Request)
request_mock.url.path = "/v1/chat/completions"
request_mock.url = MagicMock()
request_mock.url.__str__.return_value = "http://localhost/v1/chat/completions"
request_mock.method = "POST"
request_mock.query_params = {}
request_mock.headers = {"Content-Type": "application/json"}
request_mock.client = MagicMock()
request_mock.client.host = "127.0.0.1"
data = {
"model": "gpt-3.5-turbo",
"metadata": {"user_api_key_user_id": "victim"},
}
user_api_key_dict = UserAPIKeyAuth(
api_key="hashed-key",
user_id="real-user",
metadata={},
team_metadata={},
spend=0.0,
max_budget=100.0,
model_max_budget={},
team_spend=0.0,
team_max_budget=200.0,
)
updated = await add_litellm_data_to_request(
data=data,
request=request_mock,
user_api_key_dict=user_api_key_dict,
proxy_config=MagicMock(),
general_settings={},
version="test-version",
)
snapshot_body = updated["proxy_server_request"]["body"]
assert snapshot_body is not None
snapshot_metadata = snapshot_body.get("metadata") or {}
assert "user_api_key_user_id" not in snapshot_metadata or (snapshot_metadata["user_api_key_user_id"] != "victim")
@pytest.mark.asyncio
async def test_add_litellm_data_to_request_body_snapshot_excludes_secret_fields():
"""Security: proxy_server_request['body'] must never contain secret_fields
because that dict holds raw HTTP headers including Authorization Bearer
tokens. The body snapshot is persisted in spend logs and other audit trails,
so leaking secret_fields there exposes user credentials.
secret_fields must still be available on the live ``data`` dict for
downstream consumers (MCP, Responses API) that legitimately need raw headers.
"""
from litellm.proxy.litellm_pre_call_utils import add_litellm_data_to_request
request_mock = MagicMock(spec=Request)
request_mock.url.path = "/v1/chat/completions"
request_mock.url = MagicMock()
request_mock.url.__str__.return_value = "http://localhost/v1/chat/completions"
request_mock.method = "POST"
request_mock.query_params = {}
request_mock.headers = {
"Content-Type": "application/json",
"Authorization": "Bearer sk-super-secret-token",
}
request_mock.client = MagicMock()
request_mock.client.host = "127.0.0.1"
data = {
"model": "gpt-3.5-turbo",
"messages": [{"role": "user", "content": "hello"}],
}
user_api_key_dict = UserAPIKeyAuth(
api_key="hashed-key",
user_id="test-user",
metadata={},
team_metadata={},
spend=0.0,
max_budget=100.0,
model_max_budget={},
team_spend=0.0,
team_max_budget=200.0,
)
updated = await add_litellm_data_to_request(
data=data,
request=request_mock,
user_api_key_dict=user_api_key_dict,
proxy_config=MagicMock(),
general_settings={},
version="test-version",
)
# secret_fields must exist on the live data dict
assert "secret_fields" in updated, "secret_fields must still be present on the live data dict"
assert "raw_headers" in updated["secret_fields"]
# But the body snapshot must NOT contain secret_fields
snapshot_body = updated["proxy_server_request"]["body"]
assert "secret_fields" not in snapshot_body, (
"secret_fields must be excluded from proxy_server_request['body'] "
"to prevent Authorization tokens from leaking into spend logs"
)
@pytest.mark.asyncio
async def test_add_litellm_data_to_request_body_snapshot_excludes_proxy_server_request():
"""Regression: the body snapshot used to include the proxy_server_request
key itself, producing the path
``proxy_server_request.body.proxy_server_request.body == body``. Custom
loggers and audit consumers must not see the self-referencing structure
(independent of redaction — fires on every successful call).
"""
from litellm.proxy.litellm_pre_call_utils import add_litellm_data_to_request
request_mock = MagicMock(spec=Request)
request_mock.url.path = "/v1/chat/completions"
request_mock.url = MagicMock()
request_mock.url.__str__.return_value = "http://localhost/v1/chat/completions"
request_mock.method = "POST"
request_mock.query_params = {}
request_mock.headers = {"Content-Type": "application/json"}
request_mock.client = MagicMock()
request_mock.client.host = "127.0.0.1"
data = {
"model": "gpt-3.5-turbo",
"messages": [{"role": "user", "content": "hello"}],
"api_key": "request-key",
"proxy_server_request": {
"body": {"messages": [{"role": "user", "content": "forged"}]},
},
}
user_api_key_dict = UserAPIKeyAuth(
api_key="hashed-key",
user_id="test-user",
metadata={},
team_metadata={},
spend=0.0,
max_budget=100.0,
model_max_budget={},
team_spend=0.0,
team_max_budget=200.0,
)
updated = await add_litellm_data_to_request(
data=data,
request=request_mock,
user_api_key_dict=user_api_key_dict,
proxy_config=MagicMock(),
general_settings={},
version="test-version",
)
snapshot_body = updated["proxy_server_request"]["body"]
assert "proxy_server_request" not in snapshot_body, (
"proxy_server_request must be excluded from its own body snapshot to prevent the body from self-referencing"
)
assert "api_key" not in snapshot_body
assert updated["proxy_server_request"]["credential_fields"] == ("api_key",)
assert snapshot_body["messages"] == [{"role": "user", "content": "hello"}]
def test_initial_snapshot_refresh_clears_a_previous_guardrail_checkpoint() -> None:
from litellm.integrations.shadow_eval_logger import GuardrailRequestSnapshot
from litellm.litellm_core_utils.litellm_logging import Logging
from litellm.proxy.litellm_pre_call_utils import refresh_proxy_server_request_body_snapshot
logging_obj: Final = Logging(
model="test-model", messages=[], stream=False, call_type="acompletion",
start_time=datetime.now(), litellm_call_id="new-request", function_id="new-request",
)
logging_obj.shadow_eval_request_snapshot = GuardrailRequestSnapshot.capture(
{"messages": [{"role": "user", "content": "previous request"}]},
{"standard_logging_guardrail_information": [{"guardrail_mode": "pre_call"}]},
)
assert logging_obj.shadow_eval_request_snapshot is not None
proxy_request: Final = {"body": {}}
data: Final = {
"messages": [{"role": "user", "content": "new request"}],
"proxy_server_request": proxy_request,
"litellm_logging_obj": logging_obj,
}
refresh_proxy_server_request_body_snapshot(data)
assert logging_obj.shadow_eval_request_snapshot is None
assert proxy_request == {"body": {"messages": [{"role": "user", "content": "new request"}]}}
def test_body_snapshot_excludes_team_callback_credentials() -> None:
from litellm.proxy.litellm_pre_call_utils import refresh_proxy_server_request_body_snapshot
from litellm.types.litellm_params import TRUSTED_CALLBACK_VARS_FIELD
callback_vars: Final = {
"langfuse_public_key": "pk-lf-team",
"langfuse_secret_key": "sk-lf-team-secret",
"langfuse_host": "https://cloud.langfuse.com",
}
proxy_request: Final = {"body": None}
data: Final = {
"messages": [{"role": "user", "content": "hi"}],
"proxy_server_request": proxy_request,
"success_callback": ["langfuse"],
**callback_vars,
TRUSTED_CALLBACK_VARS_FIELD: callback_vars,
}
refresh_proxy_server_request_body_snapshot(data)
assert proxy_request == {
"body": {"messages": [{"role": "user", "content": "hi"}], "success_callback": ["langfuse"]}
}, proxy_request
@pytest.mark.asyncio
@pytest.mark.parametrize("pre_call_ran", [False, True])
async def test_post_guardrail_snapshot_preserves_logging_only_masking_in_spend_logs(
monkeypatch: pytest.MonkeyPatch, pre_call_ran: bool
) -> None:
from litellm.litellm_core_utils.litellm_logging import Logging
from litellm.proxy.guardrails.guardrail_hooks.presidio import _OPTIONAL_PresidioPIIMasking
from litellm.proxy.litellm_pre_call_utils import refresh_proxy_server_request_body_snapshot
from litellm.proxy.spend_tracking.spend_tracking_utils import _get_proxy_server_request_for_spend_logs_payload
monkeypatch.setenv("STORE_PROMPTS_IN_SPEND_LOGS", "true")
messages: Final = [{"role": "user", "content": "email probe@example.invalid"}]
metadata: Final = {
"standard_logging_guardrail_information": [{"guardrail_mode": "pre_call"}] if pre_call_ran else []
}
data: Final = {"messages": messages, "metadata": metadata, "proxy_server_request": {}}
logging_obj: Final = Logging(
model="test-model", messages=messages, stream=False, call_type="acompletion",
start_time=datetime.now(), litellm_call_id="mask-spend", function_id="mask-spend", kwargs=data,
)
data["litellm_logging_obj"] = logging_obj
refresh_proxy_server_request_body_snapshot(data, guardrails_applied=True)
logging_obj.update_messages(messages)
snapshot: Final = logging_obj.shadow_eval_request_snapshot
assert (snapshot is not None) is pre_call_ran
guardrail: Final = _OPTIONAL_PresidioPIIMasking(
mock_testing=True, logging_only=True, mock_redacted_text={"text": "email [EMAIL]", "items": []}
)
kwargs, _ = await guardrail.async_logging_hook(
kwargs=logging_obj.model_call_details, result=None, call_type="acompletion"
)
stored: Final = json.loads(_get_proxy_server_request_for_spend_logs_payload(
metadata={}, litellm_params=kwargs["litellm_params"], kwargs=kwargs,
))
assert kwargs["messages"] == [{"role": "user", "content": "email [EMAIL]"}]
assert stored["messages"] == kwargs["messages"]
if snapshot is not None:
assert snapshot.body["messages"] == [{"role": "user", "content": "email probe@example.invalid"}]
assert "probe@example.invalid" not in json.dumps(stored)
def test_refresh_proxy_server_request_body_snapshot_picks_up_guardrail_masking():
"""
Regression: proxy_server_request['body'] is snapshotted by
add_litellm_data_to_request BEFORE guardrails (e.g. Presidio PII masking) run
in pre_call_hook. Without a refresh after pre_call_hook, the persisted body
silently bypasses whatever masking the guardrail applied, so raw PII/PCI
lands in SpendLogs when store_prompts_in_spend_logs is enabled.
"""
from litellm.proxy.litellm_pre_call_utils import (
refresh_proxy_server_request_body_snapshot,
)
class _FakeLoggingObj:
"""Stands in for the live, non-JSON-serializable Logging instance that
litellm.utils.function_setup stamps onto `data` between the initial
snapshot and pre_call_hook."""
data = {
"model": "gpt-3.5-turbo",
"messages": [{"role": "user", "content": "my ssn is 123-45-6789"}],
"secret_fields": {"raw_headers": {"authorization": "Bearer sk-secret"}},
"litellm_logging_obj": _FakeLoggingObj(),
"proxy_server_request": {
"url": "http://localhost/v1/chat/completions",
"method": "POST",
"body": {
"model": "gpt-3.5-turbo",
"messages": [{"role": "user", "content": "my ssn is 123-45-6789"}],
},
},
}
# Simulate a PII-masking guardrail mutating `messages` in place, like Presidio's
# async_pre_call_hook does, after the initial snapshot was already taken.
data["messages"] = [{"role": "user", "content": "my ssn is <MASKED>"}]
refresh_proxy_server_request_body_snapshot(data)
refreshed_body = data["proxy_server_request"]["body"]
assert refreshed_body["messages"] == data["messages"]
# Still excludes secrets, self-reference, and the live logging object, same as
# the initial snapshot -- and proves the persisted body stays JSON-serializable.
assert "secret_fields" not in refreshed_body
assert "proxy_server_request" not in refreshed_body
assert "litellm_logging_obj" not in refreshed_body
assert "123-45-6789" not in json.dumps(refreshed_body)
@pytest.mark.asyncio
async def test_add_litellm_data_to_request_strips_string_encoded_admin_injection():
"""Regression: metadata arriving as a JSON string (multipart/form-data or
extra_body) must not bypass the admin-injection strip. The parse happens
AFTER receipt, so the strip has to run after the parse, not before.
"""
from litellm.proxy.litellm_pre_call_utils import add_litellm_data_to_request
request_mock = MagicMock(spec=Request)
request_mock.url.path = "/v1/chat/completions"
request_mock.url = MagicMock()
request_mock.url.__str__.return_value = "http://localhost/v1/chat/completions"
request_mock.method = "POST"
request_mock.query_params = {}
request_mock.headers = {"Content-Type": "multipart/form-data"}
request_mock.client = MagicMock()
request_mock.client.host = "127.0.0.1"
# Attacker encodes an admin-injection payload inside a JSON string.
attacker_payload = {
"user_api_key_metadata": {"disable_global_guardrails": True},
"user_api_key_team_metadata": {"disable_global_guardrails": True},
"_pipeline_managed_guardrails": ["evaded"],
}
data = {
"model": "gpt-3.5-turbo",
"metadata": json.dumps(attacker_payload),
"litellm_metadata": json.dumps(attacker_payload),
}
real_admin_metadata = {"admin_flag": "from_proxy"}
user_api_key_dict = UserAPIKeyAuth(
api_key="hashed-key",
metadata=real_admin_metadata,
team_metadata=real_admin_metadata,
spend=0.0,
max_budget=100.0,
model_max_budget={},
team_spend=0.0,
team_max_budget=200.0,
)
updated = await add_litellm_data_to_request(
data=data,
request=request_mock,
user_api_key_dict=user_api_key_dict,
proxy_config=MagicMock(),
general_settings={},
version="test-version",
)
populated = updated["metadata"]
# The real admin payload from user_api_key_dict wins.
assert populated["user_api_key_metadata"] == real_admin_metadata
assert populated["user_api_key_team_metadata"] == real_admin_metadata
assert populated.get("_pipeline_managed_guardrails") != ["evaded"]
other = updated.get("litellm_metadata") or {}
# After the strip, litellm_metadata has no admin-injection slots.
assert "user_api_key_metadata" not in other
assert "user_api_key_team_metadata" not in other
assert "_pipeline_managed_guardrails" not in other
@pytest.mark.asyncio
async def test_add_litellm_data_to_request_strips_user_control_fields():
"""Strip untrusted proxy-control fields before guardrails, logging, and headers read metadata."""
request_mock = MagicMock(spec=Request)
request_mock.url.path = "/v1/chat/completions"
request_mock.url = MagicMock()
request_mock.url.__str__.return_value = "http://localhost/v1/chat/completions"
request_mock.method = "POST"
request_mock.query_params = {}
request_mock.headers = {"Content-Type": "application/json"}
request_mock.client = MagicMock()
request_mock.client.host = "127.0.0.1"
malicious_metadata = {
"disable_global_guardrails": True,
"opted_out_global_guardrails": ["pii"],
"pillar_response_headers": {"set-cookie": "session=evil"},
"_pillar_response_headers_trusted": True,
"pillar_flagged": True,
"pillar_scanners": {"jailbreak": True},
"pillar_evidence": [{"evidence": "spoofed"}],
"pillar_session_id_response": "spoofed-session",
"applied_guardrails": ["spoofed"],
"applied_policies": ["spoofed-policy"],
"policy_sources": {"spoofed-policy": "request"},
"routing_decision": {"cause": "forged", "routed_model": "spoofed"},
"litellm_gateway_injected_cache": "forged-deployment-id",
"_session_deployment_affinity_ttl": 999999,
"internal_call_origin": "autorouter_classifier",
"_guardrail_pipelines": [{"name": "spoofed"}],
"_pipeline_managed_guardrails": ["evaded"],
"safe_user_metadata": "kept",
}
data = {
"model": "gpt-3.5-turbo",
"messages": [{"role": "user", "content": "hello"}],
"mock_response": "free response",
"mock_tool_calls": [{"id": "call_1"}],
"disable_global_guardrails": True,
"enable_prompt_caching": True,
"routing_decision": {"cause": "forged", "routed_model": "spoofed"},
"litellm_gateway_injected_cache": "forged-deployment-id",
"metadata": copy.deepcopy(malicious_metadata),
"litellm_metadata": copy.deepcopy(malicious_metadata),
"weights": {"gpt-3.5-turbo": {"forged-deployment-id": 100}},
"_router_weights": {"gpt-3.5-turbo": {"forged-deployment-id": 100}},
}
updated = await add_litellm_data_to_request(
data=data,
request=request_mock,
user_api_key_dict=UserAPIKeyAuth(api_key="hashed-key"),
proxy_config=MagicMock(),
general_settings={},
version="test-version",
)
assert "mock_response" not in updated
assert "mock_tool_calls" not in updated
assert "disable_global_guardrails" not in updated
assert "enable_prompt_caching" not in updated
assert "routing_decision" not in updated
assert "litellm_gateway_injected_cache" not in updated
assert "weights" not in updated
assert "_router_weights" not in updated
assert "weights" not in updated["proxy_server_request"]["body"]
assert "_router_weights" not in updated["proxy_server_request"]["body"]
stripped_keys = {
"disable_global_guardrails",
"opted_out_global_guardrails",
"pillar_response_headers",
"_pillar_response_headers_trusted",
"pillar_flagged",
"pillar_scanners",
"pillar_evidence",
"pillar_session_id_response",
"applied_guardrails",
"applied_policies",
"policy_sources",
"routing_decision",
"litellm_gateway_injected_cache",
"_session_deployment_affinity_ttl",
"internal_call_origin",
"_guardrail_pipelines",
"_pipeline_managed_guardrails",
}
assert "litellm_metadata" not in updated
for stripped_key in stripped_keys:
assert stripped_key not in updated["metadata"]
assert updated["metadata"]["safe_user_metadata"] == "kept"
requester_metadata = updated["metadata"]["requester_metadata"]
for stripped_key in stripped_keys:
assert stripped_key not in requester_metadata
snapshot_body = updated["proxy_server_request"]["body"]
assert "mock_response" not in snapshot_body
assert "mock_tool_calls" not in snapshot_body
assert "pillar_response_headers" not in snapshot_body["metadata"]
@pytest.mark.asyncio
@pytest.mark.parametrize(
"key_value, expected",
[(True, True), (False, False), ("yes", None), (None, None)],
)
async def test_key_metadata_enable_prompt_caching_promoted_to_request_root(key_value, expected):
"""Key metadata enable_prompt_caching is stamped onto the request root (bools only), even when the client spoofs it."""
request_mock = MagicMock(spec=Request)
request_mock.url.path = "/v1/chat/completions"
request_mock.url = MagicMock()
request_mock.url.__str__.return_value = "http://localhost/v1/chat/completions"
request_mock.method = "POST"
request_mock.query_params = {}
request_mock.headers = {"Content-Type": "application/json"}
request_mock.client = MagicMock()
request_mock.client.host = "127.0.0.1"
data = {
"model": "claude-sonnet-4-5",
"messages": [{"role": "user", "content": "hello"}],
"enable_prompt_caching": "spoofed-by-client",
}
key_metadata = {} if key_value is None else {"enable_prompt_caching": key_value}
updated = await add_litellm_data_to_request(
data=data,
request=request_mock,
user_api_key_dict=UserAPIKeyAuth(api_key="hashed-key", metadata=key_metadata),
proxy_config=MagicMock(),
general_settings={},
version="test-version",
)
assert updated.get("enable_prompt_caching") == expected
@pytest.mark.asyncio
@pytest.mark.parametrize(
"control_field",
[
"callbacks",
"service_callback",
"logger_fn",
"litellm_disabled_callbacks",
"_agentic_loop_depth",
"_agentic_loop_fingerprints",
"_code_interpreter_interception_active",
"_code_interpreter_interception_converted_stream",
"_code_interpreter_interception_sandbox_key",
"_headroom_interception_converted_stream",
"max_agentic_loops",
],
)
async def test_add_litellm_data_to_request_strips_callback_control_fields(
control_field,
):
"""``callbacks`` / ``service_callback`` / ``logger_fn`` get appended to
the worker-wide ``litellm.{input,success,failure,_async_*,service}_callback``
lists and ``litellm.user_logger_fn`` from inside ``function_setup`` —
one request poisons every subsequent caller in that worker.
``litellm_disabled_callbacks`` is the inverse: a request-body value
silently disables admin-configured audit/observability for the call.
None has a documented per-request use, so all four are stripped at
the proxy boundary alongside the existing internal-only fields."""
request_mock = MagicMock(spec=Request)
request_mock.url.path = "/v1/chat/completions"
request_mock.url = MagicMock()
request_mock.url.__str__.return_value = "http://localhost/v1/chat/completions"
request_mock.method = "POST"
request_mock.query_params = {}
request_mock.headers = {"Content-Type": "application/json"}
request_mock.client = MagicMock()
request_mock.client.host = "127.0.0.1"
sample_values = {
"callbacks": ["langfuse"],
"service_callback": ["langfuse"],
"litellm_disabled_callbacks": ["langfuse"],
"logger_fn": "module.func",
"_agentic_loop_depth": 5,
"_agentic_loop_fingerprints": ["forged"],
"_code_interpreter_interception_active": True,
"_code_interpreter_interception_converted_stream": True,
"_code_interpreter_interception_sandbox_key": "forged-key",
"_headroom_interception_converted_stream": True,
"max_agentic_loops": 9999,
}
sample_value = sample_values[control_field]
updated = await add_litellm_data_to_request(
data={
"model": "gpt-3.5-turbo",
"messages": [{"role": "user", "content": "hi"}],
control_field: sample_value,
},
request=request_mock,
user_api_key_dict=UserAPIKeyAuth(api_key="hashed-key"),
proxy_config=MagicMock(),
general_settings={},
version="test-version",
)
assert control_field not in updated
# The post-strip body snapshot used by audit/spend logging must also
# not retain the attacker-injected control field.
snapshot_body = updated["proxy_server_request"]["body"]
assert control_field not in snapshot_body
@pytest.mark.asyncio
@pytest.mark.parametrize("timeout_field", ["timeout", "request_timeout", "stream_timeout"])
async def test_add_litellm_data_to_request_marks_body_timeout_as_client_side(timeout_field):
"""Router._get_timeout resolves the effective timeout from any of kwargs["timeout"],
kwargs["request_timeout"], or kwargs["stream_timeout"], all settable directly in the
request body. Without recognizing all three, a caller could force a 408 on every
deployment in a fallback chain without it being flagged as caller-controlled, cooling
down deployments other tenants rely on (see cooldown_handlers._trigger_cooldown_for_failed_deployment)."""
request_mock = MagicMock(spec=Request)
request_mock.url.path = "/v1/chat/completions"
request_mock.url = MagicMock()
request_mock.url.__str__.return_value = "http://localhost/v1/chat/completions"
request_mock.method = "POST"
request_mock.query_params = {}
request_mock.headers = {"Content-Type": "application/json"}
request_mock.client = MagicMock()
request_mock.client.host = "127.0.0.1"
updated = await add_litellm_data_to_request(
data={
"model": "gpt-3.5-turbo",
"messages": [{"role": "user", "content": "hi"}],
timeout_field: 0.001,
},
request=request_mock,
user_api_key_dict=UserAPIKeyAuth(api_key="hashed-key"),
proxy_config=MagicMock(),
general_settings={},
version="test-version",
)
assert updated["client_side_timeout"] is True
@pytest.mark.asyncio
async def test_add_litellm_data_to_request_ignores_forged_client_side_timeout():
"""The client_side_timeout marker itself must never be trusted verbatim from the
request body: a caller forging client_side_timeout=True without a real timeout
override could dodge cooldown protection on an actual deployment failure."""
request_mock = MagicMock(spec=Request)
request_mock.url.path = "/v1/chat/completions"
request_mock.url = MagicMock()
request_mock.url.__str__.return_value = "http://localhost/v1/chat/completions"
request_mock.method = "POST"
request_mock.query_params = {}
request_mock.headers = {"Content-Type": "application/json"}
request_mock.client = MagicMock()
request_mock.client.host = "127.0.0.1"
updated = await add_litellm_data_to_request(
data={
"model": "gpt-3.5-turbo",
"messages": [{"role": "user", "content": "hi"}],
"client_side_timeout": True,
},
request=request_mock,
user_api_key_dict=UserAPIKeyAuth(api_key="hashed-key"),
proxy_config=MagicMock(),
general_settings={},
version="test-version",
)
assert not updated.get("client_side_timeout")
@pytest.mark.asyncio
async def test_client_side_timeout_marker_never_reaches_the_provider():
"""A proxy request with a caller-supplied timeout gets kwargs["client_side_timeout"]
stamped for the router's cooldown logic. That router-only marker must not ride
into the provider payload: unregistered kwargs are swept into extra_body /
additionalModelRequestFields, so Bedrock rejects the whole call with
`client_side_timeout: Extra inputs are not permitted`."""
from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler
request_mock = MagicMock(spec=Request)
request_mock.url.path = "/v1/chat/completions"
request_mock.url = MagicMock()
request_mock.url.__str__.return_value = "http://localhost/v1/chat/completions"
request_mock.method = "POST"
request_mock.query_params = {}
request_mock.headers = {"Content-Type": "application/json"}
request_mock.client = MagicMock()
request_mock.client.host = "127.0.0.1"
updated = await add_litellm_data_to_request(
data={
"model": "bedrock/us.anthropic.claude-sonnet-5",
"messages": [{"role": "user", "content": "hi"}],
"max_tokens": 10,
"timeout": 30,
},
request=request_mock,
user_api_key_dict=UserAPIKeyAuth(api_key="hashed-key"),
proxy_config=MagicMock(),
general_settings={},
version="test-version",
)
assert updated["client_side_timeout"] is True
converse_response = MagicMock()
converse_response.status_code = 200
converse_response.headers = {}
converse_response.json.return_value = {
"output": {"message": {"role": "assistant", "content": [{"text": "ok"}]}},
"stopReason": "end_turn",
"usage": {"inputTokens": 10, "outputTokens": 5, "totalTokens": 15},
}
converse_response.text = json.dumps(converse_response.json.return_value)
client = AsyncHTTPHandler()
with patch.object(client, "post", return_value=converse_response) as mock_post:
await litellm.acompletion(
**updated,
aws_access_key_id="fake-access-key",
aws_secret_access_key="fake-secret-key",
aws_region_name="us-east-1",
client=client,
)
mock_post.assert_called_once()
assert mock_post.call_args.kwargs["url"].endswith("/converse")
provider_body = json.loads(mock_post.call_args.kwargs["data"])
assert "client_side_timeout" not in json.dumps(provider_body), provider_body
@pytest.mark.asyncio
async def test_add_litellm_data_to_request_allows_client_mock_response_with_admin_opt_in():
request_mock = MagicMock(spec=Request)
request_mock.url.path = "/v1/chat/completions"
request_mock.url = MagicMock()
request_mock.url.__str__.return_value = "http://localhost/v1/chat/completions"
request_mock.method = "POST"
request_mock.query_params = {}
request_mock.headers = {"Content-Type": "application/json"}
request_mock.client = MagicMock()
request_mock.client.host = "127.0.0.1"
updated = await add_litellm_data_to_request(
data={
"model": "gpt-3.5-turbo",
"messages": [{"role": "user", "content": "hello"}],
"mock_response": "allowed mock",
"mock_tool_calls": [{"id": "call_1"}],
},
request=request_mock,
user_api_key_dict=UserAPIKeyAuth(
api_key="hashed-key",
metadata={"allow_client_mock_response": True},
),
proxy_config=MagicMock(),
general_settings={},
version="test-version",
)
assert updated["mock_response"] == "allowed mock"
assert updated["mock_tool_calls"] == [{"id": "call_1"}]
@pytest.mark.asyncio
async def test_add_litellm_data_to_request_strips_client_redaction_bypass_controls():
request_mock = MagicMock(spec=Request)
request_mock.url.path = "/v1/chat/completions"
request_mock.url = MagicMock()
request_mock.url.__str__.return_value = "http://localhost/v1/chat/completions"
request_mock.method = "POST"
request_mock.query_params = {}
request_mock.headers = {
"Content-Type": "application/json",
"litellm-disable-message-redaction": "true",
}
request_mock.client = MagicMock()
request_mock.client.host = "127.0.0.1"
original_turn_off_message_logging = litellm.turn_off_message_logging
litellm.turn_off_message_logging = True
try:
updated = await add_litellm_data_to_request(
data={
"model": "gpt-3.5-turbo",
"messages": [{"role": "user", "content": "hello"}],
"turn_off_message_logging": False,
"metadata": {
"headers": {"litellm-disable-message-redaction": "true"},
"turn_off_message_logging": False,
},
"litellm_metadata": json.dumps(
{
"headers": {"LiteLLM-Disable-Message-Redaction": "true"},
"turn_off_message_logging": "false",
}
),
"litellm_params": {
"metadata": {"turn_off_message_logging": False},
},
},
request=request_mock,
user_api_key_dict=UserAPIKeyAuth(api_key="hashed-key"),
proxy_config=MagicMock(),
general_settings={},
version="test-version",
)
finally:
litellm.turn_off_message_logging = original_turn_off_message_logging
assert "turn_off_message_logging" not in updated
assert "turn_off_message_logging" not in (updated.get("litellm_params") or {}).get("metadata", {})
assert "turn_off_message_logging" not in updated["metadata"]
assert "turn_off_message_logging" not in (updated.get("litellm_metadata") or {})
assert "litellm-disable-message-redaction" not in {header.lower() for header in updated["metadata"]["headers"]}
assert "litellm-disable-message-redaction" not in {
header.lower() for header in updated["metadata"]["requester_metadata"].get("headers", {})
}
assert "litellm-disable-message-redaction" not in {
header.lower() for header in updated["proxy_server_request"]["headers"]
}
assert "litellm-disable-message-redaction" not in {
header.lower() for header in updated["proxy_server_request"]["body"]["metadata"]["headers"]
}
assert "litellm-disable-message-redaction" not in {
header.lower() for header in (updated.get("litellm_metadata") or {}).get("headers", {})
}
@pytest.mark.parametrize(
"admin_metadata_kwargs",
[
{
"metadata": {
"logging": [
{
"callback_name": "langfuse",
"callback_type": "success_and_failure",
"callback_vars": {"turn_off_message_logging": False},
}
]
}
},
{
"team_metadata": {
"logging": [
{
"callback_name": "langfuse",
"callback_type": "success_and_failure",
"callback_vars": {"turn_off_message_logging": False},
}
]
}
},
],
)
@pytest.mark.asyncio
async def test_add_litellm_data_to_request_admin_callback_vars_turn_off_message_logging_overrides_global(
admin_metadata_kwargs,
):
from litellm.litellm_core_utils.initialize_dynamic_callback_params import (
initialize_standard_callback_dynamic_params,
)
from litellm.litellm_core_utils.redact_messages import should_redact_message_logging
request_mock = MagicMock(spec=Request)
request_mock.url.path = "/v1/chat/completions"
request_mock.url = MagicMock()
request_mock.url.__str__.return_value = "http://localhost/v1/chat/completions"
request_mock.method = "POST"
request_mock.query_params = {}
request_mock.headers = {"Content-Type": "application/json"}
request_mock.client = MagicMock()
request_mock.client.host = "127.0.0.1"
original_turn_off_message_logging = litellm.turn_off_message_logging
litellm.turn_off_message_logging = True
try:
updated = await add_litellm_data_to_request(
data={
"model": "gpt-3.5-turbo",
"messages": [{"role": "user", "content": "hello"}],
},
request=request_mock,
user_api_key_dict=UserAPIKeyAuth(api_key="hashed-key", **admin_metadata_kwargs),
proxy_config=MagicMock(),
general_settings={},
version="test-version",
)
assert updated.get("turn_off_message_logging") == "False"
dynamic_params = initialize_standard_callback_dynamic_params(updated)
assert dynamic_params.get("turn_off_message_logging") == "False"
assert should_redact_message_logging({"standard_callback_dynamic_params": dynamic_params}) is False
finally:
litellm.turn_off_message_logging = original_turn_off_message_logging
@pytest.mark.parametrize(
"admin_metadata_kwargs",
[
{
"metadata": {
"logging": [
{
"callback_name": "langfuse",
"callback_type": "success_and_failure",
"callback_vars": {"turn_off_message_logging": True},
}
]
}
},
{
"team_metadata": {
"logging": [
{
"callback_name": "langfuse",
"callback_type": "success_and_failure",
"callback_vars": {"turn_off_message_logging": True},
}
]
}
},
],
)
@pytest.mark.asyncio
async def test_add_litellm_data_to_request_admin_callback_vars_turn_off_message_logging_enables_redaction_when_global_off(
admin_metadata_kwargs,
):
from litellm.litellm_core_utils.initialize_dynamic_callback_params import (
initialize_standard_callback_dynamic_params,
)
from litellm.litellm_core_utils.redact_messages import should_redact_message_logging
request_mock = MagicMock(spec=Request)
request_mock.url.path = "/v1/chat/completions"
request_mock.url = MagicMock()
request_mock.url.__str__.return_value = "http://localhost/v1/chat/completions"
request_mock.method = "POST"
request_mock.query_params = {}
request_mock.headers = {"Content-Type": "application/json"}
request_mock.client = MagicMock()
request_mock.client.host = "127.0.0.1"
original_turn_off_message_logging = litellm.turn_off_message_logging
litellm.turn_off_message_logging = False
try:
updated = await add_litellm_data_to_request(
data={
"model": "gpt-3.5-turbo",
"messages": [{"role": "user", "content": "hello"}],
},
request=request_mock,
user_api_key_dict=UserAPIKeyAuth(api_key="hashed-key", **admin_metadata_kwargs),
proxy_config=MagicMock(),
general_settings={},
version="test-version",
)
assert updated.get("turn_off_message_logging") == "True"
dynamic_params = initialize_standard_callback_dynamic_params(updated)
assert dynamic_params.get("turn_off_message_logging") == "True"
assert should_redact_message_logging({"standard_callback_dynamic_params": dynamic_params}) is True
finally:
litellm.turn_off_message_logging = original_turn_off_message_logging
@pytest.mark.parametrize(
"auth_kwargs",
[
{"metadata": {"allow_client_message_redaction_opt_out": True}},
{"team_metadata": {"allow_client_message_redaction_opt_out": True}},
],
)
@pytest.mark.asyncio
async def test_add_litellm_data_to_request_allows_redaction_opt_out_with_admin_opt_in(
auth_kwargs,
):
request_mock = MagicMock(spec=Request)
request_mock.url.path = "/v1/chat/completions"
request_mock.url = MagicMock()
request_mock.url.__str__.return_value = "http://localhost/v1/chat/completions"
request_mock.method = "POST"
request_mock.query_params = {}
request_mock.headers = {
"Content-Type": "application/json",
"litellm-disable-message-redaction": "true",
}
request_mock.client = MagicMock()
request_mock.client.host = "127.0.0.1"
original_turn_off_message_logging = litellm.turn_off_message_logging
litellm.turn_off_message_logging = True
try:
updated = await add_litellm_data_to_request(
data={
"model": "gpt-3.5-turbo",
"messages": [{"role": "user", "content": "hello"}],
"turn_off_message_logging": False,
"metadata": {
"headers": {"litellm-disable-message-redaction": "true"},
"turn_off_message_logging": False,
},
"litellm_metadata": json.dumps({"headers": {"LiteLLM-Disable-Message-Redaction": "true"}}),
},
request=request_mock,
user_api_key_dict=UserAPIKeyAuth(api_key="hashed-key", **auth_kwargs),
proxy_config=MagicMock(),
general_settings={},
version="test-version",
)
finally:
litellm.turn_off_message_logging = original_turn_off_message_logging
assert updated["turn_off_message_logging"] is False
assert updated["metadata"]["turn_off_message_logging"] is False
assert "litellm-disable-message-redaction" in {header.lower() for header in updated["metadata"]["headers"]}
assert "litellm-disable-message-redaction" in {
header.lower() for header in updated["metadata"]["requester_metadata"].get("headers", {})
}
assert "litellm-disable-message-redaction" in {
header.lower() for header in updated["proxy_server_request"]["headers"]
}
assert "litellm-disable-message-redaction" in {
header.lower() for header in updated["proxy_server_request"]["body"]["metadata"]["headers"]
}
assert "litellm_metadata" not in updated
@pytest.mark.asyncio
async def test_add_litellm_data_to_request_honors_header_tags():
"""Header-supplied tags flow through to request metadata."""
from litellm.proxy.litellm_pre_call_utils import add_litellm_data_to_request
request_mock = MagicMock(spec=Request)
request_mock.url.path = "/v1/chat/completions"
request_mock.url = MagicMock()
request_mock.url.__str__.return_value = "http://localhost/v1/chat/completions"
request_mock.method = "POST"
request_mock.query_params = {}
request_mock.headers = {
"Content-Type": "application/json",
"x-litellm-tags": "production,ab-test",
}
request_mock.client = MagicMock()
request_mock.client.host = "127.0.0.1"
data = {"model": "gpt-3.5-turbo"}
user_api_key_dict = UserAPIKeyAuth(
api_key="hashed-key",
metadata={},
team_metadata={},
spend=0.0,
max_budget=100.0,
model_max_budget={},
team_spend=0.0,
team_max_budget=200.0,
)
updated = await add_litellm_data_to_request(
data=data,
request=request_mock,
user_api_key_dict=user_api_key_dict,
proxy_config=MagicMock(),
general_settings={},
version="test-version",
)
assert updated["metadata"].get("tags") == ["production", "ab-test"]
@pytest.mark.asyncio
async def test_add_litellm_data_to_request_preserves_caller_metadata_tags():
"""Caller-supplied metadata.tags are preserved and reach the router."""
from litellm.proxy.litellm_pre_call_utils import add_litellm_data_to_request
request_mock = MagicMock(spec=Request)
request_mock.url.path = "/v1/chat/completions"
request_mock.url = MagicMock()
request_mock.url.__str__.return_value = "http://localhost/v1/chat/completions"
request_mock.method = "POST"
request_mock.query_params = {}
request_mock.headers = {"Content-Type": "application/json"}
request_mock.client = MagicMock()
request_mock.client.host = "127.0.0.1"
data = {
"model": "gpt-3.5-turbo",
"metadata": {"tags": ["caller-tag"]},
}
user_api_key_dict = UserAPIKeyAuth(
api_key="hashed-key",
metadata={},
team_metadata={},
spend=0.0,
max_budget=100.0,
model_max_budget={},
team_spend=0.0,
team_max_budget=200.0,
)
updated = await add_litellm_data_to_request(
data=data,
request=request_mock,
user_api_key_dict=user_api_key_dict,
proxy_config=MagicMock(),
general_settings={},
version="test-version",
)
assert updated["metadata"].get("tags") == ["caller-tag"]
@pytest.mark.asyncio
async def test_add_litellm_data_to_request_unions_caller_header_tags_with_static_key_tags():
"""Caller-supplied `x-litellm-tags` must union with static key-level
tags, not overwrite them."""
from litellm.proxy.litellm_pre_call_utils import add_litellm_data_to_request
request_mock = MagicMock(spec=Request)
request_mock.url.path = "/v1/chat/completions"
request_mock.url = MagicMock()
request_mock.url.__str__.return_value = "http://localhost/v1/chat/completions"
request_mock.method = "POST"
request_mock.query_params = {}
request_mock.headers = {
"Content-Type": "application/json",
"x-litellm-tags": "tenant:1681",
}
request_mock.client = MagicMock()
request_mock.client.host = "127.0.0.1"
data = {"model": "gpt-3.5-turbo"}
user_api_key_dict = UserAPIKeyAuth(
api_key="hashed-key",
metadata={"tags": ["team:platform", "env:prod"]},
team_metadata={},
spend=0.0,
max_budget=100.0,
model_max_budget={},
team_spend=0.0,
team_max_budget=200.0,
)
updated = await add_litellm_data_to_request(
data=data,
request=request_mock,
user_api_key_dict=user_api_key_dict,
proxy_config=MagicMock(),
general_settings={},
version="test-version",
)
final_tags = updated["metadata"].get("tags") or []
assert "team:platform" in final_tags
assert "env:prod" in final_tags
assert "tenant:1681" in final_tags
@pytest.mark.asyncio
async def test_add_litellm_data_to_request_unions_caller_header_tags_with_static_team_tags():
"""Same union behavior must hold for team-level static tags."""
from litellm.proxy.litellm_pre_call_utils import add_litellm_data_to_request
request_mock = MagicMock(spec=Request)
request_mock.url.path = "/v1/chat/completions"
request_mock.url = MagicMock()
request_mock.url.__str__.return_value = "http://localhost/v1/chat/completions"
request_mock.method = "POST"
request_mock.query_params = {}
request_mock.headers = {
"Content-Type": "application/json",
"x-litellm-tags": "tenant:42",
}
request_mock.client = MagicMock()
request_mock.client.host = "127.0.0.1"
data = {"model": "gpt-3.5-turbo"}
user_api_key_dict = UserAPIKeyAuth(
api_key="hashed-key",
metadata={},
team_metadata={"tags": ["team:eng", "owner:platform"]},
spend=0.0,
max_budget=100.0,
model_max_budget={},
team_spend=0.0,
team_max_budget=200.0,
)
updated = await add_litellm_data_to_request(
data=data,
request=request_mock,
user_api_key_dict=user_api_key_dict,
proxy_config=MagicMock(),
general_settings={},
version="test-version",
)
final_tags = updated["metadata"].get("tags") or []
assert "team:eng" in final_tags
assert "owner:platform" in final_tags
assert "tenant:42" in final_tags
@pytest.mark.asyncio
async def test_add_litellm_data_to_request_unions_dedups_overlapping_caller_and_static_tags():
"""A tag that appears in both the static set and the caller header
must show up exactly once in the merged list."""
from litellm.proxy.litellm_pre_call_utils import add_litellm_data_to_request
request_mock = MagicMock(spec=Request)
request_mock.url.path = "/v1/chat/completions"
request_mock.url = MagicMock()
request_mock.url.__str__.return_value = "http://localhost/v1/chat/completions"
request_mock.method = "POST"
request_mock.query_params = {}
request_mock.headers = {
"Content-Type": "application/json",
"x-litellm-tags": "env:prod,tenant:7",
}
request_mock.client = MagicMock()
request_mock.client.host = "127.0.0.1"
data = {"model": "gpt-3.5-turbo"}
user_api_key_dict = UserAPIKeyAuth(
api_key="hashed-key",
metadata={"tags": ["env:prod", "team:platform"]},
team_metadata={},
spend=0.0,
max_budget=100.0,
model_max_budget={},
team_spend=0.0,
team_max_budget=200.0,
)
updated = await add_litellm_data_to_request(
data=data,
request=request_mock,
user_api_key_dict=user_api_key_dict,
proxy_config=MagicMock(),
general_settings={},
version="test-version",
)
final_tags = updated["metadata"].get("tags") or []
assert final_tags.count("env:prod") == 1
assert "team:platform" in final_tags
assert "tenant:7" in final_tags
@pytest.mark.asyncio
async def test_add_litellm_data_to_request_user_spend_and_budget():
from litellm.proxy.litellm_pre_call_utils import add_litellm_data_to_request
request_mock = MagicMock(spec=Request)
request_mock.url.path = "/v1/completions"
request_mock.url = MagicMock()
request_mock.url.__str__.return_value = "http://localhost/v1/completions"
request_mock.method = "POST"
request_mock.query_params = {}
request_mock.headers = {"Content-Type": "application/json"}
request_mock.client = MagicMock()
request_mock.client.host = "127.0.0.1"
data = {
"model": "gpt-3.5-turbo",
"messages": [{"role": "user", "content": "Hello"}],
}
user_api_key_dict = UserAPIKeyAuth(
api_key="hashed-key",
metadata={},
team_metadata={},
user_spend=150.0,
user_max_budget=500.0,
)
updated_data = await add_litellm_data_to_request(
data=data,
request=request_mock,
user_api_key_dict=user_api_key_dict,
proxy_config=MagicMock(),
general_settings={},
version="test-version",
)
metadata = updated_data.get("metadata", {})
assert metadata["user_api_key_user_spend"] == 150.0
assert metadata["user_api_key_user_max_budget"] == 500.0
@pytest.mark.asyncio
async def test_add_litellm_data_to_request_audio_transcription_multipart():
from litellm.proxy.litellm_pre_call_utils import add_litellm_data_to_request
# Setup request mock for /v1/audio/transcriptions
request_mock = MagicMock(spec=Request)
request_mock.url.path = "/v1/audio/transcriptions"
request_mock.url = MagicMock()
request_mock.url.__str__.return_value = "http://localhost/v1/audio/transcriptions"
request_mock.method = "POST"
request_mock.query_params = {}
request_mock.headers = {
"Content-Type": "multipart/form-data",
"Authorization": "Bearer sk-1234",
}
request_mock.client = MagicMock()
request_mock.client.host = "127.0.0.1"
# Simulate multipart data (metadata as string)
metadata_dict = {"tags": ["jobID:214590dsff09fds", "taskName:run_page_classification"]}
stringified_metadata = json.dumps(metadata_dict)
data = {
"model": "fake-openai-endpoint",
"metadata": stringified_metadata, # Simulating multipart-form field
"file": b"Fake audio bytes",
}
user_api_key_dict = UserAPIKeyAuth(
api_key="hashed-key",
metadata={},
team_metadata={},
spend=0.0,
max_budget=100.0,
model_max_budget={},
team_spend=0.0,
team_max_budget=200.0,
)
updated_data = await add_litellm_data_to_request(
data=data,
request=request_mock,
user_api_key_dict=user_api_key_dict,
proxy_config=MagicMock(),
general_settings={},
version="test-version",
)
# Assert metadata was parsed correctly
metadata_field = updated_data.get("metadata", {})
litellm_metadata = updated_data.get("litellm_metadata", {})
assert isinstance(metadata_field, dict)
assert "tags" in metadata_field
assert metadata_field["tags"] == [
"jobID:214590dsff09fds",
"taskName:run_page_classification",
]
@pytest.mark.asyncio
async def test_add_litellm_data_to_request_disabled_callbacks():
"""
Test that litellm_disabled_callbacks from key metadata is properly added to the request data.
"""
from litellm.proxy.litellm_pre_call_utils import add_litellm_data_to_request
# Setup mock request
request_mock = MagicMock(spec=Request)
request_mock.url.path = "/chat/completions"
request_mock.url = MagicMock()
request_mock.url.__str__.return_value = "http://localhost/chat/completions"
request_mock.method = "POST"
request_mock.query_params = {}
request_mock.headers = {"Content-Type": "application/json"}
request_mock.client = MagicMock()
request_mock.client.host = "127.0.0.1"
# Setup user API key with disabled callbacks in metadata
user_api_key_dict = UserAPIKeyAuth(
api_key="test_api_key",
user_id="test_user_id",
org_id="test_org_id",
metadata={"litellm_disabled_callbacks": ["langfuse", "langsmith", "datadog"]},
)
# Setup request data
data = {
"model": "gpt-3.5-turbo",
"messages": [{"role": "user", "content": "Hello"}],
}
# Setup proxy config
proxy_config = MagicMock()
# Call add_litellm_data_to_request
result = await add_litellm_data_to_request(
data=data,
request=request_mock,
user_api_key_dict=user_api_key_dict,
proxy_config=proxy_config,
)
# Verify that litellm_disabled_callbacks was added to the request data
assert "litellm_disabled_callbacks" in result
assert result["litellm_disabled_callbacks"] == ["langfuse", "langsmith", "datadog"]
# Verify that other data is still present
assert "model" in result
assert result["model"] == "gpt-3.5-turbo"
assert "messages" in result
@pytest.mark.asyncio
async def test_add_litellm_data_to_request_disabled_callbacks_empty():
"""
Test that litellm_disabled_callbacks is not added when it's empty.
"""
from litellm.proxy.litellm_pre_call_utils import add_litellm_data_to_request
# Setup mock request
request_mock = MagicMock(spec=Request)
request_mock.url.path = "/chat/completions"
request_mock.url = MagicMock()
request_mock.url.__str__.return_value = "http://localhost/chat/completions"
request_mock.method = "POST"
request_mock.query_params = {}
request_mock.headers = {"Content-Type": "application/json"}
request_mock.client = MagicMock()
request_mock.client.host = "127.0.0.1"
# Setup user API key with empty disabled callbacks
user_api_key_dict = UserAPIKeyAuth(
api_key="test_api_key",
user_id="test_user_id",
org_id="test_org_id",
metadata={"litellm_disabled_callbacks": []},
)
# Setup request data
data = {
"model": "gpt-3.5-turbo",
"messages": [{"role": "user", "content": "Hello"}],
}
# Setup proxy config
proxy_config = MagicMock()
# Call add_litellm_data_to_request
result = await add_litellm_data_to_request(
data=data,
request=request_mock,
user_api_key_dict=user_api_key_dict,
proxy_config=proxy_config,
)
# Verify that litellm_disabled_callbacks is not added when empty
assert "litellm_disabled_callbacks" not in result
# Verify that other data is still present
assert "model" in result
assert result["model"] == "gpt-3.5-turbo"
assert "messages" in result
@pytest.mark.asyncio
async def test_add_litellm_data_to_request_disabled_callbacks_not_present():
"""
Test that litellm_disabled_callbacks is not added when it's not present in metadata.
"""
from litellm.proxy.litellm_pre_call_utils import add_litellm_data_to_request
# Setup mock request
request_mock = MagicMock(spec=Request)
request_mock.url.path = "/chat/completions"
request_mock.url = MagicMock()
request_mock.url.__str__.return_value = "http://localhost/chat/completions"
request_mock.method = "POST"
request_mock.query_params = {}
request_mock.headers = {"Content-Type": "application/json"}
request_mock.client = MagicMock()
request_mock.client.host = "127.0.0.1"
# Setup user API key without disabled callbacks in metadata
user_api_key_dict = UserAPIKeyAuth(
api_key="test_api_key",
user_id="test_user_id",
org_id="test_org_id",
metadata={}, # No litellm_disabled_callbacks
)
# Setup request data
data = {
"model": "gpt-3.5-turbo",
"messages": [{"role": "user", "content": "Hello"}],
}
# Setup proxy config
proxy_config = MagicMock()
# Call add_litellm_data_to_request
result = await add_litellm_data_to_request(
data=data,
request=request_mock,
user_api_key_dict=user_api_key_dict,
proxy_config=proxy_config,
)
# Verify that litellm_disabled_callbacks is not added when not present
assert "litellm_disabled_callbacks" not in result
# Verify that other data is still present
assert "model" in result
assert result["model"] == "gpt-3.5-turbo"
assert "messages" in result
@pytest.mark.asyncio
async def test_add_litellm_data_to_request_disabled_callbacks_invalid_type():
"""
Test that litellm_disabled_callbacks is not added when it's not a list.
"""
from litellm.proxy.litellm_pre_call_utils import add_litellm_data_to_request
# Setup mock request
request_mock = MagicMock(spec=Request)
request_mock.url.path = "/chat/completions"
request_mock.url = MagicMock()
request_mock.url.__str__.return_value = "http://localhost/chat/completions"
request_mock.method = "POST"
request_mock.query_params = {}
request_mock.headers = {"Content-Type": "application/json"}
request_mock.client = MagicMock()
request_mock.client.host = "127.0.0.1"
# Setup user API key with invalid disabled callbacks type
user_api_key_dict = UserAPIKeyAuth(
api_key="test_api_key",
user_id="test_user_id",
org_id="test_org_id",
metadata={"litellm_disabled_callbacks": "not_a_list"}, # Should be a list
)
# Setup request data
data = {
"model": "gpt-3.5-turbo",
"messages": [{"role": "user", "content": "Hello"}],
}
# Setup proxy config
proxy_config = MagicMock()
# Call add_litellm_data_to_request
result = await add_litellm_data_to_request(
data=data,
request=request_mock,
user_api_key_dict=user_api_key_dict,
proxy_config=proxy_config,
)
# Verify that litellm_disabled_callbacks is not added when invalid type
assert "litellm_disabled_callbacks" not in result
# Verify that other data is still present
assert "model" in result
assert result["model"] == "gpt-3.5-turbo"
assert "messages" in result
@pytest.mark.asyncio
async def test_add_litellm_data_to_request_disabled_callbacks_with_logging_settings():
"""
Test that litellm_disabled_callbacks works correctly alongside logging settings.
"""
from litellm.proxy.litellm_pre_call_utils import add_litellm_data_to_request
# Setup mock request
request_mock = MagicMock(spec=Request)
request_mock.url.path = "/chat/completions"
request_mock.url = MagicMock()
request_mock.url.__str__.return_value = "http://localhost/chat/completions"
request_mock.method = "POST"
request_mock.query_params = {}
request_mock.headers = {"Content-Type": "application/json"}
request_mock.client = MagicMock()
request_mock.client.host = "127.0.0.1"
# Setup user API key with both logging settings and disabled callbacks
user_api_key_dict = UserAPIKeyAuth(
api_key="test_api_key",
user_id="test_user_id",
org_id="test_org_id",
metadata={
"logging": [
{
"callback_name": "langfuse",
"callback_type": "success",
"callback_vars": {},
}
],
"litellm_disabled_callbacks": ["langsmith", "datadog"],
},
)
# Setup request data
data = {
"model": "gpt-3.5-turbo",
"messages": [{"role": "user", "content": "Hello"}],
}
# Setup proxy config
proxy_config = MagicMock()
# Call add_litellm_data_to_request
result = await add_litellm_data_to_request(
data=data,
request=request_mock,
user_api_key_dict=user_api_key_dict,
proxy_config=proxy_config,
)
# Verify that both logging settings and disabled callbacks are handled correctly
assert "litellm_disabled_callbacks" in result
assert result["litellm_disabled_callbacks"] == ["langsmith", "datadog"]
# Verify that other data is still present
assert "model" in result
assert result["model"] == "gpt-3.5-turbo"
assert "messages" in result
def test_key_dynamic_logging_settings():
"""
Test KeyAndTeamLoggingSettings.get_key_dynamic_logging_settings method with arize and langfuse callbacks
"""
# Test with arize logging
key_with_arize = UserAPIKeyAuth(
api_key="test-key",
metadata={"logging": [{"callback_name": "arize", "callback_type": "success"}]},
team_metadata={},
)
result = KeyAndTeamLoggingSettings.get_key_dynamic_logging_settings(key_with_arize)
assert result == [{"callback_name": "arize", "callback_type": "success"}]
# Test with langfuse logging
key_with_langfuse = UserAPIKeyAuth(
api_key="test-key",
metadata={"logging": [{"callback_name": "langfuse", "callback_type": "success"}]},
team_metadata={},
)
result = KeyAndTeamLoggingSettings.get_key_dynamic_logging_settings(key_with_langfuse)
assert result == [{"callback_name": "langfuse", "callback_type": "success"}]
# Test with no logging metadata
key_without_logging = UserAPIKeyAuth(api_key="test-key", metadata={}, team_metadata={})
result = KeyAndTeamLoggingSettings.get_key_dynamic_logging_settings(key_without_logging)
assert result is None
def test_team_dynamic_logging_settings():
"""
Test KeyAndTeamLoggingSettings.get_team_dynamic_logging_settings method with arize and langfuse callbacks
"""
# Test with arize team logging
key_with_team_arize = UserAPIKeyAuth(
api_key="test-key",
metadata={},
team_metadata={"logging": [{"callback_name": "arize", "callback_type": "failure"}]},
)
result = KeyAndTeamLoggingSettings.get_team_dynamic_logging_settings(key_with_team_arize)
assert result == [{"callback_name": "arize", "callback_type": "failure"}]
# Test with langfuse team logging
key_with_team_langfuse = UserAPIKeyAuth(
api_key="test-key",
metadata={},
team_metadata={"logging": [{"callback_name": "langfuse", "callback_type": "success"}]},
)
result = KeyAndTeamLoggingSettings.get_team_dynamic_logging_settings(key_with_team_langfuse)
assert result == [{"callback_name": "langfuse", "callback_type": "success"}]
# Test with no team logging metadata
key_without_team_logging = UserAPIKeyAuth(api_key="test-key", metadata={}, team_metadata={})
result = KeyAndTeamLoggingSettings.get_team_dynamic_logging_settings(key_without_team_logging)
assert result is None
def test_key_dynamic_logging_settings_decrypts_callback_vars(monkeypatch):
"""Encrypted callback_vars on the key are decrypted before downstream use."""
from litellm.proxy.common_utils.callback_utils import encrypt_callback_vars
monkeypatch.setenv("LITELLM_SALT_KEY", "test-salt-32-bytes-aaaaaaaaaaaaaa")
encrypted_metadata = encrypt_callback_vars(
{
"logging": [
{
"callback_name": "langfuse",
"callback_type": "success",
"callback_vars": {
"langfuse_public_key": "pk-real",
"langfuse_secret_key": "sk-real",
},
}
]
}
)
cv_on_disk = encrypted_metadata["logging"][0]["callback_vars"]
assert cv_on_disk["langfuse_secret_key"] != "sk-real" # sanity: stored encrypted
key = UserAPIKeyAuth(api_key="t", metadata=encrypted_metadata, team_metadata={})
result = KeyAndTeamLoggingSettings.get_key_dynamic_logging_settings(key)
cv = result[0]["callback_vars"]
assert cv["langfuse_secret_key"] == "sk-real"
assert cv["langfuse_public_key"] == "pk-real"
def test_team_dynamic_logging_settings_decrypts_callback_vars(monkeypatch):
"""Encrypted callback_vars on the team are decrypted before downstream use."""
from litellm.proxy.common_utils.callback_utils import encrypt_callback_vars
monkeypatch.setenv("LITELLM_SALT_KEY", "test-salt-32-bytes-aaaaaaaaaaaaaa")
encrypted_team = encrypt_callback_vars(
{
"logging": [
{
"callback_name": "langfuse",
"callback_type": "failure",
"callback_vars": {"langfuse_secret_key": "team-sk-real"},
}
]
}
)
key = UserAPIKeyAuth(api_key="t", metadata={}, team_metadata=encrypted_team)
result = KeyAndTeamLoggingSettings.get_team_dynamic_logging_settings(key)
assert result[0]["callback_vars"]["langfuse_secret_key"] == "team-sk-real"
def test_get_dynamic_logging_metadata_with_arize_team_logging():
"""
Test _get_dynamic_logging_metadata function with arize team logging and dynamic parameters
"""
# Setup user with arize team logging including callback_vars
user_api_key_dict = UserAPIKeyAuth(
api_key="test-key",
metadata={},
team_metadata={
"logging": [
{
"callback_name": "arize",
"callback_type": "success",
"callback_vars": {
"arize_api_key": "test_arize_api_key",
"arize_space_id": "test_arize_space_id",
},
}
]
},
)
# Mock proxy_config (not used in this test path since we have team dynamic logging)
mock_proxy_config = MagicMock()
# Call the function
result = _get_dynamic_logging_metadata(user_api_key_dict=user_api_key_dict, proxy_config=mock_proxy_config)
# Verify the result
assert result is not None
assert isinstance(result, TeamCallbackMetadata)
assert result.success_callback == ["arize"]
assert result.callback_vars is not None
assert result.callback_vars["arize_api_key"] == "test_arize_api_key"
assert result.callback_vars["arize_space_id"] == "test_arize_space_id"
def test_add_team_callback_rejects_env_reference():
with pytest.raises(PydanticValidationError) as exc_info:
AddTeamCallback(
callback_name="langfuse",
callback_type="success",
callback_vars={"langfuse_secret_key": "os.environ/LANGFUSE_SECRET_KEY_TEMP"},
)
assert "os.environ/" in str(exc_info.value)
def test_get_dynamic_logging_metadata_ignores_env_reference_from_key_metadata(
monkeypatch,
):
monkeypatch.setenv("LANGFUSE_SECRET_KEY_TEMP", "server-side-secret")
monkeypatch.setattr(
litellm.utils,
"get_secret",
lambda *args, **kwargs: pytest.fail("get_secret should not be called"),
)
user_api_key_dict = UserAPIKeyAuth(
api_key="test-key",
metadata={
"logging": [
{
"callback_name": "langfuse",
"callback_type": "success",
"callback_vars": {
"langfuse_secret_key": "os.environ/LANGFUSE_SECRET_KEY_TEMP",
},
}
]
},
team_metadata={},
)
result = _get_dynamic_logging_metadata(user_api_key_dict=user_api_key_dict, proxy_config=MagicMock())
assert result is None
def test_get_num_retries_from_request():
"""
Test LiteLLMProxyRequestSetup._get_num_retries_from_request method
"""
# Test case 1: Header is present with valid integer string
headers_with_retries = {"x-litellm-num-retries": "3"}
result = LiteLLMProxyRequestSetup._get_num_retries_from_request(headers_with_retries)
assert result == 3
# Test case 2: Header is not present
headers_without_retries = {"Content-Type": "application/json"}
result = LiteLLMProxyRequestSetup._get_num_retries_from_request(headers_without_retries)
assert result is None
# Test case 3: Empty headers dictionary
empty_headers = {}
result = LiteLLMProxyRequestSetup._get_num_retries_from_request(empty_headers)
assert result is None
# Test case 4: Header present with zero value
headers_with_zero = {"x-litellm-num-retries": "0"}
result = LiteLLMProxyRequestSetup._get_num_retries_from_request(headers_with_zero)
assert result == 0
# Test case 5: Header present with large number
headers_with_large_number = {"x-litellm-num-retries": "100"}
result = LiteLLMProxyRequestSetup._get_num_retries_from_request(headers_with_large_number)
assert result == 100
# Test case 6: Multiple headers with num retries header
headers_multiple = {
"Content-Type": "application/json",
"x-litellm-num-retries": "5",
"Authorization": "Bearer token",
}
result = LiteLLMProxyRequestSetup._get_num_retries_from_request(headers_multiple)
assert result == 5
# Test case 7: Header present with invalid value (should raise ValueError when int() is called)
headers_with_invalid = {"x-litellm-num-retries": "invalid"}
with pytest.raises(ValueError, match="invalid literal for int\\(\\) with base"):
LiteLLMProxyRequestSetup._get_num_retries_from_request(headers_with_invalid)
# Test case 8: Header present with float string (should raise ValueError when int() is called)
headers_with_float = {"x-litellm-num-retries": "3.5"}
with pytest.raises(ValueError, match="invalid literal for int\\(\\) with base"):
LiteLLMProxyRequestSetup._get_num_retries_from_request(headers_with_float)
# Test case 9: Header present with negative number
headers_with_negative = {"x-litellm-num-retries": "-1"}
result = LiteLLMProxyRequestSetup._get_num_retries_from_request(headers_with_negative)
assert result == -1
def test_get_keepalive_seconds_from_request():
"""
Test LiteLLMProxyRequestSetup._get_keepalive_seconds_from_request method
"""
# Header present with valid float string
headers_with_keepalive = {"x-litellm-keepalive-seconds": "15"}
result = LiteLLMProxyRequestSetup._get_keepalive_seconds_from_request(headers_with_keepalive)
assert result == 15.0
# Header not present
result = LiteLLMProxyRequestSetup._get_keepalive_seconds_from_request({"Content-Type": "application/json"})
assert result is None
# Empty headers dictionary
result = LiteLLMProxyRequestSetup._get_keepalive_seconds_from_request({})
assert result is None
# Header present with a fractional value
result = LiteLLMProxyRequestSetup._get_keepalive_seconds_from_request({"x-litellm-keepalive-seconds": "1.5"})
assert result == 1.5
# Header present with invalid value raises ValueError, matching the other
# x-litellm-* numeric header helpers (_get_timeout_from_request, etc.)
with pytest.raises(ValueError, match="could not convert string to float: 'not-a-number"):
LiteLLMProxyRequestSetup._get_keepalive_seconds_from_request({"x-litellm-keepalive-seconds": "not-a-number"})
def test_add_litellm_data_for_backend_llm_call_merges_keepalive_seconds_header():
"""
The x-litellm-keepalive-seconds header must be merged into the data dict
that add_litellm_data_to_request later data.update()s onto the request body,
the same way x-litellm-timeout/x-litellm-num-retries already are.
"""
result = LiteLLMProxyRequestSetup.add_litellm_data_for_backend_llm_call(
headers={"x-litellm-keepalive-seconds": "20"},
request_data={},
user_api_key_dict=UserAPIKeyAuth(api_key="sk-test"),
)
assert result.get("keepalive_seconds") == 20.0
def test_add_user_api_key_auth_to_request_metadata():
"""
Test that add_user_api_key_auth_to_request_metadata properly adds user API key authentication data to request metadata
"""
# Setup test data
data = {
"model": "gpt-3.5-turbo",
"messages": [{"role": "user", "content": "Hello"}],
"litellm_metadata": {}, # This will be the metadata variable name
}
user_api_key_dict = UserAPIKeyAuth(
api_key="hashed-test-key-123",
user_id="test-user-123",
org_id="test-org-456",
team_id="test-team-789",
key_alias="test-key-alias",
user_email="test@example.com",
team_alias="test-team-alias",
end_user_id="test-end-user-123",
request_route="/chat/completions",
end_user_max_budget=500.0,
)
metadata_variable_name = "litellm_metadata"
# Call the function
result = LiteLLMProxyRequestSetup.add_user_api_key_auth_to_request_metadata(
data=data,
user_api_key_dict=user_api_key_dict,
_metadata_variable_name=metadata_variable_name,
)
# Verify the metadata was properly added
metadata = result[metadata_variable_name]
# Check that user API key information was added
assert metadata["user_api_key_hash"] == "hashed-test-key-123"
assert metadata["user_api_key_alias"] == "test-key-alias"
assert metadata["user_api_key_team_id"] == "test-team-789"
assert metadata["user_api_key_user_id"] == "test-user-123"
assert metadata["user_api_key_org_id"] == "test-org-456"
assert metadata["user_api_key_team_alias"] == "test-team-alias"
assert metadata["user_api_key_end_user_id"] == "test-end-user-123"
assert metadata["user_api_key_user_email"] == "test@example.com"
assert metadata["user_api_key_request_route"] == "/chat/completions"
# Check that the hashed API key was added
assert metadata["user_api_key"] == "hashed-test-key-123"
# Check that end user max budget was added
assert metadata["user_api_end_user_max_budget"] == 500.0
# Verify original data is preserved
assert result["model"] == "gpt-3.5-turbo"
assert result["messages"] == [{"role": "user", "content": "Hello"}]
def _auth_with_callback_credentials() -> UserAPIKeyAuth:
return UserAPIKeyAuth(
api_key="hashed-test-key-123",
key_alias="test-key-alias",
team_id="test-team-789",
team_alias="test-team-alias",
metadata={
"logging": [{"callback_name": "langfuse", "callback_vars": {"langfuse_secret_key": "sk-KEY-CANARY"}}],
"rpm_limit_type": "guaranteed_throughput",
},
team_metadata={
"callback_settings": {"langfuse": {"callback_vars": {"langfuse_secret_key": "sk-TEAM-CANARY"}}},
"secret_manager_settings": {"vault_token": "vt-TEAM-CANARY"},
"model_rpm_limit": {"gpt-4": 10},
},
project_metadata={
"logging": [{"callback_vars": {"langfuse_secret_key": "sk-PROJECT-CANARY"}}],
"project_tier": "gold",
},
organization_metadata={
"secret_manager_settings": {"vault_token": "vt-ORG-CANARY"},
"org_tier": "platinum",
},
)
def test_stamped_auth_object_carries_no_callback_credentials():
"""
Regression (LIT-5487): the UserAPIKeyAuth stamped into request metadata reaches every
raw-metadata logging integration, so it must not carry team/key callback credentials.
"""
user_api_key_dict = _auth_with_callback_credentials()
otel_span = object()
user_api_key_dict.parent_otel_span = otel_span
user_api_key_dict.budget_reservation = {"amount": 1.0}
user_api_key_dict.via_virtual_key = True
data = {"litellm_metadata": {}}
result = LiteLLMProxyRequestSetup.add_user_api_key_auth_to_request_metadata(
data=data,
user_api_key_dict=user_api_key_dict,
_metadata_variable_name="litellm_metadata",
)
stamped = result["litellm_metadata"]["user_api_key_auth"]
emitted = json.dumps(
{
"metadata": stamped.metadata,
"team_metadata": stamped.team_metadata,
"project_metadata": stamped.project_metadata,
"organization_metadata": stamped.organization_metadata,
},
default=str,
)
assert "sk-KEY-CANARY" not in emitted
assert "sk-TEAM-CANARY" not in emitted
assert "vt-TEAM-CANARY" not in emitted
assert "sk-PROJECT-CANARY" not in emitted
assert "vt-ORG-CANARY" not in emitted
# consumers keep the type and the non-credential slots they read
assert isinstance(stamped, UserAPIKeyAuth)
assert stamped.key_alias == "test-key-alias"
assert stamped.team_id == "test-team-789"
assert stamped.team_alias == "test-team-alias"
assert stamped.api_key == "hashed-test-key-123"
assert stamped.metadata["rpm_limit_type"] == "guaranteed_throughput"
assert stamped.team_metadata["model_rpm_limit"] == {"gpt-4": 10}
assert stamped.project_metadata["project_tier"] == "gold"
assert stamped.organization_metadata["org_tier"] == "platinum"
# server-only markers are excluded from model_dump, so rebuilding the object
# instead of copying it would silently drop them
assert stamped.via_virtual_key is True
assert stamped.budget_reservation == {"amount": 1.0}
assert stamped.parent_otel_span is otel_span
def test_stamping_does_not_mutate_the_cached_auth_object():
"""
Regression (LIT-5487): UserAPIKeyAuth is cached and model_copy is shallow, so stripping
in place would poison the shared dicts and silently kill team callbacks fleet-wide.
"""
user_api_key_dict = _auth_with_callback_credentials()
metadata_before = copy.deepcopy(user_api_key_dict.metadata)
team_metadata_before = copy.deepcopy(user_api_key_dict.team_metadata)
LiteLLMProxyRequestSetup.add_user_api_key_auth_to_request_metadata(
data={"litellm_metadata": {}},
user_api_key_dict=user_api_key_dict,
_metadata_variable_name="litellm_metadata",
)
assert user_api_key_dict.metadata == metadata_before
assert user_api_key_dict.team_metadata == team_metadata_before
def test_management_endpoint_metadata_drops_callback_credentials():
"""
Regression (LIT-5487): user_api_key_auth_metadata is part of StandardLoggingPayload, so a
callback_settings-shaped team must not push credentials into it.
"""
data = {"litellm_metadata": {}}
result = LiteLLMProxyRequestSetup.add_management_endpoint_metadata_to_request_metadata(
data=data,
management_endpoint_metadata={
"callback_settings": {"langfuse": {"callback_vars": {"langfuse_secret_key": "sk-TEAM-CANARY"}}},
"secret_manager_settings": {"vault_token": "vt-TEAM-CANARY"},
"logging": [{"callback_vars": {"langfuse_secret_key": "sk-LOGGING-CANARY"}}],
"other_field": "value",
},
_metadata_variable_name="litellm_metadata",
)
auth_metadata = result["litellm_metadata"]["user_api_key_auth_metadata"]
emitted = json.dumps(auth_metadata, default=str)
assert "sk-TEAM-CANARY" not in emitted
assert "vt-TEAM-CANARY" not in emitted
assert "sk-LOGGING-CANARY" not in emitted
assert auth_metadata["other_field"] == "value"
@pytest.mark.parametrize(
"data, model_group_settings, expected_headers_added",
[
# Test case 1: Model is in forward_client_headers_to_llm_api list
(
{"model": "gpt-4", "messages": [{"role": "user", "content": "Hello"}]},
MagicMock(forward_client_headers_to_llm_api=["gpt-4"]),
True,
),
# Test case 2: Model is not in forward_client_headers_to_llm_api list
(
{"model": "claude-3", "messages": [{"role": "user", "content": "Hello"}]},
MagicMock(forward_client_headers_to_llm_api=["gpt-4"]),
False,
),
# Test case 3: Model group settings is None
(
{"model": "gpt-4", "messages": [{"role": "user", "content": "Hello"}]},
None,
False,
),
# Test case 4: forward_client_headers_to_llm_api is None
(
{"model": "gpt-4", "messages": [{"role": "user", "content": "Hello"}]},
MagicMock(forward_client_headers_to_llm_api=None),
False,
),
# Test case 5: Data has no model
(
{"messages": [{"role": "user", "content": "Hello"}]},
MagicMock(forward_client_headers_to_llm_api=["gpt-4"]),
False,
),
# Test case 6: Model is None
(
{"model": None, "messages": [{"role": "user", "content": "Hello"}]},
MagicMock(forward_client_headers_to_llm_api=["gpt-4"]),
False,
),
],
)
def test_add_headers_to_llm_call_by_model_group(data, model_group_settings, expected_headers_added):
"""
Test LiteLLMProxyRequestSetup.add_headers_to_llm_call_by_model_group method
This tests various scenarios:
1. When model is in the forward_client_headers_to_llm_api list
2. When model is not in the list
3. When model_group_settings is None
4. When forward_client_headers_to_llm_api is None
5. When data has no model
6. When model is None
"""
import litellm
# Setup test headers and user API key
headers = {
"Authorization": "Bearer token123",
"User-Agent": "test-client/1.0",
"X-Custom-Header": "custom-value",
}
user_api_key_dict = UserAPIKeyAuth(api_key="test-key", user_id="test-user", org_id="test-org")
# Mock the model_group_settings
original_model_group_settings = getattr(litellm, "model_group_settings", None)
litellm.model_group_settings = model_group_settings
try:
# Mock the add_headers_to_llm_call method to return expected headers
expected_returned_headers = {
"X-LiteLLM-User": "test-user",
"X-LiteLLM-Org": "test-org",
}
with patch.object(
LiteLLMProxyRequestSetup,
"add_headers_to_llm_call",
return_value=expected_returned_headers if expected_headers_added else {},
) as mock_add_headers:
# Make a copy of original data to verify it's not mutated unexpectedly
original_data = copy.deepcopy(data)
# Call the method under test
result = LiteLLMProxyRequestSetup.add_headers_to_llm_call_by_model_group(
data=data, headers=headers, user_api_key_dict=user_api_key_dict
)
# Verify the result
assert result is not None
assert isinstance(result, dict)
if expected_headers_added:
# Verify that add_headers_to_llm_call was called
mock_add_headers.assert_called_once_with(headers, user_api_key_dict)
# Verify that headers were added to the data
assert "headers" in result
assert result["headers"] == expected_returned_headers
else:
# Verify that add_headers_to_llm_call was not called
mock_add_headers.assert_not_called()
# Verify that no headers were added
assert "headers" not in result or result.get("headers") is None
# Verify that original data fields are preserved
for key, value in original_data.items():
if key != "headers": # headers might be added
assert result[key] == value
finally:
# Restore original model_group_settings
litellm.model_group_settings = original_model_group_settings
def test_add_headers_to_llm_call_by_model_group_empty_headers_returned():
"""
Test that when add_headers_to_llm_call returns empty dict, no headers are added to data
"""
import litellm
# Setup test data
data = {"model": "gpt-4", "messages": [{"role": "user", "content": "Hello"}]}
headers = {"Authorization": "Bearer token123"}
user_api_key_dict = UserAPIKeyAuth(api_key="test-key")
# Mock model_group_settings with model in the list
mock_settings = MagicMock(forward_client_headers_to_llm_api=["gpt-4"])
original_model_group_settings = getattr(litellm, "model_group_settings", None)
litellm.model_group_settings = mock_settings
try:
with patch.object(
LiteLLMProxyRequestSetup,
"add_headers_to_llm_call",
return_value={}, # Return empty dict
) as mock_add_headers:
result = LiteLLMProxyRequestSetup.add_headers_to_llm_call_by_model_group(
data=data, headers=headers, user_api_key_dict=user_api_key_dict
)
# Verify that add_headers_to_llm_call was called
mock_add_headers.assert_called_once_with(headers, user_api_key_dict)
# Verify that no headers were added since returned headers were empty
assert "headers" not in result
# Verify original data is preserved
assert result["model"] == "gpt-4"
assert result["messages"] == [{"role": "user", "content": "Hello"}]
finally:
# Restore original model_group_settings
litellm.model_group_settings = original_model_group_settings
def test_add_headers_to_llm_call_by_model_group_existing_headers_in_data():
"""
Test that existing headers in data are overwritten when new headers are added
"""
import litellm
# Setup test data with existing headers
data = {
"model": "gpt-4",
"messages": [{"role": "user", "content": "Hello"}],
"headers": {"Existing-Header": "existing-value"},
}
headers = {"Authorization": "Bearer token123"}
user_api_key_dict = UserAPIKeyAuth(api_key="test-key")
# Mock model_group_settings with model in the list
mock_settings = MagicMock(forward_client_headers_to_llm_api=["gpt-4"])
original_model_group_settings = getattr(litellm, "model_group_settings", None)
litellm.model_group_settings = mock_settings
try:
new_headers = {"X-LiteLLM-User": "test-user"}
with patch.object(
LiteLLMProxyRequestSetup,
"add_headers_to_llm_call",
return_value=new_headers,
) as mock_add_headers:
result = LiteLLMProxyRequestSetup.add_headers_to_llm_call_by_model_group(
data=data, headers=headers, user_api_key_dict=user_api_key_dict
)
# Verify that add_headers_to_llm_call was called
mock_add_headers.assert_called_once_with(headers, user_api_key_dict)
# Verify that headers were overwritten
assert "headers" in result
assert result["headers"] == new_headers
assert result["headers"] != {"Existing-Header": "existing-value"}
# Verify original data is preserved
assert result["model"] == "gpt-4"
assert result["messages"] == [{"role": "user", "content": "Hello"}]
finally:
# Restore original model_group_settings
litellm.model_group_settings = original_model_group_settings
from typing import Optional
from fastapi.responses import Response
from litellm.integrations.custom_logger import CustomLogger
from litellm.proxy.common_request_processing import ProxyBaseLLMRequestProcessing
from litellm.proxy.utils import ProxyLogging
from litellm.types.utils import StandardLoggingPayload
class TestCustomLogger(CustomLogger):
def __init__(self):
self.standard_logging_object: Optional[StandardLoggingPayload] = None
super().__init__()
async def async_log_success_event(self, kwargs, response_obj, start_time, end_time):
print(f"SUCCESS CALLBACK CALLED! kwargs keys: {list(kwargs.keys())}")
self.standard_logging_object = kwargs.get("standard_logging_object")
print(f"Captured standard_logging_object: {self.standard_logging_object}")
async def async_log_failure_event(self, kwargs, response_obj, start_time, end_time):
print(f"FAILURE CALLBACK CALLED! kwargs keys: {list(kwargs.keys())}")
@pytest.mark.asyncio
async def test_add_litellm_metadata_from_request_headers():
"""
Test that add_litellm_metadata_from_request_headers properly adds litellm metadata from request headers,
makes an LLM request using base_process_llm_request, sleeps for 3 seconds, and checks standard_logging_payload has spend_logs_metadata from headers
Relevant issue: https://github.com/BerriAI/litellm/issues/14008
"""
# Set up test logger
litellm._turn_on_debug()
test_logger = TestCustomLogger()
original_callbacks = litellm.callbacks
litellm.callbacks = [test_logger]
try:
# Prepare test data (ensure no streaming, add mock_response and api_key to route to litellm.acompletion)
headers = {
"x-litellm-spend-logs-metadata": '{"user_id": "12345", "project_id": "proj_abc", "request_type": "chat_completion", "timestamp": "2025-09-02T10:30:00Z"}'
}
data = {
"model": "gpt-4",
"messages": [{"role": "user", "content": "Hello"}],
"stream": False,
"mock_response": "Hi",
"api_key": "fake-key",
}
# Create mock request with headers
mock_request = MagicMock(spec=Request)
mock_request.headers = headers
mock_request.url.path = "/chat/completions"
# Create mock response
mock_fastapi_response = MagicMock(spec=Response)
# Create mock user API key dict
mock_user_api_key_dict = UserAPIKeyAuth(
api_key="test-key",
user_id="test-user",
org_id="test-org",
metadata={"allow_client_mock_response": True},
)
# Create mock proxy logging object
mock_proxy_logging_obj = MagicMock(spec=ProxyLogging)
# Create async functions for the hooks
async def mock_during_call_hook(*args, **kwargs):
return None
async def mock_pre_call_hook(*args, **kwargs):
return data
async def mock_post_call_success_hook(*args, **kwargs):
# Return the response unchanged
return kwargs.get("response", args[2] if len(args) > 2 else None)
mock_proxy_logging_obj.during_call_hook = mock_during_call_hook
mock_proxy_logging_obj.pre_call_hook = mock_pre_call_hook
mock_proxy_logging_obj.post_call_success_hook = mock_post_call_success_hook
# Create mock proxy config
mock_proxy_config = MagicMock()
# Create mock general settings
general_settings = {}
# Create mock select_data_generator with correct signature
def mock_select_data_generator(response=None, user_api_key_dict=None, request_data=None):
async def mock_generator():
yield "data: " + json.dumps({"choices": [{"delta": {"content": "Hello"}}]}) + "\n\n"
yield "data: [DONE]\n\n"
return mock_generator()
# Create the processor
processor = ProxyBaseLLMRequestProcessing(data=data)
# Call base_process_llm_request (it will use the mock_response="Hi" parameter)
result = await processor.base_process_llm_request(
request=mock_request,
fastapi_response=mock_fastapi_response,
user_api_key_dict=mock_user_api_key_dict,
route_type="acompletion",
proxy_logging_obj=mock_proxy_logging_obj,
general_settings=general_settings,
proxy_config=mock_proxy_config,
select_data_generator=mock_select_data_generator,
llm_router=None,
model="gpt-4",
is_streaming_request=False,
)
# Sleep for 3 seconds to allow logging to complete
await asyncio.sleep(3)
# Check if standard_logging_object was set
assert test_logger.standard_logging_object is not None, (
"standard_logging_object should be populated after LLM request"
)
# Verify the logging object contains expected metadata
standard_logging_obj = test_logger.standard_logging_object
print(f"Standard logging object captured: {json.dumps(standard_logging_obj, indent=4, default=str)}")
SPEND_LOGS_METADATA = standard_logging_obj["metadata"]["spend_logs_metadata"]
assert SPEND_LOGS_METADATA == dict(json.loads(headers["x-litellm-spend-logs-metadata"])), (
"spend_logs_metadata should be the same as the headers"
)
finally:
litellm.callbacks = original_callbacks
@pytest.mark.asyncio
async def test_anthropic_messages_standard_logging_object_matches_fixture():
"""
Regression: /v1/messages calls routed to non-Anthropic providers should keep
call_type=anthropic_messages in standard logging payloads.
"""
litellm._turn_on_debug()
test_logger = TestCustomLogger()
original_callbacks = litellm.callbacks
litellm.callbacks = [test_logger]
try:
data = {
"model": "gemini/gemini-2.5-flash",
"messages": [{"role": "user", "content": "Hi."}],
"stream": False,
"mock_response": "Hello! How can I help you today?",
"api_key": "fake-key",
"max_tokens": 4096,
}
mock_request = MagicMock(spec=Request)
mock_request.headers = {"user-agent": "PostmanRuntime/7.53.0"}
mock_request.url.path = "/v1/messages"
mock_request.url = MagicMock()
mock_request.url.__str__.return_value = "http://localhost/v1/messages"
mock_request.method = "POST"
mock_request.query_params = {}
mock_request.client = MagicMock()
mock_request.client.host = "127.0.0.1"
mock_fastapi_response = MagicMock(spec=Response)
mock_user_api_key_dict = UserAPIKeyAuth(
api_key="test-key",
user_id="default_user_id",
metadata={"allow_client_mock_response": True},
)
mock_proxy_logging_obj = MagicMock(spec=ProxyLogging)
async def mock_during_call_hook(*args, **kwargs):
return None
async def mock_pre_call_hook(*args, **kwargs):
return data
async def mock_post_call_success_hook(*args, **kwargs):
return kwargs.get("response", args[2] if len(args) > 2 else None)
mock_proxy_logging_obj.during_call_hook = mock_during_call_hook
mock_proxy_logging_obj.pre_call_hook = mock_pre_call_hook
mock_proxy_logging_obj.post_call_success_hook = mock_post_call_success_hook
processor = ProxyBaseLLMRequestProcessing(data=data)
await processor.base_process_llm_request(
request=mock_request,
fastapi_response=mock_fastapi_response,
user_api_key_dict=mock_user_api_key_dict,
route_type="anthropic_messages",
proxy_logging_obj=mock_proxy_logging_obj,
general_settings={},
proxy_config=MagicMock(),
select_data_generator=None,
llm_router=None,
model="gemini/gemini-2.5-flash",
is_streaming_request=False,
)
await asyncio.sleep(3)
assert test_logger.standard_logging_object is not None
actual = test_logger.standard_logging_object
expected = {
"call_type": "anthropic_messages",
"status": "success",
"model": "gemini/gemini-2.5-flash",
}
# Compare only stable fields from the saved proxy log snapshot.
actual_projection = {
"call_type": actual.get("call_type"),
"status": actual.get("status"),
"model": actual.get("model"),
}
assert actual_projection == expected
assert actual.get("call_type") == "anthropic_messages"
finally:
litellm.callbacks = original_callbacks
def test_add_litellm_metadata_from_request_headers_x_litellm_trace_id_sets_chain_id():
"""x-litellm-trace-id sets both metadata and top-level litellm_session_id/litellm_trace_id for call chaining."""
headers = {"x-litellm-trace-id": "foo"}
data = {"metadata": {}}
LiteLLMProxyRequestSetup.add_litellm_metadata_from_request_headers(
headers=headers, data=data, _metadata_variable_name="metadata"
)
assert data["metadata"]["trace_id"] == "foo"
assert data["metadata"]["session_id"] == "foo"
assert data["litellm_session_id"] == "foo"
assert data["litellm_trace_id"] == "foo"
def test_add_litellm_metadata_from_request_headers_x_litellm_session_id_sets_chain_id():
"""x-litellm-session-id sets both metadata and top-level litellm_session_id/litellm_trace_id for call chaining."""
headers = {"x-litellm-session-id": "bar"}
data = {"metadata": {}}
LiteLLMProxyRequestSetup.add_litellm_metadata_from_request_headers(
headers=headers, data=data, _metadata_variable_name="metadata"
)
assert data["metadata"]["trace_id"] == "bar"
assert data["metadata"]["session_id"] == "bar"
assert data["litellm_session_id"] == "bar"
assert data["litellm_trace_id"] == "bar"
def test_add_litellm_metadata_from_request_headers_both_headers_trace_id_precedence():
"""When both x-litellm-trace-id and x-litellm-session-id are present, trace-id takes precedence for chain_id."""
headers = {
"x-litellm-trace-id": "trace-value",
"x-litellm-session-id": "session-value",
}
data = {"metadata": {}}
LiteLLMProxyRequestSetup.add_litellm_metadata_from_request_headers(
headers=headers, data=data, _metadata_variable_name="metadata"
)
assert data["metadata"]["trace_id"] == "trace-value"
assert data["metadata"]["session_id"] == "trace-value"
assert data["litellm_session_id"] == "trace-value"
assert data["litellm_trace_id"] == "trace-value"
def test_add_litellm_metadata_from_request_headers_generic_session_id_header():
"""A generic x-<vendor>-session-id header is used when no explicit litellm header is set."""
headers = {"x-claude-code-session-id": "e96634a3-fa28-4083-b354-55542e2dca01"}
data = {"metadata": {}}
LiteLLMProxyRequestSetup.add_litellm_metadata_from_request_headers(
headers=headers, data=data, _metadata_variable_name="metadata"
)
assert data["metadata"]["session_id"] == "e96634a3-fa28-4083-b354-55542e2dca01"
assert data["litellm_session_id"] == "e96634a3-fa28-4083-b354-55542e2dca01"
assert data["litellm_trace_id"] == "e96634a3-fa28-4083-b354-55542e2dca01"
def test_add_litellm_metadata_from_anthropic_user_id_sets_session_id():
data = {"metadata": {"user_id": "user_abc123_account__session_e96634a3-fa28-4083-b354-55542e2dca01"}}
LiteLLMProxyRequestSetup.add_litellm_metadata_from_request_headers(
headers={}, data=data, _metadata_variable_name="metadata"
)
assert data["metadata"]["session_id"] == "e96634a3-fa28-4083-b354-55542e2dca01"
assert data["litellm_session_id"] == "e96634a3-fa28-4083-b354-55542e2dca01"
assert "litellm_trace_id" not in data
def test_add_litellm_metadata_from_anthropic_user_id_dict_sets_session_id():
data = {
"metadata": {
"user_id": {
"device_id": "device",
"account_uuid": "account",
"session_id": "sess_4f8c1d2a-1234",
}
}
}
LiteLLMProxyRequestSetup.add_litellm_metadata_from_request_headers(
headers={}, data=data, _metadata_variable_name="metadata"
)
assert data["metadata"]["user_id"] == "sess_4f8c1d2a-1234"
assert data["metadata"]["session_id"] == "sess_4f8c1d2a-1234"
assert data["litellm_session_id"] == "sess_4f8c1d2a-1234"
assert "litellm_trace_id" not in data
def test_add_litellm_metadata_from_headers_session_id_beats_anthropic_user_id():
data = {
"metadata": {
"user_id": "user_abc123_account__session_body-session-id",
}
}
LiteLLMProxyRequestSetup.add_litellm_metadata_from_request_headers(
headers={"x-litellm-session-id": "header-session-id"},
data=data,
_metadata_variable_name="metadata",
)
assert data["metadata"]["session_id"] == "header-session-id"
assert data["litellm_session_id"] == "header-session-id"
assert data["litellm_trace_id"] == "header-session-id"
def test_add_litellm_metadata_from_headers_session_id_beats_anthropic_user_id_dict():
data = {
"metadata": {
"user_id": {
"session_id": "body-session-id",
}
}
}
LiteLLMProxyRequestSetup.add_litellm_metadata_from_request_headers(
headers={"x-litellm-session-id": "header-session-id"},
data=data,
_metadata_variable_name="metadata",
)
assert data["metadata"]["session_id"] == "header-session-id"
assert data["litellm_session_id"] == "header-session-id"
assert data["litellm_trace_id"] == "header-session-id"
@pytest.mark.parametrize(
"user_id",
[
"user_abc123_account__session_",
"user_abc123_account_",
"user_abc123_account__session_invalid!",
],
)
def test_add_litellm_metadata_from_anthropic_user_id_ignores_invalid_session_id(user_id: str):
data = {"metadata": {"user_id": user_id}}
LiteLLMProxyRequestSetup.add_litellm_metadata_from_request_headers(
headers={}, data=data, _metadata_variable_name="metadata"
)
assert data == {"metadata": {"user_id": user_id}}
@pytest.mark.parametrize(
"user_id",
[
{},
{"session_id": 123},
{"session_id": "invalid session id"},
{"session_id": ""},
],
)
def test_add_litellm_metadata_from_anthropic_user_id_dict_ignores_invalid_session_id(
user_id: object,
):
data = {"metadata": {"user_id": user_id}}
LiteLLMProxyRequestSetup.add_litellm_metadata_from_request_headers(
headers={}, data=data, _metadata_variable_name="metadata"
)
assert data == {"metadata": {"user_id": user_id}}
def test_add_litellm_metadata_from_request_headers_explicit_header_beats_generic():
"""Explicit x-litellm-trace-id wins over a generic x-*-session-id header."""
headers = {
"x-litellm-trace-id": "explicit-trace-id-value",
"x-claude-code-session-id": "e96634a3-fa28-4083-b354-55542e2dca01",
}
data = {"metadata": {}}
LiteLLMProxyRequestSetup.add_litellm_metadata_from_request_headers(
headers=headers, data=data, _metadata_variable_name="metadata"
)
assert data["litellm_session_id"] == "explicit-trace-id-value"
assert data["litellm_trace_id"] == "explicit-trace-id-value"
def test_get_chain_id_from_headers_generic_vendor_session_id():
"""get_chain_id_from_headers picks up any x-<vendor>-session-id with a valid value."""
from litellm.proxy.litellm_pre_call_utils import get_chain_id_from_headers
assert (
get_chain_id_from_headers({"x-claude-code-session-id": "e96634a3-fa28-4083-b354-55542e2dca01"})
== "e96634a3-fa28-4083-b354-55542e2dca01"
)
# Short / non-alphanumeric values should be ignored
assert get_chain_id_from_headers({"x-foo-session-id": "short"}) is None
assert get_chain_id_from_headers({"x-foo-session-id": "has spaces!!"}) is None
# Explicit headers still take precedence
assert (
get_chain_id_from_headers(
{
"x-litellm-trace-id": "explicit-id-value",
"x-claude-code-session-id": "e96634a3-fa28-4083-b354-55542e2dca01",
}
)
== "explicit-id-value"
)
CODEX_USER_AGENT = "codex_cli_rs/0.62.0 (Mac OS 25.5.0; arm64) Apple_Terminal"
CODEX_SESSION_UUID = "0199f0c2-8b41-7c3e-9a52-6d1f4b8e2a77"
@pytest.mark.parametrize(
"user_agent",
[
"codex-tui",
"codex-tui/0.149.0 (Mac OS 26.5.1; arm64) ghostty/1.3.1 (codex-tui; 0.149.0)",
"codex_cli_rs/0.62.0 (Mac OS 25.5.0; arm64) Apple_Terminal",
"codex_exec/0.62.0 (Linux 6.1; x86_64) unknown",
"codex_vscode/0.62.0 (Mac OS 26.5.1; arm64) vscode/1.99.0",
"Codex CLI/1.0",
],
)
def test_is_codex_user_agent_accepts_every_first_party_originator(user_agent: str):
"""Codex ships several originators sharing only the `codex` stem, and the TUI
sends a bare `codex-tui` with no version, so matching one spelling misses real clients."""
from litellm.proxy.litellm_pre_call_utils import is_codex_user_agent
assert is_codex_user_agent(user_agent) is True
@pytest.mark.parametrize(
"user_agent",
["codexify/1.0", "mycodex-tui/1.0", "curl/8.7.1", "claude-cli/2.1.0 (external, cli)", ""],
)
def test_is_codex_user_agent_rejects_non_codex_clients(user_agent: str):
from litellm.proxy.litellm_pre_call_utils import is_codex_user_agent
assert is_codex_user_agent(user_agent) is False
def test_get_chain_id_from_headers_codex_tui_user_agent():
"""The real Codex TUI user agent must group turns, not just the codex_cli_rs spelling."""
from litellm.proxy.litellm_pre_call_utils import get_chain_id_from_headers
ua = "codex-tui/0.149.0 (Mac OS 26.5.1; arm64) ghostty/1.3.1 (codex-tui; 0.149.0)"
assert get_chain_id_from_headers({"user-agent": ua, "session-id": CODEX_SESSION_UUID}) == CODEX_SESSION_UUID
assert (
get_chain_id_from_headers({"user-agent": "codex-tui", "session-id": CODEX_SESSION_UUID}) == CODEX_SESSION_UUID
)
@pytest.mark.parametrize(
"header",
["session-id", "session_id", "thread-id", "conversation_id", "Session-Id"],
)
def test_get_chain_id_from_headers_codex_unprefixed_session_id(header: str):
"""Codex sends its conversation uuid unprefixed, so the x-<vendor>-session-id regex misses it."""
from litellm.proxy.litellm_pre_call_utils import get_chain_id_from_headers
assert get_chain_id_from_headers({"user-agent": CODEX_USER_AGENT, header: CODEX_SESSION_UUID}) == CODEX_SESSION_UUID
@pytest.mark.parametrize(
"user_agent",
["curl/8.7.1", "claude-cli/2.1.0 (external, cli)", "OpenAI/Python 1.0.0"],
)
def test_get_chain_id_from_headers_unprefixed_session_id_requires_codex(user_agent: str):
"""An unprefixed session-id from a non-Codex caller must not group traces.
The name is generic enough that two unrelated callers could collide on a value
and have their sessions merged, so the bare-header path is Codex-only.
"""
from litellm.proxy.litellm_pre_call_utils import get_chain_id_from_headers
assert get_chain_id_from_headers({"user-agent": user_agent, "session-id": CODEX_SESSION_UUID}) is None
assert get_chain_id_from_headers({"session-id": CODEX_SESSION_UUID}) is None
def test_get_chain_id_from_headers_codex_prefers_session_over_thread():
from litellm.proxy.litellm_pre_call_utils import get_chain_id_from_headers
assert (
get_chain_id_from_headers(
{
"user-agent": CODEX_USER_AGENT,
"thread-id": "e96634a3-fa28-4083-b354-55542e2dca01",
"session-id": CODEX_SESSION_UUID,
}
)
== CODEX_SESSION_UUID
)
def test_get_chain_id_from_headers_codex_ignores_implausible_value():
from litellm.proxy.litellm_pre_call_utils import get_chain_id_from_headers
assert get_chain_id_from_headers({"user-agent": CODEX_USER_AGENT, "session-id": "short"}) is None
assert get_chain_id_from_headers({"user-agent": CODEX_USER_AGENT, "session-id": "has spaces!!"}) is None
def test_get_chain_id_from_headers_explicit_beats_codex_header():
from litellm.proxy.litellm_pre_call_utils import get_chain_id_from_headers
assert (
get_chain_id_from_headers(
{
"user-agent": CODEX_USER_AGENT,
"x-litellm-trace-id": "explicit-id-value",
"session-id": CODEX_SESSION_UUID,
}
)
== "explicit-id-value"
)
def test_add_litellm_metadata_groups_codex_turns_into_one_session():
"""Every turn of a Codex session must log under one session id, not a fresh per-call trace id."""
headers = {"user-agent": CODEX_USER_AGENT, "session-id": CODEX_SESSION_UUID}
turns = [{"litellm_metadata": {}}, {"litellm_metadata": {}}]
for turn in turns:
LiteLLMProxyRequestSetup.add_litellm_metadata_from_request_headers(
headers=headers, data=turn, _metadata_variable_name="litellm_metadata"
)
for turn in turns:
assert turn["litellm_session_id"] == CODEX_SESSION_UUID
assert turn["litellm_trace_id"] == CODEX_SESSION_UUID
assert turn["litellm_metadata"]["session_id"] == CODEX_SESSION_UUID
OPENCODE_SESSION_ID = "ses_f91e6e825ffeuhlu5EbglxjAN2"
OPENCODE_HEADERS = {
"x-session-affinity": OPENCODE_SESSION_ID,
"X-Session-Id": OPENCODE_SESSION_ID,
"User-Agent": "opencode/1.18.28",
}
def test_add_litellm_metadata_groups_opencode_turns_into_one_session():
"""Every turn of an opencode session must land on metadata.session_id, which is what
DeploymentAffinityCheck reads for session pinning, instead of a fresh per-call id."""
turns = [{"metadata": {}}, {"metadata": {}}]
for turn in turns:
LiteLLMProxyRequestSetup.add_litellm_metadata_from_request_headers(
headers=OPENCODE_HEADERS, data=turn, _metadata_variable_name="metadata"
)
for turn in turns:
assert turn["metadata"]["session_id"] == OPENCODE_SESSION_ID
assert turn["metadata"]["trace_id"] == OPENCODE_SESSION_ID
assert turn["litellm_session_id"] == OPENCODE_SESSION_ID
assert turn["litellm_trace_id"] == OPENCODE_SESSION_ID
@pytest.mark.parametrize("value", ["short", "has spaces!!", ""])
def test_get_chain_id_from_headers_bare_session_id_ignores_implausible_value(value: str):
from litellm.proxy.litellm_pre_call_utils import get_chain_id_from_headers
assert get_chain_id_from_headers({"x-session-id": value}) is None
@pytest.mark.parametrize(
"other_header",
[
"x-litellm-trace-id",
"x-litellm-session-id",
"x-claude-code-session-id",
"x-parent-session-id",
],
)
def test_get_chain_id_from_headers_bare_session_id_loses_to_more_specific_header(other_header: str):
"""opencode subagent calls carry x-parent-session-id next to X-Session-Id; explicit and
vendor-scoped headers must keep winning over the bare header."""
from litellm.proxy.litellm_pre_call_utils import get_chain_id_from_headers
assert (
get_chain_id_from_headers(
{
"x-session-id": OPENCODE_SESSION_ID,
other_header: "e96634a3-fa28-4083-b354-55542e2dca01",
}
)
== "e96634a3-fa28-4083-b354-55542e2dca01"
)
def test_trace_id_from_traceparent_valid():
from litellm.proxy.litellm_pre_call_utils import _trace_id_from_traceparent
assert (
_trace_id_from_traceparent("00-4bf92f3577b34da6a3ce929d0e0e4736-00f067aa0ba902b7-01")
== "4bf92f3577b34da6a3ce929d0e0e4736"
)
# Case-insensitive, normalized to lowercase
assert (
_trace_id_from_traceparent("00-4BF92F3577B34DA6A3CE929D0E0E4736-00f067aa0ba902b7-01")
== "4bf92f3577b34da6a3ce929d0e0e4736"
)
@pytest.mark.parametrize(
"traceparent",
[
"not-a-traceparent",
"00-tooshort-00f067aa0ba902b7-01",
"00-4bf92f3577b34da6a3ce929d0e0e4736-00f067aa0ba902b7", # missing flags segment
"00-4bf92f3577g34da6a3ce929d0e0e4736-00f067aa0ba902b7-01", # non-hex char
"00-00000000000000000000000000000000-00f067aa0ba902b7-01", # all-zero trace-id, invalid per spec
"",
],
)
def test_trace_id_from_traceparent_rejects_malformed(traceparent: str):
from litellm.proxy.litellm_pre_call_utils import _trace_id_from_traceparent
assert _trace_id_from_traceparent(traceparent) is None
def test_session_id_from_baggage_valid():
from litellm.proxy.litellm_pre_call_utils import _session_id_from_baggage
assert _session_id_from_baggage("session.id=abc-123,user.id=42") == "abc-123"
assert _session_id_from_baggage("user.id=42, session.id=xyz-789") == "xyz-789"
@pytest.mark.parametrize(
"baggage",
[
"user.id=42",
"",
"session.id=",
],
)
def test_session_id_from_baggage_absent_or_empty(baggage: str):
from litellm.proxy.litellm_pre_call_utils import _session_id_from_baggage
assert _session_id_from_baggage(baggage) is None
def test_add_litellm_metadata_from_request_headers_traceparent_sets_trace_id_only():
"""A bare traceparent header (no litellm-specific headers) sets litellm_trace_id
from its trace-id component and leaves litellm_session_id unset."""
headers = {"traceparent": "00-4bf92f3577b34da6a3ce929d0e0e4736-00f067aa0ba902b7-01"}
data = {"metadata": {}}
LiteLLMProxyRequestSetup.add_litellm_metadata_from_request_headers(
headers=headers, data=data, _metadata_variable_name="metadata"
)
assert data["litellm_trace_id"] == "4bf92f3577b34da6a3ce929d0e0e4736"
assert data["metadata"]["trace_id"] == "4bf92f3577b34da6a3ce929d0e0e4736"
assert "litellm_session_id" not in data
def test_add_litellm_metadata_from_request_headers_baggage_sets_session_id_only():
"""A bare baggage header (no litellm-specific headers) sets litellm_session_id
from its session.id entry and leaves litellm_trace_id unset."""
headers = {"baggage": "session.id=baggage-session-42,user.id=7"}
data = {"metadata": {}}
LiteLLMProxyRequestSetup.add_litellm_metadata_from_request_headers(
headers=headers, data=data, _metadata_variable_name="metadata"
)
assert data["litellm_session_id"] == "baggage-session-42"
assert data["metadata"]["session_id"] == "baggage-session-42"
assert "litellm_trace_id" not in data
def test_add_litellm_metadata_from_request_headers_baggage_session_id_not_logged_raw(caplog):
"""The raw baggage session.id value must never reach the debug log line -
it isn't sanitized until set_session_id() runs much later in
Logging.__init__(), so logging it here would let a caller with control
characters or terminal escape sequences forge plaintext log output."""
import logging
poisoned = "poisoned\x1b[31mFAKE_RED_TEXT\x1b[0m"
headers = {"baggage": f"session.id={poisoned}"}
data = {"metadata": {}}
with caplog.at_level(logging.DEBUG, logger="LiteLLM Proxy"):
LiteLLMProxyRequestSetup.add_litellm_metadata_from_request_headers(
headers=headers, data=data, _metadata_variable_name="metadata"
)
assert data["litellm_session_id"] == poisoned
assert not any(poisoned in record.getMessage() for record in caplog.records)
def test_add_litellm_metadata_from_request_headers_traceparent_and_baggage_together():
"""traceparent and baggage are resolved independently - trace_id and
session_id do not have to be the same value, unlike the chain_id path."""
headers = {
"traceparent": "00-4bf92f3577b34da6a3ce929d0e0e4736-00f067aa0ba902b7-01",
"baggage": "session.id=baggage-session-42",
}
data = {"metadata": {}}
LiteLLMProxyRequestSetup.add_litellm_metadata_from_request_headers(
headers=headers, data=data, _metadata_variable_name="metadata"
)
assert data["litellm_trace_id"] == "4bf92f3577b34da6a3ce929d0e0e4736"
assert data["litellm_session_id"] == "baggage-session-42"
def test_add_litellm_metadata_from_request_headers_explicit_trace_id_beats_traceparent():
"""x-litellm-trace-id must win over a traceparent header carrying a
different trace-id - explicit litellm headers are always highest priority."""
headers = {
"x-litellm-trace-id": "explicit-trace-id-value",
"traceparent": "00-4bf92f3577b34da6a3ce929d0e0e4736-00f067aa0ba902b7-01",
}
data = {"metadata": {}}
LiteLLMProxyRequestSetup.add_litellm_metadata_from_request_headers(
headers=headers, data=data, _metadata_variable_name="metadata"
)
assert data["litellm_trace_id"] == "explicit-trace-id-value"
assert data["litellm_session_id"] == "explicit-trace-id-value"
def _otel_span_with_trace_id(trace_id: int) -> NonRecordingSpan:
return NonRecordingSpan(SpanContext(trace_id=trace_id, span_id=0x00F067AA0BA902B7, is_remote=False))
def _request_mock_without_trace_headers() -> MagicMock:
request_mock: Final = MagicMock(spec=Request)
request_mock.url = MagicMock()
request_mock.url.path = "/v1/chat/completions"
request_mock.url.__str__.return_value = "http://localhost/v1/chat/completions"
request_mock.method = "POST"
request_mock.query_params = {}
request_mock.headers = {"Content-Type": "application/json"}
request_mock.client = MagicMock()
request_mock.client.host = "127.0.0.1"
return request_mock
@pytest.mark.asyncio
async def test_add_litellm_data_to_request_defaults_trace_id_to_otel_server_span():
"""With OTel on and a client that sends no trace headers, the request's
litellm_trace_id (and so the spend log session_id) must be the W3C trace-id
of the proxy's server span, so a trace in the OTel backend can be looked up
in the Logs UI and vice versa."""
otel_trace_id: Final = 0x4BF92F3577B34DA6A3CE929D0E0E4736
user_api_key_dict: Final = UserAPIKeyAuth(
api_key="hashed-key", parent_otel_span=_otel_span_with_trace_id(otel_trace_id)
)
data: Final = await add_litellm_data_to_request(
data={"model": "gpt-5.6", "messages": [{"role": "user", "content": "hi"}]},
request=_request_mock_without_trace_headers(),
user_api_key_dict=user_api_key_dict,
proxy_config=MagicMock(),
general_settings={},
)
assert data["litellm_trace_id"] == format(otel_trace_id, "032x")
assert data["metadata"]["trace_id"] == format(otel_trace_id, "032x")
assert "litellm_session_id" not in data
@pytest.mark.asyncio
async def test_add_litellm_data_to_request_falls_back_to_request_state_otel_span():
"""Custom auth hooks return a UserAPIKeyAuth without parent_otel_span even
though user_api_key_auth already opened the server span on request.state,
so the fallback must read the span from there or custom-auth requests would
keep getting an unrelated session id."""
otel_trace_id: Final = 0x4BF92F3577B34DA6A3CE929D0E0E4736
request_mock: Final = _request_mock_without_trace_headers()
request_mock.state.parent_otel_span = _otel_span_with_trace_id(otel_trace_id)
data: Final = await add_litellm_data_to_request(
data={"model": "gpt-5.6"},
request=request_mock,
user_api_key_dict=UserAPIKeyAuth(api_key="hashed-key", parent_otel_span=None),
proxy_config=MagicMock(),
general_settings={},
)
assert data["litellm_trace_id"] == format(otel_trace_id, "032x")
assert data["metadata"]["trace_id"] == format(otel_trace_id, "032x")
@pytest.mark.asyncio
async def test_add_litellm_data_to_request_otel_span_does_not_override_caller_trace_id():
"""A caller's own trace identity (x-litellm-trace-id header or body
metadata.trace_id) keeps priority over the OTel server span's trace-id."""
span: Final = _otel_span_with_trace_id(0x4BF92F3577B34DA6A3CE929D0E0E4736)
header_request: Final = _request_mock_without_trace_headers()
header_request.headers = {"Content-Type": "application/json", "x-litellm-trace-id": "caller-trace"}
from_header: Final = await add_litellm_data_to_request(
data={"model": "gpt-5.6"},
request=header_request,
user_api_key_dict=UserAPIKeyAuth(api_key="hashed-key", parent_otel_span=span),
proxy_config=MagicMock(),
general_settings={},
)
assert from_header["litellm_trace_id"] == "caller-trace"
assert from_header["metadata"]["trace_id"] == "caller-trace"
from_body: Final = await add_litellm_data_to_request(
data={"model": "gpt-5.6", "metadata": {"trace_id": "body-trace"}},
request=_request_mock_without_trace_headers(),
user_api_key_dict=UserAPIKeyAuth(api_key="hashed-key", parent_otel_span=span),
proxy_config=MagicMock(),
general_settings={},
)
assert "litellm_trace_id" not in from_body
assert from_body["metadata"]["trace_id"] == "body-trace"
@pytest.mark.asyncio
@pytest.mark.parametrize("path", ["/v1/responses", "/v1/messages"])
async def test_add_litellm_data_to_request_otel_span_does_not_override_body_trace_id_on_litellm_metadata_routes(path):
"""On routes that keep LiteLLM state in litellm_metadata, the caller's body
metadata.trace_id is only promoted into litellm_metadata later in the
pipeline, so the OTel fallback must look at the requester metadata too or
it would claim the slot first and the caller's id would be lost."""
request_mock: Final = _request_mock_without_trace_headers()
request_mock.url.path = path
request_mock.url.__str__.return_value = f"http://localhost{path}"
data: Final = await add_litellm_data_to_request(
data={"model": "gpt-5.6", "metadata": {"trace_id": "body-trace"}},
request=request_mock,
user_api_key_dict=UserAPIKeyAuth(
api_key="hashed-key", parent_otel_span=_otel_span_with_trace_id(0x4BF92F3577B34DA6A3CE929D0E0E4736)
),
proxy_config=MagicMock(),
general_settings={},
)
assert "litellm_trace_id" not in data
assert data["litellm_metadata"]["trace_id"] == "body-trace"
@pytest.mark.asyncio
@pytest.mark.parametrize("empty_trace_id", [None, ""])
async def test_add_litellm_data_to_request_otel_span_fills_empty_body_trace_id(empty_trace_id):
"""A serialized-but-empty litellm_trace_id in the body (null or "") carries
no identity, so it must not block the OTel server span fallback."""
otel_trace_id: Final = 0x4BF92F3577B34DA6A3CE929D0E0E4736
data: Final = await add_litellm_data_to_request(
data={"model": "gpt-5.6", "litellm_trace_id": empty_trace_id},
request=_request_mock_without_trace_headers(),
user_api_key_dict=UserAPIKeyAuth(
api_key="hashed-key", parent_otel_span=_otel_span_with_trace_id(otel_trace_id)
),
proxy_config=MagicMock(),
general_settings={},
)
assert data["litellm_trace_id"] == format(otel_trace_id, "032x")
assert data["metadata"]["trace_id"] == format(otel_trace_id, "032x")
@pytest.mark.asyncio
@pytest.mark.parametrize("parent_otel_span", [None, "invalid_span", "not_a_span", "plain_string"])
async def test_add_litellm_data_to_request_no_trace_id_without_valid_otel_span(parent_otel_span):
"""No OTel span (OTel off), a span with an invalid context, an object that
only quacks like a span, or a value that is not a span at all (custom auth
is typed loosely and can hand back anything) must leave litellm_trace_id
unset, and never fail the request, so downstream keeps generating its own id."""
span: Final = {
"invalid_span": INVALID_SPAN,
"not_a_span": MagicMock(),
"plain_string": "not-a-span",
}.get(parent_otel_span)
data: Final = await add_litellm_data_to_request(
data={"model": "gpt-5.6"},
request=_request_mock_without_trace_headers(),
user_api_key_dict=UserAPIKeyAuth(api_key="hashed-key", parent_otel_span=span),
proxy_config=MagicMock(),
general_settings={},
)
assert "litellm_trace_id" not in data
assert "trace_id" not in data["metadata"]
def test_add_litellm_metadata_from_request_headers_anthropic_metadata_beats_baggage():
"""The existing Anthropic metadata.user_id session_id path must win over a
baggage session.id fallback."""
data = {
"metadata": {
"user_id": "user_abc123_account__session_e96634a3-fa28-4083-b354-55542e2dca01",
}
}
LiteLLMProxyRequestSetup.add_litellm_metadata_from_request_headers(
headers={"baggage": "session.id=baggage-session-42"},
data=data,
_metadata_variable_name="metadata",
)
assert data["litellm_session_id"] == "e96634a3-fa28-4083-b354-55542e2dca01"
assert "litellm_trace_id" not in data
def test_get_internal_user_header_from_mapping_returns_expected_header():
mappings = [
{"header_name": "X-OpenWebUI-User-Id", "litellm_user_role": "internal_user"},
{"header_name": "X-OpenWebUI-User-Email", "litellm_user_role": "customer"},
]
header_name = LiteLLMProxyRequestSetup.get_internal_user_header_from_mapping(mappings)
assert header_name == "X-OpenWebUI-User-Id"
def test_get_internal_user_header_from_mapping_none_when_absent():
mappings = [{"header_name": "X-OpenWebUI-User-Email", "litellm_user_role": "customer"}]
header_name = LiteLLMProxyRequestSetup.get_internal_user_header_from_mapping(mappings)
assert header_name is None
single = {"header_name": "X-Only-Customer", "litellm_user_role": "customer"}
header_name = LiteLLMProxyRequestSetup.get_internal_user_header_from_mapping(single)
assert header_name is None
def test_add_internal_user_from_user_mapping_sets_user_id_when_header_present():
user_api_key_dict = UserAPIKeyAuth(api_key="test-key")
headers = {"X-OpenWebUI-User-Id": "internal-user-123"}
general_settings = {
"user_header_mappings": [
{
"header_name": "X-OpenWebUI-User-Id",
"litellm_user_role": "internal_user",
},
{"header_name": "X-OpenWebUI-User-Email", "litellm_user_role": "customer"},
]
}
result = LiteLLMProxyRequestSetup.add_internal_user_from_user_mapping(general_settings, user_api_key_dict, headers)
assert result is user_api_key_dict
assert user_api_key_dict.user_id == "internal-user-123"
def test_add_internal_user_from_user_mapping_no_header_or_mapping_returns_unchanged():
user_api_key_dict = UserAPIKeyAuth(api_key="test-key")
result = LiteLLMProxyRequestSetup.add_internal_user_from_user_mapping(
None, user_api_key_dict, {"X-OpenWebUI-User-Id": "abc"}
)
assert result is user_api_key_dict
assert user_api_key_dict.user_id is None
general_settings = {
"user_header_mappings": [{"header_name": "X-OpenWebUI-User-Id", "litellm_user_role": "internal_user"}]
}
result = LiteLLMProxyRequestSetup.add_internal_user_from_user_mapping(
general_settings, user_api_key_dict, {"Other": "value"}
)
assert result is user_api_key_dict
assert user_api_key_dict.user_id is None
def test_get_sanitized_user_information_from_key_includes_guardrails_metadata():
"""
Test that get_sanitized_user_information_from_key includes guardrails field from key metadata in the returned payload
"""
user_api_key_dict = UserAPIKeyAuth(
api_key="test-key-hash",
key_alias="test-alias",
user_id="test-user",
metadata={"guardrails": ["presidio", "aporia"], "other_field": "value"},
)
result = LiteLLMProxyRequestSetup.get_sanitized_user_information_from_key(user_api_key_dict=user_api_key_dict)
assert result["user_api_key_auth_metadata"] is not None
assert "guardrails" in result["user_api_key_auth_metadata"]
assert result["user_api_key_auth_metadata"]["guardrails"] == ["presidio", "aporia"]
assert result["user_api_key_auth_metadata"]["other_field"] == "value"
def test_user_and_team_spend_and_budget_flow_to_standard_logging_metadata():
"""
Full flow: UserAPIKeyAuth -> get_sanitized_user_information_from_key ->
get_standard_logging_metadata. User-level and team-level spend + max budget
must reach the StandardLoggingPayload metadata that custom loggers receive,
alongside the key-level values
"""
from litellm.litellm_core_utils.litellm_logging import StandardLoggingPayloadSetup
user_api_key_dict = UserAPIKeyAuth(
api_key="test-key-hash",
spend=1.5,
max_budget=10.0,
user_id="test-user",
user_spend=25.5,
user_max_budget=100.0,
team_id="test-team",
team_spend=250.75,
team_max_budget=1000.0,
)
sanitized = LiteLLMProxyRequestSetup.get_sanitized_user_information_from_key(user_api_key_dict=user_api_key_dict)
assert sanitized["user_api_key_spend"] == 1.5
assert sanitized["user_api_key_max_budget"] == 10.0
assert sanitized["user_api_key_user_spend"] == 25.5
assert sanitized["user_api_key_user_max_budget"] == 100.0
assert sanitized["user_api_key_team_spend"] == 250.75
assert sanitized["user_api_key_team_max_budget"] == 1000.0
logging_metadata = StandardLoggingPayloadSetup.get_standard_logging_metadata(dict(sanitized))
assert logging_metadata["user_api_key_user_spend"] == 25.5
assert logging_metadata["user_api_key_user_max_budget"] == 100.0
assert logging_metadata["user_api_key_team_spend"] == 250.75
assert logging_metadata["user_api_key_team_max_budget"] == 1000.0
def test_user_and_team_spend_and_budget_default_to_none_in_standard_logging_metadata():
"""
Keys with no user or team level budgets report None for the new fields in the
StandardLoggingPayload metadata instead of raising
"""
from litellm.litellm_core_utils.litellm_logging import StandardLoggingPayloadSetup
user_api_key_dict = UserAPIKeyAuth(api_key="test-key-hash")
sanitized = LiteLLMProxyRequestSetup.get_sanitized_user_information_from_key(user_api_key_dict=user_api_key_dict)
logging_metadata = StandardLoggingPayloadSetup.get_standard_logging_metadata(dict(sanitized))
assert logging_metadata["user_api_key_user_spend"] is None
assert logging_metadata["user_api_key_user_max_budget"] is None
assert logging_metadata["user_api_key_team_spend"] is None
assert logging_metadata["user_api_key_team_max_budget"] is None
@pytest.mark.asyncio
async def test_team_guardrails_append_to_key_guardrails():
"""
Test that team guardrails are appended to key guardrails instead of overriding them.
Team guardrails should only be added if they are not already present in key guardrails.
"""
request_mock = MagicMock(spec=Request)
request_mock.url.path = "/chat/completions"
request_mock.url = MagicMock()
request_mock.url.__str__.return_value = "http://localhost/chat/completions"
request_mock.method = "POST"
request_mock.query_params = {}
request_mock.headers = {"Content-Type": "application/json"}
request_mock.client = MagicMock()
request_mock.client.host = "127.0.0.1"
data = {
"model": "gpt-3.5-turbo",
"messages": [{"role": "user", "content": "test"}],
}
user_api_key_dict = UserAPIKeyAuth(
api_key="test-key",
metadata={"guardrails": ["key-guardrail-1", "key-guardrail-2"]},
team_metadata={"guardrails": ["team-guardrail-1", "key-guardrail-1"]},
)
with patch("litellm.proxy.utils._premium_user_check"):
updated_data = await add_litellm_data_to_request(
data=data,
request=request_mock,
user_api_key_dict=user_api_key_dict,
proxy_config=MagicMock(),
general_settings={},
version="test-version",
)
metadata = updated_data.get("metadata", {})
guardrails = metadata.get("guardrails", [])
assert "key-guardrail-1" in guardrails
assert "key-guardrail-2" in guardrails
assert "team-guardrail-1" in guardrails
assert guardrails.count("key-guardrail-1") == 1
@pytest.mark.asyncio
async def test_request_guardrails_do_not_override_key_guardrails():
"""
Test that request-level guardrails do not override key-level guardrails.
Key guardrails should be preserved when request contains guardrails (including empty array).
"""
request_mock = MagicMock(spec=Request)
request_mock.url.path = "/chat/completions"
request_mock.url = MagicMock()
request_mock.url.__str__.return_value = "http://localhost/chat/completions"
request_mock.method = "POST"
request_mock.query_params = {}
request_mock.headers = {"Content-Type": "application/json"}
request_mock.client = MagicMock()
request_mock.client.host = "127.0.0.1"
user_api_key_dict = UserAPIKeyAuth(
api_key="test-key",
metadata={"guardrails": ["key-guardrail-1"]},
team_metadata={},
)
# Test case: Request with empty guardrails should not result in empty guardrails
data_with_empty = {
"model": "gpt-3.5-turbo",
"messages": [{"role": "user", "content": "test"}],
"guardrails": [],
}
with patch("litellm.proxy.utils._premium_user_check"):
updated_data_empty = await add_litellm_data_to_request(
data=data_with_empty,
request=request_mock,
user_api_key_dict=user_api_key_dict,
proxy_config=MagicMock(),
general_settings={},
version="test-version",
)
_metadata = updated_data_empty.get("metadata", {})
requested_guardrails = _metadata.get("guardrails", [])
assert "guardrails" not in updated_data_empty
assert "key-guardrail-1" in requested_guardrails
assert len(requested_guardrails) == 1
@pytest.mark.asyncio
async def test_project_guardrails_merge_with_key_and_team():
"""
Test that project guardrails are merged with key and team guardrails (union semantics).
All three levels should contribute to the final guardrails list without duplicates.
"""
request_mock = MagicMock(spec=Request)
request_mock.url.path = "/chat/completions"
request_mock.url = MagicMock()
request_mock.url.__str__.return_value = "http://localhost/chat/completions"
request_mock.method = "POST"
request_mock.query_params = {}
request_mock.headers = {"Content-Type": "application/json"}
request_mock.client = MagicMock()
request_mock.client.host = "127.0.0.1"
data = {
"model": "gpt-3.5-turbo",
"messages": [{"role": "user", "content": "test"}],
}
user_api_key_dict = UserAPIKeyAuth(
api_key="test-key",
metadata={"guardrails": ["key-guardrail-1"]},
team_metadata={"guardrails": ["team-guardrail-1", "key-guardrail-1"]},
project_metadata={"guardrails": ["project-guardrail-1", "team-guardrail-1"]},
)
with patch("litellm.proxy.utils._premium_user_check"):
updated_data = await add_litellm_data_to_request(
data=data,
request=request_mock,
user_api_key_dict=user_api_key_dict,
proxy_config=MagicMock(),
general_settings={},
version="test-version",
)
metadata = updated_data.get("metadata", {})
guardrails = metadata.get("guardrails", [])
# All three sources contribute
assert "key-guardrail-1" in guardrails
assert "team-guardrail-1" in guardrails
assert "project-guardrail-1" in guardrails
# No duplicates
assert guardrails.count("key-guardrail-1") == 1
assert guardrails.count("team-guardrail-1") == 1
@pytest.mark.asyncio
async def test_project_guardrails_only():
"""
Test that project guardrails work when key and team have no guardrails configured.
"""
request_mock = MagicMock(spec=Request)
request_mock.url.path = "/chat/completions"
request_mock.url = MagicMock()
request_mock.url.__str__.return_value = "http://localhost/chat/completions"
request_mock.method = "POST"
request_mock.query_params = {}
request_mock.headers = {"Content-Type": "application/json"}
request_mock.client = MagicMock()
request_mock.client.host = "127.0.0.1"
data = {
"model": "gpt-3.5-turbo",
"messages": [{"role": "user", "content": "test"}],
}
user_api_key_dict = UserAPIKeyAuth(
api_key="test-key",
metadata={},
team_metadata={},
project_metadata={"guardrails": ["project-guardrail-1", "project-guardrail-2"]},
)
with patch("litellm.proxy.utils._premium_user_check"):
updated_data = await add_litellm_data_to_request(
data=data,
request=request_mock,
user_api_key_dict=user_api_key_dict,
proxy_config=MagicMock(),
general_settings={},
version="test-version",
)
metadata = updated_data.get("metadata", {})
guardrails = metadata.get("guardrails", [])
assert "project-guardrail-1" in guardrails
assert "project-guardrail-2" in guardrails
assert len(guardrails) == 2
def test_update_model_if_key_alias_exists():
"""
Test that _update_model_if_key_alias_exists properly updates the model when a key alias exists.
"""
# Test case 1: Key alias exists and matches model
data = {"model": "modelAlias", "messages": [{"role": "user", "content": "Hello"}]}
user_api_key_dict = UserAPIKeyAuth(
api_key="test-key",
aliases={"modelAlias": "xai/grok-4-fast-non-reasoning"},
)
_update_model_if_key_alias_exists(data=data, user_api_key_dict=user_api_key_dict)
assert data["model"] == "xai/grok-4-fast-non-reasoning"
# Test case 2: Key alias doesn't exist
data = {
"model": "unknown-model",
"messages": [{"role": "user", "content": "Hello"}],
}
user_api_key_dict = UserAPIKeyAuth(
api_key="test-key",
aliases={"modelAlias": "xai/grok-4-fast-non-reasoning"},
)
original_model = data["model"]
_update_model_if_key_alias_exists(data=data, user_api_key_dict=user_api_key_dict)
assert data["model"] == original_model # Should remain unchanged
# Test case 3: Model is None
data = {"model": None, "messages": [{"role": "user", "content": "Hello"}]}
user_api_key_dict = UserAPIKeyAuth(
api_key="test-key",
aliases={"modelAlias": "xai/grok-4-fast-non-reasoning"},
)
_update_model_if_key_alias_exists(data=data, user_api_key_dict=user_api_key_dict)
assert data["model"] is None # Should remain None
# Test case 4: Model key doesn't exist in data
data = {"messages": [{"role": "user", "content": "Hello"}]}
user_api_key_dict = UserAPIKeyAuth(
api_key="test-key",
aliases={"modelAlias": "xai/grok-4-fast-non-reasoning"},
)
_update_model_if_key_alias_exists(data=data, user_api_key_dict=user_api_key_dict)
assert "model" not in data # Should not add model if it doesn't exist
# Test case 5: Multiple aliases, matching one
data = {"model": "alias1", "messages": [{"role": "user", "content": "Hello"}]}
user_api_key_dict = UserAPIKeyAuth(
api_key="test-key",
aliases={
"alias1": "model1",
"alias2": "model2",
"alias3": "model3",
},
)
_update_model_if_key_alias_exists(data=data, user_api_key_dict=user_api_key_dict)
assert data["model"] == "model1"
# Test case 6: Empty aliases dict
data = {"model": "modelAlias", "messages": [{"role": "user", "content": "Hello"}]}
user_api_key_dict = UserAPIKeyAuth(api_key="test-key", aliases={})
original_model = data["model"]
_update_model_if_key_alias_exists(data=data, user_api_key_dict=user_api_key_dict)
assert data["model"] == original_model # Should remain unchanged
@pytest.mark.asyncio
async def test_embedding_header_forwarding_with_model_group():
"""
Test that headers are properly forwarded for embedding requests when
forward_client_headers_to_llm_api is configured for the model group.
This test verifies the fix for embedding endpoints not forwarding headers
similar to how chat completion endpoints do.
"""
import importlib
import litellm.proxy.litellm_pre_call_utils as pre_call_utils_module
# Reload the module to ensure it has a fresh reference to litellm
# This is necessary because conftest.py reloads litellm at module scope,
# which can cause the module's litellm reference to become stale
importlib.reload(pre_call_utils_module)
# Re-import the function after reload to get the fresh version
from litellm.proxy.litellm_pre_call_utils import add_litellm_data_to_request
# Setup mock request for embeddings
request_mock = MagicMock(spec=Request)
request_mock.url.path = "/v1/embeddings"
request_mock.url = MagicMock()
request_mock.url.__str__.return_value = "http://localhost/v1/embeddings"
request_mock.method = "POST"
request_mock.query_params = {}
request_mock.headers = {
"Content-Type": "application/json",
"X-Custom-Header": "custom-value",
"X-Request-ID": "test-request-123",
"Authorization": "Bearer sk-test-key",
}
request_mock.client = MagicMock()
request_mock.client.host = "127.0.0.1"
# Setup embedding request data
data = {
"model": "local-openai/text-embedding-3-small",
"input": ["Text to embed"],
}
# Setup user API key
user_api_key_dict = UserAPIKeyAuth(
api_key="test-key",
user_id="test-user",
org_id="test-org",
)
# Mock model_group_settings to enable header forwarding for the model
# Use string-based patch to ensure we patch the current sys.modules['litellm']
# This avoids issues with module reloading during parallel test execution
mock_settings = MagicMock(forward_client_headers_to_llm_api=["local-openai/*"])
with patch("litellm.model_group_settings", mock_settings):
# Call add_litellm_data_to_request which includes header forwarding logic
updated_data = await add_litellm_data_to_request(
data=data,
request=request_mock,
user_api_key_dict=user_api_key_dict,
proxy_config=MagicMock(),
general_settings={},
version="test-version",
)
# Verify that headers were added to the request data
assert "headers" in updated_data, "Headers should be added to embedding request"
# Verify that only x- prefixed headers (except x-stainless) were forwarded
forwarded_headers = updated_data["headers"]
assert "X-Custom-Header" in forwarded_headers, "X-Custom-Header should be forwarded"
assert forwarded_headers["X-Custom-Header"] == "custom-value"
assert "X-Request-ID" in forwarded_headers, "X-Request-ID should be forwarded"
assert forwarded_headers["X-Request-ID"] == "test-request-123"
# Verify that authorization header was NOT forwarded (sensitive header)
assert "Authorization" not in forwarded_headers, "Authorization header should not be forwarded"
# Verify that Content-Type was NOT forwarded (doesn't start with x-)
assert "Content-Type" not in forwarded_headers, "Content-Type should not be forwarded"
# Verify original data fields are preserved
assert updated_data["model"] == "local-openai/text-embedding-3-small"
assert updated_data["input"] == ["Text to embed"]
@pytest.mark.asyncio
async def test_embedding_header_forwarding_without_model_group_config():
"""
Test that headers are NOT forwarded for embedding requests when
the model is not in the forward_client_headers_to_llm_api list.
"""
import litellm
# Setup mock request for embeddings
request_mock = MagicMock(spec=Request)
request_mock.url.path = "/v1/embeddings"
request_mock.url = MagicMock()
request_mock.url.__str__.return_value = "http://localhost/v1/embeddings"
request_mock.method = "POST"
request_mock.query_params = {}
request_mock.headers = {
"Content-Type": "application/json",
"X-Custom-Header": "custom-value",
}
request_mock.client = MagicMock()
request_mock.client.host = "127.0.0.1"
# Setup embedding request data with a model NOT in the forward list
data = {
"model": "text-embedding-ada-002",
"input": ["Text to embed"],
}
user_api_key_dict = UserAPIKeyAuth(
api_key="test-key",
user_id="test-user",
)
# Mock model_group_settings with a different model in the forward list
mock_settings = MagicMock(forward_client_headers_to_llm_api=["gpt-4", "claude-*"])
original_model_group_settings = getattr(litellm, "model_group_settings", None)
litellm.model_group_settings = mock_settings
try:
updated_data = await add_litellm_data_to_request(
data=data,
request=request_mock,
user_api_key_dict=user_api_key_dict,
proxy_config=MagicMock(),
general_settings={},
version="test-version",
)
# Verify that headers were NOT added since model is not in forward list
assert "headers" not in updated_data or updated_data.get("headers") is None, (
"Headers should not be forwarded for models not in forward_client_headers_to_llm_api list"
)
# Verify original data fields are preserved
assert updated_data["model"] == "text-embedding-ada-002"
assert updated_data["input"] == ["Text to embed"]
finally:
# Restore original model_group_settings
litellm.model_group_settings = original_model_group_settings
@pytest.mark.asyncio
async def test_add_guardrails_from_policy_engine():
"""
Test that add_guardrails_from_policy_engine adds guardrails from matching policies
and tracks applied policies in metadata.
"""
from litellm.proxy.policy_engine.attachment_registry import get_attachment_registry
from litellm.proxy.policy_engine.policy_registry import get_policy_registry
from litellm.types.proxy.policy_engine import (
Policy,
PolicyAttachment,
PolicyGuardrails,
)
# Setup test data
data = {
"model": "gpt-4",
"messages": [{"role": "user", "content": "Hello"}],
"metadata": {},
}
user_api_key_dict = UserAPIKeyAuth(
api_key="test-key",
team_alias="healthcare-team",
key_alias="my-key",
)
# Setup mock policies in the registry (policies define WHAT guardrails to apply)
policy_registry = get_policy_registry()
policy_registry._policies = {
"global-baseline": Policy(
guardrails=PolicyGuardrails(add=["pii_blocker"]),
),
"healthcare": Policy(
guardrails=PolicyGuardrails(add=["hipaa_audit"]),
),
}
policy_registry._initialized = True
# Setup attachments in the attachment registry (attachments define WHERE policies apply)
attachment_registry = get_attachment_registry()
attachment_registry._attachments = [
PolicyAttachment(policy="global-baseline", scope="*"), # applies to all
PolicyAttachment(policy="healthcare", teams=["healthcare-team"]), # applies to healthcare team
]
attachment_registry._initialized = True
# Call the function
await add_guardrails_from_policy_engine(
data=data,
metadata_variable_name="metadata",
user_api_key_dict=user_api_key_dict,
)
# Verify guardrails were added
assert "guardrails" in data["metadata"]
assert "pii_blocker" in data["metadata"]["guardrails"]
assert "hipaa_audit" in data["metadata"]["guardrails"]
# Verify applied policies were tracked
assert "applied_policies" in data["metadata"]
assert "global-baseline" in data["metadata"]["applied_policies"]
assert "healthcare" in data["metadata"]["applied_policies"]
# Clean up registries
policy_registry._policies = {}
policy_registry._initialized = False
attachment_registry._attachments = []
attachment_registry._initialized = False
def test_match_and_track_policies_preserves_attachment_and_request_body_order():
from litellm.proxy.policy_engine.attachment_registry import AttachmentRegistry
from litellm.types.proxy.policy_engine import Policy, PolicyMatchContext
attachment_policy_names = [f"attachment-policy-{index}" for index in range(8)]
request_body_policy_names = ["body-policy-1", "body-policy-2"]
policy_names = [*attachment_policy_names, *request_body_policy_names]
policies = {policy_name: Policy() for policy_name in policy_names}
attachment_registry = AttachmentRegistry()
attachment_registry.load_attachments(
[{"policy": policy_name, "scope": "*"} for policy_name in attachment_policy_names]
)
applied_policy_names, _ = _match_and_track_policies(
data={"metadata": {}},
context=PolicyMatchContext(model="gpt-4"),
request_body_policies=request_body_policy_names,
policies_override=policies,
attachment_registry_override=attachment_registry,
)
assert applied_policy_names == policy_names
def test_match_and_track_policies_keeps_condition_missing_child_alongside_unconditional_sibling():
from litellm.proxy.policy_engine.attachment_registry import AttachmentRegistry
from litellm.types.proxy.policy_engine import (
Policy,
PolicyCondition,
PolicyGuardrails,
PolicyMatchContext,
)
policies = {
"baseline": Policy(guardrails=PolicyGuardrails(add=["baseline_guardrail"])),
"parent": Policy(guardrails=PolicyGuardrails(add=["pii_blocker"])),
"child": Policy(
inherit="parent",
guardrails=PolicyGuardrails(add=["child_guard"]),
condition=PolicyCondition(model="claude.*"),
),
}
attachment_registry = AttachmentRegistry()
attachment_registry.load_attachments(
[
{"policy": "baseline", "scope": "*"},
{"policy": "child", "scope": "*"},
]
)
data = {"metadata": {}}
applied_policy_names, _ = _match_and_track_policies(
data=data,
context=PolicyMatchContext(model="gpt-5.5"),
request_body_policies=[],
policies_override=policies,
attachment_registry_override=attachment_registry,
)
assert applied_policy_names == ["baseline", "child"]
assert data["metadata"]["applied_policies"] == ["baseline", "child"]
@pytest.mark.asyncio
async def test_add_guardrails_from_policy_engine_keeps_a_policy_added_guardrail_its_pipeline_also_steps():
from litellm.proxy.policy_engine.attachment_registry import get_attachment_registry
from litellm.proxy.policy_engine.policy_registry import get_policy_registry
from litellm.types.proxy.policy_engine import (
GuardrailPipeline,
PipelineStep,
Policy,
PolicyAttachment,
PolicyGuardrails,
)
data = {"model": "gpt-4", "messages": [{"role": "user", "content": "Hello"}], "metadata": {}}
policy_registry = get_policy_registry()
policy_registry._policies = {
"response-governance": Policy(
guardrails=PolicyGuardrails(add=["pii_blocker"]),
pipeline=GuardrailPipeline(mode="post_call", steps=[PipelineStep(guardrail="pii_blocker")]),
),
}
policy_registry._initialized = True
attachment_registry = get_attachment_registry()
attachment_registry._attachments = [PolicyAttachment(policy="response-governance", scope="*")]
attachment_registry._initialized = True
try:
await add_guardrails_from_policy_engine(
data=data,
metadata_variable_name="metadata",
user_api_key_dict=UserAPIKeyAuth(api_key="test-key"),
)
finally:
policy_registry._policies = {}
policy_registry._initialized = False
attachment_registry._attachments = []
attachment_registry._initialized = False
assert data["metadata"]["guardrails"] == ["pii_blocker"]
assert data["metadata"]["_pipeline_managed_guardrails"] == {"pii_blocker"}
assert [pipeline.mode for _policy_name, pipeline in data["metadata"]["_guardrail_pipelines"]] == ["post_call"]
@pytest.mark.asyncio
async def test_add_guardrails_from_policy_engine_applies_inherited_parent_guardrail_when_child_condition_misses():
from litellm.proxy.policy_engine.attachment_registry import get_attachment_registry
from litellm.proxy.policy_engine.policy_registry import get_policy_registry
from litellm.types.proxy.policy_engine import (
Policy,
PolicyAttachment,
PolicyCondition,
PolicyGuardrails,
)
data = {"model": "gpt-5.5", "messages": [{"role": "user", "content": "Hello"}], "metadata": {}}
policy_registry = get_policy_registry()
policy_registry._policies = {
"parent": Policy(guardrails=PolicyGuardrails(add=["pii_blocker"])),
"child": Policy(
inherit="parent",
guardrails=PolicyGuardrails(add=["child_guard"]),
condition=PolicyCondition(model="claude.*"),
),
}
policy_registry._initialized = True
attachment_registry = get_attachment_registry()
attachment_registry._attachments = [PolicyAttachment(policy="child", scope="*")]
attachment_registry._initialized = True
try:
await add_guardrails_from_policy_engine(
data=data,
metadata_variable_name="metadata",
user_api_key_dict=UserAPIKeyAuth(api_key="test-key"),
)
finally:
policy_registry._policies = {}
policy_registry._initialized = False
attachment_registry._attachments = []
attachment_registry._initialized = False
assert "pii_blocker" in data["metadata"]["guardrails"]
assert "child_guard" not in data["metadata"]["guardrails"]
@pytest.mark.asyncio
async def test_add_guardrails_from_policy_engine_accepts_dynamic_policies_and_pops_from_data():
"""
Test that add_guardrails_from_policy_engine accepts dynamic 'policies' from the request body
and removes them to prevent forwarding to the LLM provider.
This is critical because 'policies' is a LiteLLM proxy-specific parameter that should
not be sent to the actual LLM API (e.g., OpenAI, Anthropic, etc.).
"""
from litellm.proxy.policy_engine.policy_registry import get_policy_registry
# Setup test data with 'policies' in the request body
data = {
"model": "gpt-4",
"messages": [{"role": "user", "content": "Hello"}],
"policies": [
"PII-POLICY-GLOBAL",
"HIPAA-POLICY",
], # Dynamic policies - should be accepted and removed
"metadata": {},
}
user_api_key_dict = UserAPIKeyAuth(
api_key="test-key",
team_alias="test-team",
key_alias="test-key",
)
# Initialize empty policy registry (we're just testing the accept and pop behavior)
policy_registry = get_policy_registry()
policy_registry._policies = {}
policy_registry._initialized = False
# Call the function - should accept dynamic policies and not raise an error
await add_guardrails_from_policy_engine(
data=data,
metadata_variable_name="metadata",
user_api_key_dict=user_api_key_dict,
)
# Verify that 'policies' was removed from the request body
assert "policies" not in data, (
"'policies' should be removed from request body to prevent forwarding to LLM provider"
)
# Verify that other fields are preserved
assert "model" in data
assert data["model"] == "gpt-4"
assert "messages" in data
assert data["messages"] == [{"role": "user", "content": "Hello"}]
assert "metadata" in data
@pytest.mark.asyncio
async def test_api_created_global_policy_applies_to_new_key_without_restart():
"""
Regression: policies created at runtime via policy builder must apply
immediately when attached globally, even if the server started with no
initialized policy config.
"""
from litellm.proxy.policy_engine.attachment_registry import get_attachment_registry
from litellm.proxy.policy_engine.policy_registry import get_policy_registry
from litellm.types.proxy.policy_engine import (
Policy,
PolicyAttachment,
PolicyGuardrails,
)
data = {
"model": "gpt-4",
"messages": [{"role": "user", "content": "Hello"}],
"metadata": {},
}
user_api_key_dict = UserAPIKeyAuth(api_key="test-key")
policy_registry = get_policy_registry()
attachment_registry = get_attachment_registry()
policy_registry._policies = {}
policy_registry._policies_by_id = {}
policy_registry._initialized = False
attachment_registry._attachments = []
attachment_registry._initialized = False
try:
policy_registry.add_policy(
"runtime-global-policy",
Policy(guardrails=PolicyGuardrails(add=["runtime-guardrail"])),
)
attachment_registry.add_attachment(PolicyAttachment(policy="runtime-global-policy", scope="*"))
await add_guardrails_from_policy_engine(
data=data,
metadata_variable_name="metadata",
user_api_key_dict=user_api_key_dict,
)
assert "runtime-guardrail" in data["metadata"]["guardrails"]
assert "runtime-global-policy" in data["metadata"]["applied_policies"]
finally:
policy_registry._policies = {}
policy_registry._policies_by_id = {}
policy_registry._initialized = False
attachment_registry._attachments = []
attachment_registry._initialized = False
@pytest.mark.asyncio
async def test_add_guardrails_from_policy_engine_policy_version_by_id():
"""
Test that add_guardrails_from_policy_engine executes a specific policy version
when policy_<uuid> is passed in the request body.
"""
from litellm.proxy.policy_engine.attachment_registry import get_attachment_registry
from litellm.proxy.policy_engine.policy_registry import get_policy_registry
from litellm.types.proxy.policy_engine import Policy, PolicyGuardrails
policy_version_uuid = "12345678-1234-5678-1234-567812345678"
policy_version_ref = f"policy_{policy_version_uuid}"
# Policy from the specific version (e.g. published) - different guardrail than production
published_version_policy = Policy(
guardrails=PolicyGuardrails(add=["published_version_guardrail"]),
)
data = {
"model": "gpt-4",
"messages": [{"role": "user", "content": "Hello"}],
"policies": [policy_version_ref],
"metadata": {},
}
user_api_key_dict = UserAPIKeyAuth(
api_key="test-key",
team_alias="test-team",
key_alias="test-key",
)
policy_registry = get_policy_registry()
policy_registry._policies = {}
policy_registry._initialized = True
attachment_registry = get_attachment_registry()
attachment_registry._attachments = []
attachment_registry._initialized = True
with patch.object(
policy_registry,
"get_policy_by_id_for_request",
return_value=("test-policy-from-version", published_version_policy),
):
await add_guardrails_from_policy_engine(
data=data,
metadata_variable_name="metadata",
user_api_key_dict=user_api_key_dict,
)
# Verify guardrails from the specific version were applied
assert "metadata" in data
assert "guardrails" in data["metadata"]
assert "published_version_guardrail" in data["metadata"]["guardrails"]
assert "policies" not in data
# Clean up
policy_registry._policies = {}
policy_registry._initialized = False
@pytest.mark.asyncio
async def test_bearer_token_not_in_debug_logs():
"""
E2E regression test for the client-reported JWT leak.
Calls add_litellm_data_to_request with a Bearer token in the request
headers and captures all debug log output. Asserts the raw token never
appears in any log message — covering the exact paths the client reported:
- "Request Headers: ..."
- "receiving data: ..."
- "[PROXY] returned data from litellm_pre_call_utils: ..."
"""
import logging
from io import StringIO
from litellm.proxy.litellm_pre_call_utils import add_litellm_data_to_request
from litellm.proxy.proxy_server import ProxyConfig
secret_token = "eyJhbGciOiJSUzI1NiIsInR5cCI6IkpXVCJ9.eyJzdWIiOiIxMjM0NTY3ODkwIn0.fakesignature"
mock_request = MagicMock(spec=Request)
mock_request.headers = {
"authorization": f"Bearer {secret_token}",
"content-type": "application/json",
}
mock_request.url = MagicMock()
mock_request.url.__str__ = lambda self: "http://localhost:4000/v1/chat/completions"
mock_request.method = "POST"
mock_request.query_params = {}
data = {
"model": "gpt-4",
"messages": [{"role": "user", "content": "hi"}],
}
user_api_key_dict = UserAPIKeyAuth(api_key="sk-1234")
# Capture all debug log output from the proxy logger
log_capture = StringIO()
log_handler = logging.StreamHandler(log_capture)
log_handler.setLevel(logging.DEBUG)
logger = logging.getLogger("LiteLLM Proxy")
logger.addHandler(log_handler)
original_level = logger.level
logger.setLevel(logging.DEBUG)
try:
with (
patch("litellm.proxy.proxy_server.llm_router", None),
patch("litellm.proxy.proxy_server.premium_user", True),
):
await add_litellm_data_to_request(
data=data,
request=mock_request,
user_api_key_dict=user_api_key_dict,
proxy_config=ProxyConfig(),
general_settings={},
)
finally:
logger.removeHandler(log_handler)
logger.setLevel(original_level)
log_output = log_capture.getvalue()
assert secret_token not in log_output, (
f"Bearer token leaked in debug logs. Found token in log output:\n{log_output[:500]}"
)
# ============================================================================
# Tests for credential overrides from model_config (team/project metadata)
# ============================================================================
@pytest.fixture()
def setup_test_credentials():
"""Populate litellm.credential_list with test credentials and enable feature flag, clean up after."""
original = litellm.credential_list[:]
original_flag = litellm.enable_model_config_credential_overrides
litellm.enable_model_config_credential_overrides = True
litellm.credential_list.extend(
[
CredentialItem(
credential_name="hotel-azure-eastus",
credential_info={},
credential_values={
"api_base": "https://hotel-eastus.openai.azure.com/",
"api_key": "key-hotel-eastus",
},
),
CredentialItem(
credential_name="hotel-azure-westus",
credential_info={},
credential_values={
"api_base": "https://hotel-westus.openai.azure.com/",
"api_key": "key-hotel-westus",
},
),
CredentialItem(
credential_name="hotel-rec-azure",
credential_info={},
credential_values={
"api_base": "https://hotel-rec-app.openai.azure.com/",
"api_key": "key-hotel-rec",
},
),
CredentialItem(
credential_name="hotel-rec-vision",
credential_info={},
credential_values={
"api_base": "https://hotel-rec-vision.openai.azure.com/",
"api_key": "key-hotel-rec-vision",
"api_version": "2024-06-01",
},
),
CredentialItem(
credential_name="flight-azure-centralus",
credential_info={},
credential_values={
"api_base": "https://flight-centralus.openai.azure.com/",
"api_key": "key-flight-centralus",
},
),
]
)
yield
litellm.credential_list[:] = original
litellm.enable_model_config_credential_overrides = original_flag
# --- Unit tests for _extract_credential_from_entry ---
def test_extract_credential_from_entry_azure():
entry = {"azure": {"litellm_credentials": "my-cred"}}
assert _extract_credential_from_entry(entry) == "my-cred"
def test_extract_credential_from_entry_no_credential():
entry = {"azure": {"some_other_key": "value"}}
assert _extract_credential_from_entry(entry) is None
def test_extract_credential_from_entry_empty():
assert _extract_credential_from_entry({}) is None
def test_extract_credential_from_entry_non_dict_value():
entry = {"azure": "not-a-dict"}
assert _extract_credential_from_entry(entry) is None
def test_extract_credential_from_entry_non_dict_entry():
"""Non-dict entry (e.g. string) should return None, not crash."""
assert _extract_credential_from_entry("my-cred-name") is None
assert _extract_credential_from_entry(["a", "list"]) is None
assert _extract_credential_from_entry(42) is None
# --- Unit tests for _resolve_credential_from_model_config ---
def test_resolve_project_model_specific_wins():
project_config = {
"gpt-4": {"azure": {"litellm_credentials": "proj-gpt4"}},
"defaultconfig": {"azure": {"litellm_credentials": "proj-default"}},
}
team_config = {
"gpt-4": {"azure": {"litellm_credentials": "team-gpt4"}},
"defaultconfig": {"azure": {"litellm_credentials": "team-default"}},
}
result = _resolve_credential_from_model_config("gpt-4", project_config, team_config)
assert result == "proj-gpt4"
def test_resolve_project_default_wins_over_team():
project_config = {
"defaultconfig": {"azure": {"litellm_credentials": "proj-default"}},
}
team_config = {
"gpt-4": {"azure": {"litellm_credentials": "team-gpt4"}},
"defaultconfig": {"azure": {"litellm_credentials": "team-default"}},
}
result = _resolve_credential_from_model_config("gpt-4", project_config, team_config)
assert result == "proj-default"
def test_resolve_team_model_specific_wins_over_team_default():
team_config = {
"gpt-4": {"azure": {"litellm_credentials": "team-gpt4"}},
"defaultconfig": {"azure": {"litellm_credentials": "team-default"}},
}
result = _resolve_credential_from_model_config("gpt-4", None, team_config)
assert result == "team-gpt4"
def test_resolve_team_default_used_as_fallback():
team_config = {
"defaultconfig": {"azure": {"litellm_credentials": "team-default"}},
}
result = _resolve_credential_from_model_config("gpt-3.5", None, team_config)
assert result == "team-default"
def test_resolve_no_match_returns_none():
result = _resolve_credential_from_model_config("gpt-4", None, None)
assert result is None
def test_resolve_empty_configs_returns_none():
result = _resolve_credential_from_model_config("gpt-4", {}, {})
assert result is None
def test_resolve_model_not_in_any_config():
project_config = {"gpt-4": {"azure": {"litellm_credentials": "x"}}}
result = _resolve_credential_from_model_config("gpt-3.5", project_config, None)
assert result is None
# --- Integration tests for _apply_credential_overrides_from_model_config ---
def test_apply_overrides_project_model_specific(setup_test_credentials):
"""Scenario 2: Hotel Rec App -> gpt-4-vision -> project model-specific."""
data = {"model": "gpt-4-vision"}
user_api_key_dict = UserAPIKeyAuth(
api_key="test-key",
team_metadata={
"model_config": {
"defaultconfig": {"azure": {"litellm_credentials": "hotel-azure-eastus"}},
"gpt-4": {"azure": {"litellm_credentials": "hotel-azure-westus"}},
}
},
project_metadata={
"model_config": {
"defaultconfig": {"azure": {"litellm_credentials": "hotel-rec-azure"}},
"gpt-4-vision": {"azure": {"litellm_credentials": "hotel-rec-vision"}},
}
},
)
_apply_credential_overrides_from_model_config(data=data, user_api_key_dict=user_api_key_dict)
assert data["api_base"] == "https://hotel-rec-vision.openai.azure.com/"
assert data["api_key"] == "key-hotel-rec-vision"
assert data["api_version"] == "2024-06-01"
def test_apply_overrides_project_default(setup_test_credentials):
"""Scenario 1: Hotel Rec App -> gpt-4 -> project default."""
data = {"model": "gpt-4"}
user_api_key_dict = UserAPIKeyAuth(
api_key="test-key",
team_metadata={
"model_config": {
"defaultconfig": {"azure": {"litellm_credentials": "hotel-azure-eastus"}},
"gpt-4": {"azure": {"litellm_credentials": "hotel-azure-westus"}},
}
},
project_metadata={
"model_config": {
"defaultconfig": {"azure": {"litellm_credentials": "hotel-rec-azure"}},
"gpt-4-vision": {"azure": {"litellm_credentials": "hotel-rec-vision"}},
}
},
)
_apply_credential_overrides_from_model_config(data=data, user_api_key_dict=user_api_key_dict)
assert data["api_base"] == "https://hotel-rec-app.openai.azure.com/"
assert data["api_key"] == "key-hotel-rec"
def test_apply_overrides_team_model_specific(setup_test_credentials):
"""Scenario 4: Hotel Review App -> gpt-4 -> team model-specific."""
data = {"model": "gpt-4"}
user_api_key_dict = UserAPIKeyAuth(
api_key="test-key",
team_metadata={
"model_config": {
"defaultconfig": {"azure": {"litellm_credentials": "hotel-azure-eastus"}},
"gpt-4": {"azure": {"litellm_credentials": "hotel-azure-westus"}},
}
},
project_metadata={},
)
_apply_credential_overrides_from_model_config(data=data, user_api_key_dict=user_api_key_dict)
assert data["api_base"] == "https://hotel-westus.openai.azure.com/"
assert data["api_key"] == "key-hotel-westus"
def test_apply_overrides_team_default(setup_test_credentials):
"""Scenario 3: Hotel Review App -> gpt-3.5 -> team default."""
data = {"model": "gpt-3.5"}
user_api_key_dict = UserAPIKeyAuth(
api_key="test-key",
team_metadata={
"model_config": {
"defaultconfig": {"azure": {"litellm_credentials": "hotel-azure-eastus"}},
"gpt-4": {"azure": {"litellm_credentials": "hotel-azure-westus"}},
}
},
project_metadata={},
)
_apply_credential_overrides_from_model_config(data=data, user_api_key_dict=user_api_key_dict)
assert data["api_base"] == "https://hotel-eastus.openai.azure.com/"
assert data["api_key"] == "key-hotel-eastus"
def test_apply_overrides_no_config(setup_test_credentials):
"""Scenario 6: No model_config anywhere -> data unchanged."""
data = {"model": "gpt-4"}
user_api_key_dict = UserAPIKeyAuth(
api_key="test-key",
team_metadata={},
project_metadata={},
)
_apply_credential_overrides_from_model_config(data=data, user_api_key_dict=user_api_key_dict)
assert "api_base" not in data
assert "api_key" not in data
def test_apply_overrides_clientside_credentials_take_precedence(
setup_test_credentials,
):
"""Clientside api_base/api_key in data should block model_config override."""
data = {
"model": "gpt-4",
"api_base": "https://my-custom-endpoint.openai.azure.com/",
"api_key": "my-custom-key",
}
user_api_key_dict = UserAPIKeyAuth(
api_key="test-key",
team_metadata={"model_config": {"defaultconfig": {"azure": {"litellm_credentials": "hotel-azure-eastus"}}}},
)
_apply_credential_overrides_from_model_config(data=data, user_api_key_dict=user_api_key_dict)
assert data["api_base"] == "https://my-custom-endpoint.openai.azure.com/"
assert data["api_key"] == "my-custom-key"
def test_apply_overrides_missing_credential_name(setup_test_credentials):
"""model_config references a credential that doesn't exist -> no override."""
data = {"model": "gpt-4"}
user_api_key_dict = UserAPIKeyAuth(
api_key="test-key",
team_metadata={"model_config": {"gpt-4": {"azure": {"litellm_credentials": "nonexistent-credential"}}}},
)
_apply_credential_overrides_from_model_config(data=data, user_api_key_dict=user_api_key_dict)
assert "api_base" not in data
assert "api_key" not in data
def test_apply_overrides_api_version_only_if_present(setup_test_credentials):
"""api_version should only be set if the credential contains it."""
data = {"model": "gpt-3.5"}
user_api_key_dict = UserAPIKeyAuth(
api_key="test-key",
team_metadata={"model_config": {"defaultconfig": {"azure": {"litellm_credentials": "hotel-azure-eastus"}}}},
)
_apply_credential_overrides_from_model_config(data=data, user_api_key_dict=user_api_key_dict)
assert data["api_base"] == "https://hotel-eastus.openai.azure.com/"
assert data["api_key"] == "key-hotel-eastus"
assert "api_version" not in data
def test_apply_overrides_no_model_in_data(setup_test_credentials):
"""No model in request data -> skip override."""
data = {"messages": [{"role": "user", "content": "hello"}]}
user_api_key_dict = UserAPIKeyAuth(
api_key="test-key",
team_metadata={"model_config": {"defaultconfig": {"azure": {"litellm_credentials": "some-cred"}}}},
)
_apply_credential_overrides_from_model_config(data=data, user_api_key_dict=user_api_key_dict)
assert "api_base" not in data
def test_apply_overrides_none_metadata(setup_test_credentials):
"""None metadata on both team and project -> skip override."""
data = {"model": "gpt-4"}
user_api_key_dict = UserAPIKeyAuth(
api_key="test-key",
team_metadata=None,
project_metadata=None,
)
_apply_credential_overrides_from_model_config(data=data, user_api_key_dict=user_api_key_dict)
assert "api_base" not in data
def test_apply_overrides_clientside_api_version_preserved(setup_test_credentials):
"""Clientside api_version should not be overwritten by credential."""
data = {"model": "gpt-4-vision", "api_version": "2025-01-01"}
user_api_key_dict = UserAPIKeyAuth(
api_key="test-key",
team_metadata={"model_config": {"gpt-4-vision": {"azure": {"litellm_credentials": "hotel-rec-vision"}}}},
)
_apply_credential_overrides_from_model_config(data=data, user_api_key_dict=user_api_key_dict)
# api_base and api_key should be set from credential
assert data["api_base"] == "https://hotel-rec-vision.openai.azure.com/"
assert data["api_key"] == "key-hotel-rec-vision"
# api_version should be preserved from the request, not overwritten
assert data["api_version"] == "2025-01-01"
def test_resolve_non_dict_model_config_ignored():
"""Non-dict model_config (e.g. string) should be safely skipped."""
result = _resolve_credential_from_model_config("gpt-4", "not-a-dict", None)
assert result is None
result = _resolve_credential_from_model_config("gpt-4", None, ["also", "not", "a", "dict"])
assert result is None
# Valid config still works alongside invalid one
result = _resolve_credential_from_model_config(
"gpt-4",
"invalid",
{"gpt-4": {"azure": {"litellm_credentials": "valid-cred"}}},
)
assert result == "valid-cred"
def test_resolve_pre_alias_model_name_fallback():
"""model_config keyed on pre-alias name should match after alias resolution."""
team_config = {
"gpt-4": {"azure": {"litellm_credentials": "team-gpt4"}},
}
# Post-alias name doesn't match, but pre-alias does (team scope)
result = _resolve_credential_from_model_config("azure/gpt-4-0613", None, team_config, pre_alias_model_name="gpt-4")
assert result == "team-gpt4"
# Same test for project scope
project_config = {
"gpt-4": {"azure": {"litellm_credentials": "proj-gpt4"}},
}
result = _resolve_credential_from_model_config(
"azure/gpt-4-0613", project_config, None, pre_alias_model_name="gpt-4"
)
assert result == "proj-gpt4"
def test_resolve_post_alias_name_takes_priority():
"""Post-alias (resolved) name should be tried before pre-alias name."""
team_config = {
"gpt-4": {"azure": {"litellm_credentials": "pre-alias-cred"}},
"gpt-4o-team-1": {"azure": {"litellm_credentials": "post-alias-cred"}},
}
# Team scope
result = _resolve_credential_from_model_config("gpt-4o-team-1", None, team_config, pre_alias_model_name="gpt-4")
assert result == "post-alias-cred"
# Project scope
result = _resolve_credential_from_model_config("gpt-4o-team-1", team_config, None, pre_alias_model_name="gpt-4")
assert result == "post-alias-cred"
def test_apply_overrides_with_alias(setup_test_credentials):
"""Credential override should work when model name was changed by alias."""
# Simulate: user called "my-gpt4", alias resolved to "azure/gpt-4-custom"
# model_config is keyed on "my-gpt4" (the pre-alias name)
data = {"model": "azure/gpt-4-custom"}
user_api_key_dict = UserAPIKeyAuth(
api_key="test-key",
team_metadata={
"model_config": {
"my-gpt4": {"azure": {"litellm_credentials": "hotel-azure-eastus"}},
}
},
)
_apply_credential_overrides_from_model_config(
data=data,
user_api_key_dict=user_api_key_dict,
pre_alias_model_name="my-gpt4",
)
assert data["api_base"] == "https://hotel-eastus.openai.azure.com/"
assert data["api_key"] == "key-hotel-eastus"
def test_apply_overrides_feature_flag_disabled_by_default():
"""Feature flag defaults to False — credential overrides are inert until explicitly enabled."""
assert litellm.enable_model_config_credential_overrides is False
data = {"model": "gpt-4"}
user_api_key_dict = UserAPIKeyAuth(
api_key="test-key",
team_metadata={"model_config": {"gpt-4": {"azure": {"litellm_credentials": "hotel-azure-eastus"}}}},
)
_apply_credential_overrides_from_model_config(data=data, user_api_key_dict=user_api_key_dict)
assert "api_base" not in data
assert "api_key" not in data
def test_extract_credential_provider_hint_prefers_exact_match():
"""Provider hint selects the correct provider in a multi-provider entry."""
entry = {
"openai": {"litellm_credentials": "openai-cred"},
"azure": {"litellm_credentials": "azure-cred"},
}
# With provider hint, should pick the exact match
assert _extract_credential_from_entry(entry, provider="azure") == "azure-cred"
assert _extract_credential_from_entry(entry, provider="openai") == "openai-cred"
# Without provider hint, falls back to first key (insertion order)
result = _extract_credential_from_entry(entry)
assert result in ("openai-cred", "azure-cred")
# Unknown provider falls back to first available
result = _extract_credential_from_entry(entry, provider="bedrock")
assert result in ("openai-cred", "azure-cred")
def test_resolve_provider_hint_from_model_name():
"""Provider prefix in model name (e.g. azure/gpt-4) threads through to entry extraction."""
config = {
"gpt-4": {
"openai": {"litellm_credentials": "openai-cred"},
"azure": {"litellm_credentials": "azure-cred"},
},
}
# Model name "azure/gpt-4" -> provider="azure" -> should prefer azure-cred
# But _resolve_credential_from_model_config tries "azure/gpt-4" first (no match),
# then falls to defaultconfig (no match). So we need to use pre_alias_model_name.
result = _resolve_credential_from_model_config(
"azure/gpt-4", config, None, pre_alias_model_name="gpt-4", provider="azure"
)
assert result == "azure-cred"
def test_clean_headers_preserves_x_api_key_when_byok_enabled():
"""
Regression test: when forward_llm_provider_auth_headers=True,
clean_headers() must preserve the client-supplied x-api-key header
so it can be forwarded to the upstream Anthropic API (BYOK flow).
"""
headers = Headers(
{
"x-api-key": "sk-ant-api03-client-key",
"x-litellm-api-key": "sk-proxy-virtual-key",
"content-type": "application/json",
}
)
result = clean_headers(
headers=headers,
litellm_key_header_name="x-litellm-api-key",
forward_llm_provider_auth_headers=True,
authenticated_with_header="x-litellm-api-key",
)
# x-api-key must be preserved for BYOK
assert result.get("x-api-key") == "sk-ant-api03-client-key"
# x-litellm-api-key must NOT leak to the upstream
assert "x-litellm-api-key" not in result
def test_clean_headers_strips_x_api_key_when_byok_disabled():
"""
Regression test: with forward_llm_provider_auth_headers=False (default),
x-api-key must be stripped so proxy-configured keys are not overridden
by a client-supplied one.
"""
headers = Headers(
{
"x-api-key": "sk-ant-api03-client-key",
"x-litellm-api-key": "sk-proxy-virtual-key",
}
)
result = clean_headers(
headers=headers,
litellm_key_header_name="x-litellm-api-key",
forward_llm_provider_auth_headers=False,
authenticated_with_header="x-litellm-api-key",
)
assert "x-api-key" not in result
def test_clean_headers_strips_x_api_key_when_byok_enabled_but_x_api_key_was_auth_header():
"""
Anti-replay regression: even when forward_llm_provider_auth_headers=True,
if the client authenticated TO the proxy using x-api-key (i.e., the proxy
key arrived as x-api-key), clean_headers() must NOT forward that header
upstream. Otherwise a proxy-auth key would leak to the LLM provider.
"""
headers = Headers(
{
"x-api-key": "sk-proxy-auth-key-masquerading-as-anthropic-key",
"content-type": "application/json",
}
)
result = clean_headers(
headers=headers,
litellm_key_header_name="x-litellm-api-key",
forward_llm_provider_auth_headers=True,
authenticated_with_header="x-api-key",
)
# Even with BYOK enabled, x-api-key must be stripped when it was used
# as the LiteLLM auth header (anti-replay guard).
assert "x-api-key" not in result
# ---------------------------------------------------------------------------
# Team guardrail + global policy regression tests
# ---------------------------------------------------------------------------
@pytest.mark.asyncio
async def test_move_guardrails_to_metadata_moves_include_guardrail_response_before_the_no_guardrail_early_out():
policy_registry = MagicMock()
policy_registry.is_initialized.return_value = False
user_api_key_dict = UserAPIKeyAuth(api_key="test-key")
true_data = {
"model": "gpt-4.1-mini",
"messages": [{"role": "user", "content": "Hello"}],
"metadata": {},
"include_guardrail_response": True,
}
with patch("litellm.proxy.policy_engine.policy_registry.get_policy_registry", return_value=policy_registry):
await move_guardrails_to_metadata(
data=true_data,
_metadata_variable_name="metadata",
user_api_key_dict=user_api_key_dict,
)
assert "include_guardrail_response" not in true_data
assert true_data["metadata"]["include_guardrail_response"] is True
string_data = {
"model": "gpt-4.1-mini",
"messages": [{"role": "user", "content": "Hello"}],
"metadata": {},
"include_guardrail_response": "true",
}
with patch("litellm.proxy.policy_engine.policy_registry.get_policy_registry", return_value=policy_registry):
await move_guardrails_to_metadata(
data=string_data,
_metadata_variable_name="metadata",
user_api_key_dict=user_api_key_dict,
)
assert "include_guardrail_response" not in string_data
assert string_data["metadata"]["include_guardrail_response"] is False
@pytest.mark.asyncio
async def test_team_guardrail_merges_with_global_policy():
"""
Regression: team's direct guardrail must be present alongside guardrails
resolved from a global policy (scope='*') configured by the admin.
The bug: get_guardrail_from_metadata checked litellm_metadata before
metadata. When the request contained a non-empty litellm_metadata field
(without a 'guardrails' key), the merged list in data["metadata"] was
shadowed and non-default guardrails silently received an empty
requested_guardrails list.
"""
from litellm.proxy.policy_engine.attachment_registry import get_attachment_registry
from litellm.proxy.policy_engine.policy_registry import get_policy_registry
from litellm.proxy.litellm_pre_call_utils import move_guardrails_to_metadata
from litellm.types.proxy.policy_engine import (
Policy,
PolicyAttachment,
PolicyGuardrails,
)
data = {
"model": "gpt-4",
"messages": [{"role": "user", "content": "Hello"}],
# Simulate a request that carries litellm_metadata (without guardrails)
# which previously shadowed data["metadata"]["guardrails"].
"litellm_metadata": {"some_user_field": "some_value"},
"metadata": {},
}
user_api_key_dict = UserAPIKeyAuth(
api_key="test-key",
team_metadata={"guardrails": ["team-direct-guardrail"]},
)
policy_registry = get_policy_registry()
policy_registry._policies = {
"global-policy": Policy(
guardrails=PolicyGuardrails(add=["policy-guardrail-1", "policy-guardrail-2"]),
),
}
policy_registry._initialized = True
attachment_registry = get_attachment_registry()
attachment_registry._attachments = [
PolicyAttachment(policy="global-policy", scope="*"),
]
attachment_registry._initialized = True
try:
with patch("litellm.proxy.utils._premium_user_check"):
await move_guardrails_to_metadata(
data=data,
_metadata_variable_name="metadata",
user_api_key_dict=user_api_key_dict,
)
guardrails = data["metadata"].get("guardrails", [])
assert "team-direct-guardrail" in guardrails, f"Team guardrail missing from merged list: {guardrails}"
assert "policy-guardrail-1" in guardrails, f"policy-guardrail-1 missing: {guardrails}"
assert "policy-guardrail-2" in guardrails, f"policy-guardrail-2 missing: {guardrails}"
assert len(guardrails) == len(set(guardrails)), f"Duplicates in guardrails list: {guardrails}"
# Verify get_guardrail_from_metadata returns the merged list even
# when litellm_metadata is present (the bug: it returned [] before fix)
from litellm.integrations.custom_guardrail import CustomGuardrail
class _DummyGuardrail(CustomGuardrail):
pass
dummy = _DummyGuardrail(guardrail_name="team-direct-guardrail")
returned = dummy.get_guardrail_from_metadata(data)
assert "team-direct-guardrail" in returned, (
f"get_guardrail_from_metadata shadowed by litellm_metadata; got: {returned}"
)
finally:
policy_registry._policies = {}
policy_registry._initialized = False
attachment_registry._attachments = []
attachment_registry._initialized = False
@pytest.mark.asyncio
async def test_get_guardrail_from_metadata_prefers_metadata_over_litellm_metadata():
"""
Unit test: get_guardrail_from_metadata must read from data["metadata"] first.
A non-empty data["litellm_metadata"] without a 'guardrails' key must not
shadow data["metadata"]["guardrails"].
"""
from litellm.integrations.custom_guardrail import CustomGuardrail
class _DummyGuardrail(CustomGuardrail):
pass
dummy = _DummyGuardrail(guardrail_name="my-guardrail")
data = {
"metadata": {"guardrails": ["my-guardrail", "other-guardrail"]},
"litellm_metadata": {"some_field": "some_value"}, # no 'guardrails' key
}
result = dummy.get_guardrail_from_metadata(data)
assert result == [
"my-guardrail",
"other-guardrail",
], f"Expected guardrails from metadata, got: {result}"
def test_get_guardrail_from_metadata_reads_litellm_metadata_when_no_metadata():
"""
get_guardrail_from_metadata must still read from litellm_metadata when
data["metadata"] has no 'guardrails' key (thread/assistant endpoint path).
"""
from litellm.integrations.custom_guardrail import CustomGuardrail
class _DummyGuardrail(CustomGuardrail):
pass
dummy = _DummyGuardrail(guardrail_name="my-guardrail")
data = {
"metadata": {"requester_metadata": {"user": "alice"}}, # no guardrails key
"litellm_metadata": {"guardrails": ["my-guardrail"]},
}
result = dummy.get_guardrail_from_metadata(data)
assert result == ["my-guardrail"], f"Expected guardrails from litellm_metadata fallback, got: {result}"
def _build_request_mock_with_headers(headers: dict) -> Request:
request_mock = MagicMock(spec=Request)
request_mock.url = MagicMock()
request_mock.url.path = "/v1/chat/completions"
request_mock.url.__str__.return_value = "http://localhost/v1/chat/completions"
request_mock.method = "POST"
request_mock.query_params = {}
request_mock.headers = headers
request_mock.client = MagicMock()
request_mock.client.host = "127.0.0.1"
request_mock.state = MagicMock()
request_mock.state._cached_headers = None
return request_mock
class TestApplyClientTagPolicyPreAuth:
"""Tests for ``LiteLLMProxyRequestSetup.apply_client_tag_policy_pre_auth``.
Regression coverage for the bug where ``x-litellm-tags`` header was
invisible to ``_tag_max_budget_check`` because the merge happened
post-auth in ``add_litellm_data_to_request``.
"""
def test_merges_header_tags_into_metadata(self):
request_mock = _build_request_mock_with_headers({"x-litellm-tags": "tenant:acme,env:prod"})
data = {"model": "gpt-3.5-turbo"}
user_api_key_dict = UserAPIKeyAuth(
api_key="hashed-key",
metadata={},
team_metadata={},
)
LiteLLMProxyRequestSetup.apply_client_tag_policy_pre_auth(
request=request_mock,
request_data=data,
user_api_key_dict=user_api_key_dict,
)
assert data["metadata"]["tags"] == ["tenant:acme", "env:prod"]
def test_unions_header_tags_with_existing_metadata_tags(self):
request_mock = _build_request_mock_with_headers({"x-litellm-tags": "tenant:acme,env:prod"})
data = {
"model": "gpt-3.5-turbo",
"metadata": {"tags": ["env:prod", "team:platform"]},
}
user_api_key_dict = UserAPIKeyAuth(
api_key="hashed-key",
metadata={},
team_metadata={},
)
LiteLLMProxyRequestSetup.apply_client_tag_policy_pre_auth(
request=request_mock,
request_data=data,
user_api_key_dict=user_api_key_dict,
)
# Existing tags first, dedupe header tags
assert data["metadata"]["tags"] == ["env:prod", "team:platform", "tenant:acme"]
def test_preserves_body_tags(self):
# Pre-auth must NOT touch body-supplied tags. _tag_max_budget_check
# (inside common_checks) enforces per-tag budgets on whatever tags
# it sees in request_data, including body tags. The helper only
# adds header tags to metadata.tags.
request_mock = _build_request_mock_with_headers({"x-litellm-tags": "tenant:acme"})
data = {
"model": "gpt-3.5-turbo",
"tags": ["root-tag"],
"litellm_metadata": {"tags": ["litellm-meta-tag"]},
}
user_api_key_dict = UserAPIKeyAuth(
api_key="hashed-key",
metadata={},
team_metadata={},
)
LiteLLMProxyRequestSetup.apply_client_tag_policy_pre_auth(
request=request_mock,
request_data=data,
user_api_key_dict=user_api_key_dict,
)
assert data["tags"] == ["root-tag"]
# litellm_metadata is the active metadata key (it's present), so
# header tags merge into it and union with existing tags there.
assert data["litellm_metadata"]["tags"] == [
"litellm-meta-tag",
"tenant:acme",
]
def test_uses_litellm_metadata_when_present(self):
request_mock = _build_request_mock_with_headers({"x-litellm-tags": "tenant:acme"})
data = {
"model": "gpt-3.5-turbo",
"litellm_metadata": {"foo": "bar"},
}
user_api_key_dict = UserAPIKeyAuth(
api_key="hashed-key",
metadata={},
team_metadata={},
)
LiteLLMProxyRequestSetup.apply_client_tag_policy_pre_auth(
request=request_mock,
request_data=data,
user_api_key_dict=user_api_key_dict,
)
# get_metadata_variable_name_from_kwargs returns "litellm_metadata"
# when present, so header tags should land there to be visible to
# _tag_max_budget_check.
assert data["litellm_metadata"]["tags"] == ["tenant:acme"]
assert "tags" not in data.get("metadata", {})
def test_no_header_no_mutation(self):
request_mock = _build_request_mock_with_headers({})
data = {"model": "gpt-3.5-turbo"}
user_api_key_dict = UserAPIKeyAuth(
api_key="hashed-key",
metadata={},
team_metadata={},
)
LiteLLMProxyRequestSetup.apply_client_tag_policy_pre_auth(
request=request_mock,
request_data=data,
user_api_key_dict=user_api_key_dict,
)
assert "metadata" not in data or "tags" not in data["metadata"]
def test_string_metadata_tags_survive_header_merge(self):
# metadata can arrive as a JSON string (multipart/form-data, extra_body).
# The pre-auth merge must parse it so an over-budget body tag isn't
# silently dropped when a within-budget header tag is also present.
request_mock = _build_request_mock_with_headers({"x-litellm-tags": "free"})
data = {
"model": "gpt-3.5-turbo",
"metadata": '{"tags": ["paid"]}',
}
user_api_key_dict = UserAPIKeyAuth(
api_key="hashed-key",
metadata={},
team_metadata={},
)
LiteLLMProxyRequestSetup.apply_client_tag_policy_pre_auth(
request=request_mock,
request_data=data,
user_api_key_dict=user_api_key_dict,
)
assert isinstance(data["metadata"], dict)
assert data["metadata"]["tags"] == ["paid", "free"]
@pytest.mark.asyncio
async def test_string_metadata_does_not_bypass_tag_max_budget_check(self):
"""Regression: string metadata containing an over-budget tag must not
be silently overwritten when an x-litellm-tags header is present."""
from litellm.proxy._types import LiteLLM_BudgetTable, LiteLLM_TagTable
from litellm.proxy.auth.auth_checks import _tag_max_budget_check
from litellm.proxy.utils import ProxyLogging
request_mock = _build_request_mock_with_headers({"x-litellm-tags": "free"})
data = {
"model": "gpt-3.5-turbo",
"metadata": '{"tags": ["paid"]}',
}
user_api_key_dict = UserAPIKeyAuth(
api_key="hashed-key",
metadata={},
team_metadata={},
)
LiteLLMProxyRequestSetup.apply_client_tag_policy_pre_auth(
request=request_mock,
request_data=data,
user_api_key_dict=user_api_key_dict,
)
paid_tag = LiteLLM_TagTable(
tag_name="paid",
spend=0.0,
litellm_budget_table=LiteLLM_BudgetTable(max_budget=0.10),
)
async def mock_get_current_spend(counter_key, fallback_spend, max_budget=None, **kwargs):
if counter_key == "spend:tag:paid":
return 0.50
return fallback_spend
with (
patch(
"litellm.proxy.proxy_server.get_current_spend",
mock_get_current_spend,
),
patch(
"litellm.proxy.auth.auth_checks.get_tag_objects_batch",
new_callable=AsyncMock,
return_value={"paid": paid_tag},
),
):
with pytest.raises(litellm.BudgetExceededError) as exc_info:
await _tag_max_budget_check(
request_body=data,
prisma_client=MagicMock(),
user_api_key_cache=MagicMock(),
proxy_logging_obj=ProxyLogging(user_api_key_cache=None),
valid_token=UserAPIKeyAuth(token="test-token"),
)
assert exc_info.value.current_cost == 0.50
assert exc_info.value.max_budget == 0.10
@pytest.mark.asyncio
async def test_header_tags_visible_to_tag_max_budget_check(self):
"""End-to-end: helper + ``_tag_max_budget_check`` enforces budget on
header-supplied tags. Without the helper, this would silently pass."""
from litellm.proxy._types import LiteLLM_BudgetTable, LiteLLM_TagTable
from litellm.proxy.auth.auth_checks import _tag_max_budget_check
from litellm.proxy.utils import ProxyLogging
request_mock = _build_request_mock_with_headers({"x-litellm-tags": "tenant:acme"})
data = {"model": "gpt-3.5-turbo"}
user_api_key_dict = UserAPIKeyAuth(
api_key="hashed-key",
metadata={},
team_metadata={},
)
LiteLLMProxyRequestSetup.apply_client_tag_policy_pre_auth(
request=request_mock,
request_data=data,
user_api_key_dict=user_api_key_dict,
)
tag_object = LiteLLM_TagTable(
tag_name="tenant:acme",
spend=0.0,
litellm_budget_table=LiteLLM_BudgetTable(max_budget=0.10),
)
async def mock_get_current_spend(counter_key, fallback_spend, max_budget=None, **kwargs):
if counter_key == "spend:tag:tenant:acme":
return 0.50
return fallback_spend
with (
patch(
"litellm.proxy.proxy_server.get_current_spend",
mock_get_current_spend,
),
patch(
"litellm.proxy.auth.auth_checks.get_tag_objects_batch",
new_callable=AsyncMock,
return_value={"tenant:acme": tag_object},
),
):
with pytest.raises(litellm.BudgetExceededError) as exc_info:
await _tag_max_budget_check(
request_body=data,
prisma_client=MagicMock(),
user_api_key_cache=MagicMock(),
proxy_logging_obj=ProxyLogging(user_api_key_cache=None),
valid_token=UserAPIKeyAuth(token="test-token"),
)
assert exc_info.value.current_cost == 0.50
assert exc_info.value.max_budget == 0.10
@pytest.mark.asyncio
@pytest.mark.parametrize(
"route",
[
"/bedrock/model/us.anthropic.claude-sonnet-4-6/invoke",
"/v1/messages",
],
)
async def test_header_tags_visible_to_tag_max_budget_check_on_metadata_route(self, route):
"""Regression: on LITELLM_METADATA_ROUTES (bedrock, /v1/messages, ...),
common_checks pre-seeds ``litellm_metadata`` and writes key tags there
before ``_tag_max_budget_check`` reads from the same key. The auth wrapper
calls ``apply_client_tag_policy_pre_auth`` first, so without an earlier
pre-seed header tags land in ``metadata`` and the budget check (now
resolving to ``litellm_metadata``) silently ignores them. This test mirrors
the actual auth-time call order and verifies that an over-budget
header-supplied tag still trips ``_tag_max_budget_check``.
"""
from litellm.proxy._types import LiteLLM_BudgetTable, LiteLLM_TagTable
from litellm.proxy.auth.auth_checks import common_checks
from litellm.proxy.utils import ProxyLogging
request_mock = _build_request_mock_with_headers({"x-litellm-tags": "tenant:acme"})
data = {"model": "us.anthropic.claude-sonnet-4-6"}
valid_token = UserAPIKeyAuth(
token="test-token",
api_key="hashed-key",
metadata={},
team_metadata={},
)
LiteLLMProxyRequestSetup.pre_seed_litellm_metadata_for_route(
request_data=data,
route=route,
)
LiteLLMProxyRequestSetup.apply_client_tag_policy_pre_auth(
request=request_mock,
request_data=data,
user_api_key_dict=valid_token,
)
tag_object = LiteLLM_TagTable(
tag_name="tenant:acme",
spend=0.0,
litellm_budget_table=LiteLLM_BudgetTable(max_budget=0.10),
)
async def mock_get_current_spend(counter_key, fallback_spend, max_budget=None, **kwargs):
if counter_key == "spend:tag:tenant:acme":
return 0.50
return fallback_spend
with (
patch(
"litellm.proxy.proxy_server.prisma_client",
MagicMock(),
),
patch(
"litellm.proxy.proxy_server.get_current_spend",
mock_get_current_spend,
),
patch(
"litellm.proxy.auth.auth_checks.get_tag_objects_batch",
new_callable=AsyncMock,
return_value={"tenant:acme": tag_object},
),
):
with pytest.raises(litellm.BudgetExceededError) as exc_info:
await common_checks(
request_body=data,
team_object=None,
user_object=None,
end_user_object=None,
global_proxy_spend=None,
general_settings={},
route=route,
llm_router=None,
proxy_logging_obj=ProxyLogging(user_api_key_cache=None),
valid_token=valid_token,
request=request_mock,
)
assert exc_info.value.current_cost == 0.50
assert exc_info.value.max_budget == 0.10
assert "metadata" not in data
assert data["litellm_metadata"]["tags"] == ["tenant:acme"]
class TestApplyKeyTagsPreAuth:
def test_merges_key_tags_into_metadata(self):
data = {"model": "gpt-3.5-turbo"}
user_api_key_dict = UserAPIKeyAuth(
api_key="hashed-key",
metadata={"tags": ["engineering", "production"]},
team_metadata={},
)
LiteLLMProxyRequestSetup.apply_key_tags_pre_auth(
request_data=data,
user_api_key_dict=user_api_key_dict,
)
assert data["metadata"]["tags"] == ["engineering", "production"]
def test_unions_key_tags_with_existing_request_tags(self):
data = {
"model": "gpt-3.5-turbo",
"metadata": {"tags": ["request-tag"]},
}
user_api_key_dict = UserAPIKeyAuth(
api_key="hashed-key",
metadata={"tags": ["key-tag", "request-tag"]},
team_metadata={},
)
LiteLLMProxyRequestSetup.apply_key_tags_pre_auth(
request_data=data,
user_api_key_dict=user_api_key_dict,
)
# request-tag deduplicated; key-tag appended
assert data["metadata"]["tags"] == ["request-tag", "key-tag"]
def test_no_key_tags_no_mutation(self):
data = {"model": "gpt-3.5-turbo"}
user_api_key_dict = UserAPIKeyAuth(
api_key="hashed-key",
metadata={},
team_metadata={},
)
LiteLLMProxyRequestSetup.apply_key_tags_pre_auth(
request_data=data,
user_api_key_dict=user_api_key_dict,
)
assert "metadata" not in data or "tags" not in data.get("metadata", {})
def test_empty_key_metadata_no_mutation(self):
data = {"model": "gpt-3.5-turbo"}
user_api_key_dict = UserAPIKeyAuth(
api_key="hashed-key",
metadata={},
team_metadata={},
)
LiteLLMProxyRequestSetup.apply_key_tags_pre_auth(
request_data=data,
user_api_key_dict=user_api_key_dict,
)
assert "metadata" not in data
def test_uses_litellm_metadata_when_present(self):
data = {
"model": "gpt-3.5-turbo",
"litellm_metadata": {"foo": "bar"},
}
user_api_key_dict = UserAPIKeyAuth(
api_key="hashed-key",
metadata={"tags": ["key-tag"]},
team_metadata={},
)
LiteLLMProxyRequestSetup.apply_key_tags_pre_auth(
request_data=data,
user_api_key_dict=user_api_key_dict,
)
assert data["litellm_metadata"]["tags"] == ["key-tag"]
assert "tags" not in data.get("metadata", {})
def test_string_metadata_parsed_before_merge(self):
data = {
"model": "gpt-3.5-turbo",
"metadata": '{"tags": ["existing"]}',
}
user_api_key_dict = UserAPIKeyAuth(
api_key="hashed-key",
metadata={"tags": ["key-tag"]},
team_metadata={},
)
LiteLLMProxyRequestSetup.apply_key_tags_pre_auth(
request_data=data,
user_api_key_dict=user_api_key_dict,
)
assert isinstance(data["metadata"], dict)
assert data["metadata"]["tags"] == ["existing", "key-tag"]
@pytest.mark.asyncio
async def test_key_tags_visible_to_tag_max_budget_check(self):
from litellm.proxy._types import LiteLLM_BudgetTable, LiteLLM_TagTable
from litellm.proxy.auth.auth_checks import _tag_max_budget_check
from litellm.proxy.utils import ProxyLogging
data = {"model": "gpt-3.5-turbo"}
user_api_key_dict = UserAPIKeyAuth(
api_key="hashed-key",
metadata={"tags": ["engineering"]},
team_metadata={},
)
LiteLLMProxyRequestSetup.apply_key_tags_pre_auth(
request_data=data,
user_api_key_dict=user_api_key_dict,
)
tag_object = LiteLLM_TagTable(
tag_name="engineering",
spend=0.0,
litellm_budget_table=LiteLLM_BudgetTable(max_budget=0.10),
)
async def mock_get_current_spend(counter_key, fallback_spend, max_budget=None, **kwargs):
if counter_key == "spend:tag:engineering":
return 0.50
return fallback_spend
with (
patch(
"litellm.proxy.proxy_server.get_current_spend",
mock_get_current_spend,
),
patch(
"litellm.proxy.auth.auth_checks.get_tag_objects_batch",
new_callable=AsyncMock,
return_value={"engineering": tag_object},
),
):
with pytest.raises(litellm.BudgetExceededError) as exc_info:
await _tag_max_budget_check(
request_body=data,
prisma_client=MagicMock(),
user_api_key_cache=MagicMock(),
proxy_logging_obj=ProxyLogging(user_api_key_cache=None),
valid_token=UserAPIKeyAuth(token="test-token"),
)
assert exc_info.value.current_cost == 0.50
assert exc_info.value.max_budget == 0.10
@pytest.mark.asyncio
async def test_key_tags_within_budget_passes_check(self):
from litellm.proxy._types import LiteLLM_BudgetTable, LiteLLM_TagTable
from litellm.proxy.auth.auth_checks import _tag_max_budget_check
from litellm.proxy.utils import ProxyLogging
data = {"model": "gpt-3.5-turbo"}
user_api_key_dict = UserAPIKeyAuth(
api_key="hashed-key",
metadata={"tags": ["engineering"]},
team_metadata={},
)
LiteLLMProxyRequestSetup.apply_key_tags_pre_auth(
request_data=data,
user_api_key_dict=user_api_key_dict,
)
tag_object = LiteLLM_TagTable(
tag_name="engineering",
spend=0.05,
litellm_budget_table=LiteLLM_BudgetTable(max_budget=0.10),
)
async def mock_get_current_spend(counter_key, fallback_spend, max_budget=None, **kwargs):
if counter_key == "spend:tag:engineering":
return 0.05
return fallback_spend
with (
patch(
"litellm.proxy.proxy_server.get_current_spend",
mock_get_current_spend,
),
patch(
"litellm.proxy.auth.auth_checks.get_tag_objects_batch",
new_callable=AsyncMock,
return_value={"engineering": tag_object},
),
):
await _tag_max_budget_check(
request_body=data,
prisma_client=MagicMock(),
user_api_key_cache=MagicMock(),
proxy_logging_obj=ProxyLogging(user_api_key_cache=None),
valid_token=UserAPIKeyAuth(token="test-token"),
)
# ============================================================================
# Tests for #27516: provider hint resolution from deployment when the
# user-facing model name has no provider prefix.
# ============================================================================
def test_resolve_provider_from_deployment_uses_litellm_params_model():
"""When custom_llm_provider is unset, fall back to the prefix of model."""
router = MagicMock()
deployment = MagicMock()
deployment.litellm_params.model = "bedrock/us.anthropic.claude-sonnet-4-6"
deployment.litellm_params.custom_llm_provider = None
router.get_deployment_by_model_group_name.return_value = deployment
assert _resolve_provider_from_deployment(router, "claude-sonnet-4.6") == "bedrock"
def test_resolve_provider_from_deployment_prefers_custom_llm_provider():
"""Explicit custom_llm_provider on the deployment wins over model prefix."""
router = MagicMock()
deployment = MagicMock()
deployment.litellm_params.model = "us.anthropic.claude-sonnet-4-6"
deployment.litellm_params.custom_llm_provider = "bedrock"
router.get_deployment_by_model_group_name.return_value = deployment
assert _resolve_provider_from_deployment(router, "claude-sonnet-4.6") == "bedrock"
def test_resolve_provider_from_deployment_no_match():
"""No deployment for the model group -> None."""
router = MagicMock()
router.get_deployment_by_model_group_name.return_value = None
assert _resolve_provider_from_deployment(router, "unknown-model") is None
def test_resolve_provider_from_deployment_router_raises():
"""Router exceptions must not propagate — fall back to None."""
router = MagicMock()
router.get_deployment_by_model_group_name.side_effect = RuntimeError("boom")
assert _resolve_provider_from_deployment(router, "claude-sonnet-4.6") is None
def test_resolve_provider_from_deployment_falls_back_to_pre_alias():
"""If post-alias lookup fails, the pre-alias name is also tried."""
router = MagicMock()
deployment = MagicMock()
deployment.litellm_params.model = "bedrock/anthropic.claude-sonnet-4-6"
deployment.litellm_params.custom_llm_provider = None
def lookup(model_group_name):
if model_group_name == "pre-alias-name":
return deployment
return None
router.get_deployment_by_model_group_name.side_effect = lookup
result = _resolve_provider_from_deployment(router, "post-alias-name", pre_alias_model_name="pre-alias-name")
assert result == "bedrock"
def test_apply_overrides_multi_provider_default_picks_correct_provider(
setup_test_credentials,
):
"""
Regression for #27516: when defaultconfig has multiple providers and the
request model has no '/' prefix, the deployment's custom_llm_provider must
drive provider matching instead of falling through to dict insertion order.
"""
litellm.credential_list.append(
CredentialItem(
credential_name="bedrock-team-1",
credential_info={},
credential_values={"api_key": "ABSK-bedrock-key-for-team-1"},
)
)
litellm.credential_list.append(
CredentialItem(
credential_name="gemini-team-1",
credential_info={},
credential_values={"api_key": "gemini-key-for-team-1"},
)
)
data = {"model": "claude-sonnet-4.6"}
user_api_key_dict = UserAPIKeyAuth(
api_key="test-key",
team_metadata={
"model_config": {
"defaultconfig": {
# gemini comes first in insertion order — the bug picked it.
"gemini": {"litellm_credentials": "gemini-team-1"},
"bedrock": {"litellm_credentials": "bedrock-team-1"},
}
}
},
)
router = MagicMock()
deployment = MagicMock()
deployment.litellm_params.model = "us.anthropic.claude-sonnet-4-6"
deployment.litellm_params.custom_llm_provider = "bedrock"
router.get_deployment_by_model_group_name.return_value = deployment
_apply_credential_overrides_from_model_config(
data=data,
user_api_key_dict=user_api_key_dict,
llm_router=router,
)
assert data["api_key"] == "ABSK-bedrock-key-for-team-1"
def test_apply_overrides_no_router_keeps_legacy_behaviour(setup_test_credentials):
"""
Without a router, the function still works for the single-provider case
(the historical behaviour). Multi-provider configs with no '/' prefix
keep the legacy first-entry behaviour because there is no way to
disambiguate — this preserves backwards compatibility.
"""
data = {"model": "gpt-4"}
user_api_key_dict = UserAPIKeyAuth(
api_key="test-key",
team_metadata={"model_config": {"defaultconfig": {"azure": {"litellm_credentials": "hotel-azure-eastus"}}}},
)
_apply_credential_overrides_from_model_config(data=data, user_api_key_dict=user_api_key_dict, llm_router=None)
assert data["api_base"] == "https://hotel-eastus.openai.azure.com/"
assert data["api_key"] == "key-hotel-eastus"
def test_apply_overrides_provider_prefix_in_model_skips_router_lookup(
setup_test_credentials,
):
"""
When the request model already has a 'provider/...' prefix, the router
lookup must be skipped — the explicit prefix is authoritative.
"""
data = {"model": "azure/gpt-4"}
user_api_key_dict = UserAPIKeyAuth(
api_key="test-key",
team_metadata={
"model_config": {
"defaultconfig": {
"azure": {"litellm_credentials": "hotel-azure-eastus"},
"bedrock": {"litellm_credentials": "hotel-rec-azure"},
}
}
},
)
router = MagicMock()
_apply_credential_overrides_from_model_config(data=data, user_api_key_dict=user_api_key_dict, llm_router=router)
assert data["api_base"] == "https://hotel-eastus.openai.azure.com/"
assert data["api_key"] == "key-hotel-eastus"
router.get_deployment_by_model_group_name.assert_not_called()
def _make_request_mock(path: str, headers: dict) -> MagicMock:
request_mock = MagicMock(spec=Request)
request_mock.url = MagicMock()
request_mock.url.path = path
request_mock.url.__str__.return_value = f"http://localhost{path}"
request_mock.method = "POST"
request_mock.query_params = {}
request_mock.headers = headers
request_mock.client = MagicMock()
request_mock.client.host = "127.0.0.1"
return request_mock
@pytest.mark.asyncio
@pytest.mark.parametrize(
"user_agent, request_drop_params, operator_drop_params, expected_drop_params",
[
("claude-cli/2.0.69 (external, cli)", None, None, True),
("claude-cli/1.0.44 (external, sdk-py)", None, None, True),
("claude-cli/2.0.69 (external, cli)", False, None, False),
("claude-cli/2.0.69 (external, cli)", None, False, None),
("claude-cli/2.0.69 (external, cli)", None, True, None),
("codex_cli_rs/0.144.5 (Mac OS 26.4.0; arm64) WezTerm", None, None, True),
("codex_exec/0.144.5 (Mac OS 26.4.0; arm64) WarpTerminal (codex_exec; 0.144.5)", None, None, True),
("codex_vscode/0.144.5 (Mac OS 26.4.0; arm64) vscode/1.104.1", None, None, True),
("codex_exec/0.144.5 (Mac OS 26.4.0; arm64)", False, None, False),
("codex_exec/0.144.5 (Mac OS 26.4.0; arm64)", None, True, None),
("PostmanRuntime/7.53.0", None, None, None),
(None, None, None, None),
],
)
async def test_add_litellm_data_to_request_agentic_cli_drop_params(
user_agent, request_drop_params, operator_drop_params, expected_drop_params
):
"""Claude Code sends Anthropic-specific params and Codex sends
service_tier, both of which fail on providers that reject them, so those
user agents must turn on drop_params automatically, without overriding an
explicit caller value, an explicit operator-level litellm_settings value,
or affecting other clients.
"""
headers = {"Content-Type": "application/json"}
if user_agent is not None:
headers["user-agent"] = user_agent
request_mock = _make_request_mock("/v1/messages", headers)
data = {"model": "gpt-4o", "messages": [{"role": "user", "content": "hi"}]}
if request_drop_params is not None:
data["drop_params"] = request_drop_params
proxy_config = MagicMock()
proxy_config.config = (
{"litellm_settings": {"drop_params": operator_drop_params}}
if operator_drop_params is not None
else {"litellm_settings": {}}
)
updated = await add_litellm_data_to_request(
data=data,
request=request_mock,
user_api_key_dict=UserAPIKeyAuth(api_key="hashed-key"),
proxy_config=proxy_config,
general_settings={},
version="test-version",
)
assert updated.get("drop_params") == expected_drop_params
@pytest.mark.asyncio
async def test_add_litellm_data_to_request_merges_metadata_tags_on_responses_route():
"""Regression for #31584: user-supplied metadata.tags must be merged into
litellm_metadata.tags on /v1/responses so they reach SpendLogs.request_tags."""
from litellm.proxy.litellm_pre_call_utils import add_litellm_data_to_request
request_mock = MagicMock(spec=Request)
request_mock.url = MagicMock()
request_mock.url.path = "/v1/responses"
request_mock.url.__str__.return_value = "http://localhost/v1/responses"
request_mock.method = "POST"
request_mock.query_params = {}
request_mock.headers = {"Content-Type": "application/json"}
request_mock.client = MagicMock()
request_mock.client.host = "127.0.0.1"
request_mock.state = MagicMock()
data = {
"model": "gpt-4o",
"input": "hello",
"metadata": {"tags": ["cost-center-1", "team-alpha"]},
}
user_api_key_dict = UserAPIKeyAuth(
api_key="hashed-key",
metadata={},
team_metadata={},
)
updated = await add_litellm_data_to_request(
data=data,
request=request_mock,
user_api_key_dict=user_api_key_dict,
proxy_config=MagicMock(),
general_settings={},
version="test-version",
)
assert "cost-center-1" in updated["litellm_metadata"]["tags"]
assert "team-alpha" in updated["litellm_metadata"]["tags"]
@pytest.mark.asyncio
async def test_add_litellm_data_to_request_unions_metadata_tags_with_header_tags_on_responses_route():
"""On /v1/responses, tags from metadata.tags AND x-litellm-tags header
must both appear in litellm_metadata.tags."""
from litellm.proxy.litellm_pre_call_utils import add_litellm_data_to_request
request_mock = MagicMock(spec=Request)
request_mock.url = MagicMock()
request_mock.url.path = "/v1/responses"
request_mock.url.__str__.return_value = "http://localhost/v1/responses"
request_mock.method = "POST"
request_mock.query_params = {}
request_mock.headers = {
"Content-Type": "application/json",
"x-litellm-tags": "header-tag",
}
request_mock.client = MagicMock()
request_mock.client.host = "127.0.0.1"
request_mock.state = MagicMock()
data = {
"model": "gpt-4o",
"input": "hello",
"metadata": {"tags": ["body-tag"]},
}
user_api_key_dict = UserAPIKeyAuth(
api_key="hashed-key",
metadata={},
team_metadata={},
)
updated = await add_litellm_data_to_request(
data=data,
request=request_mock,
user_api_key_dict=user_api_key_dict,
proxy_config=MagicMock(),
general_settings={},
version="test-version",
)
tags = updated["litellm_metadata"]["tags"]
assert "header-tag" in tags
assert "body-tag" in tags
def _make_chat_request_mock() -> MagicMock:
return _make_request_mock("/v1/chat/completions", {"Content-Type": "application/json"})
@pytest.mark.asyncio
async def test_overwrite_user_with_key_hash_clobbers_caller_supplied_user(monkeypatch):
"""The flag exists so providers can ban by a tamper-proof id; a caller-chosen
`user` must never survive, and the raw sk- key must never be forwarded."""
from litellm.proxy._types import hash_token
monkeypatch.setattr(litellm, "overwrite_user_with_key_hash", True)
raw_key = "sk-overwrite-user-test-1234"
user_api_key_dict = UserAPIKeyAuth(api_key=raw_key)
user_api_key_dict.via_virtual_key = True
data = {"model": "gpt-4o", "user": "attacker-chosen-id"}
updated_data = await add_litellm_data_to_request(
data=data,
request=_make_chat_request_mock(),
user_api_key_dict=user_api_key_dict,
proxy_config=MagicMock(),
general_settings={},
version="test-version",
)
assert updated_data["user"] == hash_token(raw_key)
assert updated_data["user"] != "attacker-chosen-id"
assert raw_key not in updated_data["user"]
@pytest.mark.asyncio
async def test_overwrite_user_with_key_hash_sets_user_when_absent(monkeypatch):
from litellm.proxy._types import hash_token
monkeypatch.setattr(litellm, "overwrite_user_with_key_hash", True)
raw_key = "sk-overwrite-user-test-5678"
user_api_key_dict = UserAPIKeyAuth(api_key=raw_key)
user_api_key_dict.via_virtual_key = True
data = {"model": "gpt-4o"}
updated_data = await add_litellm_data_to_request(
data=data,
request=_make_chat_request_mock(),
user_api_key_dict=user_api_key_dict,
proxy_config=MagicMock(),
general_settings={},
version="test-version",
)
assert updated_data["user"] == hash_token(raw_key)
@pytest.mark.asyncio
async def test_overwrite_user_with_key_hash_disabled_preserves_caller_user():
assert litellm.overwrite_user_with_key_hash is False
user_api_key_dict = UserAPIKeyAuth(api_key="sk-overwrite-user-test-9999")
user_api_key_dict.via_virtual_key = True
data = {"model": "gpt-4o", "user": "caller-chosen-id"}
updated_data = await add_litellm_data_to_request(
data=data,
request=_make_chat_request_mock(),
user_api_key_dict=user_api_key_dict,
proxy_config=MagicMock(),
general_settings={},
version="test-version",
)
assert updated_data["user"] == "caller-chosen-id"
@pytest.mark.asyncio
async def test_overwrite_user_with_key_hash_skips_custom_auth_credential(monkeypatch):
"""Custom-auth credentials are not sk-prefixed or JWTs, so UserAPIKeyAuth stores
them raw; the stamp must skip them entirely so auth material never leaks."""
monkeypatch.setattr(litellm, "overwrite_user_with_key_hash", True)
raw_credential = "my-custom-auth-credential-abc123"
user_api_key_dict = UserAPIKeyAuth(api_key=raw_credential)
assert user_api_key_dict.api_key == raw_credential
updated_data = await add_litellm_data_to_request(
data={"model": "gpt-4o", "user": "caller-chosen-id"},
request=_make_chat_request_mock(),
user_api_key_dict=user_api_key_dict,
proxy_config=MagicMock(),
general_settings={},
version="test-version",
)
assert updated_data["user"] == "caller-chosen-id"
@pytest.mark.asyncio
async def test_overwrite_user_with_key_hash_skips_jwt_auth(monkeypatch):
"""A hashed JWT rotates on every token re-issue, so it is useless as a stable
ban id; JWT-authenticated requests are not stamped."""
from litellm.proxy._types import hash_token
monkeypatch.setattr(litellm, "overwrite_user_with_key_hash", True)
hashed_jwt = f"hashed-jwt-{hash_token('some-jwt-token')}"
user_api_key_dict = UserAPIKeyAuth(api_key=hashed_jwt)
updated_data = await add_litellm_data_to_request(
data={"model": "gpt-4o", "user": "caller-chosen-id"},
request=_make_chat_request_mock(),
user_api_key_dict=user_api_key_dict,
proxy_config=MagicMock(),
general_settings={},
version="test-version",
)
assert updated_data["user"] == "caller-chosen-id"
@pytest.mark.asyncio
async def test_overwrite_user_with_key_hash_skips_hex_shaped_custom_credential(monkeypatch):
"""A custom-auth credential that happens to be 64 hex chars is indistinguishable
from a key hash by shape alone; only the server-set via_virtual_key marker may
authorize stamping, so this raw credential must never be forwarded."""
monkeypatch.setattr(litellm, "overwrite_user_with_key_hash", True)
hex_shaped_credential = "a" * 64
user_api_key_dict = UserAPIKeyAuth(api_key=hex_shaped_credential)
assert user_api_key_dict.api_key == hex_shaped_credential
assert user_api_key_dict.via_virtual_key is False
updated_data = await add_litellm_data_to_request(
data={"model": "gpt-4o", "user": "caller-chosen-id"},
request=_make_chat_request_mock(),
user_api_key_dict=user_api_key_dict,
proxy_config=MagicMock(),
general_settings={},
version="test-version",
)
assert updated_data["user"] == "caller-chosen-id"
def test_via_virtual_key_cannot_be_forged_from_validated_input():
from_kwargs = UserAPIKeyAuth(api_key="b" * 64, via_virtual_key=True)
assert from_kwargs.via_virtual_key is False
from_dict = UserAPIKeyAuth.model_validate({"api_key": "b" * 64, "via_virtual_key": True})
assert from_dict.via_virtual_key is False
@pytest.mark.asyncio
async def test_overwrite_user_with_key_hash_stamps_master_key_alias(monkeypatch):
"""Master-key requests carry the stable alias instead of a hash (so the master
key never propagates anywhere); the alias is the stampable id for them."""
from litellm.constants import LITELLM_PROXY_MASTER_KEY_ALIAS
monkeypatch.setattr(litellm, "overwrite_user_with_key_hash", True)
user_api_key_dict = UserAPIKeyAuth(api_key=LITELLM_PROXY_MASTER_KEY_ALIAS)
user_api_key_dict.via_virtual_key = True
updated_data = await add_litellm_data_to_request(
data={"model": "gpt-4o", "user": "attacker-chosen-id"},
request=_make_chat_request_mock(),
user_api_key_dict=user_api_key_dict,
proxy_config=MagicMock(),
general_settings={},
version="test-version",
)
assert updated_data["user"] == LITELLM_PROXY_MASTER_KEY_ALIAS
@pytest.mark.asyncio
async def test_overwrite_user_with_key_hash_rejects_alias_without_marker(monkeypatch):
from litellm.constants import LITELLM_PROXY_MASTER_KEY_ALIAS
monkeypatch.setattr(litellm, "overwrite_user_with_key_hash", True)
user_api_key_dict = UserAPIKeyAuth(api_key=LITELLM_PROXY_MASTER_KEY_ALIAS)
assert user_api_key_dict.via_virtual_key is False
updated_data = await add_litellm_data_to_request(
data={"model": "gpt-4o", "user": "caller-chosen-id"},
request=_make_chat_request_mock(),
user_api_key_dict=user_api_key_dict,
proxy_config=MagicMock(),
general_settings={},
version="test-version",
)
assert updated_data["user"] == "caller-chosen-id"
def test_get_sanitized_user_information_from_key_drops_callback_config():
"""
Regression (LIT-4306): `user_api_key_auth_metadata` lands in the
StandardLoggingPayload every integration receives, so the per-key callback
config (and the integration credentials inside it) must not ride along.
Everything else - notably `priority`, which the dynamic rate limiter reads
back off this exact field - has to survive.
"""
user_api_key_dict = UserAPIKeyAuth(
api_key="test-key-hash",
metadata={
"logging": [
{
"callback_name": "langsmith",
"callback_vars": {"langsmith_api_key": "litellm_enc::ciphertext"},
}
],
"callback_settings": {"callback_vars": {"langfuse_secret_key": "litellm_enc::other"}},
"priority": "high",
},
)
result = LiteLLMProxyRequestSetup.get_sanitized_user_information_from_key(user_api_key_dict=user_api_key_dict)
auth_metadata = result["user_api_key_auth_metadata"]
assert "logging" not in auth_metadata
assert "callback_settings" not in auth_metadata
assert "litellm_enc::" not in json.dumps(auth_metadata)
assert auth_metadata["priority"] == "high"
# UserAPIKeyAuth is the live auth object; the per-key callbacks are resolved
# from it during pre-call, so it must not be mutated by building the log view
assert "logging" in (user_api_key_dict.metadata or {})
def test_team_alias_targeting_deleted_team_deployment_keeps_requested_model(monkeypatch):
"""
Regression: a team's model_aliases can point at the internal routing key
(model_name_{team_id}_{uuid}) of a team deployment that was since deleted,
e.g. after an admin replaces per-team duplicates with one gateway-level
model. Rewriting to the dead internal name made every request fail with
"no healthy deployments for model_name_..." even though the requested
public name resolves at the gateway level. The rewrite must be skipped
when the alias target has no live deployment.
"""
import litellm.proxy.litellm_pre_call_utils as pre_call_utils
from litellm.proxy.litellm_pre_call_utils import _update_model_if_team_alias_exists
monkeypatch.delenv("LITELLM_ENABLE_TEAM_STALE_ALIAS_BYPASS", raising=False)
pre_call_utils._ENABLE_TEAM_STALE_ALIAS_BYPASS = None
class _MockRouter:
model_name_to_deployment_indices = {"gpt-4": [0]}
team_model_to_deployment_indices = {}
test_data = {"model": "gpt-4"}
user_api_key_dict = UserAPIKeyAuth(
api_key="test_key",
team_id="team-1",
team_model_aliases={"gpt-4": "model_name_team-1_dead-uuid"},
)
with patch("litellm.proxy.proxy_server.llm_router", _MockRouter()):
_update_model_if_team_alias_exists(data=test_data, user_api_key_dict=user_api_key_dict)
assert test_data.get("model") == "gpt-4"
def test_team_alias_targeting_live_team_deployment_still_rewrites(monkeypatch):
import litellm.proxy.litellm_pre_call_utils as pre_call_utils
from litellm.proxy.litellm_pre_call_utils import _update_model_if_team_alias_exists
monkeypatch.delenv("LITELLM_ENABLE_TEAM_STALE_ALIAS_BYPASS", raising=False)
pre_call_utils._ENABLE_TEAM_STALE_ALIAS_BYPASS = None
class _MockRouter:
model_name_to_deployment_indices = {"model_name_team-1_live-uuid": [0]}
team_model_to_deployment_indices = {}
test_data = {"model": "gpt-4"}
user_api_key_dict = UserAPIKeyAuth(
api_key="test_key",
team_id="team-1",
team_model_aliases={"gpt-4": "model_name_team-1_live-uuid"},
)
with patch("litellm.proxy.proxy_server.llm_router", _MockRouter()):
_update_model_if_team_alias_exists(data=test_data, user_api_key_dict=user_api_key_dict)
assert test_data.get("model") == "model_name_team-1_live-uuid"
def test_warn_stale_team_alias_once_logs_once_per_key(monkeypatch):
from collections import OrderedDict
import litellm.proxy.litellm_pre_call_utils as pre_call_utils
monkeypatch.setattr(pre_call_utils, "_STALE_TEAM_ALIAS_WARNING_KEYS", OrderedDict())
with patch.object(pre_call_utils.verbose_proxy_logger, "warning") as mock_warning:
pre_call_utils._warn_stale_team_alias_once("team-1:gpt-4", "stale alias %s", "gpt-4")
pre_call_utils._warn_stale_team_alias_once("team-1:gpt-4", "stale alias %s", "gpt-4")
assert mock_warning.call_count == 1
def test_warn_stale_team_alias_once_evicts_oldest_key_beyond_cap(monkeypatch):
from collections import OrderedDict
import litellm.proxy.litellm_pre_call_utils as pre_call_utils
monkeypatch.setattr(pre_call_utils, "_STALE_TEAM_ALIAS_WARNING_KEYS", OrderedDict())
monkeypatch.setattr(pre_call_utils, "_MAX_STALE_ALIAS_WARNING_KEYS", 2)
with patch.object(pre_call_utils.verbose_proxy_logger, "warning"):
pre_call_utils._warn_stale_team_alias_once("key-1", "stale alias")
pre_call_utils._warn_stale_team_alias_once("key-2", "stale alias")
pre_call_utils._warn_stale_team_alias_once("key-3", "stale alias")
assert list(pre_call_utils._STALE_TEAM_ALIAS_WARNING_KEYS) == ["key-2", "key-3"]
_OAUTH_TOKEN = "Bearer sk-ant-oat01-regression-token-lit5108"
def _all_header_dicts(data: dict, metadata_variable_name: str) -> list[dict]:
metadata = data.get(metadata_variable_name) or {}
proxy_server_request = data["proxy_server_request"]
body = proxy_server_request["body"]
return [
metadata.get("headers") or {},
(metadata.get("requester_metadata") or {}).get("headers") or {},
proxy_server_request["headers"],
(body.get(metadata_variable_name) or {}).get("headers") or {},
]
@pytest.mark.asyncio
@pytest.mark.parametrize(
"path, metadata_variable_name",
[
("/v1/messages", "litellm_metadata"),
("/v1/chat/completions", "metadata"),
],
)
async def test_add_litellm_data_to_request_redacts_oauth_header_from_logging_copies(path, metadata_variable_name):
"""The Anthropic subscription token is forwarded upstream but never handed to logging."""
request_mock = _make_request_mock(
path,
{
"Content-Type": "application/json",
"anthropic-version": "2023-06-01",
"Authorization": _OAUTH_TOKEN,
"x-litellm-api-key": "Bearer sk-virtual-key",
},
)
updated = await add_litellm_data_to_request(
data={"model": "anthropic-claude", "messages": [{"role": "user", "content": "hello"}]},
request=request_mock,
user_api_key_dict=UserAPIKeyAuth(api_key="hashed-key"),
proxy_config=MagicMock(),
general_settings={"forward_client_headers_to_llm_api": True},
version="test-version",
)
for header_dict in _all_header_dicts(updated, metadata_variable_name):
assert header_dict.get("Authorization") != _OAUTH_TOKEN
assert "sk-ant-oat01" not in json.dumps(header_dict)
assert "sk-ant-oat01" not in json.dumps(updated["proxy_server_request"], default=repr)
assert updated["proxy_server_request"]["headers"] is updated[metadata_variable_name]["headers"]
from litellm.litellm_core_utils.get_provider_specific_headers import (
ProviderSpecificHeaderUtils,
)
assert (
ProviderSpecificHeaderUtils.get_provider_specific_headers(
provider_specific_header=updated["provider_specific_header"],
custom_llm_provider="anthropic",
)["Authorization"]
== _OAUTH_TOKEN
)
@pytest.mark.asyncio
@pytest.mark.parametrize(
"path, metadata_variable_name",
[
("/v1/messages", "litellm_metadata"),
("/v1/chat/completions", "metadata"),
],
)
async def test_add_litellm_data_to_request_stamps_used_client_oauth_token(path, metadata_variable_name):
"""A seat-billed request and a configured-key request must land in spend logs differing on exactly
the credential flag, and the flag must never carry the token itself."""
async def metadata_for(client_headers: dict) -> dict:
request_mock = _make_request_mock(path, {"Content-Type": "application/json", **client_headers})
updated = await add_litellm_data_to_request(
data={"model": "anthropic-claude", "messages": [{"role": "user", "content": "hello"}]},
request=request_mock,
user_api_key_dict=UserAPIKeyAuth(api_key="hashed-key"),
proxy_config=MagicMock(),
general_settings={"forward_client_headers_to_llm_api": True},
version="test-version",
)
return updated[metadata_variable_name]
def spend_log_row_metadata(request_metadata: dict) -> dict:
row = get_logging_payload(
kwargs={
"model": "claude-sonnet-5",
"custom_llm_provider": "anthropic",
"litellm_params": {"metadata": request_metadata},
},
response_obj={},
start_time=datetime.now(timezone.utc),
end_time=datetime.now(timezone.utc),
)
return json.loads(row["metadata"])
seat_row = spend_log_row_metadata(
await metadata_for({"Authorization": _OAUTH_TOKEN, "x-litellm-api-key": "Bearer sk-virtual-key"})
)
key_row = spend_log_row_metadata(await metadata_for({"Authorization": "Bearer sk-virtual-key"}))
assert seat_row["used_client_oauth_token"] is True
assert key_row["used_client_oauth_token"] is False
differing_keys = {key for key in seat_row.keys() | key_row.keys() if seat_row.get(key) != key_row.get(key)}
assert differing_keys == {"used_client_oauth_token"}
assert "sk-ant-oat01" not in json.dumps(seat_row, default=repr)
@pytest.mark.asyncio
async def test_add_litellm_data_to_request_keeps_every_forwarded_credential_out_of_logging_copies():
"""Credentials kept for transport must not survive anywhere under proxy_server_request."""
secrets = {
"x-api-key": "sk-byok-provider-key-lit5108",
"cookie": "litellm_jwt=session-token-lit5108",
"proxy-authorization": "Bearer proxy-token-lit5108",
}
request_mock = _make_request_mock(
"/v1/chat/completions",
{
"Content-Type": "application/json",
"x-litellm-api-key": "Bearer sk-virtual-key",
**secrets,
},
)
updated = await add_litellm_data_to_request(
data={"model": "gpt-4o", "messages": [{"role": "user", "content": "hello"}]},
request=request_mock,
user_api_key_dict=UserAPIKeyAuth(api_key="hashed-key"),
proxy_config=MagicMock(),
general_settings={
"forward_llm_provider_auth_headers": True,
"forward_client_headers_to_llm_api": True,
},
version="test-version",
)
assert updated["api_key"] == secrets["x-api-key"]
assert updated["headers"]["x-api-key"] == secrets["x-api-key"]
logged = json.dumps(updated["proxy_server_request"], default=repr)
for value in secrets.values():
assert value not in logged
@pytest.mark.parametrize(
"header, expected_redacted",
[
("Authorization", True),
("X-Api-Key", True),
("x-goog-api-key", True),
("Ocp-Apim-Subscription-Key", True),
("API-Key", True),
("Cookie", True),
("Proxy-Authorization", True),
("anthropic-version", False),
("user-agent", False),
],
)
def test_redact_credential_headers_classifies_each_header(header, expected_redacted):
from litellm.proxy.litellm_pre_call_utils import redact_credential_headers
headers = {header: "secret-value"}
redacted = redact_credential_headers(headers)
assert redacted[header] == ("***REDACTED***" if expected_redacted else "secret-value")
assert headers[header] == "secret-value"
@pytest.mark.asyncio
async def test_add_litellm_data_to_request_debug_log_does_not_print_credentials():
"""The request-header debug line carries values the stdout secret filter does not match."""
import litellm.proxy.litellm_pre_call_utils as pre_call_utils
request_mock = _make_request_mock(
"/v1/chat/completions",
{
"Content-Type": "application/json",
"Ocp-Apim-Subscription-Key": "apim-plaintext-token-lit5108",
"x-litellm-api-key": "Bearer sk-virtual-key",
},
)
with patch.object(pre_call_utils.verbose_proxy_logger, "debug") as mock_debug:
await add_litellm_data_to_request(
data={"model": "gpt-4o", "messages": [{"role": "user", "content": "hello"}]},
request=request_mock,
user_api_key_dict=UserAPIKeyAuth(api_key="hashed-key"),
proxy_config=MagicMock(),
general_settings={"forward_llm_provider_auth_headers": True},
version="test-version",
)
logged = " ".join(str(call) for call in mock_debug.call_args_list)
assert "apim-plaintext-token-lit5108" not in logged
def _callback_credential_request_mock() -> MagicMock:
request_mock = MagicMock(spec=Request)
request_mock.url = MagicMock()
request_mock.url.path = "/v1/chat/completions"
request_mock.url.__str__.return_value = "http://localhost/v1/chat/completions"
request_mock.method = "POST"
request_mock.query_params = {}
request_mock.headers = {"Content-Type": "application/json"}
request_mock.client = MagicMock()
request_mock.client.host = "127.0.0.1"
return request_mock
_DATADOG_TEAM_KEY = UserAPIKeyAuth(
api_key="hashed-key",
team_id="team-1",
team_metadata={
"logging": [
{
"callback_name": "datadog",
"callback_type": "success",
"callback_vars": {"dd_api_key": "team-dd-key"},
}
]
},
)
@pytest.mark.asyncio
async def test_add_litellm_data_to_request_strips_caller_supplied_callback_credentials():
"""
The team admin sets dd_api_key only; a caller pairing its own dd_site with that key
would ship the team's Datadog credential to a host it controls.
"""
caller_destinations = {"dd_site": "attacker.example.com", "dd_agent_host": "attacker.example.com"}
data = {
"model": "gpt-3.5-turbo",
"messages": [{"role": "user", "content": "hello"}],
**caller_destinations,
"gcs_bucket_name": "attacker-bucket",
TRUSTED_CALLBACK_VARS_FIELD: {"dd_site": "smuggled.example.com"},
"metadata": {**caller_destinations, "safe_user_metadata": "kept"},
"litellm_metadata": dict(caller_destinations),
"litellm_params": {"metadata": dict(caller_destinations)},
}
updated = await add_litellm_data_to_request(
data=data,
request=_callback_credential_request_mock(),
user_api_key_dict=_DATADOG_TEAM_KEY,
proxy_config=MagicMock(),
general_settings={},
version="test-version",
)
assert "dd_site" not in updated
assert "dd_agent_host" not in updated
assert "gcs_bucket_name" not in updated
assert updated["dd_api_key"] == "team-dd-key"
assert updated[TRUSTED_CALLBACK_VARS_FIELD] == {"dd_api_key": "team-dd-key"}
assert "litellm_metadata" not in updated
assert "dd_site" not in updated["metadata"]
assert "dd_agent_host" not in updated["metadata"]
assert "dd_site" not in updated["litellm_params"]["metadata"]
assert updated["metadata"]["safe_user_metadata"] == "kept"
@pytest.mark.asyncio
async def test_add_litellm_data_to_request_strips_caller_supplied_callback_credentials_with_clientside_creds_allowed():
"""`allow_client_side_credentials` opens the auth-layer ban; the strip must still hold."""
data = {
"model": "gpt-3.5-turbo",
"messages": [{"role": "user", "content": "hello"}],
"dd_site": "attacker.example.com",
}
updated = await add_litellm_data_to_request(
data=data,
request=_callback_credential_request_mock(),
user_api_key_dict=_DATADOG_TEAM_KEY,
proxy_config=MagicMock(),
general_settings={"allow_client_side_credentials": True},
version="test-version",
)
assert "dd_site" not in updated
assert updated[TRUSTED_CALLBACK_VARS_FIELD] == {"dd_api_key": "team-dd-key"}
@pytest.mark.asyncio
async def test_add_litellm_data_to_request_omits_trusted_callback_vars_without_team_callbacks():
"""Without team/key callback settings the trusted field must not exist for a callback to read."""
data = {
"model": "gpt-3.5-turbo",
"messages": [{"role": "user", "content": "hello"}],
TRUSTED_CALLBACK_VARS_FIELD: {"dd_api_key": "caller-key", "dd_site": "attacker.example.com"},
}
updated = await add_litellm_data_to_request(
data=data,
request=_callback_credential_request_mock(),
user_api_key_dict=UserAPIKeyAuth(api_key="hashed-key"),
proxy_config=MagicMock(),
general_settings={},
version="test-version",
)
assert TRUSTED_CALLBACK_VARS_FIELD not in updated
def test_trusted_callback_vars_never_reach_the_provider():
"""
The stamped field rides the request body, so it has to be a recognised litellm param;
otherwise the OpenAI param builder sweeps it into extra_body and the provider 400s.
"""
from litellm.utils import get_non_default_completion_params
non_default = get_non_default_completion_params(
{
"model": "gpt-4",
TRUSTED_CALLBACK_VARS_FIELD: {"dd_api_key": "team-dd-key"},
"some_provider_param": "kept",
}
)
assert TRUSTED_CALLBACK_VARS_FIELD not in non_default
assert non_default["some_provider_param"] == "kept"
@pytest.mark.asyncio
async def test_key_level_callback_vars_survive_the_strip():
"""
Key-level callbacks configure their own destination and credentials, and they replace
team settings rather than merging with them, so only the request body is untrusted.
"""
key_with_datadog_callback = UserAPIKeyAuth(
api_key="hashed-key",
metadata={
"logging": [
{
"callback_name": "datadog",
"callback_type": "success",
"callback_vars": {"dd_api_key": "key-dd-key", "dd_site": "us5.datadoghq.com"},
}
]
},
)
data = {
"model": "gpt-3.5-turbo",
"messages": [{"role": "user", "content": "hello"}],
"dd_site": "attacker.example.com",
}
updated = await add_litellm_data_to_request(
data=data,
request=_callback_credential_request_mock(),
user_api_key_dict=key_with_datadog_callback,
proxy_config=MagicMock(),
general_settings={},
version="test-version",
)
assert updated[TRUSTED_CALLBACK_VARS_FIELD] == {"dd_api_key": "key-dd-key", "dd_site": "us5.datadoghq.com"}
assert updated["dd_site"] == "us5.datadoghq.com"
class TestPromotedTraceControlFields:
"""LIT-5137: caller metadata trace fields must reach litellm_metadata."""
def _make_request(self, path: str) -> MagicMock:
request = MagicMock(spec=Request)
request.url = MagicMock()
request.url.path = path
request.url.__str__.return_value = f"http://localhost{path}"
request.method = "POST"
request.query_params = {}
request.headers = {"Content-Type": "application/json"}
request.client = MagicMock()
request.client.host = "127.0.0.1"
return request
async def _run(self, path: str, data: dict, headers: dict | None = None) -> dict:
request = self._make_request(path)
if headers is not None:
request.headers = {"Content-Type": "application/json", **headers}
return await add_litellm_data_to_request(
data=data,
request=request,
user_api_key_dict=UserAPIKeyAuth(api_key="hashed-key"),
proxy_config=MagicMock(),
general_settings={},
version="test-version",
)
def test_returns_litellm_metadata_for_responses_route(self):
assert _get_metadata_variable_name(self._make_request("/v1/responses")) == "litellm_metadata"
def test_promotes_trace_prefixed_and_allow_listed_fields(self):
requester_metadata = {
"trace_id": "trace-1",
"trace_name": "name-1",
"trace_user_id": "user-1",
"trace_metadata": {"tenant_id": "tenant-1"},
"trace_version": "v1",
"trace_release": "r1",
"session_id": "session-1",
"mask_input": True,
"mask_output": True,
}
promoted = _promoted_trace_control_fields(
requester_metadata=requester_metadata,
litellm_metadata={},
)
assert dict(promoted) == requester_metadata
def test_does_not_promote_unlisted_trace_prefixed_fields(self):
"""trace_public flips a trace to publicly readable, so the allow-list is explicit."""
promoted = _promoted_trace_control_fields(
requester_metadata={"trace_id": "trace-1", "trace_public": True, "trace_tags": ["a"]},
litellm_metadata={},
)
assert dict(promoted) == {"trace_id": "trace-1"}
def test_does_not_promote_non_trace_fields(self):
promoted = _promoted_trace_control_fields(
requester_metadata={
"trace_id": "trace-1",
"tags": ["free-tier"],
"user_api_key": "forged",
"user_api_key_user_id": "forged-user",
"spend_logs_metadata": {"forged": True},
"guardrails": ["disabled"],
"debug_langfuse": True,
"session": "not-session-id",
"existing_trace_id": "victim-trace",
"update_trace_keys": ["input", "output"],
},
litellm_metadata={},
)
assert dict(promoted) == {"trace_id": "trace-1"}
def test_does_not_promote_trace_mutation_controls(self):
"""existing_trace_id + update_trace_keys let a caller overwrite any trace in the project."""
promoted = _promoted_trace_control_fields(
requester_metadata={
"trace_id": "trace-1",
"existing_trace_id": "someone-elses-trace",
"update_trace_keys": ["input", "output"],
},
litellm_metadata={},
)
assert dict(promoted) == {"trace_id": "trace-1"}
def test_existing_litellm_metadata_value_wins(self):
promoted = _promoted_trace_control_fields(
requester_metadata={"trace_id": "from-body", "session_id": "from-body", "trace_name": "from-body"},
litellm_metadata={"trace_id": "from-header", "session_id": "from-header"},
)
assert dict(promoted) == {"trace_name": "from-body"}
def test_empty_requester_metadata_promotes_nothing(self):
assert _promoted_trace_control_fields(requester_metadata={}, litellm_metadata={}) == ()
@pytest.mark.asyncio
async def test_responses_route_end_to_end(self):
caller_metadata = {
"trace_id": "22662678-30c1-41a1-a24b-216d6e5fb83d",
"session_id": "218af06c-28a2-4705-8a0a-5f9970d39326",
"trace_user_id": "user-123",
"trace_metadata": {"tenant_id": "tenant-1"},
"mask_input": True,
}
updated = await self._run(
"/v1/responses",
{"model": "gpt-4.1-mini", "input": "say resp", "metadata": copy.deepcopy(caller_metadata)},
)
litellm_metadata = updated["litellm_metadata"]
for key, value in caller_metadata.items():
assert litellm_metadata[key] == value
assert updated["metadata"] == caller_metadata
@pytest.mark.asyncio
async def test_messages_route_end_to_end(self):
updated = await self._run(
"/v1/messages",
{
"model": "claude-sonnet-4-5",
"max_tokens": 32,
"messages": [{"role": "user", "content": "hi"}],
"metadata": {"trace_id": "msg-trace-1", "session_id": "msg-session-1"},
},
)
assert updated["litellm_metadata"]["trace_id"] == "msg-trace-1"
assert updated["litellm_metadata"]["session_id"] == "msg-session-1"
@pytest.mark.asyncio
async def test_session_id_header_beats_body_metadata(self):
updated = await self._run(
"/v1/responses",
{"model": "gpt-4.1-mini", "input": "say resp", "metadata": {"session_id": "from-body"}},
headers={"x-litellm-session-id": "from-header-12345678"},
)
assert updated["litellm_metadata"]["session_id"] == "from-header-12345678"
@pytest.mark.asyncio
async def test_forged_user_api_key_fields_are_not_promoted(self):
updated = await self._run(
"/v1/responses",
{
"model": "gpt-4.1-mini",
"input": "say resp",
"metadata": {"trace_id": "trace-1", "user_api_key_user_id": "forged", "spend_logs_metadata": {"a": 1}},
},
)
litellm_metadata = updated["litellm_metadata"]
assert litellm_metadata["trace_id"] == "trace-1"
assert litellm_metadata.get("user_api_key_user_id") != "forged"
@pytest.mark.asyncio
async def test_chat_completions_route_is_untouched(self):
updated = await self._run(
"/v1/chat/completions",
{
"model": "gpt-4.1-mini",
"messages": [{"role": "user", "content": "hello"}],
"metadata": {"trace_id": "trace-1", "session_id": "session-1"},
},
)
assert "litellm_metadata" not in updated
assert updated["metadata"]["trace_id"] == "trace-1"
assert updated["metadata"]["session_id"] == "session-1"
@pytest.mark.asyncio
async def test_add_litellm_data_to_request_inherited_tags_excludes_caller_tags():
"""inherited_tags must carry only what key/team/project policy contributed,
never anything the caller's own request (header/body) supplied, even when the
caller resubmits the identical value -- it's a snapshot taken before the
caller's own tags are merged in, not a set difference against caller_tags.
tag_based_routing.py's allow_fail_open relies on this so a caller can't strip
an inherited constraint's protection by resubmitting its exact value."""
request_mock = MagicMock(spec=Request)
request_mock.url.path = "/v1/chat/completions"
request_mock.url = MagicMock()
request_mock.url.__str__.return_value = "http://localhost/v1/chat/completions"
request_mock.method = "POST"
request_mock.query_params = {}
request_mock.headers = {"Content-Type": "application/json"}
request_mock.client = MagicMock()
request_mock.client.host = "127.0.0.1"
data = {
"model": "gpt-3.5-turbo",
# Caller resubmits the exact value the key policy also contributes.
"tags": ["key-supplied"],
}
user_api_key_dict = UserAPIKeyAuth(
api_key="hashed-key",
user_id="real-user",
metadata={"tags": ["key-supplied"]},
team_metadata={"tags": ["team-supplied"]},
spend=0.0,
max_budget=100.0,
model_max_budget={},
team_spend=0.0,
team_max_budget=200.0,
)
updated = await add_litellm_data_to_request(
data=data,
request=request_mock,
user_api_key_dict=user_api_key_dict,
proxy_config=MagicMock(),
general_settings={},
version="test-version",
)
assert set(updated["metadata"]["tags"]) == {"key-supplied", "team-supplied"}
assert set(updated["metadata"]["inherited_tags"]) == {"key-supplied", "team-supplied"}
assert tuple(updated["metadata"]["caller_tags"]) == ("key-supplied",)
@pytest.mark.asyncio
async def test_add_litellm_data_to_request_inherited_tags_survives_pre_auth_header_merge():
"""Regression: apply_client_tag_policy_pre_auth (run from user_api_key_auth,
for _tag_max_budget_check) merges the caller's x-litellm-tags header into the
same metadata.tags list this function later reads from -- before this
function ever runs. A snapshot-based inherited_tags would misattribute that
caller-controlled value as policy-backed; inherited_tags must instead be read
directly from key/team/project metadata, immune to whatever the pre-auth pass
already merged into "tags"."""
request_mock = MagicMock(spec=Request)
request_mock.url.path = "/v1/chat/completions"
request_mock.url = MagicMock()
request_mock.url.__str__.return_value = "http://localhost/v1/chat/completions"
request_mock.method = "POST"
request_mock.query_params = {}
request_mock.headers = {"Content-Type": "application/json", "x-litellm-tags": "caller-invented-tag"}
request_mock.client = MagicMock()
request_mock.client.host = "127.0.0.1"
data: dict = {"model": "gpt-3.5-turbo"}
user_api_key_dict = UserAPIKeyAuth(
api_key="hashed-key",
user_id="real-user",
metadata={"tags": ["key-supplied"]},
team_metadata={},
spend=0.0,
max_budget=100.0,
model_max_budget={},
team_spend=0.0,
team_max_budget=200.0,
)
# Simulate the real request pipeline: the pre-auth merge runs first, on the
# same data dict, before add_litellm_data_to_request is ever called.
LiteLLMProxyRequestSetup.apply_client_tag_policy_pre_auth(
request=request_mock,
request_data=data,
user_api_key_dict=user_api_key_dict,
)
assert data["metadata"]["tags"] == ["caller-invented-tag"]
updated = await add_litellm_data_to_request(
data=data,
request=request_mock,
user_api_key_dict=user_api_key_dict,
proxy_config=MagicMock(),
general_settings={},
version="test-version",
)
assert set(updated["metadata"]["tags"]) == {"caller-invented-tag", "key-supplied"}
assert updated["metadata"]["inherited_tags"] == ("key-supplied",)
assert updated["metadata"]["caller_tags"] == ("caller-invented-tag",)
@pytest.mark.asyncio
async def test_add_litellm_data_to_request_caller_tags_excludes_key_and_team_tags():
"""caller_tags must carry only what the caller itself sent (header + body
tags), never anything merged in from key/team metadata, even though the
merged "tags" field (used for matching) legitimately contains all three."""
request_mock = MagicMock(spec=Request)
request_mock.url.path = "/v1/chat/completions"
request_mock.url = MagicMock()
request_mock.url.__str__.return_value = "http://localhost/v1/chat/completions"
request_mock.method = "POST"
request_mock.query_params = {}
request_mock.headers = {"Content-Type": "application/json"}
request_mock.client = MagicMock()
request_mock.client.host = "127.0.0.1"
data = {
"model": "gpt-3.5-turbo",
"tags": ["caller-supplied"],
}
user_api_key_dict = UserAPIKeyAuth(
api_key="hashed-key",
user_id="real-user",
metadata={"tags": ["key-supplied"]},
team_metadata={"tags": ["team-supplied"]},
spend=0.0,
max_budget=100.0,
model_max_budget={},
team_spend=0.0,
team_max_budget=200.0,
)
updated = await add_litellm_data_to_request(
data=data,
request=request_mock,
user_api_key_dict=user_api_key_dict,
proxy_config=MagicMock(),
general_settings={},
version="test-version",
)
assert set(updated["metadata"]["tags"]) == {"caller-supplied", "key-supplied", "team-supplied"}
assert tuple(updated["metadata"]["caller_tags"]) == ("caller-supplied",)
@pytest.mark.asyncio
async def test_add_litellm_data_to_request_caller_tags_includes_header_tags():
"""The x-litellm-tags header is as much a caller-controlled input as the
body's "tags" field; both must land in caller_tags."""
request_mock = MagicMock(spec=Request)
request_mock.url.path = "/v1/chat/completions"
request_mock.url = MagicMock()
request_mock.url.__str__.return_value = "http://localhost/v1/chat/completions"
request_mock.method = "POST"
request_mock.query_params = {}
request_mock.headers = {"Content-Type": "application/json", "x-litellm-tags": "header-tag"}
request_mock.client = MagicMock()
request_mock.client.host = "127.0.0.1"
data = {"model": "gpt-3.5-turbo"}
user_api_key_dict = UserAPIKeyAuth(
api_key="hashed-key",
user_id="real-user",
metadata={"tags": ["key-supplied"]},
team_metadata={},
spend=0.0,
max_budget=100.0,
model_max_budget={},
team_spend=0.0,
team_max_budget=200.0,
)
updated = await add_litellm_data_to_request(
data=data,
request=request_mock,
user_api_key_dict=user_api_key_dict,
proxy_config=MagicMock(),
general_settings={},
version="test-version",
)
assert set(updated["metadata"]["tags"]) == {"header-tag", "key-supplied"}
assert tuple(updated["metadata"]["caller_tags"]) == ("header-tag",)
@pytest.mark.asyncio
async def test_add_litellm_data_to_request_caller_tags_empty_when_caller_sends_nothing():
"""caller_tags must be present (an empty tuple), not absent, when the caller
supplied no tags of their own -- an empty-but-present value tells
tag_based_routing.py's allow_fail_open that any required/excluded tag on the
request is entirely inherited, not that no origin information is available.
"""
request_mock = MagicMock(spec=Request)
request_mock.url.path = "/v1/chat/completions"
request_mock.url = MagicMock()
request_mock.url.__str__.return_value = "http://localhost/v1/chat/completions"
request_mock.method = "POST"
request_mock.query_params = {}
request_mock.headers = {"Content-Type": "application/json"}
request_mock.client = MagicMock()
request_mock.client.host = "127.0.0.1"
data = {"model": "gpt-3.5-turbo"}
user_api_key_dict = UserAPIKeyAuth(
api_key="hashed-key",
user_id="real-user",
metadata={"tags": ["key-supplied"]},
team_metadata={},
spend=0.0,
max_budget=100.0,
model_max_budget={},
team_spend=0.0,
team_max_budget=200.0,
)
updated = await add_litellm_data_to_request(
data=data,
request=request_mock,
user_api_key_dict=user_api_key_dict,
proxy_config=MagicMock(),
general_settings={},
version="test-version",
)
assert updated["metadata"]["tags"] == ["key-supplied"]
assert updated["metadata"]["caller_tags"] == ()
OAUTH_TOKEN = "Bearer sk-ant-oat01-fake-subscription-token-for-testing-0123456789"
GOOGLE_ACCESS_TOKEN = "Bearer ya29.fake-google-access-token-for-testing"
BEDROCK_API_KEY = "ABSKQmVkcm9ja0FQSUtleUZvclRlc3Rpbmc="
CROSS_ACCOUNT_AUTHORIZATION = "Bearer deliberately-configured-pass-through-token"
SIGV4_PREFIX = "AWS4-HMAC-SHA256"
AUTHORIZATION_HEADER_CASINGS = ["authorization", "Authorization", "AUTHORIZATION"]
LEAK_TARGET_PROVIDERS = ["bedrock", "bedrock_converse", "bedrock_mantle", "vertex_ai"]
BEDROCK_ENDPOINT = (
"https://bedrock-runtime.us-west-2.amazonaws.com/model/us.anthropic.claude-sonnet-4-5-20250929-v1:0/invoke"
)
BEDROCK_REGION = "us-west-2"
BEDROCK_REQUEST_DATA = {"messages": [{"role": "user", "content": "Say OK"}], "max_tokens": 32}
SIGV4_OPTIONAL_PARAMS = {
"aws_access_key_id": "AKIAIOSFODNN7EXAMPLE",
"aws_secret_access_key": "wJalrXUtnFEMI/K7MDENG/bPxRfiCYEXAMPLEKEY",
"aws_region_name": BEDROCK_REGION,
}
def _client_headers(authorization_header_name: str | None = "authorization") -> dict:
headers = {
"content-type": "application/json",
"anthropic-version": "2023-06-01",
"user-agent": "claude-cli/2.1.239",
}
if authorization_header_name is not None:
headers[authorization_header_name] = OAUTH_TOKEN
return headers
def _headers_forwarded_to(client_headers: dict, custom_llm_provider: str) -> dict:
data: dict = {}
add_provider_specific_headers_to_request(data=data, headers=client_headers)
return ProviderSpecificHeaderUtils.get_provider_specific_headers(
provider_specific_header=data.get("provider_specific_header"),
custom_llm_provider=custom_llm_provider,
)
def _authorization_values(headers) -> list:
return [value for name, value in headers.items() if name.lower() == "authorization"]
def _signed_headers_for_bedrock(request_headers: dict, api_key: str | None = None) -> dict:
with patch.dict(os.environ, {"AWS_BEARER_TOKEN_BEDROCK": ""}):
signed_headers, _ = BaseAWSLLM()._sign_request(
service_name="bedrock",
headers=request_headers,
optional_params=SIGV4_OPTIONAL_PARAMS,
request_data=BEDROCK_REQUEST_DATA,
api_base=BEDROCK_ENDPOINT,
api_key=api_key,
)
return signed_headers
def _signed_headers_component(signature: str, component: str) -> str:
for part in signature.removeprefix(SIGV4_PREFIX).split(","):
name, _, value = part.strip().partition("=")
if name == component:
return value
raise AssertionError(f"{component} missing from {signature}")
@pytest.mark.parametrize("authorization_header_name", AUTHORIZATION_HEADER_CASINGS)
@pytest.mark.parametrize("custom_llm_provider", LEAK_TARGET_PROVIDERS)
def test_oauth_credential_is_never_forwarded_to_bedrock_or_vertex(authorization_header_name, custom_llm_provider):
"""
A client's Anthropic OAuth credential is meaningless to AWS and Google, and sending it
there both breaks the request and hands a third-party cloud a credential it should
never hold. It must not survive the pre-call path for any non-Anthropic provider.
"""
forwarded = _headers_forwarded_to(_client_headers(authorization_header_name), custom_llm_provider)
assert _authorization_values(forwarded) == []
assert OAUTH_TOKEN not in forwarded.values()
@pytest.mark.parametrize("authorization_header_name", AUTHORIZATION_HEADER_CASINGS)
def test_oauth_credential_still_reaches_anthropic_unchanged(authorization_header_name):
forwarded = _headers_forwarded_to(_client_headers(authorization_header_name), "anthropic")
assert forwarded[authorization_header_name] == OAUTH_TOKEN
assert _authorization_values(forwarded) == [OAUTH_TOKEN]
def test_oauth_credential_entry_is_scoped_to_anthropic_alone():
data: dict = {}
add_provider_specific_headers_to_request(data=data, headers=_client_headers())
scoped_headers = data["provider_specific_header"]
if not isinstance(scoped_headers, list):
scoped_headers = [scoped_headers]
credential_entries = [entry for entry in scoped_headers if OAUTH_TOKEN in entry["extra_headers"].values()]
assert [entry["custom_llm_provider"] for entry in credential_entries] == ["anthropic"]
@pytest.mark.parametrize("custom_llm_provider", ["anthropic", "bedrock", "bedrock_mantle", "vertex_ai"])
def test_client_anthropic_api_headers_reach_every_anthropic_messages_provider(custom_llm_provider):
client_headers = {
"anthropic-beta": "claude-code-20250219,interleaved-thinking-2025-05-14",
"anthropic-version": "2023-06-01",
"user-agent": "claude-cli/2.1.239",
}
forwarded = _headers_forwarded_to(client_headers, custom_llm_provider)
assert forwarded == {
"anthropic-beta": "claude-code-20250219,interleaved-thinking-2025-05-14",
"anthropic-version": "2023-06-01",
}
def test_client_anthropic_api_headers_stay_off_openai_compatible_providers():
forwarded = _headers_forwarded_to({"anthropic-beta": "claude-code-20250219"}, "openai")
assert forwarded == {}
@pytest.mark.parametrize("authorization_header_name", AUTHORIZATION_HEADER_CASINGS)
def test_add_provider_specific_headers_reports_a_forwarded_oauth_credential(authorization_header_name):
assert add_provider_specific_headers_to_request(data={}, headers=_client_headers(authorization_header_name)) is True
@pytest.mark.parametrize(
"headers",
[
_client_headers(None),
{"content-type": "application/json", "authorization": "Bearer sk-a-normal-key"},
{"anthropic-beta": "claude-code-20250219", "authorization": "Bearer sk-ant-api03-a-configured-key"},
],
)
def test_add_provider_specific_headers_reports_no_oauth_credential_without_a_forwarded_token(headers):
assert add_provider_specific_headers_to_request(data={}, headers=headers) is False
def test_no_provider_specific_header_when_client_sends_nothing_anthropic():
data: dict = {}
add_provider_specific_headers_to_request(
data=data, headers={"content-type": "application/json", "authorization": "Bearer sk-a-normal-key"}
)
assert "provider_specific_header" not in data
def test_bedrock_sigv4_signature_survives_a_client_oauth_header():
forwarded = _headers_forwarded_to(_client_headers(), "bedrock")
signed = _signed_headers_for_bedrock({"Content-Type": "application/json", **forwarded})
authorizations = _authorization_values(signed)
assert len(authorizations) == 1
assert authorizations[0].startswith(SIGV4_PREFIX)
assert signed["X-Amz-Date"]
def test_bedrock_sigv4_signing_is_unchanged_by_the_client_oauth_header():
without_oauth = _signed_headers_for_bedrock(
{"Content-Type": "application/json", **_headers_forwarded_to(_client_headers(None), "bedrock")}
)
with_oauth = _signed_headers_for_bedrock(
{"Content-Type": "application/json", **_headers_forwarded_to(_client_headers(), "bedrock")}
)
assert without_oauth["Authorization"].startswith(SIGV4_PREFIX)
assert _signed_headers_component(with_oauth["Authorization"], "SignedHeaders") == (
_signed_headers_component(without_oauth["Authorization"], "SignedHeaders")
)
def test_bedrock_get_request_headers_keeps_the_sigv4_signature():
forwarded = _headers_forwarded_to(_client_headers(), "bedrock")
with patch.dict(os.environ, {"AWS_BEARER_TOKEN_BEDROCK": ""}):
prepped = BaseAWSLLM().get_request_headers(
credentials=Credentials(
SIGV4_OPTIONAL_PARAMS["aws_access_key_id"],
SIGV4_OPTIONAL_PARAMS["aws_secret_access_key"],
),
aws_region_name=BEDROCK_REGION,
extra_headers=forwarded,
endpoint_url=BEDROCK_ENDPOINT,
data=json.dumps(BEDROCK_REQUEST_DATA),
headers={"Content-Type": "application/json", **forwarded},
)
authorizations = _authorization_values(prepped.headers)
assert len(authorizations) == 1
assert authorizations[0].startswith(SIGV4_PREFIX)
def test_bedrock_api_key_deployment_keeps_its_own_bearer_token():
forwarded = _headers_forwarded_to(_client_headers(), "bedrock")
signed = _signed_headers_for_bedrock({"Content-Type": "application/json", **forwarded}, api_key=BEDROCK_API_KEY)
assert _authorization_values(signed) == [f"Bearer {BEDROCK_API_KEY}"]
def test_deliberately_configured_authorization_still_overrides_sigv4():
signed = _signed_headers_for_bedrock(
{"Content-Type": "application/json", "Authorization": CROSS_ACCOUNT_AUTHORIZATION}
)
assert _authorization_values(signed) == [CROSS_ACCOUNT_AUTHORIZATION]
def test_vertex_sends_exactly_one_authorization_header():
forwarded = _headers_forwarded_to(_client_headers(), "vertex_ai")
vertex_request_headers = {
"content-type": "application/json",
"Authorization": GOOGLE_ACCESS_TOKEN,
}
vertex_request_headers.update(forwarded)
assert _authorization_values(vertex_request_headers) == [GOOGLE_ACCESS_TOKEN]
@pytest.mark.asyncio
async def test_newrelic_team_callback_vars_reach_trusted_field():
"""A key with a newrelic team callback stamps its vars into the proxy-owned
trusted field, and a caller-supplied newrelic_api_key in the body is
stripped rather than merged."""
key_with_newrelic_callback = UserAPIKeyAuth(
api_key="hashed-key",
metadata={
"logging": [
{
"callback_name": "newrelic",
"callback_type": "success",
"callback_vars": {"newrelic_api_key": "team-nr-key", "newrelic_region": "eu"},
}
]
},
)
data = {
"model": "gpt-3.5-turbo",
"messages": [{"role": "user", "content": "hello"}],
"newrelic_api_key": "attacker-key",
}
updated = await add_litellm_data_to_request(
data=data,
request=_callback_credential_request_mock(),
user_api_key_dict=key_with_newrelic_callback,
proxy_config=MagicMock(),
general_settings={},
version="test-version",
)
assert updated[TRUSTED_CALLBACK_VARS_FIELD] == {
"newrelic_api_key": "team-nr-key",
"newrelic_region": "eu",
}
assert updated["success_callback"] == ["newrelic"]
from litellm.litellm_core_utils.initialize_dynamic_callback_params import (
initialize_standard_callback_dynamic_params,
)
params = initialize_standard_callback_dynamic_params(updated)
assert params.get("newrelic_api_key") == "team-nr-key"
assert params.get("newrelic_region") == "eu"
from litellm.integrations.otel.presets import dynamic_otlp_endpoint, dynamic_otlp_headers
assert dynamic_otlp_headers("newrelic", params) == {"api-key": "team-nr-key"}
assert dynamic_otlp_endpoint("newrelic", params) == "https://otlp.eu01.nr-data.net"
from litellm.utils import get_non_default_completion_params
forwarded = get_non_default_completion_params(updated)
assert not any(param.startswith("newrelic_") for param in forwarded)
assert TRUSTED_CALLBACK_VARS_FIELD not in forwarded
def test_newrelic_vars_scoped_to_newrelic_callback_entry():
"""New Relic routing reads these vars from the trusted overlay with no
callback-name check, so a team that puts newrelic_* under a different
callback's vars must not have them enter the shared bag (and so never
exports to New Relic). Vars under a real newrelic entry are kept."""
from litellm.proxy._types import AddTeamCallback
from litellm.proxy.litellm_pre_call_utils import convert_key_logging_metadata_to_callback
smuggled = convert_key_logging_metadata_to_callback(
AddTeamCallback(
callback_name="langfuse",
callback_type="success",
callback_vars={
"langfuse_public_key": "pk",
"newrelic_api_key": "SMUGGLED",
"newrelic_region": "eu",
},
),
None,
)
assert smuggled.callback_vars == {"langfuse_public_key": "pk"}
legit = convert_key_logging_metadata_to_callback(
AddTeamCallback(
callback_name="newrelic",
callback_type="success",
callback_vars={"newrelic_api_key": "REAL", "newrelic_region": "us"},
),
None,
)
assert legit.callback_vars == {"newrelic_api_key": "REAL", "newrelic_region": "us"}
def _reserved_stamp_request(path: str) -> MagicMock:
request_mock = MagicMock(spec=Request)
request_mock.url = MagicMock()
request_mock.url.path = path
request_mock.url.__str__.return_value = f"http://localhost{path}"
request_mock.method = "POST"
request_mock.query_params = {}
request_mock.headers = {"Content-Type": "application/json"}
request_mock.client = MagicMock()
request_mock.client.host = "127.0.0.1"
return request_mock
def _reserved_stamp_key(key_metadata: dict | None = None) -> UserAPIKeyAuth:
return UserAPIKeyAuth(
api_key="hashed-key",
metadata=key_metadata or {},
team_metadata={},
spend=0.0,
max_budget=100.0,
model_max_budget={},
team_spend=0.0,
team_max_budget=200.0,
)
_PLANTED_STAMPS = {
"attempted_fallbacks": 99,
"original_model_group": "spoofed-group",
"request_retry_count": -100,
"_client_output_ceiling": {"api_base": "https://attacker.example"},
ROUTER_USAGE_COUNTED_TOKENS_METADATA_KEY: 10**9,
"client_key": "client_value",
}
@pytest.mark.asyncio
async def test_add_litellm_data_to_request_strips_router_reserved_stamps_from_both_buckets() -> None:
"""attempted_fallbacks and original_model_group are router-written facts the spend row
reads back; a client planting them in either bucket is dropped at the boundary so the
router never sees a reserved key it did not write."""
from litellm.proxy.litellm_pre_call_utils import add_litellm_data_to_request
data = {
"model": "gpt-3.5-turbo",
"messages": [{"role": "user", "content": "hi"}],
"metadata": dict(_PLANTED_STAMPS),
"litellm_metadata": dict(_PLANTED_STAMPS),
}
updated = await add_litellm_data_to_request(
data=data,
request=_reserved_stamp_request("/v1/chat/completions"),
user_api_key_dict=_reserved_stamp_key(),
proxy_config=MagicMock(),
general_settings={},
version="test-version",
)
assert "litellm_metadata" not in updated
assert "attempted_fallbacks" not in updated["metadata"]
assert "original_model_group" not in updated["metadata"]
assert "_client_output_ceiling" not in updated["metadata"]
assert "request_retry_count" not in updated["metadata"]
assert ROUTER_USAGE_COUNTED_TOKENS_METADATA_KEY not in updated["metadata"]
assert updated["metadata"]["client_key"] == "client_value"
@pytest.mark.asyncio
async def test_add_litellm_data_to_request_strips_router_reserved_stamps_from_json_string_litellm_metadata() -> None:
from litellm.proxy.litellm_pre_call_utils import add_litellm_data_to_request
data = {
"model": "gpt-3.5-turbo",
"messages": [{"role": "user", "content": "hi"}],
"litellm_metadata": json.dumps(_PLANTED_STAMPS),
}
updated = await add_litellm_data_to_request(
data=data,
request=_reserved_stamp_request("/v1/chat/completions"),
user_api_key_dict=_reserved_stamp_key(),
proxy_config=MagicMock(),
general_settings={},
version="test-version",
)
assert "litellm_metadata" not in updated
assert "attempted_fallbacks" not in updated["metadata"]
assert "original_model_group" not in updated["metadata"]
assert "request_retry_count" not in updated["metadata"]
assert updated["metadata"]["client_key"] == "client_value"
@pytest.mark.asyncio
async def test_add_litellm_data_to_request_strips_router_reserved_stamps_despite_pricing_override_opt_in() -> None:
"""The pricing strip is gated on allow_client_pricing_override; the reserved-stamp strip
is not, because no key or team setting makes a client-written fallback count valid."""
from litellm.proxy.litellm_pre_call_utils import add_litellm_data_to_request
data = {
"model": "gpt-3.5-turbo",
"messages": [{"role": "user", "content": "hi"}],
"litellm_metadata": {**_PLANTED_STAMPS, "model_info": {"input_cost_per_token": 0.0}},
}
updated = await add_litellm_data_to_request(
data=data,
request=_reserved_stamp_request("/v1/chat/completions"),
user_api_key_dict=_reserved_stamp_key({"allow_client_pricing_override": True}),
proxy_config=MagicMock(),
general_settings={},
version="test-version",
)
assert "litellm_metadata" not in updated
assert updated["metadata"]["model_info"] == {"input_cost_per_token": 0.0}
assert "attempted_fallbacks" not in updated["metadata"]
assert "original_model_group" not in updated["metadata"]
assert "request_retry_count" not in updated["metadata"]
@pytest.mark.asyncio
async def test_add_litellm_data_to_request_strips_router_reserved_stamps_on_responses_route():
"""On the Responses family the proxy-owned bucket is litellm_metadata and the client's
OpenAI metadata param is the sibling; both lose the reserved keys."""
from litellm.proxy.litellm_pre_call_utils import add_litellm_data_to_request
data = {
"model": "gpt-3.5-turbo",
"input": "hi",
"metadata": dict(_PLANTED_STAMPS),
"litellm_metadata": dict(_PLANTED_STAMPS),
}
updated = await add_litellm_data_to_request(
data=data,
request=_reserved_stamp_request("/v1/responses"),
user_api_key_dict=_reserved_stamp_key(),
proxy_config=MagicMock(),
general_settings={},
version="test-version",
)
for bucket in ("metadata", "litellm_metadata"):
assert "attempted_fallbacks" not in updated[bucket]
assert "original_model_group" not in updated[bucket]
assert updated[bucket]["client_key"] == "client_value"
@pytest.mark.asyncio
async def test_router_keeps_proxy_metadata_bucket_identity_after_reserved_stamp_strip():
"""Regression for the #38586 break: a client that planted a reserved key in
litellm_metadata made the router hand downstream a scrubbed copy, so the proxy's
post_call write-backs (guardrail telemetry, applied guardrails) landed in a dict the
spend row never read. After the boundary strip plus the in-place scrub, the object the
router forwards is the proxy's own request_data bucket; on chat routes that bucket is
``metadata``, since the boundary folds client ``litellm_metadata`` into it."""
from litellm.proxy.litellm_pre_call_utils import add_litellm_data_to_request
data = {
"model": "gpt-3.5-turbo",
"messages": [{"role": "user", "content": "hi"}],
"litellm_metadata": dict(_PLANTED_STAMPS),
}
request_data = await add_litellm_data_to_request(
data=data,
request=_reserved_stamp_request("/v1/chat/completions"),
user_api_key_dict=_reserved_stamp_key(),
proxy_config=MagicMock(),
general_settings={},
version="test-version",
)
proxy_bucket = request_data["metadata"]
assert "attempted_fallbacks" not in proxy_bucket
assert "original_model_group" not in proxy_bucket
router = litellm.Router(
model_list=[
{
"model_name": "gpt-3.5-turbo",
"litellm_params": {"model": "gpt-3.5-turbo", "mock_response": "hi"},
}
]
)
forwarded_buckets = []
original_acompletion = router._acompletion
async def _spy(*args, **spy_kwargs):
forwarded_buckets.append(spy_kwargs["metadata"])
return await original_acompletion(*args, **spy_kwargs)
router._acompletion = _spy
await router.acompletion(**request_data)
assert forwarded_buckets == [proxy_bucket]
assert forwarded_buckets[0] is proxy_bucket
assert proxy_bucket["attempted_fallbacks"] == 0
assert proxy_bucket.get("original_model_group") != "spoofed-group"
proxy_bucket["standard_logging_guardrail_information"] = [{"guardrail_name": "postcall-guard"}]
assert forwarded_buckets[0]["standard_logging_guardrail_information"] == [{"guardrail_name": "postcall-guard"}]
@pytest.mark.asyncio
async def test_add_litellm_data_to_request_folds_litellm_metadata_into_metadata_on_chat_routes():
data = {
"model": "gpt-3.5-turbo",
"metadata": {"tags": ["from-metadata"]},
"litellm_metadata": {"trace_id": "abc", "tags": ["from-litellm-metadata"]},
}
updated = await add_litellm_data_to_request(
data=data,
request=_make_chat_request_mock(),
user_api_key_dict=UserAPIKeyAuth(api_key="hashed-key"),
proxy_config=MagicMock(),
general_settings={},
version="test-version",
)
assert "litellm_metadata" not in updated
assert updated["metadata"]["trace_id"] == "abc"
assert updated["metadata"]["tags"] == ["from-metadata", "from-litellm-metadata"]
@pytest.mark.asyncio
async def test_add_litellm_data_to_request_keeps_litellm_metadata_on_litellm_metadata_routes():
data = {"model": "claude-sonnet-5", "litellm_metadata": {"trace_id": "abc"}}
updated = await add_litellm_data_to_request(
data=data,
request=_make_request_mock("/v1/messages", {"Content-Type": "application/json"}),
user_api_key_dict=UserAPIKeyAuth(api_key="hashed-key"),
proxy_config=MagicMock(),
general_settings={},
version="test-version",
)
assert updated["litellm_metadata"]["trace_id"] == "abc"
def _stamp_model_access_groups(matched_model_access_groups, metadata_variable_name="metadata"):
user_api_key_dict = UserAPIKeyAuth(api_key="hashed-key")
user_api_key_dict.matched_model_access_groups = matched_model_access_groups
return LiteLLMProxyRequestSetup.add_user_api_key_auth_to_request_metadata(
data={metadata_variable_name: {}},
user_api_key_dict=user_api_key_dict,
_metadata_variable_name=metadata_variable_name,
)[metadata_variable_name]
def test_matched_model_access_groups_are_stamped_into_request_metadata():
"""The post-call spend writer reads the groups off request metadata, not off UserAPIKeyAuth."""
stamped = _stamp_model_access_groups(["tier-a", "tier-b"])
assert stamped[MODEL_ACCESS_GROUP_METADATA_KEY] == ["tier-a", "tier-b"]
assert MODEL_ACCESS_GROUP_METADATA_KEY not in _stamp_model_access_groups(None)
def test_stamped_model_access_groups_survive_the_litellm_metadata_merge():
"""
The key must keep its ``user_api_key`` prefix: when a request carries both metadata dicts,
get_litellm_metadata_from_kwargs returns litellm_metadata and copies a key over from metadata
only when that substring is in its name, so an unprefixed key is silently dropped.
"""
kwargs = {
"litellm_params": {
"metadata": _stamp_model_access_groups(["tier-a"]),
"litellm_metadata": {"trace_id": "abc"},
}
}
assert get_litellm_metadata_from_kwargs(kwargs)[MODEL_ACCESS_GROUP_METADATA_KEY] == ["tier-a"]
def _request_for(path: str) -> MagicMock:
request = MagicMock(spec=Request)
request.scope = {"path": path}
request.url = MagicMock()
request.url.path = path
request.url.__str__.return_value = f"http://localhost{path}"
request.method = "POST"
request.query_params = {}
request.headers = {"Content-Type": "application/json"}
request.client = MagicMock()
request.client.host = "127.0.0.1"
return request
def _spend_log_session_id(data: dict[str, object], metadata_key: str = "metadata") -> str | None:
"""Resolve session_id the way LiteLLM_SpendLogs does, reading the omit decision stamped on the request."""
from litellm.litellm_core_utils.get_litellm_params import get_litellm_params
from litellm.litellm_core_utils.litellm_logging import StandardLoggingPayloadSetup
from litellm.proxy.spend_tracking.spend_tracking_utils import _get_session_id_for_spend_log
metadata = data[metadata_key]
assert isinstance(metadata, dict)
litellm_params = get_litellm_params(
litellm_session_id=str(data["litellm_session_id"]) if "litellm_session_id" in data else None,
litellm_trace_id=str(data["litellm_trace_id"]) if "litellm_trace_id" in data else None,
metadata=metadata,
)
trace_id = StandardLoggingPayloadSetup.get_standard_logging_payload_trace_id(
logging_obj=SimpleNamespace(litellm_trace_id="per-call-random-trace-id"),
litellm_params=litellm_params,
)
return _get_session_id_for_spend_log(
kwargs={},
metadata=metadata,
standard_logging_payload={"trace_id": trace_id},
omit_when_missing=bool(metadata.get(SESSION_ID_OMITTED_METADATA_KEY)),
)
@pytest.mark.asyncio
@pytest.mark.parametrize("request_correlation_in_logs", [False, True])
async def test_missing_session_id_generate_makes_spend_log_and_callback_session_ids_agree(
monkeypatch: pytest.MonkeyPatch, request_correlation_in_logs: bool
):
"""Without a session header, SpendLogs.session_id and the metadata.session_id that Langfuse logs
must be the same generated id, so cross-referencing the two by session_id works. The id is marked
as generated so affinity consumers (Fireworks x-session-affinity, router session pins) skip it."""
monkeypatch.setattr(litellm, "request_correlation_in_logs", request_correlation_in_logs)
data = {"model": "gpt-4o", "messages": [{"role": "user", "content": "hi"}]}
updated = await add_litellm_data_to_request(
data=data,
request=_request_for("/v1/chat/completions"),
user_api_key_dict=UserAPIKeyAuth(api_key="hashed-key"),
proxy_config=MagicMock(),
general_settings={"missing_session_id": "generate"},
)
callback_session_id = updated["metadata"]["session_id"]
assert isinstance(callback_session_id, str) and len(callback_session_id) == 36
assert _spend_log_session_id(updated) == callback_session_id
assert updated["metadata"][SESSION_ID_GENERATED_METADATA_KEY] is True
assert (
get_fireworks_session_id({"litellm_session_id": updated["litellm_session_id"], "metadata": updated["metadata"]})
is None
)
@pytest.mark.asyncio
async def test_missing_session_id_unset_keeps_legacy_divergence():
updated = await add_litellm_data_to_request(
data={"model": "gpt-4o", "messages": []},
request=_request_for("/v1/chat/completions"),
user_api_key_dict=UserAPIKeyAuth(api_key="hashed-key"),
proxy_config=MagicMock(),
general_settings={},
)
assert "session_id" not in updated["metadata"]
assert "litellm_session_id" not in updated
assert _spend_log_session_id(updated) == "per-call-random-trace-id"
@pytest.mark.asyncio
async def test_missing_session_id_omit_leaves_spend_log_session_id_null():
"""Under `omit` a traceparent still becomes the trace id but never a session id, so SpendLogs and
Langfuse agree on having no session."""
request = _request_for("/v1/chat/completions")
request.headers = {"traceparent": "00-4bf92f3577b34da6a3ce929d0e0e4736-00f067aa0ba902b7-01"}
updated = await add_litellm_data_to_request(
data={"model": "gpt-4o", "messages": []},
request=request,
user_api_key_dict=UserAPIKeyAuth(api_key="hashed-key"),
proxy_config=MagicMock(),
general_settings={"missing_session_id": "omit"},
)
assert "session_id" not in updated["metadata"]
assert updated["litellm_trace_id"] == "4bf92f3577b34da6a3ce929d0e0e4736"
assert updated["metadata"][SESSION_ID_OMITTED_METADATA_KEY] is True
assert _spend_log_session_id(updated) is None
@pytest.mark.asyncio
async def test_missing_session_id_omit_keeps_client_supplied_session_id():
request = _request_for("/v1/chat/completions")
request.headers = {"x-litellm-session-id": "client-session-1"}
updated = await add_litellm_data_to_request(
data={"model": "gpt-4o", "messages": []},
request=request,
user_api_key_dict=UserAPIKeyAuth(api_key="hashed-key"),
proxy_config=MagicMock(),
general_settings={"missing_session_id": "omit"},
)
assert updated["metadata"]["session_id"] == "client-session-1"
assert _spend_log_session_id(updated) == "client-session-1"
@pytest.mark.asyncio
@pytest.mark.parametrize(
"client_body",
[
{"model": "gpt-4o", "messages": [], "litellm_session_id": "cust-sess-1"},
{"model": "gpt-4o", "messages": [], "litellm_session_id": "cust-sess-1", "metadata": {"trace_id": "trace-1"}},
],
)
async def test_missing_session_id_omit_keeps_body_litellm_session_id(
monkeypatch: pytest.MonkeyPatch, client_body: dict[str, object]
):
from litellm.litellm_core_utils.get_litellm_params import get_litellm_params
from litellm.litellm_core_utils.litellm_logging import StandardLoggingPayloadSetup
monkeypatch.setattr(litellm, "request_correlation_in_logs", True)
updated = await add_litellm_data_to_request(
data=client_body,
request=_request_for("/v1/chat/completions"),
user_api_key_dict=UserAPIKeyAuth(api_key="hashed-key"),
proxy_config=MagicMock(),
general_settings={"missing_session_id": "omit"},
)
callback_session_id = StandardLoggingPayloadSetup.get_standard_logging_payload_session_id(
logging_obj=SimpleNamespace(litellm_session_id=""),
litellm_params=get_litellm_params(litellm_session_id="cust-sess-1", metadata=updated["metadata"]),
)
assert callback_session_id == "cust-sess-1"
assert updated["metadata"]["session_id"] == "cust-sess-1"
assert _spend_log_session_id(updated) == "cust-sess-1"
@pytest.mark.asyncio
async def test_missing_session_id_omit_body_litellm_session_id_does_not_override_metadata_session_id():
updated = await add_litellm_data_to_request(
data={
"model": "gpt-4o",
"messages": [],
"litellm_session_id": "cust-sess-1",
"metadata": {"session_id": "meta-sess-1"},
},
request=_request_for("/v1/chat/completions"),
user_api_key_dict=UserAPIKeyAuth(api_key="hashed-key"),
proxy_config=MagicMock(),
general_settings={"missing_session_id": "omit"},
)
assert updated["metadata"]["session_id"] == "meta-sess-1"
assert _spend_log_session_id(updated) == "meta-sess-1"
@pytest.mark.asyncio
@pytest.mark.parametrize("path", ["/v1/responses", "/v1/messages"])
async def test_missing_session_id_omit_keeps_metadata_session_id_on_litellm_metadata_routes(path: str):
updated = await add_litellm_data_to_request(
data={
"model": "gpt-4o",
"input": "hi",
"litellm_session_id": "cust-sess-1",
"metadata": {"session_id": "meta-sess-1"},
},
request=_request_for(path),
user_api_key_dict=UserAPIKeyAuth(api_key="hashed-key"),
proxy_config=MagicMock(),
general_settings={"missing_session_id": "omit"},
)
assert updated["litellm_metadata"]["session_id"] == "meta-sess-1"
assert _spend_log_session_id(updated, "litellm_metadata") == "meta-sess-1"
@pytest.mark.asyncio
@pytest.mark.parametrize("path", ["/v1/responses", "/v1/messages"])
async def test_missing_session_id_omit_keeps_body_litellm_session_id_on_litellm_metadata_routes(path: str):
updated = await add_litellm_data_to_request(
data={"model": "gpt-4o", "input": "hi", "litellm_session_id": "cust-sess-1"},
request=_request_for(path),
user_api_key_dict=UserAPIKeyAuth(api_key="hashed-key"),
proxy_config=MagicMock(),
general_settings={"missing_session_id": "omit"},
)
assert updated["litellm_metadata"]["session_id"] == "cust-sess-1"
assert _spend_log_session_id(updated, "litellm_metadata") == "cust-sess-1"
@pytest.mark.asyncio
async def test_missing_session_id_omit_ignores_empty_body_litellm_session_id():
updated = await add_litellm_data_to_request(
data={"model": "gpt-4o", "messages": [], "litellm_session_id": ""},
request=_request_for("/v1/chat/completions"),
user_api_key_dict=UserAPIKeyAuth(api_key="hashed-key"),
proxy_config=MagicMock(),
general_settings={"missing_session_id": "omit"},
)
assert "session_id" not in updated["metadata"]
assert _spend_log_session_id(updated) is None
@pytest.mark.asyncio
async def test_missing_session_id_generate_reuses_traceparent_trace_id():
"""A W3C traceparent already decides SpendLogs.session_id, so the callback session id must reuse it."""
request = _request_for("/v1/chat/completions")
request.headers = {"traceparent": "00-4bf92f3577b34da6a3ce929d0e0e4736-00f067aa0ba902b7-01"}
updated = await add_litellm_data_to_request(
data={"model": "gpt-4o", "messages": []},
request=request,
user_api_key_dict=UserAPIKeyAuth(api_key="hashed-key"),
proxy_config=MagicMock(),
general_settings={"missing_session_id": "generate"},
)
assert updated["metadata"]["session_id"] == "4bf92f3577b34da6a3ce929d0e0e4736"
assert _spend_log_session_id(updated) == "4bf92f3577b34da6a3ce929d0e0e4736"
@pytest.mark.asyncio
@pytest.mark.parametrize("policy", ["generate", "reject"])
async def test_missing_session_id_policy_keeps_client_supplied_session_id(policy: str):
request = _request_for("/v1/chat/completions")
request.headers = {"x-litellm-session-id": "client-session-1"}
updated = await add_litellm_data_to_request(
data={"model": "gpt-4o", "messages": []},
request=request,
user_api_key_dict=UserAPIKeyAuth(api_key="hashed-key"),
proxy_config=MagicMock(),
general_settings={"missing_session_id": policy},
)
assert updated["litellm_session_id"] == "client-session-1"
assert updated["metadata"]["session_id"] == "client-session-1"
assert _spend_log_session_id(updated) == "client-session-1"
assert SESSION_ID_GENERATED_METADATA_KEY not in updated["metadata"]
assert (
get_fireworks_session_id({"litellm_session_id": "client-session-1", "metadata": updated["metadata"]})
== "client-session-1"
)
@pytest.mark.asyncio
async def test_missing_session_id_reject_accepts_body_metadata_session_id():
updated = await add_litellm_data_to_request(
data={"model": "gpt-4o", "messages": [], "metadata": {"session_id": "body-session-1"}},
request=_request_for("/v1/chat/completions"),
user_api_key_dict=UserAPIKeyAuth(api_key="hashed-key"),
proxy_config=MagicMock(),
general_settings={"missing_session_id": "reject"},
)
assert updated["metadata"]["session_id"] == "body-session-1"
@pytest.mark.asyncio
async def test_missing_session_id_reject_returns_400_without_session_id():
with pytest.raises(ProxyException) as exc_info:
await add_litellm_data_to_request(
data={"model": "gpt-4o", "messages": []},
request=_request_for("/v1/chat/completions"),
user_api_key_dict=UserAPIKeyAuth(api_key="hashed-key"),
proxy_config=MagicMock(),
general_settings={"missing_session_id": "reject"},
)
assert exc_info.value.code == "400"
assert exc_info.value.param == "session_id"
@pytest.mark.asyncio
@pytest.mark.parametrize("path", ["/mcp/", "/mcp/tools", "/key/health"])
async def test_missing_session_id_policy_skips_non_inference_routes(path: str):
updated = await add_litellm_data_to_request(
data={"model": "gpt-4o"},
request=_request_for(path),
user_api_key_dict=UserAPIKeyAuth(api_key="hashed-key"),
proxy_config=MagicMock(),
general_settings={"missing_session_id": "reject"},
)
assert "session_id" not in updated["metadata"]
@pytest.mark.asyncio
async def test_missing_session_id_unknown_value_is_ignored():
updated = await add_litellm_data_to_request(
data={"model": "gpt-4o", "messages": []},
request=_request_for("/v1/chat/completions"),
user_api_key_dict=UserAPIKeyAuth(api_key="hashed-key"),
proxy_config=MagicMock(),
general_settings={"missing_session_id": "typo"},
)
assert "session_id" not in updated["metadata"]
@pytest.mark.asyncio
@pytest.mark.parametrize(
"path, sent_in, metadata_key, general_settings",
[
("/v1/chat/completions", "metadata", "metadata", {}),
("/v1/chat/completions", "litellm_metadata", "metadata", {}),
("/v1/chat/completions", "litellm_metadata", "metadata", {"missing_session_id": "generate"}),
("/v1/messages", "litellm_metadata", "litellm_metadata", {}),
("/v1/messages", "metadata", "litellm_metadata", {}),
("/mcp/tools", "metadata", "metadata", {"missing_session_id": "omit"}),
("/mcp/tools", "litellm_metadata", "metadata", {"missing_session_id": "omit"}),
],
)
async def test_client_supplied_omit_marker_never_reaches_the_spend_log(
path: str, sent_in: str, metadata_key: str, general_settings: dict[str, str]
):
"""The omit marker is proxy-owned: only the pre-call policy may set it. A caller that sends it in either
metadata bucket, including the one later merged into the route's bucket, must not be able to null out
SpendLogs.session_id on a request the proxy did not omit."""
updated = await add_litellm_data_to_request(
data={"model": "gpt-4o", "messages": [], sent_in: {SESSION_ID_OMITTED_METADATA_KEY: True}},
request=_request_for(path),
user_api_key_dict=UserAPIKeyAuth(api_key="hashed-key"),
proxy_config=MagicMock(),
general_settings=general_settings,
)
assert SESSION_ID_OMITTED_METADATA_KEY not in updated[metadata_key]
assert _spend_log_session_id(updated, metadata_key) == (
updated[metadata_key]["session_id"]
if general_settings.get("missing_session_id") == "generate"
else "per-call-random-trace-id"
)
def test_default_team_settings_bool_turn_off_message_logging_redacts():
from litellm.proxy.proxy_server import ProxyConfig
pc = ProxyConfig()
pc.config = {
"litellm_settings": {
"default_team_settings": [
{
"team_id": "team-redact",
"success_callback": ["gcs_bucket"],
"failure_callback": ["gcs_bucket"],
"turn_off_message_logging": True,
}
]
}
}
callback_metadata = LiteLLMProxyRequestSetup.add_team_based_callbacks_from_config(
team_id="team-redact",
proxy_config=pc,
)
assert callback_metadata is not None
assert callback_metadata.success_callback == ["gcs_bucket"]
assert callback_metadata.callback_vars == {"turn_off_message_logging": "True"}
assert (
_get_turn_off_message_logging_from_dynamic_params(
{"standard_callback_dynamic_params": dict(callback_metadata.callback_vars)}
)
is True
)
def test_add_user_api_key_auth_to_request_metadata_attributes_a_cli_session_to_its_alias():
data = {"model": "gpt-5.4-nano", "messages": [{"role": "user", "content": "hi"}], "litellm_metadata": {}}
session = UserAPIKeyAuth(
api_key="cli-session-Qm7xJ2kP9sLw4vT1nR8yAa",
key_alias="cli-session-alice",
user_id="alice",
is_session_token=True,
)
metadata = LiteLLMProxyRequestSetup.add_user_api_key_auth_to_request_metadata(
data=data, user_api_key_dict=session, _metadata_variable_name="litellm_metadata"
)["litellm_metadata"]
assert metadata["user_api_key"] == "cli-session-alice"
assert metadata["user_api_key_hash"] == "cli-session-alice"
assert metadata["user_api_key_alias"] == "cli-session-alice"
assert "Qm7xJ2kP9sLw4vT1nR8yAa" not in (metadata["user_api_key"], metadata["user_api_key_hash"])
def test_add_user_api_key_auth_to_request_metadata_keeps_the_hashed_token_for_virtual_keys():
from litellm.proxy._types import hash_token
data = {"model": "gpt-5.4-nano", "messages": [], "litellm_metadata": {}}
hashed = hash_token("sk-virtual-key")
virtual_key = UserAPIKeyAuth(api_key=hashed, key_alias="cli-session-alice", user_id="alice")
metadata = LiteLLMProxyRequestSetup.add_user_api_key_auth_to_request_metadata(
data=data, user_api_key_dict=virtual_key, _metadata_variable_name="litellm_metadata"
)["litellm_metadata"]
assert metadata["user_api_key"] == hashed
assert metadata["user_api_key_hash"] == hashed
@pytest.mark.asyncio
@pytest.mark.parametrize("path", ["/mcp-rest/tools/call", "/v1/responses", "/v1/chat/completions"])
@pytest.mark.parametrize("custom_auth", ["x-mcp-auth", "x-private-mcp-token"])
async def test_mcp_credentials_only_removed_from_logging_copies(path: str, custom_auth: str):
from litellm.types.mcp_server.mcp_server_manager import MCPServer
metadata_name: Final = "litellm_metadata" if path == "/v1/responses" else "metadata"
secrets: Final = {
"X-MCP-Deepwiki-Authorization": "upstream-sentinel",
custom_auth: "client-auth-sentinel",
"x-service-token": "configured-secret-sentinel",
}
attribution: Final = {"x-app-id": "app-a", "x-nuid": "user-a", "x-user-id": "identity-a"}
request: Final = _make_request_mock(path, {"Content-Type": "application/json", **secrets, **attribution})
request.headers = Headers(request.headers)
settings: Final = {"mcp_client_side_auth_header_name": custom_auth, "user_header_name": "x-user-id"}
server: Final = MCPServer(
server_id="header-test", name="header-test", transport="http", url="https://example.com/mcp",
extra_headers=["x-service-token", "x-user-id"],
)
with (
patch("litellm.proxy.proxy_server.general_settings", settings),
patch.dict(
"litellm.proxy._experimental.mcp_server.mcp_server_manager.global_mcp_server_manager.config_mcp_servers",
{"header-test": server}, clear=True,
),
):
updated: Final = await add_litellm_data_to_request(
data={"model": "test-model", "messages": [{"role": "user", "content": "hello"}]},
request=request, user_api_key_dict=UserAPIKeyAuth(api_key="hashed-key"),
proxy_config=MagicMock(), general_settings=settings, version="test",
)
for header_dict in _all_header_dicts(updated, metadata_name):
assert not any(value in json.dumps(header_dict) for value in secrets.values())
assert updated[metadata_name]["headers"] == updated["proxy_server_request"]["headers"]
for name, value in attribution.items():
assert updated[metadata_name]["headers"][name] == value
for name, value in secrets.items():
assert updated["secret_fields"]["raw_headers"][name.lower()] == value
assert request.headers[name] == value
def test_signoz_callback_vars_are_scoped_to_the_signoz_callback():
from litellm.proxy._types import AddTeamCallback
from litellm.proxy.litellm_pre_call_utils import convert_key_logging_metadata_to_callback
under_signoz = convert_key_logging_metadata_to_callback(
data=AddTeamCallback(
callback_name="signoz",
callback_type="success",
callback_vars={"signoz_ingestion_key": "team-key", "signoz_ingestion_endpoint": "https://ingest.eu.signoz.cloud:443"},
),
team_callback_settings_obj=None,
)
assert under_signoz.callback_vars == {
"signoz_ingestion_key": "team-key",
"signoz_ingestion_endpoint": "https://ingest.eu.signoz.cloud:443",
}
under_other = convert_key_logging_metadata_to_callback(
data=AddTeamCallback(
callback_name="langfuse",
callback_type="success",
callback_vars={"signoz_ingestion_key": "team-key", "langfuse_host": "https://cloud.langfuse.com"},
),
team_callback_settings_obj=None,
)
assert under_other.callback_vars == {"langfuse_host": "https://cloud.langfuse.com"}