mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-10 03:28:53 +00:00
fix(proxy): scope user activity by owned API keys
This commit is contained in:
parent
071cb49d32
commit
17ed4c4f7a
4 changed files with 203 additions and 24 deletions
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue