diff --git a/litellm/proxy/management_endpoints/common_daily_activity.py b/litellm/proxy/management_endpoints/common_daily_activity.py index bf8e7bc15fb..c4db06e8982 100644 --- a/litellm/proxy/management_endpoints/common_daily_activity.py +++ b/litellm/proxy/management_endpoints/common_daily_activity.py @@ -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, diff --git a/litellm/proxy/management_endpoints/internal_user_endpoints.py b/litellm/proxy/management_endpoints/internal_user_endpoints.py index 6509977f7ff..da228f8b9c5 100644 --- a/litellm/proxy/management_endpoints/internal_user_endpoints.py +++ b/litellm/proxy/management_endpoints/internal_user_endpoints.py @@ -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, ) diff --git a/tests/test_litellm/proxy/management_endpoints/test_common_daily_activity.py b/tests/test_litellm/proxy/management_endpoints/test_common_daily_activity.py index baaf3f4ba2f..7157d871afd 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_common_daily_activity.py +++ b/tests/test_litellm/proxy/management_endpoints/test_common_daily_activity.py @@ -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() diff --git a/tests/test_litellm/proxy/management_endpoints/test_internal_user_endpoints.py b/tests/test_litellm/proxy/management_endpoints/test_internal_user_endpoints.py index f684e2040bd..9b13cb2a37d 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_internal_user_endpoints.py +++ b/tests/test_litellm/proxy/management_endpoints/test_internal_user_endpoints.py @@ -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