mirror of
https://github.com/tiennm99/litellm.git
synced 2026-08-14 12:26:25 +00:00
Merge pull request #13741 from BerriAI/litellm_dev_08_18_2025_p1
Refactor - forward model group headers - reuse same logic as global header forwarding
This commit is contained in:
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:
|
||||
|
||||
+4
-9
@@ -93,9 +93,6 @@ from litellm.router_utils.fallback_event_handlers import (
|
||||
get_fallback_model_group,
|
||||
run_async_fallback,
|
||||
)
|
||||
from litellm.router_utils.forward_clientside_headers_by_model_group import (
|
||||
ForwardClientSideHeadersByModelGroup,
|
||||
)
|
||||
from litellm.router_utils.get_retry_from_policy import (
|
||||
get_num_retries_from_retry_policy as _get_num_retries_from_retry_policy,
|
||||
)
|
||||
@@ -624,9 +621,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 +887,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 +4316,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
|
||||
@@ -606,7 +606,6 @@ def test_get_dynamic_logging_metadata_with_arize_team_logging():
|
||||
assert result.callback_vars["arize_space_id"] == "test_arize_space_id"
|
||||
|
||||
|
||||
|
||||
def test_get_num_retries_from_request():
|
||||
"""
|
||||
Test LiteLLMProxyRequestSetup._get_num_retries_from_request method
|
||||
@@ -668,6 +667,7 @@ def test_get_num_retries_from_request():
|
||||
)
|
||||
assert result == -1
|
||||
|
||||
|
||||
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
|
||||
@@ -676,9 +676,9 @@ def test_add_user_api_key_auth_to_request_metadata():
|
||||
data = {
|
||||
"model": "gpt-3.5-turbo",
|
||||
"messages": [{"role": "user", "content": "Hello"}],
|
||||
"litellm_metadata": {} # This will be the metadata variable name
|
||||
"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",
|
||||
@@ -689,21 +689,21 @@ def test_add_user_api_key_auth_to_request_metadata():
|
||||
team_alias="test-team-alias",
|
||||
end_user_id="test-end-user-123",
|
||||
request_route="/chat/completions",
|
||||
end_user_max_budget=500.0
|
||||
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
|
||||
_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"
|
||||
@@ -714,13 +714,224 @@ def test_add_user_api_key_auth_to_request_metadata():
|
||||
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"}]
|
||||
assert result["messages"] == [{"role": "user", "content": "Hello"}]
|
||||
|
||||
|
||||
@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
|
||||
|
||||
@@ -896,148 +896,6 @@ async def test_router_ageneric_api_call_with_fallbacks_helper():
|
||||
assert router.fail_calls["gpt-3.5-turbo"] == initial_fail_count + 1
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_router_forward_client_headers_by_model_group():
|
||||
"""
|
||||
Test that router.forward_client_headers_by_model_group returns the correct response
|
||||
"""
|
||||
from unittest.mock import MagicMock, patch
|
||||
|
||||
from litellm.types.router import ModelGroupSettings
|
||||
|
||||
litellm.model_group_settings = ModelGroupSettings(
|
||||
forward_client_headers_to_llm_api=[
|
||||
"gpt-3.5-turbo-allow",
|
||||
"openai/*",
|
||||
"gpt-3.5-turbo-custom",
|
||||
]
|
||||
)
|
||||
|
||||
router = litellm.Router(
|
||||
model_list=[
|
||||
{
|
||||
"model_name": "gpt-3.5-turbo-allow",
|
||||
"litellm_params": {
|
||||
"model": "gpt-3.5-turbo",
|
||||
},
|
||||
},
|
||||
{
|
||||
"model_name": "gpt-3.5-turbo-disallow",
|
||||
"litellm_params": {
|
||||
"model": "gpt-3.5-turbo",
|
||||
},
|
||||
},
|
||||
{
|
||||
"model_name": "openai/*",
|
||||
"litellm_params": {
|
||||
"model": "openai/*",
|
||||
},
|
||||
},
|
||||
{
|
||||
"model_name": "openai/gpt-4o-mini",
|
||||
"litellm_params": {
|
||||
"model": "openai/gpt-4o-mini",
|
||||
},
|
||||
},
|
||||
],
|
||||
model_group_alias={
|
||||
"gpt-3.5-turbo-custom": "gpt-3.5-turbo-disallow",
|
||||
},
|
||||
)
|
||||
|
||||
## Scenario 1: Direct model name
|
||||
with patch.object(
|
||||
litellm.main, "completion", return_value=MagicMock()
|
||||
) as mock_completion:
|
||||
await router.acompletion(
|
||||
model="gpt-3.5-turbo-allow",
|
||||
messages=[{"role": "user", "content": "Hello, world!"}],
|
||||
mock_response="Hello, world!",
|
||||
secret_fields={"raw_headers": {"test": "test"}},
|
||||
)
|
||||
|
||||
mock_completion.assert_called_once()
|
||||
print(mock_completion.call_args.kwargs["headers"])
|
||||
|
||||
## Scenario 2: Wildcard model name
|
||||
with patch.object(
|
||||
litellm.main, "completion", return_value=MagicMock()
|
||||
) as mock_completion:
|
||||
await router.acompletion(
|
||||
model="openai/gpt-3.5-turbo",
|
||||
messages=[{"role": "user", "content": "Hello, world!"}],
|
||||
mock_response="Hello, world!",
|
||||
secret_fields={"raw_headers": {"test": "test"}},
|
||||
)
|
||||
|
||||
mock_completion.assert_called_once()
|
||||
print(mock_completion.call_args.kwargs["headers"])
|
||||
|
||||
## Scenario 3: Not in model_group_settings
|
||||
with patch.object(
|
||||
litellm.main, "completion", return_value=MagicMock()
|
||||
) as mock_completion:
|
||||
await router.acompletion(
|
||||
model="openai/gpt-4o-mini",
|
||||
messages=[{"role": "user", "content": "Hello, world!"}],
|
||||
mock_response="Hello, world!",
|
||||
secret_fields={"raw_headers": {"test": "test"}},
|
||||
)
|
||||
|
||||
mock_completion.assert_called_once()
|
||||
assert mock_completion.call_args.kwargs.get("headers") is None
|
||||
|
||||
## Scenario 4: Model group alias
|
||||
with patch.object(
|
||||
litellm.main, "completion", return_value=MagicMock()
|
||||
) as mock_completion:
|
||||
await router.acompletion(
|
||||
model="gpt-3.5-turbo-custom",
|
||||
messages=[{"role": "user", "content": "Hello, world!"}],
|
||||
mock_response="Hello, world!",
|
||||
secret_fields={"raw_headers": {"test": "test"}},
|
||||
)
|
||||
|
||||
mock_completion.assert_called_once()
|
||||
print(mock_completion.call_args.kwargs["headers"])
|
||||
|
||||
|
||||
def test_router_apply_default_settings():
|
||||
"""
|
||||
Test that Router.apply_default_settings() adds the expected default pre-call checks
|
||||
"""
|
||||
router = litellm.Router(
|
||||
model_list=[
|
||||
{
|
||||
"model_name": "gpt-3.5-turbo",
|
||||
"litellm_params": {"model": "gpt-3.5-turbo"},
|
||||
}
|
||||
],
|
||||
)
|
||||
|
||||
# Apply default settings
|
||||
result = router.apply_default_settings()
|
||||
|
||||
# Verify the method returns None
|
||||
assert result is None
|
||||
|
||||
# Verify that the forward_client_headers_by_model_group pre-call check was added
|
||||
# Check if any callback is of the ForwardClientHeadersByModelGroupCheck type
|
||||
has_forward_headers_check = False
|
||||
for callback in litellm.callbacks:
|
||||
print(callback)
|
||||
print(f"callback.__class__: {callback.__class__}")
|
||||
if hasattr(
|
||||
callback, "__class__"
|
||||
) and "ForwardClientSideHeadersByModelGroup" in str(callback.__class__):
|
||||
has_forward_headers_check = True
|
||||
break
|
||||
|
||||
assert (
|
||||
has_forward_headers_check
|
||||
), "Expected ForwardClientSideHeadersByModelGroup to be added to callbacks"
|
||||
|
||||
|
||||
def test_router_get_model_access_groups_team_only_models():
|
||||
"""
|
||||
Test that Router.get_model_access_groups returns the correct response for team-only models
|
||||
|
||||
Reference in New Issue
Block a user