mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-11 03:38:38 +00:00
fix(litellm_pre_call_utils.py): forward headers by model group at litellm pre call utils level
do it at the proxy level instead of router - allows reusing same forwarding logic as global forwarding
This commit is contained in:
parent
422447b7f1
commit
06d05c691d
6 changed files with 39 additions and 96 deletions
File diff suppressed because one or more lines are too long
|
|
@ -3,7 +3,7 @@ model_list:
|
|||
litellm_params:
|
||||
model: openai/fake
|
||||
api_key: fake-key
|
||||
api_base: https://exampleopenaiendpoint-production.up.railway.app/
|
||||
api_base: https://webhook.site/4feb0d46-4b23-468c-bf55-7008b5deb36d
|
||||
- model_name: gpt-5-mini
|
||||
litellm_params:
|
||||
model: azure/gpt-5-mini
|
||||
|
|
@ -20,3 +20,7 @@ litellm_settings:
|
|||
type: redis
|
||||
ttl: 600
|
||||
supported_call_types: ["acompletion", "completion"]
|
||||
|
||||
model_group_settings:
|
||||
forward_client_headers_to_llm_api:
|
||||
- fake-openai-endpoint
|
||||
|
|
|
|||
|
|
@ -384,6 +384,29 @@ class LiteLLMProxyRequestSetup:
|
|||
|
||||
return returned_headers
|
||||
|
||||
@staticmethod
|
||||
def add_headers_to_llm_call_by_model_group(
|
||||
data: dict, headers: dict, user_api_key_dict: UserAPIKeyAuth
|
||||
) -> dict:
|
||||
"""
|
||||
Add headers to the LLM call by model group
|
||||
"""
|
||||
data_model = data.get("model")
|
||||
if (
|
||||
data_model is not None
|
||||
and litellm.model_group_settings is not None
|
||||
and litellm.model_group_settings.forward_client_headers_to_llm_api
|
||||
is not None
|
||||
and data_model
|
||||
in litellm.model_group_settings.forward_client_headers_to_llm_api
|
||||
):
|
||||
_headers = LiteLLMProxyRequestSetup.add_headers_to_llm_call(
|
||||
headers, user_api_key_dict
|
||||
)
|
||||
if _headers != {}:
|
||||
data["headers"] = _headers
|
||||
return data
|
||||
|
||||
@staticmethod
|
||||
def add_litellm_data_for_backend_llm_call(
|
||||
*,
|
||||
|
|
@ -439,7 +462,7 @@ class LiteLLMProxyRequestSetup:
|
|||
user_api_key_request_route=user_api_key_dict.request_route,
|
||||
)
|
||||
return user_api_key_logged_metadata
|
||||
|
||||
|
||||
@staticmethod
|
||||
def add_user_api_key_auth_to_request_metadata(
|
||||
data: dict,
|
||||
|
|
@ -457,9 +480,7 @@ class LiteLLMProxyRequestSetup:
|
|||
data[_metadata_variable_name].update(user_api_key_logged_metadata)
|
||||
data[_metadata_variable_name][
|
||||
"user_api_key"
|
||||
] = (
|
||||
user_api_key_dict.api_key
|
||||
) # this is just the hashed token
|
||||
] = user_api_key_dict.api_key # this is just the hashed token
|
||||
|
||||
data[_metadata_variable_name]["user_api_end_user_max_budget"] = getattr(
|
||||
user_api_key_dict, "end_user_max_budget", None
|
||||
|
|
@ -624,6 +645,11 @@ async def add_litellm_data_to_request( # noqa: PLR0915
|
|||
)
|
||||
)
|
||||
|
||||
# check for forwardable headers
|
||||
data = LiteLLMProxyRequestSetup.add_headers_to_llm_call_by_model_group(
|
||||
data=data, headers=_headers, user_api_key_dict=user_api_key_dict
|
||||
)
|
||||
|
||||
# Parse user info from headers
|
||||
user = LiteLLMProxyRequestSetup.get_user_from_headers(_headers, general_settings)
|
||||
if user is not None:
|
||||
|
|
|
|||
|
|
@ -624,9 +624,7 @@ class Router:
|
|||
Apply the default settings to the router.
|
||||
"""
|
||||
|
||||
default_pre_call_checks: OptionalPreCallChecks = [
|
||||
"forward_client_headers_by_model_group",
|
||||
]
|
||||
default_pre_call_checks: OptionalPreCallChecks = []
|
||||
self.add_optional_pre_call_checks(default_pre_call_checks)
|
||||
return None
|
||||
|
||||
|
|
@ -892,8 +890,6 @@ class Router:
|
|||
)
|
||||
elif pre_call_check == "responses_api_deployment_check":
|
||||
_callback = ResponsesApiDeploymentCheck()
|
||||
elif pre_call_check == "forward_client_headers_by_model_group":
|
||||
_callback = ForwardClientSideHeadersByModelGroup()
|
||||
if _callback is not None:
|
||||
if self.optional_callbacks is None:
|
||||
self.optional_callbacks = []
|
||||
|
|
@ -4323,7 +4319,9 @@ class Router:
|
|||
"deployment", None
|
||||
) # stable name - works for wildcard routes as well
|
||||
# Get model_group and id from kwargs like the sync version does
|
||||
model_group = kwargs["litellm_params"]["metadata"].get("model_group", None)
|
||||
model_group = kwargs["litellm_params"]["metadata"].get(
|
||||
"model_group", None
|
||||
)
|
||||
model_info = kwargs["litellm_params"].get("model_info", {}) or {}
|
||||
id = model_info.get("id", None)
|
||||
if model_group is None or id is None:
|
||||
|
|
|
|||
|
|
@ -1,84 +0,0 @@
|
|||
from typing import Any, Dict, Optional, TypedDict
|
||||
|
||||
from litellm.types.utils import CallTypes
|
||||
|
||||
from ..integrations.custom_logger import CustomLogger
|
||||
|
||||
|
||||
class PotentialModelGroups(TypedDict):
|
||||
deployment_model_name: Optional[str]
|
||||
model_group_alias: Optional[str]
|
||||
|
||||
|
||||
class ForwardClientSideHeadersByModelGroup(CustomLogger):
|
||||
def get_potential_model_groups_from_kwargs(
|
||||
self, kwargs: Dict[str, Any]
|
||||
) -> Optional[PotentialModelGroups]:
|
||||
"""
|
||||
Get the model group from the kwargs.
|
||||
|
||||
Returns the potential model groups from the kwargs.
|
||||
- deployment_model_name (useful for wildcard model names)
|
||||
- model_group_alias (if the model is an alias)
|
||||
"""
|
||||
metadata = kwargs.get("litellm_metadata") or kwargs.get("metadata")
|
||||
if metadata is None:
|
||||
return None
|
||||
deployment_model_name = metadata.get("deployment_model_name", None)
|
||||
model_group_alias = metadata.get("model_group_alias", None)
|
||||
return {
|
||||
"deployment_model_name": deployment_model_name,
|
||||
"model_group_alias": model_group_alias,
|
||||
}
|
||||
|
||||
def filter_headers(self, headers: Dict[str, Any]) -> Dict[str, Any]:
|
||||
"""
|
||||
Filter the headers to only include the headers that are forwarded to the LLM API.
|
||||
|
||||
E.g. passing 'connection': 'keep-alive' will cause the request to hang, and not be acknowledged on the other side.
|
||||
"""
|
||||
return {
|
||||
k: v
|
||||
for k, v in headers.items()
|
||||
if k.lower() not in ["connection", "content-length"]
|
||||
}
|
||||
|
||||
async def async_pre_call_deployment_hook(
|
||||
self, kwargs: Dict[str, Any], call_type: Optional[CallTypes]
|
||||
) -> Optional[dict]:
|
||||
"""
|
||||
if kwargs["proxy_server_request"]["headers"] is not None:
|
||||
and kwargs["forward_client_headers_to_llm_api"] is not None:
|
||||
|
||||
add the headers to the request
|
||||
kwargs["headers"].update(kwargs["proxy_server_request"]["headers"])
|
||||
"""
|
||||
import litellm
|
||||
|
||||
if litellm.model_group_settings is None:
|
||||
return None
|
||||
|
||||
potential_model_groups = self.get_potential_model_groups_from_kwargs(kwargs)
|
||||
|
||||
if potential_model_groups is None:
|
||||
return None
|
||||
|
||||
if (
|
||||
"secret_fields" in kwargs
|
||||
and kwargs["secret_fields"]["raw_headers"] is not None
|
||||
and isinstance(kwargs["secret_fields"]["raw_headers"], dict)
|
||||
):
|
||||
for model_group in potential_model_groups.values():
|
||||
if model_group is None:
|
||||
continue
|
||||
if (
|
||||
litellm.model_group_settings.forward_client_headers_to_llm_api
|
||||
is not None
|
||||
and model_group
|
||||
in litellm.model_group_settings.forward_client_headers_to_llm_api
|
||||
):
|
||||
kwargs.setdefault("headers", {}).update(
|
||||
self.filter_headers(kwargs["secret_fields"]["raw_headers"])
|
||||
)
|
||||
|
||||
return kwargs
|
||||
Loading…
Add table
Reference in a new issue