mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-02 02:11:58 +00:00
Co-authored-by: yuneng <yuneng@berri.ai> Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
8592 lines
327 KiB
Python
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"}
|