mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-06 02:48:13 +00:00
fix(auto-router): allow member edits with saved classifier credentials
This commit is contained in:
parent
ba0e4e2d23
commit
bbaa1423ab
2 changed files with 42 additions and 9 deletions
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue