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:
Mateo Wang 2026-09-02 11:23:22 -07:00 • committed by GitHub
commit a43228ef72
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
8 changed files with 388 additions and 48 deletions

View file

@ -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))

View file

@ -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(

View file

@ -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):

View file

@ -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

View file

@ -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

View file

@ -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."""

View file

@ -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"),
)

View file

@ -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;