fix(auto-router): retain saved classifier transport during member inference

This commit is contained in:
Tin Chi Lo 2026-09-28 13:08:38 -07:00
parent bbaa1423ab
commit 60b4e2a530
4 changed files with 59 additions and 14 deletions

View file

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

View file

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

View file

@ -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 [],
)

View file

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