Merge pull request #35288 from BerriAI/litellm_/jovial-heisenberg-68775c
Some checks are pending
CodSpeed Benchmarks / benchmarks (push) Waiting to run
UI Unit Tests / ui-unit-tests (push) Waiting to run
GitHub Actions Security Analysis / zizmor (push) Waiting to run

fix(tag-management): drop unsupported prisma select kwarg from key lookup
This commit is contained in:
yuneng-jiang 2026-07-31 17:32:36 -07:00 • committed by GitHub
commit b5cfc2ca00
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
2 changed files with 93 additions and 21 deletions

View file

@ -96,7 +96,6 @@ class _VerificationTokenTableClient(Protocol):
async def find_many(
self,
where: Mapping[str, object] | None = None,
select: Mapping[str, object] | None = None,
) -> "Sequence[PrismaVerificationToken]": ...
@ -157,7 +156,6 @@ async def _get_internal_user_api_keys(
key_records = await _table(VerificationTokenRepository(prisma_client)).find_many(
where={"user_id": user_id},
select={"token": True},
)
user_api_keys.update(key_record.token for key_record in key_records if getattr(key_record, "token", None))

View file

@ -1,17 +1,20 @@
import inspect
import json
import os
import sys
from typing import Any, Dict, Optional
from collections.abc import Sequence
from typing import Optional
import pytest
from fastapi import HTTPException
from fastapi.testclient import TestClient
from prisma.actions import LiteLLM_VerificationTokenActions
sys.path.insert(
0, os.path.abspath("../../../..")
) # Adds the parent directory to the system path
from unittest.mock import patch
from unittest.mock import Mock, patch
import litellm
from litellm.proxy._types import LitellmUserRoles, UserAPIKeyAuth
@ -21,6 +24,28 @@ from litellm.types.tag_management import TagDeleteRequest, TagInfoRequest, TagNe
client = TestClient(app)
class FakeVerificationTokenTable:
"""Stand-in for ``prisma_client.db.litellm_verificationtoken``.
``AsyncMock`` swallows any keyword argument, so a plain mock cannot catch a
call that the generated prisma client would reject at runtime. This double
binds every call against the real ``find_many`` signature, so passing an
unsupported kwarg (e.g. ``select``) raises the same ``TypeError`` the proxy
surfaces as an HTTP 500.
"""
def __init__(self, records: Sequence[Mock]):
self._records = tuple(records)
self.calls: list[dict[str, object]] = []
async def find_many(self, **kwargs: object) -> tuple[Mock, ...]:
inspect.signature(LiteLLM_VerificationTokenActions.find_many).bind(
self, **kwargs
)
self.calls.append(kwargs)
return self._records
@pytest.mark.asyncio
async def test_create_and_get_tag():
"""
@ -380,6 +405,7 @@ async def test_list_tags_no_dynamic_tags():
app.dependency_overrides.clear()
@pytest.mark.asyncio
async def test_internal_user_list_tags_only_returns_tags_used_by_their_keys():
"""
Internal users can view tag usage, but the tag list must be scoped to tags
@ -404,9 +430,8 @@ async def test_internal_user_list_tags_only_returns_tags_used_by_their_keys():
owned_key_record = Mock()
owned_key_record.token = "owned-key"
mock_db.litellm_verificationtoken.find_many = AsyncMock(
return_value=[owned_key_record]
)
fake_token_table = FakeVerificationTokenTable([owned_key_record])
mock_db.litellm_verificationtoken = fake_token_table
mock_db.litellm_dailytagspend.group_by = AsyncMock(
return_value=[
@ -446,10 +471,9 @@ async def test_internal_user_list_tags_only_returns_tags_used_by_their_keys():
"stored-owned-tag",
"dynamic-owned-tag",
]
mock_db.litellm_verificationtoken.find_many.assert_awaited_once_with(
where={"user_id": "internal-user-123"},
select={"token": True},
)
assert fake_token_table.calls == [
{"where": {"user_id": "internal-user-123"}}
]
mock_db.litellm_dailytagspend.group_by.assert_awaited_once_with(
by=["tag"],
where={
@ -468,6 +492,54 @@ async def test_internal_user_list_tags_only_returns_tags_used_by_their_keys():
app.dependency_overrides.clear()
@pytest.mark.asyncio
async def test_internal_user_list_tags_does_not_500_on_unsupported_prisma_kwarg():
"""
Regression: /tag/list returned 500 for every internal user because the
non-admin branch looked up the caller's keys with
``find_many(select={"token": True})``, and the generated prisma client has no
``select`` kwarg. This reproduces the reported case exactly: a freshly created
internal user with no tag spend yet, which must get an empty 200 rather than
"LiteLLM_VerificationTokenActions.find_many() got an unexpected keyword
argument 'select'".
"""
from unittest.mock import AsyncMock, Mock
from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
mock_user_auth = UserAPIKeyAuth(
api_key="new-user-key",
user_id="brand-new-internal-user",
user_role=LitellmUserRoles.INTERNAL_USER,
)
app.dependency_overrides[user_api_key_auth] = lambda: mock_user_auth
try:
with patch("litellm.proxy.proxy_server.prisma_client") as mock_prisma:
mock_db = Mock()
mock_prisma.db = mock_db
key_record = Mock()
key_record.token = "new-user-key"
fake_token_table = FakeVerificationTokenTable([key_record])
mock_db.litellm_verificationtoken = fake_token_table
mock_db.litellm_dailytagspend.group_by = AsyncMock(return_value=[])
mock_db.litellm_tagtable.find_many = AsyncMock(return_value=[])
response = client.get(
"/tag/list", headers={"Authorization": "Bearer new-user-key"}
)
assert response.status_code == 200, response.text
assert response.json() == []
assert fake_token_table.calls == [
{"where": {"user_id": "brand-new-internal-user"}}
]
finally:
app.dependency_overrides.clear()
@pytest.mark.asyncio
async def test_list_tags_with_date_range_filters_dynamic_tags():
"""
@ -537,9 +609,8 @@ async def test_internal_user_tag_daily_activity_is_scoped_to_their_keys():
owned_key_record = Mock()
owned_key_record.token = "owned-key"
mock_db.litellm_verificationtoken.find_many = AsyncMock(
return_value=[owned_key_record]
)
fake_token_table = FakeVerificationTokenTable([owned_key_record])
mock_db.litellm_verificationtoken = fake_token_table
mock_get_daily_activity.return_value = "daily-activity-response"
result = await get_tag_daily_activity(
@ -549,6 +620,7 @@ async def test_internal_user_tag_daily_activity_is_scoped_to_their_keys():
)
assert result == "daily-activity-response"
assert fake_token_table.calls == [{"where": {"user_id": "internal-user-123"}}]
mock_get_daily_activity.assert_awaited_once()
assert mock_get_daily_activity.await_args.kwargs["api_key"] == ["owned-key"]
@ -583,9 +655,8 @@ async def test_internal_user_tag_daily_activity_rejects_unowned_api_key_filter()
owned_key_record = Mock()
owned_key_record.token = "owned-key"
mock_db.litellm_verificationtoken.find_many = AsyncMock(
return_value=[owned_key_record]
)
fake_token_table = FakeVerificationTokenTable([owned_key_record])
mock_db.litellm_verificationtoken = fake_token_table
result = await get_tag_daily_activity(
start_date="2025-01-01",
end_date="2025-01-31",
@ -593,6 +664,7 @@ async def test_internal_user_tag_daily_activity_rejects_unowned_api_key_filter()
user_api_key_dict=mock_user_auth,
)
assert fake_token_table.calls == [{"where": {"user_id": "internal-user-123"}}]
assert result.results == []
assert result.metadata.total_spend == 0
assert result.metadata.total_api_requests == 0
@ -626,7 +698,8 @@ async def test_internal_user_tag_daily_activity_scopes_to_current_key_without_us
):
mock_db = Mock()
mock_prisma.db = mock_db
mock_db.litellm_verificationtoken.find_many = AsyncMock(return_value=[])
fake_token_table = FakeVerificationTokenTable([])
mock_db.litellm_verificationtoken = fake_token_table
mock_get_daily_activity.return_value = "daily-activity-response"
result = await get_tag_daily_activity(
@ -636,7 +709,7 @@ async def test_internal_user_tag_daily_activity_scopes_to_current_key_without_us
)
assert result == "daily-activity-response"
mock_db.litellm_verificationtoken.find_many.assert_not_awaited()
assert fake_token_table.calls == []
mock_get_daily_activity.assert_awaited_once()
assert mock_get_daily_activity.await_args.kwargs["api_key"] == [
"current-owned-key"
@ -669,7 +742,8 @@ async def test_internal_user_tag_daily_activity_without_any_scoped_keys_returns_
):
mock_db = Mock()
mock_prisma.db = mock_db
mock_db.litellm_verificationtoken.find_many = AsyncMock(return_value=[])
fake_token_table = FakeVerificationTokenTable([])
mock_db.litellm_verificationtoken = fake_token_table
result = await get_tag_daily_activity(
start_date="2025-01-01",
@ -680,7 +754,7 @@ async def test_internal_user_tag_daily_activity_without_any_scoped_keys_returns_
assert result.results == []
assert result.metadata.total_spend == 0
assert result.metadata.total_api_requests == 0
mock_db.litellm_verificationtoken.find_many.assert_not_awaited()
assert fake_token_table.calls == []
mock_get_daily_activity.assert_not_awaited()