mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-06 02:48:13 +00:00
fix(auto-router): retain saved classifier transport during member inference
This commit is contained in:
parent
bbaa1423ab
commit
60b4e2a530
4 changed files with 59 additions and 14 deletions
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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 [],
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue