mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-04 02:31:27 +00:00
Revert "fix(proxy): skip v1 team filter when user row is missing"
This reverts commit 74e1fbd77a.
This commit is contained in:
parent
74e1fbd77a
commit
2c08b393dc
2 changed files with 13 additions and 114 deletions
|
|
@ -10998,17 +10998,12 @@ async def get_all_team_and_direct_access_models(
|
|||
_model["model_info"]["direct_access"] = True
|
||||
|
||||
## FILTER OUT MODELS THAT ARE NOT IN DIRECT_ACCESS_MODELS OR ACCESS_VIA_TEAM_IDS - only show user models they can call
|
||||
should_filter_by_user_access = (
|
||||
user_api_key_dict.user_role == LitellmUserRoles.PROXY_ADMIN
|
||||
or user_teams is not None
|
||||
)
|
||||
if should_filter_by_user_access:
|
||||
all_models = [
|
||||
_model
|
||||
for _model in all_models
|
||||
if _model.get("model_info", {}).get("direct_access", False)
|
||||
or _model.get("model_info", {}).get("access_via_team_ids", [])
|
||||
]
|
||||
all_models = [
|
||||
_model
|
||||
for _model in all_models
|
||||
if _model.get("model_info", {}).get("direct_access", False)
|
||||
or _model.get("model_info", {}).get("access_via_team_ids", [])
|
||||
]
|
||||
return all_models
|
||||
|
||||
|
||||
|
|
@ -12455,19 +12450,14 @@ def _filter_v1_model_info_deployments(
|
|||
]
|
||||
|
||||
|
||||
async def _should_apply_v1_team_access_filter(
|
||||
def _should_apply_v1_team_access_filter(
|
||||
user_api_key_dict: UserAPIKeyAuth,
|
||||
prisma_client: PrismaClient,
|
||||
) -> bool:
|
||||
"""Team membership filtering requires admin role or a DB-backed user row."""
|
||||
if user_api_key_dict.user_role == LitellmUserRoles.PROXY_ADMIN:
|
||||
return True
|
||||
if user_api_key_dict.user_id is None:
|
||||
return False
|
||||
user_db_object = await UserRepository(prisma_client).table.find_unique(
|
||||
where={"user_id": user_api_key_dict.user_id}
|
||||
"""Team membership filtering requires a resolvable user or admin role."""
|
||||
return (
|
||||
user_api_key_dict.user_role == LitellmUserRoles.PROXY_ADMIN
|
||||
or user_api_key_dict.user_id is not None
|
||||
)
|
||||
return user_db_object is not None
|
||||
|
||||
|
||||
def _translate_model_name_for_response(model: dict) -> dict:
|
||||
|
|
@ -12649,9 +12639,8 @@ async def model_info_v1( # noqa: PLR0915
|
|||
llm_router=llm_router,
|
||||
)
|
||||
|
||||
if prisma_client is not None and await _should_apply_v1_team_access_filter(
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
prisma_client=prisma_client,
|
||||
if prisma_client is not None and _should_apply_v1_team_access_filter(
|
||||
user_api_key_dict=user_api_key_dict
|
||||
):
|
||||
all_models = await get_all_team_and_direct_access_models(
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
|
|
|
|||
|
|
@ -270,12 +270,6 @@ async def test_model_info_v1_restricted_key_filters_after_team_enrichment(monkey
|
|||
router.get_model_names.return_value = ["gpt-4", "team-claude-sonnet"]
|
||||
router.get_model_access_groups.return_value = {}
|
||||
|
||||
class MockUserRepository:
|
||||
def __init__(self, prisma_client):
|
||||
self.table = MagicMock()
|
||||
self.table.find_unique = AsyncMock(return_value=MagicMock(model_dump=lambda: {"user_id": "user-1", "teams": []}))
|
||||
|
||||
monkeypatch.setattr(ps, "UserRepository", MockUserRepository)
|
||||
monkeypatch.setattr(ps, "user_model", None)
|
||||
monkeypatch.setattr(ps, "llm_model_list", router.model_list)
|
||||
monkeypatch.setattr(ps, "llm_router", router)
|
||||
|
|
@ -320,87 +314,3 @@ async def test_model_info_v1_restricted_key_filters_after_team_enrichment(monkey
|
|||
resp = await ps.model_info_v1(user_api_key_dict=caller, litellm_model_id=None)
|
||||
|
||||
assert [m["model_name"] for m in resp["data"]] == ["gpt-4"]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_model_info_v1_missing_db_user_returns_deployments(monkeypatch):
|
||||
"""Keys with user_id but no DB row must not lose their model list."""
|
||||
deployment = {
|
||||
"model_name": "gpt-4",
|
||||
"litellm_params": {"model": "gpt-4"},
|
||||
"model_info": {"id": "global-id-1", "db_model": False},
|
||||
}
|
||||
router = MagicMock()
|
||||
router.model_list = [deployment]
|
||||
router.get_model_names.return_value = ["gpt-4"]
|
||||
router.get_model_access_groups.return_value = {}
|
||||
|
||||
class MockUserRepository:
|
||||
def __init__(self, prisma_client):
|
||||
self.table = MagicMock()
|
||||
self.table.find_unique = AsyncMock(return_value=None)
|
||||
|
||||
get_team_access = AsyncMock()
|
||||
monkeypatch.setattr(ps, "UserRepository", MockUserRepository)
|
||||
monkeypatch.setattr(ps, "user_model", None)
|
||||
monkeypatch.setattr(ps, "llm_model_list", router.model_list)
|
||||
monkeypatch.setattr(ps, "llm_router", router)
|
||||
monkeypatch.setattr(ps, "prisma_client", MagicMock())
|
||||
monkeypatch.setattr(ps, "get_all_team_and_direct_access_models", get_team_access)
|
||||
monkeypatch.setattr(
|
||||
ps, "_enrich_model_info_with_litellm_data", lambda model, **kw: model
|
||||
)
|
||||
import litellm.proxy.agent_endpoints.model_list_helpers as mlh
|
||||
|
||||
monkeypatch.setattr(
|
||||
mlh,
|
||||
"append_agents_to_model_info",
|
||||
AsyncMock(side_effect=lambda models, **kw: models),
|
||||
)
|
||||
|
||||
caller = UserAPIKeyAuth(
|
||||
user_id="missing-user",
|
||||
user_role=LitellmUserRoles.INTERNAL_USER,
|
||||
models=[],
|
||||
team_models=[],
|
||||
)
|
||||
resp = await ps.model_info_v1(user_api_key_dict=caller, litellm_model_id=None)
|
||||
|
||||
assert [m["model_name"] for m in resp["data"]] == ["gpt-4"]
|
||||
get_team_access.assert_not_called()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_get_all_team_and_direct_access_models_missing_user_skips_filter(
|
||||
monkeypatch,
|
||||
):
|
||||
"""Defense in depth: unresolved user_id must not empty the model list."""
|
||||
models = [
|
||||
{
|
||||
"model_name": "gpt-4",
|
||||
"litellm_params": {"model": "gpt-4"},
|
||||
"model_info": {"id": "global-id-1"},
|
||||
}
|
||||
]
|
||||
|
||||
class MockUserRepository:
|
||||
def __init__(self, prisma_client):
|
||||
self.table = MagicMock()
|
||||
self.table.find_unique = AsyncMock(return_value=None)
|
||||
|
||||
monkeypatch.setattr(ps, "UserRepository", MockUserRepository)
|
||||
|
||||
caller = UserAPIKeyAuth(
|
||||
user_id="missing-user",
|
||||
user_role=LitellmUserRoles.INTERNAL_USER,
|
||||
models=[],
|
||||
team_models=[],
|
||||
)
|
||||
result = await ps.get_all_team_and_direct_access_models(
|
||||
user_api_key_dict=caller,
|
||||
prisma_client=MagicMock(),
|
||||
llm_router=MagicMock(),
|
||||
all_models=models,
|
||||
)
|
||||
|
||||
assert result == models
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue