mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-07 02:59:05 +00:00
Merge pull request #39249 from BerriAI/litellm_router_settings_reject_unknown_keys
fix: apply optional_pre_call_checks and reject unsupported router settings on /config/update
This commit is contained in:
commit
a43228ef72
8 changed files with 388 additions and 48 deletions
|
|
@ -9,6 +9,38 @@ DEFAULT_HEALTH_CHECK_PROMPT: Final = str(os.getenv("DEFAULT_HEALTH_CHECK_PROMPT"
|
|||
AZURE_DEFAULT_RESPONSES_API_VERSION: Final = str(os.getenv("AZURE_DEFAULT_RESPONSES_API_VERSION", "preview"))
|
||||
ROUTER_MAX_FALLBACKS: Final = int(os.getenv("ROUTER_MAX_FALLBACKS", 5))
|
||||
ROUTER_FALLBACK_ERROR_DETAIL_MAX_CHARS: Final = 2000
|
||||
RUNTIME_UPDATABLE_ROUTER_SETTINGS: Final[frozenset[str]] = frozenset(
|
||||
{
|
||||
"routing_strategy_args",
|
||||
"routing_strategy",
|
||||
"routing_groups",
|
||||
"allowed_fails",
|
||||
"cooldown_time",
|
||||
"num_retries",
|
||||
"timeout",
|
||||
"max_retries",
|
||||
"retry_after",
|
||||
"fallbacks",
|
||||
"context_window_fallbacks",
|
||||
"retry_policy",
|
||||
"model_group_retry_policy",
|
||||
"model_group_alias",
|
||||
"enable_weighted_failover",
|
||||
"enable_tag_filtering",
|
||||
"tag_routing_prefix",
|
||||
"optional_pre_call_checks",
|
||||
}
|
||||
)
|
||||
ROUTER_SETTINGS_MANAGED_OUTSIDE_CONFIG: Final[frozenset[str]] = frozenset(
|
||||
{
|
||||
"model_list",
|
||||
"search_tools",
|
||||
"assistants_config",
|
||||
"router_general_settings",
|
||||
"ignore_invalid_deployments",
|
||||
"fallback_access_check",
|
||||
}
|
||||
)
|
||||
DEFAULT_BATCH_SIZE: Final = int(os.getenv("DEFAULT_BATCH_SIZE", 512))
|
||||
DEFAULT_FLUSH_INTERVAL_SECONDS: Final = int(os.getenv("DEFAULT_FLUSH_INTERVAL_SECONDS", 5))
|
||||
DEFAULT_S3_FLUSH_INTERVAL_SECONDS: Final = int(os.getenv("DEFAULT_S3_FLUSH_INTERVAL_SECONDS", 10))
|
||||
|
|
|
|||
|
|
@ -39,7 +39,7 @@ from typing import (
|
|||
import anyio
|
||||
import websockets
|
||||
import websockets.exceptions
|
||||
from pydantic import BaseModel, Json, JsonValue, ValidationError
|
||||
from pydantic import BaseModel, Json, JsonValue, TypeAdapter, ValidationError
|
||||
from typing_extensions import NotRequired, ReadOnly, assert_never
|
||||
|
||||
from litellm._uuid import uuid
|
||||
|
|
@ -60,6 +60,7 @@ from litellm.constants import (
|
|||
LITELLM_SETTINGS_SAFE_DB_OVERRIDES,
|
||||
LITELLM_UI_ALLOW_HEADERS,
|
||||
LITELLM_UI_SESSION_DURATION,
|
||||
RUNTIME_UPDATABLE_ROUTER_SETTINGS,
|
||||
)
|
||||
from litellm.litellm_core_utils.litellm_logging import (
|
||||
_init_custom_logger_compatible_class,
|
||||
|
|
@ -253,6 +254,7 @@ from litellm.constants import (
|
|||
PROXY_BUDGET_RESCHEDULER_MAX_TIME,
|
||||
PROXY_BUDGET_RESCHEDULER_MIN_TIME,
|
||||
PROXY_CONFIG_RELOAD_INTERVAL_SECONDS,
|
||||
ROUTER_SETTINGS_MANAGED_OUTSIDE_CONFIG,
|
||||
USER_SPEND_ALERTS_JOB_ID,
|
||||
WEEKLY_SPEND_REPORT_JOB_ID,
|
||||
)
|
||||
|
|
@ -5713,13 +5715,9 @@ class ProxyConfig:
|
|||
router_settings: Final = config.get("router_settings", None)
|
||||
|
||||
if router_settings and isinstance(router_settings, dict):
|
||||
# model list and search_tools already set
|
||||
exclude_args: Final = {
|
||||
"model_list",
|
||||
"search_tools",
|
||||
}
|
||||
|
||||
available_args: Final = [x for x in litellm.Router.get_valid_args() if x not in exclude_args]
|
||||
available_args: Final = [
|
||||
x for x in litellm.Router.get_valid_args() if x not in ROUTER_SETTINGS_MANAGED_OUTSIDE_CONFIG
|
||||
]
|
||||
|
||||
for k, v in router_settings.items():
|
||||
if k in available_args:
|
||||
|
|
@ -16218,6 +16216,7 @@ async def invitation_delete(
|
|||
)
|
||||
async def update_config(
|
||||
config_info: ConfigYAML,
|
||||
request: Request,
|
||||
user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth),
|
||||
):
|
||||
"""
|
||||
|
|
@ -16233,6 +16232,26 @@ async def update_config(
|
|||
if user_api_key_dict.user_role != LitellmUserRoles.PROXY_ADMIN:
|
||||
raise HTTPException(status_code=403, detail="Only proxy admins can update config")
|
||||
|
||||
request_body: Final[Mapping[str, JsonValue]] = TypeAdapter(Mapping[str, JsonValue]).validate_python(
|
||||
await request.json()
|
||||
)
|
||||
raw_router_settings: Final = request_body.get("router_settings")
|
||||
if isinstance(raw_router_settings, dict):
|
||||
supported_router_settings: Final = RUNTIME_UPDATABLE_ROUTER_SETTINGS | (
|
||||
frozenset(litellm.Router.get_valid_args()) - ROUTER_SETTINGS_MANAGED_OUTSIDE_CONFIG
|
||||
)
|
||||
unsupported_router_settings: Final = sorted(set(raw_router_settings) - supported_router_settings)
|
||||
if unsupported_router_settings:
|
||||
raise HTTPException(
|
||||
status_code=400,
|
||||
detail={
|
||||
"error": (
|
||||
f"Unsupported router settings: {', '.join(unsupported_router_settings)} "
|
||||
"are not valid router settings"
|
||||
)
|
||||
},
|
||||
)
|
||||
|
||||
if prisma_client is None:
|
||||
raise Exception("No DB Connected")
|
||||
|
||||
|
|
@ -16334,11 +16353,19 @@ async def update_config(
|
|||
)
|
||||
|
||||
# router_settings: merge existing + request, request wins.
|
||||
if config_info.router_settings is not None:
|
||||
if isinstance(raw_router_settings, dict):
|
||||
existing = await _read_section("router_settings")
|
||||
before_router_settings: Final = copy.deepcopy(existing)
|
||||
updates = config_info.router_settings.dict(exclude_none=True)
|
||||
new_router_settings: Final = {**existing, **updates}
|
||||
typed_router_settings: Final = (
|
||||
config_info.router_settings.dict(exclude_none=True) if config_info.router_settings is not None else {}
|
||||
)
|
||||
raw_router_settings_without_none: Final = {
|
||||
key: value
|
||||
for key, value in raw_router_settings.items()
|
||||
if key not in typed_router_settings and value is not None
|
||||
}
|
||||
router_settings_updates: Final = {**typed_router_settings, **raw_router_settings_without_none}
|
||||
new_router_settings: Final = {**existing, **router_settings_updates}
|
||||
await _upsert_section("router_settings", new_router_settings)
|
||||
asyncio.create_task(
|
||||
create_config_audit_log(
|
||||
|
|
|
|||
|
|
@ -50,6 +50,7 @@ from litellm.constants import (
|
|||
DEFAULT_HEALTH_CHECK_INTERVAL,
|
||||
DEFAULT_HEALTH_CHECK_STALENESS_MULTIPLIER,
|
||||
DEFAULT_MAX_LRU_CACHE_SIZE,
|
||||
RUNTIME_UPDATABLE_ROUTER_SETTINGS,
|
||||
SESSION_DEPLOYMENT_AFFINITY_TTL_METADATA_KEY,
|
||||
)
|
||||
from litellm.integrations.custom_logger import CustomLogger
|
||||
|
|
@ -354,6 +355,13 @@ _PreRoutingStrategyT = TypeVar("_PreRoutingStrategyT")
|
|||
_ALIAS_PARAMS_NEVER_FORWARDED: Final = frozenset({"model", "api_base", "api_key", "api_version"})
|
||||
_ALIAS_MARKER_FORWARDED_PARAMS_KWARG: Final = "_alias_marker_forwarded_params"
|
||||
|
||||
_RUNTIME_TOGGLEABLE_PRE_CALL_CHECKS: Final[Mapping[str, type[CustomLogger]]] = MappingProxyType(
|
||||
{
|
||||
"prompt_caching": PromptCachingDeploymentCheck,
|
||||
"enforce_model_rate_limits": ModelRateLimitingCheck,
|
||||
}
|
||||
)
|
||||
|
||||
|
||||
def _stream_chunks_have_generated_content(chunks: Sequence[ModelResponseStream]) -> bool:
|
||||
for chunk in chunks:
|
||||
|
|
@ -2072,11 +2080,39 @@ class Router:
|
|||
if _callback is None:
|
||||
continue
|
||||
|
||||
if self.optional_callbacks is not None and any(
|
||||
isinstance(callback, type(_callback)) for callback in self.optional_callbacks
|
||||
):
|
||||
continue
|
||||
if self.optional_callbacks is None:
|
||||
self.optional_callbacks = []
|
||||
self.optional_callbacks.append(_callback)
|
||||
litellm.logging_callback_manager.add_litellm_callback(_callback)
|
||||
|
||||
def set_optional_pre_call_checks(self, optional_pre_call_checks: OptionalPreCallChecks | None) -> None:
|
||||
if optional_pre_call_checks is None:
|
||||
return
|
||||
requested: Final = frozenset(optional_pre_call_checks)
|
||||
for name, callback_cls in _RUNTIME_TOGGLEABLE_PRE_CALL_CHECKS.items():
|
||||
if name not in requested:
|
||||
self._remove_optional_callbacks_of_type(callback_cls)
|
||||
self.add_optional_pre_call_checks(optional_pre_call_checks)
|
||||
|
||||
def _remove_optional_callbacks_of_type(self, callback_cls: type[CustomLogger]) -> None:
|
||||
if self.optional_callbacks is None or not any(type(cb) is callback_cls for cb in self.optional_callbacks):
|
||||
return
|
||||
self.optional_callbacks = [cb for cb in self.optional_callbacks if type(cb) is not callback_cls]
|
||||
if any(
|
||||
router is not self and any(type(cb) is callback_cls for cb in (router.optional_callbacks or []))
|
||||
for router in tuple(_live_routers)
|
||||
):
|
||||
return
|
||||
for cb in tuple(litellm.callbacks):
|
||||
if type(cb) is callback_cls:
|
||||
litellm.logging_callback_manager.remove_callback_from_list_by_object(
|
||||
litellm.callbacks, cb, require_self=False
|
||||
)
|
||||
|
||||
def print_deployment(self, deployment: dict):
|
||||
"""
|
||||
returns a copy of the deployment with the api key masked
|
||||
|
|
@ -11351,27 +11387,6 @@ class Router:
|
|||
"""
|
||||
Update the router settings.
|
||||
"""
|
||||
# only the following settings are allowed to be configured
|
||||
_allowed_settings: Final = [
|
||||
"routing_strategy_args",
|
||||
"routing_strategy",
|
||||
"routing_groups",
|
||||
"allowed_fails",
|
||||
"cooldown_time",
|
||||
"num_retries",
|
||||
"timeout",
|
||||
"max_retries",
|
||||
"retry_after",
|
||||
"fallbacks",
|
||||
"context_window_fallbacks",
|
||||
"retry_policy",
|
||||
"model_group_retry_policy",
|
||||
"model_group_alias",
|
||||
"enable_weighted_failover",
|
||||
"enable_tag_filtering",
|
||||
"tag_routing_prefix",
|
||||
]
|
||||
|
||||
_int_settings: Final = [
|
||||
"timeout",
|
||||
"num_retries",
|
||||
|
|
@ -11384,13 +11399,15 @@ class Router:
|
|||
rebuild_routing_groups = False
|
||||
relink_lar1_from_args = False
|
||||
for var in kwargs:
|
||||
if var in _allowed_settings:
|
||||
if var in RUNTIME_UPDATABLE_ROUTER_SETTINGS:
|
||||
if var in _int_settings:
|
||||
_casted_value = int(kwargs[var])
|
||||
setattr(self, var, _casted_value)
|
||||
elif var == "routing_groups":
|
||||
self._routing_groups_input = kwargs[var]
|
||||
rebuild_routing_groups = True
|
||||
elif var == "optional_pre_call_checks":
|
||||
self.set_optional_pre_call_checks(kwargs[var])
|
||||
elif var == "retry_policy":
|
||||
value = kwargs[var]
|
||||
if isinstance(value, dict):
|
||||
|
|
|
|||
|
|
@ -106,6 +106,20 @@ class RetryPolicy(BaseModel):
|
|||
InternalServerErrorRetries: int | None = None
|
||||
|
||||
|
||||
OptionalPreCallChecks = list[
|
||||
Literal[
|
||||
"prompt_caching",
|
||||
"router_budget_limiting",
|
||||
"responses_api_deployment_check",
|
||||
"deployment_affinity",
|
||||
"session_affinity",
|
||||
"forward_client_headers_by_model_group",
|
||||
"enforce_model_rate_limits",
|
||||
"encrypted_content_affinity",
|
||||
]
|
||||
]
|
||||
|
||||
|
||||
class UpdateRouterConfig(BaseModel):
|
||||
"""
|
||||
Set of params that you can modify via `router.update_settings()`.
|
||||
|
|
@ -128,6 +142,7 @@ class UpdateRouterConfig(BaseModel):
|
|||
model_group_alias: dict[str, str | dict] | None = {}
|
||||
enable_tag_filtering: bool | None = None
|
||||
tag_routing_prefix: str | None = None
|
||||
optional_pre_call_checks: OptionalPreCallChecks | None = None
|
||||
|
||||
model_config = ConfigDict(protected_namespaces=())
|
||||
|
||||
|
|
@ -869,20 +884,6 @@ class FallbackAccessCheck(Protocol):
|
|||
async def __call__(self, *, model: str, request_kwargs: Mapping[str, object], llm_router: "Router") -> bool: ...
|
||||
|
||||
|
||||
OptionalPreCallChecks = list[
|
||||
Literal[
|
||||
"prompt_caching",
|
||||
"router_budget_limiting",
|
||||
"responses_api_deployment_check",
|
||||
"deployment_affinity",
|
||||
"session_affinity",
|
||||
"forward_client_headers_by_model_group",
|
||||
"enforce_model_rate_limits",
|
||||
"encrypted_content_affinity",
|
||||
]
|
||||
]
|
||||
|
||||
|
||||
class LiteLLM_RouterFileObject(TypedDict, total=False):
|
||||
"""
|
||||
Tracking the litellm params hash, used for mapping the file id to the right model
|
||||
|
|
|
|||
|
|
@ -3076,7 +3076,9 @@ async def test_update_config_success_callback_normalization():
|
|||
admin_user = UserAPIKeyAuth(
|
||||
user_role=LitellmUserRoles.PROXY_ADMIN, api_key="sk-test"
|
||||
)
|
||||
await proxy_server.update_config(config_update, user_api_key_dict=admin_user)
|
||||
request = MagicMock()
|
||||
request.json = AsyncMock(return_value={"litellm_settings": {"success_callback": ["SQS", "sQs"]}})
|
||||
await proxy_server.update_config(config_update, request=request, user_api_key_dict=admin_user)
|
||||
|
||||
assert (
|
||||
"litellm_settings" in upserted
|
||||
|
|
|
|||
|
|
@ -60,6 +60,141 @@ def test_config_update_happy_admin(client, auth_as, mock_prisma, monkeypatch):
|
|||
assert normalize(response.json()) == {"message": "Config updated successfully"}
|
||||
|
||||
|
||||
def test_config_update_persists_optional_pre_call_checks(client, auth_as, mock_prisma, monkeypatch):
|
||||
from litellm.proxy import proxy_server as ps
|
||||
from litellm.proxy._types import LitellmUserRoles
|
||||
|
||||
table = _install_litellm_config(mock_prisma)
|
||||
monkeypatch.setattr(ps, "prisma_client", mock_prisma)
|
||||
fake_proxy_config = MagicMock()
|
||||
fake_proxy_config.add_deployment = AsyncMock()
|
||||
monkeypatch.setattr(ps, "proxy_config", fake_proxy_config)
|
||||
|
||||
with auth_as(LitellmUserRoles.PROXY_ADMIN):
|
||||
response = client.post(
|
||||
"/config/update",
|
||||
json={"router_settings": {"optional_pre_call_checks": ["prompt_caching"]}},
|
||||
)
|
||||
|
||||
assert response.status_code == 200
|
||||
persisted = json.loads(table.upsert.call_args.kwargs["data"]["create"]["param_value"])
|
||||
assert persisted["optional_pre_call_checks"] == ["prompt_caching"]
|
||||
|
||||
|
||||
def test_config_update_persists_model_group_affinity_config(client, auth_as, mock_prisma, monkeypatch):
|
||||
from litellm.proxy import proxy_server as ps
|
||||
from litellm.proxy._types import LitellmUserRoles
|
||||
|
||||
table = _install_litellm_config(mock_prisma)
|
||||
monkeypatch.setattr(ps, "prisma_client", mock_prisma)
|
||||
fake_proxy_config = MagicMock()
|
||||
fake_proxy_config.add_deployment = AsyncMock()
|
||||
monkeypatch.setattr(ps, "proxy_config", fake_proxy_config)
|
||||
|
||||
model_group_affinity_config = {"gpt-4": ["session_affinity"]}
|
||||
with auth_as(LitellmUserRoles.PROXY_ADMIN):
|
||||
response = client.post(
|
||||
"/config/update",
|
||||
json={"router_settings": {"model_group_affinity_config": model_group_affinity_config}},
|
||||
)
|
||||
|
||||
assert response.status_code == 200
|
||||
persisted = json.loads(table.upsert.call_args.kwargs["data"]["create"]["param_value"])
|
||||
assert persisted["model_group_affinity_config"] == model_group_affinity_config
|
||||
|
||||
|
||||
def test_config_update_persists_disable_cooldowns(client, auth_as, mock_prisma, monkeypatch):
|
||||
from litellm.proxy import proxy_server as ps
|
||||
from litellm.proxy._types import LitellmUserRoles
|
||||
|
||||
table = _install_litellm_config(mock_prisma)
|
||||
monkeypatch.setattr(ps, "prisma_client", mock_prisma)
|
||||
fake_proxy_config = MagicMock()
|
||||
fake_proxy_config.add_deployment = AsyncMock()
|
||||
monkeypatch.setattr(ps, "proxy_config", fake_proxy_config)
|
||||
|
||||
with auth_as(LitellmUserRoles.PROXY_ADMIN):
|
||||
response = client.post(
|
||||
"/config/update",
|
||||
json={"router_settings": {"disable_cooldowns": True}},
|
||||
)
|
||||
|
||||
assert response.status_code == 200
|
||||
persisted = json.loads(table.upsert.call_args.kwargs["data"]["create"]["param_value"])
|
||||
assert persisted["disable_cooldowns"] is True
|
||||
|
||||
|
||||
def test_config_update_rejects_assistants_config(client, auth_as, mock_prisma, monkeypatch):
|
||||
from litellm.proxy import proxy_server as ps
|
||||
from litellm.proxy._types import LitellmUserRoles
|
||||
|
||||
table = _install_litellm_config(mock_prisma)
|
||||
monkeypatch.setattr(ps, "prisma_client", mock_prisma)
|
||||
|
||||
with auth_as(LitellmUserRoles.PROXY_ADMIN):
|
||||
response = client.post(
|
||||
"/config/update",
|
||||
json={"router_settings": {"assistants_config": {"enabled": True}}},
|
||||
)
|
||||
|
||||
assert response.status_code == 400
|
||||
assert "assistants_config" in response.json()["error"]["message"]
|
||||
table.upsert.assert_not_called()
|
||||
|
||||
|
||||
def test_config_update_rejects_router_general_settings(client, auth_as, mock_prisma, monkeypatch):
|
||||
from litellm.proxy import proxy_server as ps
|
||||
from litellm.proxy._types import LitellmUserRoles
|
||||
|
||||
table = _install_litellm_config(mock_prisma)
|
||||
monkeypatch.setattr(ps, "prisma_client", mock_prisma)
|
||||
|
||||
with auth_as(LitellmUserRoles.PROXY_ADMIN):
|
||||
response = client.post(
|
||||
"/config/update",
|
||||
json={"router_settings": {"router_general_settings": {"async_only_mode": True}}},
|
||||
)
|
||||
|
||||
assert response.status_code == 400
|
||||
assert "router_general_settings" in response.json()["error"]["message"]
|
||||
table.upsert.assert_not_called()
|
||||
|
||||
|
||||
def test_config_update_rejects_unknown_router_setting(client, auth_as, mock_prisma, monkeypatch):
|
||||
from litellm.proxy import proxy_server as ps
|
||||
from litellm.proxy._types import LitellmUserRoles
|
||||
|
||||
table = _install_litellm_config(mock_prisma)
|
||||
monkeypatch.setattr(ps, "prisma_client", mock_prisma)
|
||||
|
||||
with auth_as(LitellmUserRoles.PROXY_ADMIN):
|
||||
response = client.post(
|
||||
"/config/update",
|
||||
json={"router_settings": {"optional_precall_checks": ["prompt_caching"]}},
|
||||
)
|
||||
|
||||
assert response.status_code == 400
|
||||
assert "optional_precall_checks" in response.json()["error"]["message"]
|
||||
table.upsert.assert_not_called()
|
||||
|
||||
|
||||
def test_config_update_unknown_router_setting_non_admin_forbidden(client, auth_as, mock_prisma, monkeypatch):
|
||||
from litellm.proxy import proxy_server as ps
|
||||
from litellm.proxy._types import LitellmUserRoles
|
||||
|
||||
_install_litellm_config(mock_prisma)
|
||||
monkeypatch.setattr(ps, "prisma_client", mock_prisma)
|
||||
|
||||
with auth_as(LitellmUserRoles.INTERNAL_USER):
|
||||
response = client.post(
|
||||
"/config/update",
|
||||
json={"router_settings": {"optional_precall_checks": ["prompt_caching"]}},
|
||||
)
|
||||
|
||||
assert response.status_code == 403
|
||||
assert "admin" in response.json()["error"]["message"].lower()
|
||||
|
||||
|
||||
def test_config_update_non_admin_forbidden(client, auth_as, mock_prisma, monkeypatch):
|
||||
"""POST /config/update by a non-admin caller is rejected; the error
|
||||
surfaces as a ProxyException with the admin-only message."""
|
||||
|
|
|
|||
|
|
@ -21,6 +21,7 @@ This file pins both halves of the fix.
|
|||
|
||||
import json
|
||||
from dataclasses import dataclass
|
||||
from typing import Final
|
||||
from unittest.mock import AsyncMock, MagicMock
|
||||
|
||||
import pytest
|
||||
|
|
@ -28,8 +29,19 @@ from pydantic import ValidationError
|
|||
|
||||
|
||||
import litellm
|
||||
from litellm.router_strategy.budget_limiter import RouterBudgetLimiting
|
||||
from litellm.router_utils.pre_call_checks.model_rate_limit_check import ModelRateLimitingCheck
|
||||
from litellm.router_utils.pre_call_checks.prompt_caching_deployment_check import PromptCachingDeploymentCheck
|
||||
from litellm.types.router import RetryPolicy, UpdateRouterConfig
|
||||
|
||||
|
||||
@pytest.fixture(autouse=True)
|
||||
def isolate_litellm_callbacks():
|
||||
callbacks_before: Final = litellm.callbacks.copy()
|
||||
yield
|
||||
litellm.callbacks = callbacks_before # test-quality-ok: required callback-state restoration fixture
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# UpdateRouterConfig schema membership (LIT-3152 part 1)
|
||||
# ---------------------------------------------------------------------------
|
||||
|
|
@ -100,6 +112,114 @@ def _build_router() -> litellm.Router:
|
|||
)
|
||||
|
||||
|
||||
def test_update_settings_adds_optional_pre_call_check_once():
|
||||
router = _build_router()
|
||||
|
||||
router.update_settings(num_retries=7, optional_pre_call_checks=["prompt_caching"])
|
||||
router.update_settings(optional_pre_call_checks=["prompt_caching"])
|
||||
|
||||
prompt_caching_callbacks = [
|
||||
callback for callback in router.optional_callbacks if isinstance(callback, PromptCachingDeploymentCheck)
|
||||
]
|
||||
assert len(prompt_caching_callbacks) == 1
|
||||
assert router.num_retries == 7
|
||||
|
||||
|
||||
def test_update_settings_clears_omitted_toggleable_pre_call_checks():
|
||||
router = _build_router()
|
||||
|
||||
router.update_settings(optional_pre_call_checks=["prompt_caching"])
|
||||
router.update_settings(optional_pre_call_checks=[])
|
||||
|
||||
assert not any(isinstance(callback, PromptCachingDeploymentCheck) for callback in (router.optional_callbacks or []))
|
||||
assert not any(isinstance(callback, PromptCachingDeploymentCheck) for callback in litellm.callbacks)
|
||||
|
||||
|
||||
def test_set_optional_pre_call_checks_reconciles_callback_types():
|
||||
router = _build_router()
|
||||
|
||||
router.set_optional_pre_call_checks(["prompt_caching"])
|
||||
router.set_optional_pre_call_checks([])
|
||||
|
||||
assert not any(isinstance(callback, PromptCachingDeploymentCheck) for callback in (router.optional_callbacks or []))
|
||||
assert not any(isinstance(callback, PromptCachingDeploymentCheck) for callback in litellm.callbacks)
|
||||
|
||||
|
||||
def test_remove_optional_pre_call_check_removes_local_and_global_callbacks():
|
||||
router = _build_router()
|
||||
|
||||
router.set_optional_pre_call_checks(["prompt_caching"])
|
||||
router._remove_optional_callbacks_of_type(PromptCachingDeploymentCheck)
|
||||
|
||||
assert not any(type(callback) is PromptCachingDeploymentCheck for callback in (router.optional_callbacks or []))
|
||||
assert not any(type(callback) is PromptCachingDeploymentCheck for callback in litellm.callbacks)
|
||||
|
||||
|
||||
def test_remove_optional_pre_call_check_keeps_global_callback_for_another_router():
|
||||
router_a = _build_router()
|
||||
router_b = _build_router()
|
||||
|
||||
router_a.update_settings(optional_pre_call_checks=["prompt_caching"])
|
||||
router_b.update_settings(optional_pre_call_checks=["prompt_caching"])
|
||||
|
||||
router_a.update_settings(optional_pre_call_checks=[])
|
||||
|
||||
assert not any(type(callback) is PromptCachingDeploymentCheck for callback in (router_a.optional_callbacks or []))
|
||||
assert any(type(callback) is PromptCachingDeploymentCheck for callback in (router_b.optional_callbacks or []))
|
||||
assert any(type(callback) is PromptCachingDeploymentCheck for callback in litellm.callbacks)
|
||||
|
||||
router_b.update_settings(optional_pre_call_checks=[])
|
||||
|
||||
assert not any(type(callback) is PromptCachingDeploymentCheck for callback in (router_b.optional_callbacks or []))
|
||||
assert not any(type(callback) is PromptCachingDeploymentCheck for callback in litellm.callbacks)
|
||||
|
||||
|
||||
def test_remove_optional_pre_call_check_keeps_global_callback_when_second_router_clears_first():
|
||||
router_a = _build_router()
|
||||
router_b = _build_router()
|
||||
|
||||
router_a.update_settings(optional_pre_call_checks=["prompt_caching"])
|
||||
router_b.update_settings(optional_pre_call_checks=["prompt_caching"])
|
||||
|
||||
router_b.update_settings(optional_pre_call_checks=[])
|
||||
|
||||
assert any(type(callback) is PromptCachingDeploymentCheck for callback in (router_a.optional_callbacks or []))
|
||||
assert not any(type(callback) is PromptCachingDeploymentCheck for callback in (router_b.optional_callbacks or []))
|
||||
assert any(type(callback) is PromptCachingDeploymentCheck for callback in litellm.callbacks)
|
||||
|
||||
router_a.update_settings(optional_pre_call_checks=[])
|
||||
|
||||
assert not any(type(callback) is PromptCachingDeploymentCheck for callback in litellm.callbacks)
|
||||
|
||||
|
||||
def test_update_settings_replaces_toggleable_pre_call_checks():
|
||||
router = _build_router()
|
||||
|
||||
router.update_settings(optional_pre_call_checks=["prompt_caching"])
|
||||
router.update_settings(optional_pre_call_checks=["enforce_model_rate_limits"])
|
||||
|
||||
assert not any(isinstance(callback, PromptCachingDeploymentCheck) for callback in (router.optional_callbacks or []))
|
||||
assert not any(isinstance(callback, PromptCachingDeploymentCheck) for callback in litellm.callbacks)
|
||||
assert any(isinstance(callback, ModelRateLimitingCheck) for callback in (router.optional_callbacks or []))
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_update_settings_preserves_router_budget_limiting_when_omitted(monkeypatch):
|
||||
async def _disable_periodic_sync(*args, **kwargs):
|
||||
return None
|
||||
|
||||
monkeypatch.setattr(
|
||||
"litellm.router_strategy.budget_limiter.RouterBudgetLimiting.periodic_sync_in_memory_spend_with_redis",
|
||||
_disable_periodic_sync,
|
||||
)
|
||||
router = _build_router()
|
||||
|
||||
router.add_optional_pre_call_checks(["router_budget_limiting"])
|
||||
router.update_settings(optional_pre_call_checks=[])
|
||||
|
||||
assert any(isinstance(callback, RouterBudgetLimiting) for callback in (router.optional_callbacks or []))
|
||||
|
||||
|
||||
def test_update_settings_persists_retry_policy_dict():
|
||||
"""When the proxy's ``_add_router_settings_from_db_config`` calls
|
||||
``llm_router.update_settings(retry_policy={...})`` after reading the
|
||||
|
|
@ -255,8 +375,12 @@ async def test_config_update_persists_and_reads_back_retry_policy(monkeypatch):
|
|||
RateLimitErrorRetries=7,
|
||||
)
|
||||
)
|
||||
request = MagicMock()
|
||||
request.json = AsyncMock(return_value={"router_settings": {"retry_policy": posted.model_dump()}})
|
||||
|
||||
await proxy_server.update_config(
|
||||
config_info=ConfigYAML(router_settings=posted),
|
||||
request=request,
|
||||
user_api_key_dict=UserAPIKeyAuth(user_role=LitellmUserRoles.PROXY_ADMIN, api_key="sk-1234"),
|
||||
)
|
||||
|
||||
|
|
|
|||
2
ui/litellm-dashboard/src/lib/http/schema.d.ts
generated
vendored
2
ui/litellm-dashboard/src/lib/http/schema.d.ts
generated
vendored
|
|
@ -37473,6 +37473,8 @@ export interface components {
|
|||
} | null;
|
||||
/** Num Retries */
|
||||
num_retries?: number | null;
|
||||
/** Optional Pre Call Checks */
|
||||
optional_pre_call_checks?: ("prompt_caching" | "router_budget_limiting" | "responses_api_deployment_check" | "deployment_affinity" | "session_affinity" | "forward_client_headers_by_model_group" | "enforce_model_rate_limits" | "encrypted_content_affinity")[] | null;
|
||||
/** Retry After */
|
||||
retry_after?: number | null;
|
||||
retry_policy?: components["schemas"]["RetryPolicy"] | null;
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue