From 60b4e2a530af9d1f0d1b1d83f33d9c510cd60087 Mon Sep 17 00:00:00 2001 From: Tin Chi Lo Date: Mon, 28 Sep 2026 13:08:38 -0700 Subject: [PATCH] fix(auto-router): retain saved classifier transport during member inference --- litellm/proxy/auth/auto_router_checks.py | 3 +- .../auto_router_endpoints.py | 4 ++- .../test_auto_router_endpoints.py | 33 ++++++++++++------- tests/unit/test_router/test_router.py | 33 ++++++++++++++++++- 4 files changed, 59 insertions(+), 14 deletions(-) diff --git a/litellm/proxy/auth/auto_router_checks.py b/litellm/proxy/auth/auto_router_checks.py index b83e1f3fffe..c0c2cdda002 100644 --- a/litellm/proxy/auth/auto_router_checks.py +++ b/litellm/proxy/auth/auto_router_checks.py @@ -1,6 +1,7 @@ from __future__ import annotations from collections.abc import Mapping +from types import MappingProxyType from typing import TYPE_CHECKING, Final from pydantic import TypeAdapter, ValidationError @@ -86,7 +87,7 @@ async def authorize_member_auto_router_inference( if raw_config is None: raise HTTPException(status_code=403, detail="The member auto-router configuration is invalid") default_model: Final = params.get("complexity_router_default_model") - config: Final = validate_member_auto_router_config(raw_config) + config: Final = validate_member_auto_router_config(raw_config, supplied_config=MappingProxyType({})) membership: Final = ( await get_team_membership( user_id=actor.user_id, diff --git a/litellm/proxy/management_endpoints/auto_router_endpoints.py b/litellm/proxy/management_endpoints/auto_router_endpoints.py index 9708161397a..e0b2d7df324 100644 --- a/litellm/proxy/management_endpoints/auto_router_endpoints.py +++ b/litellm/proxy/management_endpoints/auto_router_endpoints.py @@ -269,12 +269,13 @@ async def _authorize_member_dry_run_config( default_model: str | None, user_api_key_dict: UserAPIKeyAuth, team: LiteLLM_TeamTable, + supplied_config: Mapping[str, object] | None = None, ) -> UserAPIKeyAuth: from litellm.proxy.proxy_server import llm_router, prisma_client if prisma_client is None or llm_router is None: raise HTTPException(status_code=503, detail="Cannot verify auto-router model access") - validated: Final = validate_member_auto_router_config(config) + validated: Final = validate_member_auto_router_config(config, supplied_config=supplied_config) scoped_actor: Final = user_api_key_dict.model_copy( update=MappingProxyType({"team_id": team.team_id, "team_models": team.models, "org_id": team.organization_id}) ) @@ -544,6 +545,7 @@ async def preview_auto_router_routing( default_model=resolved.default_model, user_api_key_dict=user_api_key_dict, team=member_team, + supplied_config=MappingProxyType({}) if resolved.saved_model_id is not None else None, ) if member_team is not None else user_api_key_dict diff --git a/tests/test_litellm/proxy/management_endpoints/test_auto_router_endpoints.py b/tests/test_litellm/proxy/management_endpoints/test_auto_router_endpoints.py index ff3d19e8637..4e826418dd4 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_auto_router_endpoints.py +++ b/tests/test_litellm/proxy/management_endpoints/test_auto_router_endpoints.py @@ -2558,7 +2558,8 @@ async def test_jev_test_routing_authorizes_paid_evaluation_before_contacting_typ @pytest.mark.asyncio @pytest.mark.parametrize( - "case", ["allowed", "credential-free", "missing", "blocked", "key", "budget", "team", "not-router"] + "case", + ["allowed", "credential-free", "member", "member-unsaved", "missing", "blocked", "key", "budget", "team", "not-router"], ) async def test_saved_jev_probe_uses_authorized_server_configuration(monkeypatch: pytest.MonkeyPatch, case: str) -> None: router: Final = RecordingRouter("SIMPLE") @@ -2579,15 +2580,19 @@ async def test_saved_jev_probe_uses_authorized_server_configuration(monkeypatch: "model_info": { "id": "saved-jev-id", "blocked": case == "blocked", - "team_id": "owner-team" if case == "team" else None, + "team_id": ( + "owner-team" if case == "team" else "member-preview-team" if case == "member" else None + ), }, } ) ) monkeypatch.setattr(proxy_server, "llm_router", router) actor: Final = ( - _configure_member_preview(monkeypatch) - if case == "team" + _configure_member_preview( + monkeypatch, models=[*(TIERS[name][0] for name in TIERS), "saved-jev", "typesafe/jev-latest"] + ) + if case in ("team", "member", "member-unsaved") else UserAPIKeyAuth( user_role=LitellmUserRoles.PROXY_ADMIN, api_key="sk-probe", @@ -2600,8 +2605,10 @@ async def test_saved_jev_probe_uses_authorized_server_configuration(monkeypatch: request: Final = _request_from( { "prompt": "what is 2+2", - "saved_model_id": "missing-id" if case == "missing" else "saved-jev-id", - "team_id": "member-preview-team" if case == "team" else None, + "saved_model_id": ( + None if case == "member-unsaved" else "missing-id" if case == "missing" else "saved-jev-id" + ), + "team_id": "member-preview-team" if case in ("team", "member", "member-unsaved") else None, }, classifier_type="jev", jev_classifier_config=( @@ -2629,10 +2636,12 @@ async def test_saved_jev_probe_uses_authorized_server_configuration(monkeypatch: ) ) operation: Final = preview_auto_router_routing(request, actor, ROUTING_HTTP_REQUEST) - if case in ("missing", "blocked", "team", "not-router"): + if case in ("missing", "blocked", "team", "not-router", "member-unsaved"): with pytest.raises(HTTPException) as denied: await operation - assert denied.value.status_code == {"missing": 404, "blocked": 404, "team": 403, "not-router": 400}[case] + assert denied.value.status_code == { + "missing": 404, "blocked": 404, "team": 403, "not-router": 400, "member-unsaved": 400, + }[case] elif case in ("key", "budget"): with pytest.raises(ProxyException) as forbidden: await operation @@ -2645,7 +2654,7 @@ async def test_saved_jev_probe_uses_authorized_server_configuration(monkeypatch: assert result.routed_model == "cheap-model" assert evaluation.calls.last.request.headers["authorization"] == f"Bearer {stored_key}" assert stored_key not in result.model_dump_json() - assert evaluation.call_count == (1 if case in ("allowed", "credential-free") else 0) + assert evaluation.call_count == (1 if case in ("allowed", "credential-free", "member") else 0) assert router.recorded_calls == [] await handler.client.aclose() @@ -3167,13 +3176,15 @@ async def test_validate_config_gates_like_the_write_it_rehearses(monkeypatch: py assert not_their_team.value.status_code == 403 -def _configure_member_preview(monkeypatch: pytest.MonkeyPatch, *, allowed: bool = True) -> UserAPIKeyAuth: +def _configure_member_preview( + monkeypatch: pytest.MonkeyPatch, *, allowed: bool = True, models: Sequence[str] | None = None, +) -> UserAPIKeyAuth: from litellm.proxy import proxy_server from litellm.proxy._types import UI_TEAM_ID, LiteLLM_TeamTable team: Final = LiteLLM_TeamTable( team_id="member-preview-team", - models=list(TIERS[name][0] for name in TIERS), + models=list(models) if models is not None else list(TIERS[name][0] for name in TIERS), members_with_roles=[{"role": "user", "user_id": "preview-member"}], team_member_permissions=["/auto_router/manage"] if allowed else [], ) diff --git a/tests/unit/test_router/test_router.py b/tests/unit/test_router/test_router.py index 3dc96e4844b..125c5b8863d 100644 --- a/tests/unit/test_router/test_router.py +++ b/tests/unit/test_router/test_router.py @@ -17978,7 +17978,9 @@ class TestMemberAutoRouterInference: monkeypatch.setattr(proxy_server, "prisma_client", self.database) @staticmethod - def _marker(*, member: bool = True, classifier: bool = False) -> dict[str, object]: + def _marker( + *, member: bool = True, classifier: bool = False, jev: Mapping[str, object] | None = None, + ) -> dict[str, object]: target: Final = "permitted-model" if member else "restricted-model" return { "model_name": "model_name_router-team_member-router", @@ -17987,6 +17989,7 @@ class TestMemberAutoRouterInference: "complexity_router_config": { "tiers": dict.fromkeys(("SIMPLE", "MEDIUM", "COMPLEX", "REASONING"), target), "adaptive": False, **({"classifier_type": "llm", "classifier_llm_config": {"model": target}} if classifier else {}), + **({"classifier_type": "jev", "jev_classifier_config": jev} if jev is not None else {}), }, "tags": ["member" if member else "admin"], "timeout": 13.0 if member else 29.0, }, @@ -18026,6 +18029,34 @@ class TestMemberAutoRouterInference: assert response is not None return response + @pytest.mark.asyncio + @pytest.mark.parametrize("api_key", (None, "synthetic-saved-nimble-key")) + @pytest.mark.parametrize("allowed", (True, False)) + async def test_saved_nimble_transport_keeps_runtime_provider_authorization( + self, api_key: str | None, allowed: bool, respx_mock: respx.MockRouter, monkeypatch: pytest.MonkeyPatch, + ) -> None: + monkeypatch.setenv("DISABLE_AIOHTTP_TRANSPORT", "True") + evaluation: Final = respx_mock.post("https://saved-nimble.test/v1/systemone").respond(200, json={ + "answers": {"tier": {"type": "choice", "choice": "SIMPLE", "confidence": 1, "probabilities": {"SIMPLE": 1}}}, + }) + models: Final = [*self.team.models, "bespoke_nimble/nimble-latest"] + self.database.db.litellm_teamtable.find_unique.return_value = self.team.model_copy(update={"models": models}) + actor: Final = self.actor.model_copy(update={"models": models if allowed else self.actor.models}) + router: Final = self._router(self._marker(jev={ + "provider": "bespoke_nimble", "model": "nimble-latest", + "api_base": "https://saved-nimble.test", "api_key": api_key, + })) + if not allowed: + with pytest.raises(ProxyException, match="bespoke_nimble/nimble-latest"): + await self._route(router, self._request(actor=actor)) + assert evaluation.call_count == 0 + return + response: Final = await self._route(router, self._request(actor=actor)) + assert response.model == "permitted-model" + assert response.routing_decision is not None and response.routing_decision["cause"] == "jev_classifier" + assert evaluation.call_count == 1 + assert evaluation.calls.last.request.headers.get("authorization") == (f"Bearer {api_key}" if api_key else None) + @pytest.mark.asyncio @pytest.mark.parametrize("metadata_name", ("metadata", "litellm_metadata")) async def test_cached_roster_revocation_blocks_classifier_and_session_rebinding(