mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-09 03:18:44 +00:00
Merge pull request #20526 from ryan-crabbe/perf/add-litellm-data-to-request-optimizations
Perf: add_litellm_data_to_request optimizations
This commit is contained in:
commit
a8dbcb1a30
2 changed files with 79 additions and 24 deletions
|
|
@ -76,12 +76,14 @@ def _get_metadata_variable_name(request: Request) -> str:
|
|||
|
||||
For all /thread or /assistant endpoints we need to call this "litellm_metadata"
|
||||
|
||||
For ALL other endpoints we call this "metadata
|
||||
For ALL other endpoints we call this "metadata"
|
||||
"""
|
||||
if RouteChecks._is_assistants_api_request(request):
|
||||
path = request.url.path
|
||||
|
||||
if "thread" in path or "assistant" in path:
|
||||
return "litellm_metadata"
|
||||
|
||||
if any(route in request.url.path for route in LITELLM_METADATA_ROUTES):
|
||||
if any(route in path for route in LITELLM_METADATA_ROUTES):
|
||||
return "litellm_metadata"
|
||||
|
||||
return "metadata"
|
||||
|
|
@ -832,7 +834,8 @@ async def add_litellm_data_to_request( # noqa: PLR0915
|
|||
from litellm.proxy.proxy_server import llm_router, premium_user
|
||||
from litellm.types.proxy.litellm_pre_call_utils import SecretFields
|
||||
|
||||
_headers = clean_headers(
|
||||
_raw_headers: Dict[str, str] = dict(request.headers)
|
||||
_headers: Dict[str, str] = clean_headers(
|
||||
request.headers,
|
||||
litellm_key_header_name=(
|
||||
general_settings.get("litellm_key_header_name")
|
||||
|
|
@ -868,12 +871,10 @@ async def add_litellm_data_to_request( # noqa: PLR0915
|
|||
)
|
||||
)
|
||||
|
||||
data.update(
|
||||
LiteLLMProxyRequestSetup.add_litellm_metadata_from_request_headers(
|
||||
headers=_headers,
|
||||
data=data,
|
||||
_metadata_variable_name=_metadata_variable_name,
|
||||
)
|
||||
LiteLLMProxyRequestSetup.add_litellm_metadata_from_request_headers(
|
||||
headers=_headers,
|
||||
data=data,
|
||||
_metadata_variable_name=_metadata_variable_name,
|
||||
)
|
||||
|
||||
# Add headers to metadata for guardrails to access (fixes #17477)
|
||||
|
|
@ -900,7 +901,7 @@ async def add_litellm_data_to_request( # noqa: PLR0915
|
|||
if "user" not in data:
|
||||
data["user"] = user
|
||||
|
||||
data["secret_fields"] = SecretFields(raw_headers=dict(request.headers))
|
||||
data["secret_fields"] = SecretFields(raw_headers=_raw_headers)
|
||||
|
||||
## Dynamic api version (Azure OpenAI endpoints) ##
|
||||
try:
|
||||
|
|
@ -1044,10 +1045,6 @@ async def add_litellm_data_to_request( # noqa: PLR0915
|
|||
] = user_api_key_dict.user_max_budget
|
||||
|
||||
data[_metadata_variable_name]["user_api_key_metadata"] = user_api_key_dict.metadata
|
||||
_headers = dict(request.headers)
|
||||
_headers.pop(
|
||||
"authorization", None
|
||||
) # do not store the original `sk-..` api key in the db
|
||||
data[_metadata_variable_name]["headers"] = _headers
|
||||
data[_metadata_variable_name]["endpoint"] = str(request.url)
|
||||
|
||||
|
|
@ -1099,7 +1096,7 @@ async def add_litellm_data_to_request( # noqa: PLR0915
|
|||
# Check if using tag based routing
|
||||
tags = LiteLLMProxyRequestSetup.add_request_tag_to_metadata(
|
||||
llm_router=llm_router,
|
||||
headers=dict(request.headers),
|
||||
headers=_headers,
|
||||
data=data,
|
||||
)
|
||||
|
||||
|
|
@ -1128,20 +1125,13 @@ async def add_litellm_data_to_request( # noqa: PLR0915
|
|||
if disabled_callbacks and isinstance(disabled_callbacks, list):
|
||||
data["litellm_disabled_callbacks"] = disabled_callbacks
|
||||
|
||||
# Guardrails from key/team metadata
|
||||
# Guardrails from key/team metadata and policy engine
|
||||
move_guardrails_to_metadata(
|
||||
data=data,
|
||||
_metadata_variable_name=_metadata_variable_name,
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
)
|
||||
|
||||
# Guardrails from policy engine
|
||||
add_guardrails_from_policy_engine(
|
||||
data=data,
|
||||
metadata_variable_name=_metadata_variable_name,
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
)
|
||||
|
||||
# Team Model Aliases
|
||||
_update_model_if_team_alias_exists(
|
||||
data=data,
|
||||
|
|
@ -1482,6 +1472,29 @@ def move_guardrails_to_metadata(
|
|||
- Adds guardrails from policies attached to key/team metadata
|
||||
- Adds guardrails from policy engine based on team/key/model context
|
||||
"""
|
||||
# Early-out: skip all guardrails processing when nothing is configured
|
||||
key_metadata = user_api_key_dict.metadata
|
||||
team_metadata = user_api_key_dict.team_metadata
|
||||
|
||||
has_key_config = key_metadata and (
|
||||
"guardrails" in key_metadata or "policies" in key_metadata
|
||||
)
|
||||
has_team_config = team_metadata and (
|
||||
"guardrails" in team_metadata or "policies" in team_metadata
|
||||
)
|
||||
has_request_config = (
|
||||
"guardrails" in data or "guardrail_config" in data or "policies" in data
|
||||
)
|
||||
|
||||
# Only check policy engine if no local config (avoid import + registry lookup)
|
||||
if not (has_key_config or has_team_config or has_request_config):
|
||||
from litellm.proxy.policy_engine.policy_registry import get_policy_registry
|
||||
|
||||
if not get_policy_registry().is_initialized():
|
||||
# Nothing configured anywhere - clean up request body fields and return
|
||||
data.pop("policies", None)
|
||||
return
|
||||
|
||||
# Check key-level guardrails
|
||||
_add_guardrails_from_key_or_team_metadata(
|
||||
key_metadata=user_api_key_dict.metadata,
|
||||
|
|
|
|||
|
|
@ -15,6 +15,7 @@ from litellm.proxy.litellm_pre_call_utils import (
|
|||
LiteLLMProxyRequestSetup,
|
||||
_get_dynamic_logging_metadata,
|
||||
_get_enforced_params,
|
||||
_get_metadata_variable_name,
|
||||
_update_model_if_key_alias_exists,
|
||||
add_guardrails_from_policy_engine,
|
||||
add_litellm_data_to_request,
|
||||
|
|
@ -47,6 +48,47 @@ def test_check_if_token_is_service_account():
|
|||
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_get_enforced_params_for_service_account_settings():
|
||||
"""
|
||||
Test that service account enforced params are only added to service account keys
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue