mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-07 02:59:05 +00:00
fix: authorize and invalidate cached JWT team aliases
Co-authored-by: Cursor <cursoragent@cursor.com>
This commit is contained in:
parent
7d01090d5d
commit
75a7188d1f
10 changed files with 235 additions and 62 deletions
|
|
@ -18,7 +18,7 @@ from types import MappingProxyType
|
|||
from typing import TYPE_CHECKING, Any, Final, Literal, Optional, Protocol, cast
|
||||
|
||||
from fastapi import HTTPException, Request, status
|
||||
from pydantic import BaseModel
|
||||
from pydantic import BaseModel, TypeAdapter
|
||||
|
||||
import litellm
|
||||
from litellm._logging import verbose_proxy_logger
|
||||
|
|
@ -86,6 +86,7 @@ from litellm.proxy.common_utils.user_api_key_cache import (
|
|||
object_permission_cache_key,
|
||||
tag_cache_key,
|
||||
tag_registry_cache_key,
|
||||
team_model_aliases_cache_key,
|
||||
)
|
||||
from litellm.proxy.db.exception_handler import PrismaDBExceptionHandler
|
||||
from litellm.proxy.guardrails.tool_name_extraction import (
|
||||
|
|
@ -105,6 +106,7 @@ from litellm.repositories.table_repositories import (
|
|||
EndUserRepository,
|
||||
JWTKeyMappingRepository,
|
||||
ManagedVectorStoresRepository,
|
||||
ModelTableRepository,
|
||||
TagRepository,
|
||||
TeamMembershipRepository,
|
||||
)
|
||||
|
|
@ -126,6 +128,8 @@ if TYPE_CHECKING:
|
|||
else:
|
||||
Span = Any
|
||||
|
||||
_TEAM_MODEL_ALIASES_ADAPTER: Final = TypeAdapter(dict[str, str])
|
||||
|
||||
|
||||
class _PrismaDictableRow(Protocol):
|
||||
def dict(self) -> Mapping[str, object]: ...
|
||||
|
|
@ -139,6 +143,10 @@ class _PrismaModelDumpRow(Protocol):
|
|||
def model_dump(self) -> Mapping[str, object]: ...
|
||||
|
||||
|
||||
class _PrismaModelAliasesRow(Protocol):
|
||||
model_aliases: object
|
||||
|
||||
|
||||
class _PrismaTeamRow(Protocol):
|
||||
def dict(self) -> Mapping[str, object]: ...
|
||||
|
||||
|
|
@ -201,6 +209,12 @@ def _model_dump_table(repo: _PrismaTableHolder[_PrismaModelDumpRow]) -> _PrismaA
|
|||
return repo.table
|
||||
|
||||
|
||||
def _model_aliases_table(
|
||||
repo: _PrismaTableHolder[_PrismaModelAliasesRow],
|
||||
) -> _PrismaAuthTable[_PrismaModelAliasesRow]:
|
||||
return repo.table
|
||||
|
||||
|
||||
def _team_table(repo: _PrismaTableHolder[_PrismaTeamRow]) -> _PrismaAuthTable[_PrismaTeamRow]:
|
||||
return repo.table
|
||||
|
||||
|
|
@ -2435,6 +2449,37 @@ async def _get_team_object_from_db(team_id: str, prisma_client: PrismaClient) ->
|
|||
return await _team_table(TeamRepository(prisma_client)).find_unique(where={"team_id": team_id})
|
||||
|
||||
|
||||
async def get_team_model_aliases(
|
||||
model_id: int,
|
||||
prisma_client: PrismaClient,
|
||||
user_api_key_cache: UserApiKeyCache,
|
||||
) -> dict[str, str] | None: # mutable-ok: UserAPIKeyAuth.team_model_aliases requires a dict
|
||||
cache_key: Final = team_model_aliases_cache_key(model_id)
|
||||
cached: Final[object] = await user_api_key_cache.async_get_cache( # pyright: ignore[reportAny] # cache API is untyped
|
||||
cache_key
|
||||
)
|
||||
if cached is not None:
|
||||
return _TEAM_MODEL_ALIASES_ADAPTER.validate_python(cached)
|
||||
|
||||
where: Final[Mapping[str, object]] = MappingProxyType({"id": model_id})
|
||||
row: Final = await _model_aliases_table(ModelTableRepository(prisma_client)).find_unique(where=where)
|
||||
if row is None or row.model_aliases is None:
|
||||
return None
|
||||
|
||||
raw_aliases: Final = row.model_aliases
|
||||
aliases: Final = (
|
||||
_TEAM_MODEL_ALIASES_ADAPTER.validate_json(raw_aliases)
|
||||
if isinstance(raw_aliases, (str, bytes, bytearray))
|
||||
else _TEAM_MODEL_ALIASES_ADAPTER.validate_python(raw_aliases)
|
||||
)
|
||||
await user_api_key_cache.async_set_cache( # pyright: ignore[reportUnknownMemberType] # cache API has untyped kwargs
|
||||
key=cache_key,
|
||||
value=aliases,
|
||||
ttl=60,
|
||||
)
|
||||
return aliases
|
||||
|
||||
|
||||
async def _get_team_object_from_user_api_key_cache(
|
||||
team_id: str,
|
||||
prisma_client: PrismaClient,
|
||||
|
|
|
|||
|
|
@ -64,6 +64,7 @@ from .auth_checks import (
|
|||
get_role_based_models,
|
||||
get_role_based_routes,
|
||||
get_team_membership,
|
||||
get_team_model_aliases,
|
||||
get_team_object,
|
||||
get_team_object_by_alias,
|
||||
get_user_object,
|
||||
|
|
@ -1377,7 +1378,15 @@ class JWTAuthManager:
|
|||
model=requested_model,
|
||||
team_object=team_object,
|
||||
llm_router=llm_router,
|
||||
team_model_aliases=None,
|
||||
team_model_aliases=(
|
||||
await get_team_model_aliases(
|
||||
model_id=team_object.model_id,
|
||||
prisma_client=prisma_client,
|
||||
user_api_key_cache=user_api_key_cache,
|
||||
)
|
||||
if team_object.model_id is not None and prisma_client is not None
|
||||
else None
|
||||
),
|
||||
)
|
||||
):
|
||||
is_allowed = allowed_routes_check(
|
||||
|
|
@ -1914,7 +1923,15 @@ class JWTAuthManager:
|
|||
model=requested_model,
|
||||
team_object=team_object,
|
||||
llm_router=llm_router,
|
||||
team_model_aliases=None,
|
||||
team_model_aliases=(
|
||||
await get_team_model_aliases(
|
||||
model_id=team_object.model_id,
|
||||
prisma_client=prisma_client,
|
||||
user_api_key_cache=user_api_key_cache,
|
||||
)
|
||||
if team_object.model_id is not None and prisma_client is not None
|
||||
else None
|
||||
),
|
||||
)
|
||||
except ProxyException:
|
||||
continue
|
||||
|
|
|
|||
|
|
@ -18,7 +18,6 @@ import fastapi
|
|||
import orjson
|
||||
from fastapi import HTTPException, Request, WebSocket, status
|
||||
from fastapi.security.api_key import APIKeyHeader
|
||||
from pydantic import BaseModel, TypeAdapter
|
||||
|
||||
import litellm
|
||||
from litellm._logging import verbose_logger, verbose_proxy_logger
|
||||
|
|
@ -50,6 +49,7 @@ from litellm.proxy.auth.auth_checks import (
|
|||
get_end_user_object,
|
||||
get_jwt_key_mapping_object,
|
||||
get_project_object,
|
||||
get_team_model_aliases,
|
||||
get_team_object,
|
||||
get_user_object,
|
||||
is_valid_fallback_model,
|
||||
|
|
@ -107,11 +107,6 @@ except ImportError as e:
|
|||
enterprise_custom_auth = None
|
||||
|
||||
user_api_key_service_logger_obj: Final = ServiceLogging() # used for tracking latency on OTEL
|
||||
_TEAM_MODEL_ALIASES_ADAPTER: Final = TypeAdapter(dict[str, str])
|
||||
|
||||
|
||||
class _TeamModelAliasesRow(BaseModel):
|
||||
model_aliases: object = None
|
||||
|
||||
|
||||
def _normalize_public_auth_route(route: str) -> str:
|
||||
|
|
@ -1066,44 +1061,6 @@ async def _read_request_body_deferring_parse_failure(
|
|||
return populate_request_with_path_params(request_data=parsed_body, request=request), None
|
||||
|
||||
|
||||
async def _get_team_model_aliases(
|
||||
model_id: int,
|
||||
prisma_client: PrismaClient,
|
||||
user_api_key_cache: UserApiKeyCache,
|
||||
) -> dict[str, str] | None: # mutable-ok: UserAPIKeyAuth.team_model_aliases requires a dict
|
||||
cache_key: Final = f"team_model_aliases:{model_id}"
|
||||
cached: Final[object] = await user_api_key_cache.async_get_cache( # pyright: ignore[reportAny] # cache API is untyped
|
||||
cache_key
|
||||
)
|
||||
if cached is not None:
|
||||
return _TEAM_MODEL_ALIASES_ADAPTER.validate_python(cached)
|
||||
|
||||
row_value: Final[object] = ( # pyright: ignore[reportAny] # generated table result is dynamic
|
||||
await prisma_client.db.litellm_modeltable.find_unique( # pyright: ignore[reportAny] # generated table API is dynamic
|
||||
where={"id": model_id}, # mutable-ok: Prisma find_unique requires a dict
|
||||
)
|
||||
)
|
||||
if row_value is None:
|
||||
return None
|
||||
|
||||
row: Final = _TeamModelAliasesRow.model_validate(row_value, from_attributes=True)
|
||||
raw_aliases: Final = row.model_aliases
|
||||
if raw_aliases is None:
|
||||
return None
|
||||
|
||||
aliases: Final = (
|
||||
_TEAM_MODEL_ALIASES_ADAPTER.validate_json(raw_aliases)
|
||||
if isinstance(raw_aliases, (str, bytes, bytearray))
|
||||
else _TEAM_MODEL_ALIASES_ADAPTER.validate_python(raw_aliases)
|
||||
)
|
||||
await user_api_key_cache.async_set_cache( # pyright: ignore[reportUnknownMemberType] # cache API has untyped kwargs
|
||||
key=cache_key,
|
||||
value=aliases,
|
||||
ttl=60,
|
||||
)
|
||||
return aliases
|
||||
|
||||
|
||||
async def _record_unparsable_body_failure(
|
||||
user_api_key_dict: UserAPIKeyAuth,
|
||||
body_parse_exception: ProxyException,
|
||||
|
|
@ -1423,7 +1380,7 @@ async def _user_api_key_auth_builder(
|
|||
team_rpm_limit=(team_object.rpm_limit if team_object is not None else None),
|
||||
team_models=(team_object.models if team_object is not None else []),
|
||||
team_model_aliases=(
|
||||
await _get_team_model_aliases(
|
||||
await get_team_model_aliases(
|
||||
model_id=team_object.model_id,
|
||||
prisma_client=prisma_client,
|
||||
user_api_key_cache=user_api_key_cache,
|
||||
|
|
|
|||
|
|
@ -148,6 +148,10 @@ class UserApiKeyCache(DualCache):
|
|||
return await super().async_set_cache_pipeline(cache_list=normalized, local_only=local_only, **kwargs)
|
||||
|
||||
|
||||
def team_model_aliases_cache_key(model_id: int) -> str:
|
||||
return f"team_model_aliases:{model_id}"
|
||||
|
||||
|
||||
#: Value cached under ``user_object_permission_id_cache_key`` when the user links no permission row,
|
||||
#: so a human without an entitlement costs no DB read per request. Lives beside the key builder
|
||||
#: because it is part of the same cache protocol: a reader that knows the key must know this value.
|
||||
|
|
|
|||
|
|
@ -48,12 +48,16 @@ from litellm.proxy._types import (
|
|||
UserAPIKeyAuth,
|
||||
)
|
||||
from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
|
||||
from litellm.proxy.common_utils.auth_cache_invalidation_pubsub import evict_and_broadcast
|
||||
from litellm.proxy.common_utils.config_sync_pubsub import (
|
||||
coordination_redis_cache,
|
||||
publish_config_change,
|
||||
)
|
||||
from litellm.proxy.common_utils.encrypt_decrypt_utils import encrypt_value_helper
|
||||
from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache
|
||||
from litellm.proxy.common_utils.user_api_key_cache import (
|
||||
UserApiKeyCache,
|
||||
team_model_aliases_cache_key,
|
||||
)
|
||||
from litellm.proxy.management_endpoints.common_utils import _is_user_team_admin
|
||||
from litellm.proxy.management_endpoints.team_endpoints import (
|
||||
_refresh_cached_team,
|
||||
|
|
@ -1306,6 +1310,7 @@ async def _remove_unbacked_team_models(
|
|||
else await delete_team_model_alias(
|
||||
public_model_name=model_params.model_name,
|
||||
prisma_client=prisma_client,
|
||||
user_api_key_cache=user_api_key_cache,
|
||||
)
|
||||
)
|
||||
removed_alias_names: Final = {alias for alias_team_id, alias in removed_model_aliases if alias_team_id == team_id}
|
||||
|
|
@ -1671,6 +1676,7 @@ async def delete_model(
|
|||
async def delete_team_model_alias(
|
||||
public_model_name: str,
|
||||
prisma_client: PrismaClient,
|
||||
user_api_key_cache: UserApiKeyCache,
|
||||
) -> list[tuple[str, str]]:
|
||||
"""
|
||||
Delete a team model alias
|
||||
|
|
@ -1681,6 +1687,11 @@ async def delete_team_model_alias(
|
|||
- List of team id + model alias pairs that were removed
|
||||
"""
|
||||
team_model_aliases: Final = await _model_alias_table(prisma_client).find_many(include={"team": True})
|
||||
alias_cache_keys: Final = tuple(
|
||||
team_model_aliases_cache_key(team_model_alias.id)
|
||||
for team_model_alias in team_model_aliases
|
||||
if public_model_name in team_model_alias.model_aliases.values()
|
||||
)
|
||||
tasks: Final = []
|
||||
removed_model_aliases: Final[list[tuple[str, str]]] = []
|
||||
for team_model_alias in team_model_aliases:
|
||||
|
|
@ -1699,6 +1710,7 @@ async def delete_team_model_alias(
|
|||
)
|
||||
)
|
||||
await asyncio.gather(*tasks)
|
||||
await evict_and_broadcast(alias_cache_keys, user_api_key_cache)
|
||||
|
||||
return removed_model_aliases
|
||||
|
||||
|
|
|
|||
|
|
@ -90,9 +90,13 @@ from litellm.proxy.auth.auth_utils import (
|
|||
enforce_output_token_estimates_are_admin_only,
|
||||
)
|
||||
from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
|
||||
from litellm.proxy.common_utils.auth_cache_invalidation_pubsub import evict_and_broadcast
|
||||
from litellm.proxy.common_utils.callback_utils import encrypt_callback_vars
|
||||
from litellm.proxy.common_utils.json_merge_patch import apply_json_merge_patch
|
||||
from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache
|
||||
from litellm.proxy.common_utils.user_api_key_cache import (
|
||||
UserApiKeyCache,
|
||||
team_model_aliases_cache_key,
|
||||
)
|
||||
from litellm.proxy.management_endpoints.common_daily_activity import (
|
||||
get_daily_activity_aggregated,
|
||||
)
|
||||
|
|
@ -1647,6 +1651,7 @@ async def _update_model_table(
|
|||
data: UpdateTeamRequest,
|
||||
model_id: int | None,
|
||||
prisma_client: PrismaClient,
|
||||
user_api_key_cache: UserApiKeyCache,
|
||||
user_api_key_dict: UserAPIKeyAuth,
|
||||
litellm_proxy_admin_name: str,
|
||||
) -> int | None:
|
||||
|
|
@ -1673,6 +1678,11 @@ async def _update_model_table(
|
|||
)
|
||||
|
||||
_model_id = model_dict.id
|
||||
if model_id is not None:
|
||||
await evict_and_broadcast(
|
||||
(team_model_aliases_cache_key(model_id),),
|
||||
user_api_key_cache,
|
||||
)
|
||||
|
||||
return _model_id
|
||||
|
||||
|
|
@ -2222,6 +2232,7 @@ async def update_team(
|
|||
data=data,
|
||||
model_id=existing_team_row.model_id,
|
||||
prisma_client=prisma_client,
|
||||
user_api_key_cache=user_api_key_cache,
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
litellm_proxy_admin_name=litellm_proxy_admin_name,
|
||||
)
|
||||
|
|
|
|||
|
|
@ -4613,6 +4613,44 @@ async def test_find_team_with_model_access_defers_no_team_403_under_db_fallback(
|
|||
assert team_object is None
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_find_team_with_model_access_loads_aliases_before_restricted_team_selection():
|
||||
team = LiteLLM_TeamTable(
|
||||
team_id="alias-team",
|
||||
models=["FW-Kimi-K3"],
|
||||
model_id=1,
|
||||
)
|
||||
jwt_handler = JWTHandler()
|
||||
jwt_handler.litellm_jwtauth = LiteLLM_JWTAuth(enforce_team_based_model_access=True)
|
||||
|
||||
with (
|
||||
patch(
|
||||
"litellm.proxy.auth.handle_jwt.get_team_object",
|
||||
new_callable=AsyncMock,
|
||||
return_value=team,
|
||||
),
|
||||
patch(
|
||||
"litellm.proxy.auth.handle_jwt.get_team_model_aliases",
|
||||
new_callable=AsyncMock,
|
||||
return_value={"claude-opus-5": "FW-Kimi-K3"},
|
||||
) as mock_get_aliases,
|
||||
):
|
||||
team_id, team_object = await JWTAuthManager.find_team_with_model_access(
|
||||
team_ids={team.team_id},
|
||||
requested_model="claude-opus-5",
|
||||
route="/chat/completions",
|
||||
jwt_handler=jwt_handler,
|
||||
prisma_client=MagicMock(),
|
||||
user_api_key_cache=MagicMock(),
|
||||
parent_otel_span=None,
|
||||
proxy_logging_obj=MagicMock(),
|
||||
)
|
||||
|
||||
assert team_id == team.team_id
|
||||
assert team_object is team
|
||||
mock_get_aliases.assert_awaited_once()
|
||||
|
||||
|
||||
def _db_fallback_handler(litellm_jwtauth: Optional[LiteLLM_JWTAuth] = None) -> JWTHandler:
|
||||
handler = JWTHandler()
|
||||
handler.litellm_jwtauth = litellm_jwtauth or LiteLLM_JWTAuth()
|
||||
|
|
@ -4952,6 +4990,51 @@ async def test_resolve_db_team_fallback_skips_team_without_model_access():
|
|||
assert team_object is teams["allowed_team"]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_resolve_db_team_fallback_loads_aliases_before_restricted_team_selection():
|
||||
team = LiteLLM_TeamTable(
|
||||
team_id="alias-team",
|
||||
models=["FW-Kimi-K3"],
|
||||
model_id=1,
|
||||
)
|
||||
user_object = LiteLLM_UserTable(
|
||||
user_id="alias-user",
|
||||
user_role=LitellmUserRoles.INTERNAL_USER,
|
||||
teams=[team.team_id],
|
||||
)
|
||||
|
||||
with (
|
||||
patch(
|
||||
"litellm.proxy.auth.handle_jwt.get_team_object",
|
||||
new_callable=AsyncMock,
|
||||
return_value=team,
|
||||
),
|
||||
patch(
|
||||
"litellm.proxy.auth.handle_jwt.get_team_model_aliases",
|
||||
new_callable=AsyncMock,
|
||||
return_value={"claude-opus-5": "FW-Kimi-K3"},
|
||||
) as mock_get_aliases,
|
||||
):
|
||||
team_id, team_object, membership = await JWTAuthManager._resolve_db_team_fallback(
|
||||
user_object=user_object,
|
||||
user_id=None,
|
||||
requested_model="claude-opus-5",
|
||||
route="/chat/completions",
|
||||
jwt_handler=_db_fallback_handler(),
|
||||
enforce_team_based_model_access=True,
|
||||
team_id_upsert=False,
|
||||
prisma_client=MagicMock(),
|
||||
user_api_key_cache=MagicMock(),
|
||||
parent_otel_span=None,
|
||||
proxy_logging_obj=MagicMock(),
|
||||
)
|
||||
|
||||
assert team_id == team.team_id
|
||||
assert team_object is team
|
||||
assert membership is None
|
||||
mock_get_aliases.assert_awaited_once()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_resolve_db_team_fallback_enforces_team_allowed_routes():
|
||||
"""The DB-team fallback must apply the same team_allowed_routes gate as the
|
||||
|
|
|
|||
|
|
@ -31,11 +31,14 @@ from litellm.proxy._types import (
|
|||
JWTRoutingOverride,
|
||||
)
|
||||
from litellm.proxy.auth.handle_jwt import JWTHandler
|
||||
from litellm.proxy.auth.auth_checks import get_key_object, _cache_key_object
|
||||
from litellm.proxy.auth.auth_checks import (
|
||||
_cache_key_object,
|
||||
get_key_object,
|
||||
get_team_model_aliases,
|
||||
)
|
||||
from litellm.proxy.auth.route_checks import RouteChecks
|
||||
from litellm.proxy.auth.user_api_key_auth import (
|
||||
_check_key_model_budget_with_fallback,
|
||||
_get_team_model_aliases,
|
||||
_PendingAutoRegister,
|
||||
_matches_routing_override,
|
||||
_reserve_budget_after_common_checks,
|
||||
|
|
@ -1841,7 +1844,7 @@ async def test_standard_jwt_auth_propagates_user_identity_and_team_model_aliases
|
|||
find_model_table.assert_not_awaited()
|
||||
return
|
||||
|
||||
assert await _get_team_model_aliases(
|
||||
assert await get_team_model_aliases(
|
||||
model_id=model_id,
|
||||
prisma_client=prisma_client,
|
||||
user_api_key_cache=user_api_key_cache,
|
||||
|
|
@ -1863,7 +1866,7 @@ async def test_get_team_model_aliases_returns_none_when_model_row_missing():
|
|||
)
|
||||
|
||||
assert (
|
||||
await _get_team_model_aliases(
|
||||
await get_team_model_aliases(
|
||||
model_id=99,
|
||||
prisma_client=prisma_client,
|
||||
user_api_key_cache=UserApiKeyCache(),
|
||||
|
|
|
|||
|
|
@ -299,6 +299,10 @@ class TestDeleteTeamModelAlias:
|
|||
@pytest.mark.asyncio
|
||||
async def test_delete_team_model_alias_success(self):
|
||||
"""Test successful deletion of a team model alias"""
|
||||
from litellm.proxy.common_utils.user_api_key_cache import (
|
||||
UserApiKeyCache,
|
||||
team_model_aliases_cache_key,
|
||||
)
|
||||
from litellm.proxy.management_endpoints.model_management_endpoints import (
|
||||
delete_team_model_alias,
|
||||
)
|
||||
|
|
@ -328,10 +332,18 @@ class TestDeleteTeamModelAlias:
|
|||
# Create mock prisma client
|
||||
mock_prisma = MockPrismaClient(team_exists=True)
|
||||
mock_prisma.db = MockPrismaWrapper(model_aliases_list)
|
||||
user_api_key_cache = UserApiKeyCache()
|
||||
for model_id in (1, 2):
|
||||
await user_api_key_cache.async_set_cache(
|
||||
key=team_model_aliases_cache_key(model_id),
|
||||
value={"cached-alias": "cached-model"},
|
||||
)
|
||||
|
||||
# Call the function
|
||||
await delete_team_model_alias(
|
||||
public_model_name="public_model_1", prisma_client=mock_prisma
|
||||
public_model_name="public_model_1",
|
||||
prisma_client=mock_prisma,
|
||||
user_api_key_cache=user_api_key_cache,
|
||||
)
|
||||
|
||||
# Verify results
|
||||
|
|
@ -353,10 +365,16 @@ class TestDeleteTeamModelAlias:
|
|||
assert json.loads(second_update["data"]["model_aliases"]) == {
|
||||
"alias3": "public_model_3"
|
||||
}
|
||||
for model_id in (1, 2):
|
||||
assert await user_api_key_cache.async_get_cache(team_model_aliases_cache_key(model_id)) is None
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_delete_team_model_alias_no_matches(self):
|
||||
"""Test deletion when no matching model alias exists"""
|
||||
from litellm.proxy.common_utils.user_api_key_cache import (
|
||||
UserApiKeyCache,
|
||||
team_model_aliases_cache_key,
|
||||
)
|
||||
from litellm.proxy.management_endpoints.model_management_endpoints import (
|
||||
delete_team_model_alias,
|
||||
)
|
||||
|
|
@ -386,15 +404,24 @@ class TestDeleteTeamModelAlias:
|
|||
# Create mock prisma client
|
||||
mock_prisma = MockPrismaClient(team_exists=True)
|
||||
mock_prisma.db = MockPrismaWrapper(model_aliases_list)
|
||||
user_api_key_cache = UserApiKeyCache()
|
||||
cache_key = team_model_aliases_cache_key(1)
|
||||
await user_api_key_cache.async_set_cache(
|
||||
key=cache_key,
|
||||
value={"cached-alias": "cached-model"},
|
||||
)
|
||||
|
||||
# Call the function with non-existent model
|
||||
await delete_team_model_alias(
|
||||
public_model_name="non_existent_model", prisma_client=mock_prisma
|
||||
public_model_name="non_existent_model",
|
||||
prisma_client=mock_prisma,
|
||||
user_api_key_cache=user_api_key_cache,
|
||||
)
|
||||
|
||||
# Verify no updates were made
|
||||
mock_db = mock_prisma.db.litellm_modeltable
|
||||
assert len(mock_db.update_calls) == 0
|
||||
assert await user_api_key_cache.async_get_cache(cache_key) == {"cached-alias": "cached-model"}
|
||||
|
||||
|
||||
class TestClearCache:
|
||||
|
|
|
|||
|
|
@ -10261,10 +10261,19 @@ async def test_update_model_table_clears_aliases_with_empty_map():
|
|||
so existing aliases are cleared, while ``model_aliases=None`` must be a no-op that
|
||||
leaves the model table untouched.
|
||||
"""
|
||||
from litellm.proxy.common_utils.user_api_key_cache import (
|
||||
UserApiKeyCache,
|
||||
team_model_aliases_cache_key,
|
||||
)
|
||||
|
||||
model_id = 123
|
||||
user_api_key_cache = UserApiKeyCache()
|
||||
cache_key = team_model_aliases_cache_key(model_id)
|
||||
await user_api_key_cache.async_set_cache(key=cache_key, value={"old-alias": "old-model"})
|
||||
mock_prisma = MagicMock()
|
||||
mock_prisma.db.litellm_modeltable.create = AsyncMock()
|
||||
mock_prisma.db.litellm_modeltable.upsert = AsyncMock(
|
||||
return_value=MagicMock(id="model-123")
|
||||
return_value=MagicMock(id=model_id)
|
||||
)
|
||||
user_api_key_dict = UserAPIKeyAuth(
|
||||
user_role=LitellmUserRoles.PROXY_ADMIN, user_id="admin"
|
||||
|
|
@ -10272,33 +10281,38 @@ async def test_update_model_table_clears_aliases_with_empty_map():
|
|||
|
||||
returned_model_id = await _update_model_table(
|
||||
data=UpdateTeamRequest(team_id="team-1", model_aliases={}),
|
||||
model_id="model-123",
|
||||
model_id=model_id,
|
||||
prisma_client=mock_prisma,
|
||||
user_api_key_cache=user_api_key_cache,
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
litellm_proxy_admin_name="default_user_id",
|
||||
)
|
||||
|
||||
mock_prisma.db.litellm_modeltable.upsert.assert_awaited_once()
|
||||
upsert_kwargs = mock_prisma.db.litellm_modeltable.upsert.await_args.kwargs
|
||||
assert upsert_kwargs["where"] == {"id": "model-123"}
|
||||
assert upsert_kwargs["where"] == {"id": model_id}
|
||||
assert upsert_kwargs["data"]["update"]["model_aliases"] == json.dumps({})
|
||||
assert upsert_kwargs["data"]["create"]["model_aliases"] == json.dumps({})
|
||||
assert returned_model_id == "model-123"
|
||||
assert returned_model_id == model_id
|
||||
assert await user_api_key_cache.async_get_cache(cache_key) is None
|
||||
|
||||
mock_prisma.db.litellm_modeltable.create.reset_mock()
|
||||
mock_prisma.db.litellm_modeltable.upsert.reset_mock()
|
||||
await user_api_key_cache.async_set_cache(key=cache_key, value={"current-alias": "current-model"})
|
||||
|
||||
noop_model_id = await _update_model_table(
|
||||
data=UpdateTeamRequest(team_id="team-1", model_aliases=None),
|
||||
model_id="model-123",
|
||||
model_id=model_id,
|
||||
prisma_client=mock_prisma,
|
||||
user_api_key_cache=user_api_key_cache,
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
litellm_proxy_admin_name="default_user_id",
|
||||
)
|
||||
|
||||
mock_prisma.db.litellm_modeltable.create.assert_not_called()
|
||||
mock_prisma.db.litellm_modeltable.upsert.assert_not_called()
|
||||
assert noop_model_id == "model-123"
|
||||
assert noop_model_id == model_id
|
||||
assert await user_api_key_cache.async_get_cache(cache_key) == {"current-alias": "current-model"}
|
||||
|
||||
|
||||
class TestEmitTeamMembersMetric:
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue