fix(proxy): scope user activity by owned API keys

This commit is contained in:
AxelRay 2026-09-22 21:24:51 +08:00
parent 071cb49d32
commit 17ed4c4f7a
4 changed files with 203 additions and 24 deletions

View file

@ -560,6 +560,41 @@ async def get_api_key_metadata(
return await attach_user_emails(prisma_client, combined)
async def _get_deleted_keys_for_user(
prisma_client: PrismaClient,
user_id: str,
) -> Sequence["PrismaDeletedVerificationToken"]:
try:
return await DeletedVerificationTokenRepository(prisma_client).table.find_many(
where={"user_id": user_id},
order={"deleted_at": "desc"},
)
except Exception as e:
verbose_proxy_logger.warning("Failed to fetch deleted key metadata for user %s: %s", user_id, e)
return ()
async def get_user_api_key_filter(
prisma_client: PrismaClient,
user_id: str,
api_key: str | None,
) -> list[str]:
"""Return the key digests that should scope a user's activity query."""
active_keys: Final = await VerificationTokenRepository(prisma_client).table.find_many(where={"user_id": user_id})
deleted_keys: Final = await _get_deleted_keys_for_user(prisma_client, user_id)
user_api_keys: Final = list(
dict.fromkeys(
key.token
for key in (*active_keys, *deleted_keys)
if getattr(key, "token", None) and getattr(key, "user_id", None) == user_id
)
)
if api_key is None:
return user_api_keys
return [api_key] if api_key in user_api_keys else []
def _adjust_dates_for_timezone(
start_date: str,
end_date: str,

View file

@ -54,6 +54,7 @@ from litellm.proxy.management_endpoints.common_daily_activity import (
DailySpendRecord,
get_daily_activity,
get_daily_activity_aggregated,
get_user_api_key_filter,
)
from litellm.proxy.management_endpoints.common_utils import (
_is_user_team_admin,
@ -3025,29 +3026,32 @@ async def get_user_daily_activity_aggregated(
try:
is_admin: Final = _user_has_admin_view(user_api_key_dict)
if is_admin:
entity_id = user_id # None means global view, otherwise filter by user
else:
caller_user_id: Final = require_caller_user_id_for_non_admin(user_api_key_dict)
if user_id is None:
user_id = caller_user_id
if user_id != caller_user_id:
raise HTTPException(
status_code=status.HTTP_403_FORBIDDEN,
detail={"error": "Non-admin users can only view their own spend data."},
)
entity_id = user_id
caller_user_id: Final[str | None] = (
None if is_admin else require_caller_user_id_for_non_admin(user_api_key_dict)
)
requested_user_id: Final[str | None] = user_id if is_admin or user_id is not None else caller_user_id
if caller_user_id is not None and requested_user_id != caller_user_id:
raise HTTPException(
status_code=status.HTTP_403_FORBIDDEN,
detail={"error": "Non-admin users can only view their own spend data."},
)
api_key_filter: Final = (
api_key
if requested_user_id is None
else await get_user_api_key_filter(prisma_client, requested_user_id, api_key)
)
return await get_daily_activity_aggregated(
prisma_client=prisma_client,
table_name="litellm_dailyuserspend",
entity_id_field="user_id",
entity_id=entity_id,
entity_id=None,
entity_metadata_field=None,
start_date=start_date,
end_date=end_date,
model=model,
api_key=api_key,
api_key=api_key_filter,
timezone_offset_minutes=timezone,
include_current_utc_day=include_current_utc_day,
)

View file

@ -25,6 +25,7 @@ from litellm.proxy.management_endpoints.common_daily_activity import (
get_api_key_metadata,
get_daily_activity,
get_daily_activity_aggregated,
get_user_api_key_filter,
global_rollup_reconciled_through,
update_metrics,
)
@ -331,6 +332,20 @@ async def test_get_api_key_metadata_returns_active_key_metadata():
assert result["active-key-hash-123"]["team_id"] == "team-abc"
@pytest.mark.asyncio
async def test_get_user_api_key_filter_scopes_to_active_and_deleted_user_keys():
mock_prisma = MagicMock()
active_key = SimpleNamespace(token="active-key", user_id="target-user")
unrelated_key = SimpleNamespace(token="unrelated-key", user_id="other-user")
deleted_key = SimpleNamespace(token="deleted-key", user_id="target-user")
mock_prisma.db.litellm_verificationtoken.find_many = AsyncMock(return_value=[active_key, unrelated_key])
mock_prisma.db.litellm_deletedverificationtoken.find_many = AsyncMock(return_value=[deleted_key])
assert await get_user_api_key_filter(mock_prisma, "target-user", None) == ["active-key", "deleted-key"]
assert await get_user_api_key_filter(mock_prisma, "target-user", "unrelated-key") == []
@pytest.mark.asyncio
async def test_get_api_key_metadata_falls_back_to_deleted_keys():
"""Test that get_api_key_metadata should fall back to deleted keys table for missing keys."""
@ -1663,6 +1678,83 @@ async def test_get_daily_activity_aggregated_explicit_api_key_filter_scopes_both
assert set(day.breakdown.models["gpt-5"].api_key_breakdown) == {"key-1"}
@pytest.mark.asyncio
async def test_get_daily_activity_aggregated_user_key_filter_keeps_usage_with_legacy_user_id(
_aggregated_postgresql: psycopg.Connection,
):
rows: Final = [
(
"owned-row",
"legacy-user-id",
"2026-09-22",
"owned-key",
"gpt-5",
"",
"openai",
"/v1/chat/completions",
10,
7.0,
1,
1,
),
(
"deleted-row",
"another-legacy-user-id",
"2026-09-22",
"deleted-key",
"gpt-5",
"",
"openai",
"/v1/chat/completions",
10,
3.0,
1,
1,
),
(
"unrelated-row",
"other-user-id",
"2026-09-22",
"unrelated-key",
"gpt-5",
"",
"openai",
"/v1/chat/completions",
10,
100.0,
1,
1,
),
]
_seed_daily_user_spend(_aggregated_postgresql, rows)
mock_prisma = MagicMock()
mock_prisma.db.query_raw = _psycopg_query_raw(_aggregated_postgresql, [])
mock_prisma.db.litellm_verificationtoken.find_many = AsyncMock(
return_value=[
SimpleNamespace(token="owned-key", key_alias="owned", team_id=None, user_id="target-user"),
SimpleNamespace(token="deleted-key", key_alias="deleted", team_id=None, user_id="target-user"),
]
)
mock_prisma.db.litellm_usertable.find_many = AsyncMock(return_value=[])
result = await get_daily_activity_aggregated(
prisma_client=mock_prisma,
table_name="litellm_dailyuserspend",
entity_id_field="user_id",
entity_id=None,
entity_metadata_field=None,
start_date="2026-09-22",
end_date="2026-09-22",
model=None,
api_key=["owned-key", "deleted-key"],
)
assert result.metadata.total_spend == 10.0
assert set(result.results[0].breakdown.api_keys) == {"owned-key", "deleted-key"}
assert "unrelated-key" not in result.results[0].breakdown.api_keys
def _prisma_with_marker(marker: str | None) -> MagicMock:
prisma = MagicMock()
prisma.db = MagicMock()

View file

@ -2523,6 +2523,48 @@ async def test_get_user_daily_activity_aggregated_admin_global_view(monkeypatch,
)
@pytest.mark.asyncio
async def test_get_user_daily_activity_aggregated_filters_by_user_owned_keys(monkeypatch):
from unittest.mock import AsyncMock, MagicMock
from litellm.proxy.management_endpoints.internal_user_endpoints import (
get_user_daily_activity_aggregated,
)
mock_prisma_client = MagicMock()
monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma_client)
mock_response = MagicMock()
mock_get_user_api_key_filter = AsyncMock(return_value=["owned-key"])
monkeypatch.setattr(
"litellm.proxy.management_endpoints.internal_user_endpoints.get_user_api_key_filter",
mock_get_user_api_key_filter,
)
mock_get_daily_agg = AsyncMock(return_value=mock_response)
monkeypatch.setattr(
"litellm.proxy.management_endpoints.internal_user_endpoints.get_daily_activity_aggregated",
mock_get_daily_agg,
)
result = await get_user_daily_activity_aggregated(
start_date="2026-09-22",
end_date="2026-09-22",
model=None,
api_key=None,
user_id="target-user",
timezone=None,
user_api_key_dict=UserAPIKeyAuth(
user_id="admin-user",
user_role=LitellmUserRoles.PROXY_ADMIN,
),
)
assert result is mock_response
mock_get_user_api_key_filter.assert_awaited_once_with(mock_prisma_client, "target-user", None)
assert mock_get_daily_agg.call_args.kwargs["entity_id"] is None
assert mock_get_daily_agg.call_args.kwargs["api_key"] == ["owned-key"]
@pytest.mark.asyncio
async def test_get_user_daily_activity_aggregated_non_admin_cannot_view_other_users(
monkeypatch,
@ -2577,19 +2619,25 @@ async def test_get_user_daily_activity_aggregated_non_admin_cannot_view_other_us
new_callable=AsyncMock,
return_value=mock_response,
) as mock_get_daily_agg:
result = await get_user_daily_activity_aggregated(
start_date="2025-01-01",
end_date="2025-01-31",
model=None,
api_key=None,
user_id=None,
timezone=None,
user_api_key_dict=non_admin_key_dict,
)
with patch(
"litellm.proxy.management_endpoints.internal_user_endpoints.get_user_api_key_filter",
new_callable=AsyncMock,
return_value=["regular-user-key"],
):
result = await get_user_daily_activity_aggregated(
start_date="2025-01-01",
end_date="2025-01-31",
model=None,
api_key=None,
user_id=None,
timezone=None,
user_api_key_dict=non_admin_key_dict,
)
assert result is mock_response
mock_get_daily_agg.assert_called_once()
assert mock_get_daily_agg.call_args.kwargs["entity_id"] == "regular-user-123"
assert mock_get_daily_agg.call_args.kwargs["entity_id"] is None
assert mock_get_daily_agg.call_args.kwargs["api_key"] == ["regular-user-key"]
@pytest.mark.asyncio