diff --git a/litellm/proxy/management_helpers/auto_router_permissions.py b/litellm/proxy/management_helpers/auto_router_permissions.py index 111ffa7b3d7..a827ba63970 100644 --- a/litellm/proxy/management_helpers/auto_router_permissions.py +++ b/litellm/proxy/management_helpers/auto_router_permissions.py @@ -5,7 +5,7 @@ from types import MappingProxyType from typing import TYPE_CHECKING, Final, Literal from fastapi import HTTPException -from pydantic import BaseModel, ConfigDict, Field, ValidationError +from pydantic import BaseModel, ConfigDict, Field, TypeAdapter, ValidationError from typing_extensions import ReadOnly, TypedDict from litellm.models.organization import LiteLLM_OrganizationTable @@ -125,14 +125,26 @@ def authorize_member_auto_router_team( raise HTTPException(status_code=403, detail="This team does not allow you to manage your own auto routers.") -def validate_member_auto_router_config(config: Mapping[str, object]) -> RequestComplexityRouterConfig: +def validate_member_auto_router_config( + config: Mapping[str, object], *, supplied_config: Mapping[str, object] | None = None +) -> RequestComplexityRouterConfig: try: validated: Final = _MemberComplexityRouterConfig.model_validate(config) for entries in validated.tier_model_configs.values(): for entry in entries: _MemberRouterGenerationParams.model_validate(entry.litellm_params) if validated.jev_classifier_config is not None: - _MemberJevClassifierConfig.model_validate(validated.jev_classifier_config.model_dump()) + supplied: Final = config if supplied_config is None else supplied_config + supplied_jev: Final = TypeAdapter(Mapping[str, object] | None).validate_python( + supplied.get("jev_classifier_config") + ) or MappingProxyType({}) + _MemberJevClassifierConfig.model_validate( + { + **validated.jev_classifier_config.model_dump(exclude={"api_key", "api_base"}), + "api_key": supplied_jev.get("api_key"), + "api_base": supplied_jev.get("api_base"), + } + ) return validated except ValidationError as exc: location: Final = ".".join(str(part) for part in exc.errors()[0]["loc"]) @@ -347,7 +359,10 @@ async def authorize_member_auto_router_write( ) if raw_config is None: raise HTTPException(status_code=400, detail="A complexity_router_config is required.") - config: Final = validate_member_auto_router_config(effective_config if effective_config is not None else raw_config) + config: Final = validate_member_auto_router_config( + effective_config if effective_config is not None else raw_config, + supplied_config=supplied_config if supplied_config is not None else MappingProxyType({}), + ) stored_default: Final = existing.litellm_params.complexity_router_default_model if existing is not None else None default_model: Final = ( params.complexity_router_default_model diff --git a/tests/test_litellm/proxy/management_endpoints/test_model_management_endpoints.py b/tests/test_litellm/proxy/management_endpoints/test_model_management_endpoints.py index 447062af10f..abaa4f9fe92 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_model_management_endpoints.py +++ b/tests/test_litellm/proxy/management_endpoints/test_model_management_endpoints.py @@ -7658,8 +7658,19 @@ class TestTeamMemberAutoRouterWrites: @pytest.mark.asyncio @pytest.mark.parametrize("endpoint", ["patch", "legacy"]) @pytest.mark.parametrize("can_use_nimble", [True, False]) - async def test_member_partial_update_authorizes_the_persisted_classifier_provider( - self, endpoint: str, can_use_nimble: bool + @pytest.mark.parametrize( + ("stored_transport", "supplied_transport"), + [ + ({}, {}), + ({"api_base": "http://nimble.internal"}, {}), + ({"api_base": "http://nimble.internal", "api_key": "synthetic-nimble-key"}, {}), + ({"api_base": "http://nimble.internal"}, {"api_base": "https://collector.invalid"}), + ({"api_key": "synthetic-nimble-key"}, {"api_key": "synthetic-member-key"}), + ], + ) + async def test_member_partial_update_checks_saved_provider_and_submitted_transport( + self, endpoint: str, can_use_nimble: bool, + stored_transport: Mapping[str, str], supplied_transport: Mapping[str, str], ) -> None: from fastapi import HTTPException @@ -7673,7 +7684,9 @@ class TestTeamMemberAutoRouterWrites: "complexity_router_config": { "classifier_type": "jev", "tiers": {"SIMPLE": "allowed"}, - "jev_classifier_config": {"provider": "bespoke_nimble", "model": "nimble-latest"}, + "jev_classifier_config": { + "provider": "bespoke_nimble", "model": "nimble-latest", **stored_transport, + }, }, }, } @@ -7684,7 +7697,7 @@ class TestTeamMemberAutoRouterWrites: complexity_router_config={ "classifier_type": "jev", "tiers": {"SIMPLE": "allowed"}, - "jev_classifier_config": {"timeout_ms": 900}, + "jev_classifier_config": {"timeout_ms": 900, **supplied_transport}, } ), model_info=ModelInfo(id=row.model_id, team_id=team.team_id), @@ -7692,6 +7705,11 @@ class TestTeamMemberAutoRouterWrites: actor: Final = UserAPIKeyAuth(user_id="owner", user_role=LitellmUserRoles.INTERNAL_USER, models=models) with self._environment(database, row): operation: Final = patch_model(row.model_id, request, actor) if endpoint == "patch" else update_model(request, actor) + if supplied_transport: + with pytest.raises((HTTPException, ProxyException), match="Invalid member auto-router configuration"): + await operation + database.transaction.litellm_proxymodeltable.update.assert_not_awaited() + return if not can_use_nimble: with pytest.raises((HTTPException, ProxyException)) as denied: await operation @@ -7705,7 +7723,7 @@ class TestTeamMemberAutoRouterWrites: params: Final = LiteLLM_Params.model_validate_json(written) config: Final = TypeAdapter(Mapping[str, object]).validate_python(params.complexity_router_config) assert config["jev_classifier_config"] == { - "provider": "bespoke_nimble", "model": "nimble-latest", "timeout_ms": 900 + "provider": "bespoke_nimble", "model": "nimble-latest", "timeout_ms": 900, **stored_transport, } @pytest.mark.asyncio