mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-19 00:01:29 +00:00
test(proxy): share the jwt key mapping test doubles across the deletion endpoint tests
Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
parent
e79d03e604
commit
761e0775b9
4 changed files with 77 additions and 103 deletions
|
|
@ -0,0 +1,27 @@
|
|||
"""LiteLLM_JWTKeyMapping test doubles for the bulk key deletion paths."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from collections.abc import Mapping, Sequence
|
||||
from dataclasses import dataclass
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class JWTMappingRow:
|
||||
token: str
|
||||
jwt_claim_name: str
|
||||
jwt_claim_value: str
|
||||
jwt_issuer: str | None = None
|
||||
|
||||
|
||||
class CascadingJWTMappingTable:
|
||||
"""Mapping rows that LiteLLM_JWTKeyMapping_token_fkey drops when their key row is deleted."""
|
||||
|
||||
def __init__(self, rows: Sequence[JWTMappingRow]) -> None:
|
||||
self.rows: tuple[JWTMappingRow, ...] = tuple(rows)
|
||||
|
||||
async def find_many(self, where: Mapping[str, Mapping[str, Sequence[str]]]) -> list[JWTMappingRow]:
|
||||
return [row for row in self.rows if row.token in where["token"]["in"]]
|
||||
|
||||
def cascade(self, deleted_tokens: Sequence[str]) -> None:
|
||||
self.rows = tuple(row for row in self.rows if row.token not in deleted_tokens)
|
||||
|
|
@ -4,11 +4,10 @@ from types import SimpleNamespace
|
|||
from typing import Final
|
||||
|
||||
import pytest
|
||||
from fastapi.testclient import TestClient
|
||||
from fastapi import HTTPException
|
||||
from fastapi.testclient import TestClient
|
||||
from pytest_mock import MockerFixture
|
||||
|
||||
|
||||
from litellm.proxy._types import (
|
||||
LiteLLM_UserTableFiltered,
|
||||
LitellmUserRoles,
|
||||
|
|
@ -27,6 +26,10 @@ from litellm.proxy.management_endpoints.internal_user_endpoints import (
|
|||
ui_view_users,
|
||||
)
|
||||
from litellm.proxy.proxy_server import app
|
||||
from tests.test_litellm.proxy.management_endpoints.jwt_key_mapping_doubles import (
|
||||
CascadingJWTMappingTable,
|
||||
JWTMappingRow,
|
||||
)
|
||||
|
||||
client = TestClient(app)
|
||||
|
||||
|
|
@ -2679,27 +2682,6 @@ async def test_delete_user_cleans_up_created_by_invitation_links(mocker):
|
|||
assert condition[field] == {"in": ["admin-creator"]}
|
||||
|
||||
|
||||
class _JWTMappingRow:
|
||||
def __init__(self, token, jwt_claim_name, jwt_claim_value, jwt_issuer=None):
|
||||
self.token = token
|
||||
self.jwt_claim_name = jwt_claim_name
|
||||
self.jwt_claim_value = jwt_claim_value
|
||||
self.jwt_issuer = jwt_issuer
|
||||
|
||||
|
||||
class _CascadingJWTMappingTable:
|
||||
"""Mapping rows that LiteLLM_JWTKeyMapping_token_fkey drops when their key row is deleted."""
|
||||
|
||||
def __init__(self, rows):
|
||||
self.rows = rows
|
||||
|
||||
async def find_many(self, where, **kwargs):
|
||||
return [row for row in self.rows if row.token in where["token"]["in"]]
|
||||
|
||||
def cascade(self, deleted_tokens):
|
||||
self.rows = [row for row in self.rows if row.token not in deleted_tokens]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_delete_user_evicts_jwt_key_mapping_cache_of_its_keys(mocker):
|
||||
"""/user/delete bulk-deletes the user's keys without going through /key/delete, so the
|
||||
|
|
@ -2719,11 +2701,11 @@ async def test_delete_user_evicts_jwt_key_mapping_cache_of_its_keys(mocker):
|
|||
global_cache_key: Final = jwt_key_mapping_cache_key("sub", "jwt-user", None)
|
||||
issuer_cache_key: Final = jwt_key_mapping_cache_key("sub", "jwt-user", "https://issuer.example")
|
||||
unrelated_cache_key: Final = jwt_key_mapping_cache_key("sub", "other-user", None)
|
||||
jwt_table: Final = _CascadingJWTMappingTable(
|
||||
jwt_table: Final = CascadingJWTMappingTable(
|
||||
[
|
||||
_JWTMappingRow("hashed-jwt-key", "sub", "jwt-user"),
|
||||
_JWTMappingRow("hashed-issuer-key", "sub", "jwt-user", "https://issuer.example"),
|
||||
_JWTMappingRow("hashed-unrelated-key", "sub", "other-user"),
|
||||
JWTMappingRow("hashed-jwt-key", "sub", "jwt-user"),
|
||||
JWTMappingRow("hashed-issuer-key", "sub", "jwt-user", "https://issuer.example"),
|
||||
JWTMappingRow("hashed-unrelated-key", "sub", "other-user"),
|
||||
]
|
||||
)
|
||||
cache: Final = UserApiKeyCache()
|
||||
|
|
|
|||
|
|
@ -1,6 +1,5 @@
|
|||
import asyncio
|
||||
import json
|
||||
from litellm._uuid import uuid
|
||||
from types import MappingProxyType, SimpleNamespace
|
||||
from typing import Final, Mapping, Optional, cast
|
||||
from unittest.mock import AsyncMock, MagicMock, patch
|
||||
|
|
@ -9,6 +8,11 @@ import pytest
|
|||
from fastapi import HTTPException
|
||||
from fastapi.testclient import TestClient
|
||||
|
||||
from litellm._uuid import uuid
|
||||
from tests.test_litellm.proxy.management_endpoints.jwt_key_mapping_doubles import (
|
||||
CascadingJWTMappingTable,
|
||||
JWTMappingRow,
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
|
|
@ -499,9 +503,10 @@ async def test_organization_info_includes_user_email(monkeypatch):
|
|||
"""
|
||||
Test that GET /organization/info returns user_email in members list.
|
||||
"""
|
||||
from litellm.proxy._types import LiteLLM_OrganizationMembershipTable
|
||||
from datetime import datetime
|
||||
|
||||
from litellm.proxy._types import LiteLLM_OrganizationMembershipTable
|
||||
|
||||
# Simulate a membership row with a nested user object that has user_email
|
||||
raw_membership = {
|
||||
"user_id": "user_abc",
|
||||
|
|
@ -573,6 +578,10 @@ async def test_organization_member_add_rejects_unauthorized_caller(patched_org_p
|
|||
# ``organization_member_add`` catches HTTPException in its
|
||||
# catch-all and re-wraps as ProxyException with the original status
|
||||
# code preserved.
|
||||
from unittest.mock import Mock
|
||||
|
||||
from fastapi import Request
|
||||
|
||||
from litellm.proxy._types import (
|
||||
OrganizationMemberAddRequest,
|
||||
OrgMember,
|
||||
|
|
@ -581,9 +590,6 @@ async def test_organization_member_add_rejects_unauthorized_caller(patched_org_p
|
|||
from litellm.proxy.management_endpoints.organization_endpoints import (
|
||||
organization_member_add,
|
||||
)
|
||||
from unittest.mock import Mock
|
||||
|
||||
from fastapi import Request
|
||||
|
||||
data = OrganizationMemberAddRequest(
|
||||
organization_id="org-victim",
|
||||
|
|
@ -1440,27 +1446,6 @@ def test_organization_routes_reach_their_handler_with_enterprise_license(monkeyp
|
|||
)
|
||||
|
||||
|
||||
class _JWTMappingRow:
|
||||
def __init__(self, token, jwt_claim_name, jwt_claim_value, jwt_issuer=None):
|
||||
self.token = token
|
||||
self.jwt_claim_name = jwt_claim_name
|
||||
self.jwt_claim_value = jwt_claim_value
|
||||
self.jwt_issuer = jwt_issuer
|
||||
|
||||
|
||||
class _CascadingJWTMappingTable:
|
||||
"""Mapping rows that LiteLLM_JWTKeyMapping_token_fkey drops when their key row is deleted."""
|
||||
|
||||
def __init__(self, rows):
|
||||
self.rows = rows
|
||||
|
||||
async def find_many(self, where, **kwargs):
|
||||
return [row for row in self.rows if row.token in where["token"]["in"]]
|
||||
|
||||
def cascade(self, deleted_tokens):
|
||||
self.rows = [row for row in self.rows if row.token not in deleted_tokens]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_delete_organization_evicts_the_cache_of_the_keys_it_deletes(monkeypatch):
|
||||
"""/organization/delete bulk-deletes the org's keys without going through /key/delete, so the
|
||||
|
|
@ -1479,11 +1464,11 @@ async def test_delete_organization_evicts_the_cache_of_the_keys_it_deletes(monke
|
|||
jwt_key_mapping_cache_key("sub", "svc-account", "https://issuer.example"),
|
||||
)
|
||||
kept_cache_keys: Final = ("hashed-other-key", jwt_key_mapping_cache_key("sub", "other-account", None))
|
||||
kept_row: Final = _JWTMappingRow("hashed-other-key", "sub", "other-account")
|
||||
jwt_table: Final = _CascadingJWTMappingTable(
|
||||
kept_row: Final = JWTMappingRow("hashed-other-key", "sub", "other-account")
|
||||
jwt_table: Final = CascadingJWTMappingTable(
|
||||
[
|
||||
_JWTMappingRow("hashed-org-key", "sub", "svc-account"),
|
||||
_JWTMappingRow("hashed-org-key", "sub", "svc-account", "https://issuer.example"),
|
||||
JWTMappingRow("hashed-org-key", "sub", "svc-account"),
|
||||
JWTMappingRow("hashed-org-key", "sub", "svc-account", "https://issuer.example"),
|
||||
kept_row,
|
||||
]
|
||||
)
|
||||
|
|
@ -1516,4 +1501,4 @@ async def test_delete_organization_evicts_the_cache_of_the_keys_it_deletes(monke
|
|||
|
||||
assert all(cache.get_cache(key=cache_key) is None for cache_key in doomed_cache_keys)
|
||||
assert all(cache.get_cache(key=cache_key) == {"retained": True} for cache_key in kept_cache_keys)
|
||||
assert jwt_table.rows == [kept_row]
|
||||
assert jwt_table.rows == (kept_row,)
|
||||
|
|
|
|||
|
|
@ -12,8 +12,6 @@ from fastapi.testclient import TestClient
|
|||
from pydantic import ValidationError
|
||||
|
||||
from litellm._uuid import uuid
|
||||
|
||||
from litellm.proxy._types import UserAPIKeyAuth # Import UserAPIKeyAuth
|
||||
from litellm.proxy._types import (
|
||||
LiteLLM_BudgetTable,
|
||||
LiteLLM_BudgetTableFull,
|
||||
|
|
@ -33,14 +31,12 @@ from litellm.proxy._types import (
|
|||
TeamMemberAddRequest,
|
||||
TeamMemberUpdateRequest,
|
||||
UpdateTeamRequest,
|
||||
UserAPIKeyAuth, # Import UserAPIKeyAuth
|
||||
)
|
||||
from litellm.proxy.management_endpoints.team_endpoints import (
|
||||
user_api_key_auth, # Assuming this dependency is needed
|
||||
)
|
||||
from litellm.proxy.management_endpoints.team_endpoints import (
|
||||
_STRIP_DELETED_TEAM_FROM_USERS_SQL,
|
||||
GetTeamMemberPermissionsResponse,
|
||||
UpdateTeamMemberPermissionsRequest,
|
||||
_STRIP_DELETED_TEAM_FROM_USERS_SQL,
|
||||
_persist_deleted_team_records,
|
||||
_save_deleted_team_records,
|
||||
_transform_teams_to_deleted_records,
|
||||
|
|
@ -56,6 +52,7 @@ from litellm.proxy.management_endpoints.team_endpoints import (
|
|||
team_member_delete,
|
||||
team_member_update,
|
||||
update_team,
|
||||
user_api_key_auth, # Assuming this dependency is needed
|
||||
validate_team_org_change,
|
||||
)
|
||||
from litellm.proxy.management_helpers.access_group_team_sync import (
|
||||
|
|
@ -71,6 +68,10 @@ from litellm.types.proxy.management_endpoints.team_endpoints import (
|
|||
BulkTeamMemberAddResponse,
|
||||
TeamMemberAddResult,
|
||||
)
|
||||
from tests.test_litellm.proxy.management_endpoints.jwt_key_mapping_doubles import (
|
||||
CascadingJWTMappingTable,
|
||||
JWTMappingRow,
|
||||
)
|
||||
|
||||
# Setup TestClient
|
||||
client = TestClient(app)
|
||||
|
|
@ -2788,7 +2789,7 @@ async def test_upsert_team_member_budget_table_existing_budget():
|
|||
"""
|
||||
from unittest.mock import AsyncMock, MagicMock, patch
|
||||
|
||||
from litellm.proxy._types import LitellmUserRoles, LiteLLM_TeamTable, UserAPIKeyAuth
|
||||
from litellm.proxy._types import LiteLLM_TeamTable, LitellmUserRoles, UserAPIKeyAuth
|
||||
from litellm.proxy.management_endpoints.team_endpoints import (
|
||||
TeamMemberBudgetHandler,
|
||||
)
|
||||
|
|
@ -2849,7 +2850,7 @@ async def test_upsert_team_member_budget_table_no_existing_budget():
|
|||
"""
|
||||
from unittest.mock import AsyncMock, MagicMock, patch
|
||||
|
||||
from litellm.proxy._types import LitellmUserRoles, LiteLLM_TeamTable, UserAPIKeyAuth
|
||||
from litellm.proxy._types import LiteLLM_TeamTable, LitellmUserRoles, UserAPIKeyAuth
|
||||
from litellm.proxy.management_endpoints.team_endpoints import (
|
||||
TeamMemberBudgetHandler,
|
||||
)
|
||||
|
|
@ -6092,9 +6093,9 @@ async def test_new_team_standalone_validates_against_user_models(monkeypatch):
|
|||
- Team is created WITHOUT organization_id and models=['gpt-4']
|
||||
- Expected: Should fail with "Model not in allowed user models"
|
||||
"""
|
||||
import litellm
|
||||
from fastapi import Request
|
||||
|
||||
import litellm
|
||||
from litellm.proxy._types import NewTeamRequest, ProxyException, UserAPIKeyAuth
|
||||
from litellm.proxy.management_endpoints.team_endpoints import new_team
|
||||
|
||||
|
|
@ -9180,27 +9181,6 @@ async def test_team_member_delete_persists_deleted_keys(monkeypatch):
|
|||
assert cache.get_cache(key="unrelated-key") == {"retained": True}
|
||||
|
||||
|
||||
class _JWTMappingRow:
|
||||
def __init__(self, token, jwt_claim_name, jwt_claim_value, jwt_issuer=None):
|
||||
self.token = token
|
||||
self.jwt_claim_name = jwt_claim_name
|
||||
self.jwt_claim_value = jwt_claim_value
|
||||
self.jwt_issuer = jwt_issuer
|
||||
|
||||
|
||||
class _CascadingJWTMappingTable:
|
||||
"""Mapping rows that LiteLLM_JWTKeyMapping_token_fkey drops when their key row is deleted."""
|
||||
|
||||
def __init__(self, rows):
|
||||
self.rows = rows
|
||||
|
||||
async def find_many(self, where, **kwargs):
|
||||
return [row for row in self.rows if row.token in where["token"]["in"]]
|
||||
|
||||
def cascade(self, deleted_tokens):
|
||||
self.rows = [row for row in self.rows if row.token not in deleted_tokens]
|
||||
|
||||
|
||||
def _seed_jwt_mapping_cache(cache, mapping_rows):
|
||||
from litellm.proxy.auth.auth_checks import jwt_key_mapping_cache_key
|
||||
|
||||
|
|
@ -9224,11 +9204,11 @@ async def test_team_member_delete_evicts_jwt_key_mapping_cache_of_the_keys_it_de
|
|||
from litellm.proxy.management_endpoints.key_management_endpoints import LiteLLM_VerificationToken
|
||||
|
||||
doomed_rows: Final = (
|
||||
_JWTMappingRow("hashed-token-1", "sub", "user-123"),
|
||||
_JWTMappingRow("hashed-token-1", "sub", "user-123", "https://issuer.example"),
|
||||
JWTMappingRow("hashed-token-1", "sub", "user-123"),
|
||||
JWTMappingRow("hashed-token-1", "sub", "user-123", "https://issuer.example"),
|
||||
)
|
||||
kept_row: Final = _JWTMappingRow("hashed-other-key", "sub", "user-999")
|
||||
jwt_table: Final = _CascadingJWTMappingTable([*doomed_rows, kept_row])
|
||||
kept_row: Final = JWTMappingRow("hashed-other-key", "sub", "user-999")
|
||||
jwt_table: Final = CascadingJWTMappingTable([*doomed_rows, kept_row])
|
||||
|
||||
team = LiteLLM_TeamTable(
|
||||
team_id="team-1",
|
||||
|
|
@ -9271,7 +9251,7 @@ async def test_team_member_delete_evicts_jwt_key_mapping_cache_of_the_keys_it_de
|
|||
|
||||
assert all(cache.get_cache(key=cache_key) is None for cache_key in doomed_cache_keys)
|
||||
assert cache.get_cache(key=kept_cache_key) == "hashed-other-key"
|
||||
assert jwt_table.rows == [kept_row]
|
||||
assert jwt_table.rows == (kept_row,)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
|
|
@ -9286,11 +9266,11 @@ async def test_delete_team_evicts_jwt_key_mapping_cache_of_the_keys_it_deletes(
|
|||
from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache
|
||||
|
||||
doomed_rows: Final = (
|
||||
_JWTMappingRow("hashed-doomed-key", "sub", "svc-account"),
|
||||
_JWTMappingRow("hashed-doomed-key", "sub", "svc-account", "https://issuer.example"),
|
||||
JWTMappingRow("hashed-doomed-key", "sub", "svc-account"),
|
||||
JWTMappingRow("hashed-doomed-key", "sub", "svc-account", "https://issuer.example"),
|
||||
)
|
||||
kept_row: Final = _JWTMappingRow("hashed-unrelated-key", "sub", "svc-account", "https://other-issuer.example")
|
||||
jwt_table: Final = _CascadingJWTMappingTable([*doomed_rows, kept_row])
|
||||
kept_row: Final = JWTMappingRow("hashed-unrelated-key", "sub", "svc-account", "https://other-issuer.example")
|
||||
jwt_table: Final = CascadingJWTMappingTable([*doomed_rows, kept_row])
|
||||
|
||||
team = LiteLLM_TeamTable(
|
||||
team_id="team-doomed",
|
||||
|
|
@ -9346,7 +9326,7 @@ async def test_delete_team_evicts_jwt_key_mapping_cache_of_the_keys_it_deletes(
|
|||
|
||||
assert all(cache.get_cache(key=cache_key) is None for cache_key in doomed_cache_keys)
|
||||
assert cache.get_cache(key=kept_cache_key) == "hashed-unrelated-key"
|
||||
assert jwt_table.rows == [kept_row]
|
||||
assert jwt_table.rows == (kept_row,)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
|
|
@ -10570,7 +10550,7 @@ def test_new_team_request_accepts_team_member_budget_duration():
|
|||
async def test_create_team_member_budget_table_with_duration():
|
||||
"""Verify that create_team_member_budget_table passes budget_duration
|
||||
through to the new_budget call when team_member_budget_duration is provided."""
|
||||
from litellm.proxy._types import NewTeamRequest, UserAPIKeyAuth, LitellmUserRoles
|
||||
from litellm.proxy._types import LitellmUserRoles, NewTeamRequest, UserAPIKeyAuth
|
||||
from litellm.proxy.management_endpoints.team_endpoints import (
|
||||
TeamMemberBudgetHandler,
|
||||
)
|
||||
|
|
@ -11060,7 +11040,7 @@ async def test_team_member_me_matches_email_only_member(mock_db_client):
|
|||
@pytest.mark.asyncio
|
||||
async def test_team_member_me_returns_404_for_non_member(mock_db_client):
|
||||
"""A user who is not a member of the team gets 404, regardless of role."""
|
||||
from fastapi import Request, HTTPException
|
||||
from fastapi import HTTPException, Request
|
||||
|
||||
from litellm.proxy.management_endpoints.team_endpoints import team_member_me
|
||||
|
||||
|
|
@ -11094,7 +11074,7 @@ async def test_team_member_me_returns_404_for_proxy_admin_not_in_team(
|
|||
Proxy admins get 404 if they are not actually a member of the team.
|
||||
`me` only resolves for actual team members; admins use /team/info instead.
|
||||
"""
|
||||
from fastapi import Request, HTTPException
|
||||
from fastapi import HTTPException, Request
|
||||
|
||||
from litellm.proxy.management_endpoints.team_endpoints import team_member_me
|
||||
|
||||
|
|
@ -11155,7 +11135,7 @@ async def test_team_member_me_returns_defaults_when_no_membership_row(mock_db_cl
|
|||
@pytest.mark.asyncio
|
||||
async def test_team_member_me_rejects_team_key_without_user_id(mock_db_client):
|
||||
"""A team key with no user_id can't resolve 'me' — must return 400."""
|
||||
from fastapi import Request, HTTPException
|
||||
from fastapi import HTTPException, Request
|
||||
|
||||
from litellm.proxy.management_endpoints.team_endpoints import team_member_me
|
||||
|
||||
|
|
@ -11173,7 +11153,7 @@ async def test_team_member_me_rejects_team_key_without_user_id(mock_db_client):
|
|||
@pytest.mark.asyncio
|
||||
async def test_team_member_me_returns_404_for_unknown_team(mock_db_client):
|
||||
"""Unknown team_id returns 404 — propagated from get_team_object."""
|
||||
from fastapi import Request, HTTPException
|
||||
from fastapi import HTTPException, Request
|
||||
|
||||
from litellm.proxy.management_endpoints.team_endpoints import team_member_me
|
||||
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue