mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-09 22:31:41 +00:00
fix(proxy): scope cache analytics to internal user
This commit is contained in:
parent
168a0055a2
commit
491b754cf5
5 changed files with 112 additions and 7 deletions
1
.github/workflows/test-unit-proxy-db.yml
vendored
1
.github/workflows/test-unit-proxy-db.yml
vendored
|
|
@ -182,6 +182,7 @@ jobs:
|
|||
- test-group: endpoints-and-responses
|
||||
test-path: >-
|
||||
tests/proxy_unit_tests/test_blog_posts_endpoint.py
|
||||
tests/proxy_unit_tests/test_cache_analytics.py
|
||||
tests/proxy_unit_tests/test_models_fallback_endpoint.py
|
||||
tests/proxy_unit_tests/test_google_endpoint_routing.py
|
||||
tests/proxy_unit_tests/test_google_gemini_proxy_request.py
|
||||
|
|
|
|||
|
|
@ -22,14 +22,25 @@ def _parse_date(value: str, param_name: str) -> datetime:
|
|||
)
|
||||
|
||||
|
||||
def _cache_activity_user_id(user_api_key_dict: UserAPIKeyAuth) -> str | None:
|
||||
if user_api_key_dict.user_role not in (
|
||||
LitellmUserRoles.INTERNAL_USER,
|
||||
LitellmUserRoles.INTERNAL_USER_VIEW_ONLY,
|
||||
):
|
||||
return None
|
||||
if user_api_key_dict.user_id is None:
|
||||
raise HTTPException(status_code=status.HTTP_400_BAD_REQUEST, detail={"error": "No user_id found"})
|
||||
return user_api_key_dict.user_id
|
||||
|
||||
|
||||
@router.get(
|
||||
"/global/activity/cache_hits",
|
||||
tags=["Budget & Spend Tracking"],
|
||||
dependencies=[Depends(user_api_key_auth)],
|
||||
response_model=CacheActivityResponse,
|
||||
include_in_schema=False,
|
||||
)
|
||||
async def get_global_activity(
|
||||
user_api_key_dict: Annotated[UserAPIKeyAuth, Depends(user_api_key_auth)],
|
||||
start_date: Annotated[str, fastapi.Query(description="Time from which to start viewing spend")],
|
||||
end_date: Annotated[str, fastapi.Query(description="Time till which to view spend")],
|
||||
key_aliases: Annotated[
|
||||
|
|
@ -52,10 +63,13 @@ async def get_global_activity(
|
|||
},
|
||||
)
|
||||
|
||||
user_id: Final = _cache_activity_user_id(user_api_key_dict)
|
||||
|
||||
return await get_cache_activity(
|
||||
prisma_client=prisma_client,
|
||||
start_date=_parse_date(start_date, "start_date"),
|
||||
end_date=_parse_date(end_date, "end_date"),
|
||||
key_aliases=key_aliases or [],
|
||||
models=models or [],
|
||||
user_id=user_id,
|
||||
)
|
||||
|
|
|
|||
|
|
@ -69,6 +69,7 @@ GROUPS_SQL: Final = """
|
|||
OR COALESCE(vt."key_alias", 'Unnamed Key') IN (SELECT jsonb_array_elements_text($3::jsonb)))
|
||||
AND ($4::jsonb = '[]'::jsonb
|
||||
OR sl."model" IN (SELECT jsonb_array_elements_text($4::jsonb)))
|
||||
AND ($5::text IS NULL OR sl."user" = $5)
|
||||
GROUP BY 1
|
||||
ORDER BY (COUNT(*)) DESC
|
||||
"""
|
||||
|
|
@ -89,6 +90,7 @@ ERROR_BREAKDOWN_SQL: Final = """
|
|||
OR COALESCE(vt."key_alias", 'Unnamed Key') IN (SELECT jsonb_array_elements_text($3::jsonb)))
|
||||
AND ($4::jsonb = '[]'::jsonb
|
||||
OR sl."model" IN (SELECT jsonb_array_elements_text($4::jsonb)))
|
||||
AND ($5::text IS NULL OR sl."user" = $5)
|
||||
GROUP BY 1, 2, 3
|
||||
ORDER BY (COUNT(*)) DESC
|
||||
"""
|
||||
|
|
@ -100,6 +102,7 @@ KEY_ALIAS_OPTIONS_SQL: Final = """
|
|||
WHERE
|
||||
sl."startTime" >= ($1::timestamptz AT TIME ZONE 'UTC')
|
||||
AND sl."startTime" < (($2::timestamptz + INTERVAL '1 day') AT TIME ZONE 'UTC')
|
||||
AND ($3::text IS NULL OR sl."user" = $3)
|
||||
ORDER BY 1
|
||||
"""
|
||||
|
||||
|
|
@ -110,6 +113,7 @@ MODEL_OPTIONS_SQL: Final = """
|
|||
sl."startTime" >= ($1::timestamptz AT TIME ZONE 'UTC')
|
||||
AND sl."startTime" < (($2::timestamptz + INTERVAL '1 day') AT TIME ZONE 'UTC')
|
||||
AND sl."model" != ''
|
||||
AND ($3::text IS NULL OR sl."user" = $3)
|
||||
ORDER BY 1
|
||||
"""
|
||||
|
||||
|
|
@ -148,14 +152,15 @@ async def get_cache_activity(
|
|||
end_date: datetime,
|
||||
key_aliases: Sequence[str],
|
||||
models: Sequence[str],
|
||||
user_id: str | None = None,
|
||||
) -> CacheActivityResponse:
|
||||
key_aliases_json: Final = json.dumps(list(key_aliases))
|
||||
models_json: Final = json.dumps(list(models))
|
||||
group_rows, error_rows, key_alias_rows, model_rows = await asyncio.gather(
|
||||
prisma_client.db.query_raw(GROUPS_SQL, start_date, end_date, key_aliases_json, models_json),
|
||||
prisma_client.db.query_raw(ERROR_BREAKDOWN_SQL, start_date, end_date, key_aliases_json, models_json),
|
||||
prisma_client.db.query_raw(KEY_ALIAS_OPTIONS_SQL, start_date, end_date),
|
||||
prisma_client.db.query_raw(MODEL_OPTIONS_SQL, start_date, end_date),
|
||||
prisma_client.db.query_raw(GROUPS_SQL, start_date, end_date, key_aliases_json, models_json, user_id),
|
||||
prisma_client.db.query_raw(ERROR_BREAKDOWN_SQL, start_date, end_date, key_aliases_json, models_json, user_id),
|
||||
prisma_client.db.query_raw(KEY_ALIAS_OPTIONS_SQL, start_date, end_date, user_id),
|
||||
prisma_client.db.query_raw(MODEL_OPTIONS_SQL, start_date, end_date, user_id),
|
||||
)
|
||||
groups: Final = _groups_adapter.validate_python(group_rows or [])
|
||||
return CacheActivityResponse(
|
||||
|
|
|
|||
|
|
@ -12,10 +12,13 @@ from unittest.mock import AsyncMock, MagicMock
|
|||
import pytest
|
||||
from fastapi import HTTPException
|
||||
|
||||
from litellm.proxy._types import LitellmUserRoles, UserAPIKeyAuth
|
||||
from litellm.proxy.analytics_endpoints.analytics_endpoints import get_global_activity
|
||||
from litellm.proxy.analytics_endpoints.cache_activity import (
|
||||
ERROR_BREAKDOWN_SQL,
|
||||
GROUPS_SQL,
|
||||
KEY_ALIAS_OPTIONS_SQL,
|
||||
MODEL_OPTIONS_SQL,
|
||||
CacheActivityGroup,
|
||||
compute_totals,
|
||||
)
|
||||
|
|
@ -73,9 +76,19 @@ def mock_prisma(monkeypatch: pytest.MonkeyPatch) -> MagicMock:
|
|||
return prisma
|
||||
|
||||
|
||||
def auth_for_role(role: LitellmUserRoles, user_id: str | None = "proxy-admin") -> UserAPIKeyAuth:
|
||||
return UserAPIKeyAuth(user_id=user_id, user_role=role)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_returns_groups_totals_and_filter_options(mock_prisma: MagicMock):
|
||||
response = await get_global_activity(start_date="2026-07-01", end_date="2026-07-27", key_aliases=[], models=[])
|
||||
response = await get_global_activity(
|
||||
start_date="2026-07-01",
|
||||
end_date="2026-07-27",
|
||||
key_aliases=[],
|
||||
models=[],
|
||||
user_api_key_dict=auth_for_role(LitellmUserRoles.PROXY_ADMIN),
|
||||
)
|
||||
|
||||
assert [group.call_type for group in response.groups] == ["acompletion", "Unknown"]
|
||||
assert response.groups[0].api_requests == 1000
|
||||
|
|
@ -101,6 +114,7 @@ async def test_filters_are_passed_to_sql_as_json_arrays(mock_prisma: MagicMock):
|
|||
end_date="2026-07-27",
|
||||
key_aliases=["my-key"],
|
||||
models=["gpt-5.1", "claude-opus-4-8"],
|
||||
user_api_key_dict=auth_for_role(LitellmUserRoles.PROXY_ADMIN),
|
||||
)
|
||||
|
||||
filtered_calls = [
|
||||
|
|
@ -112,10 +126,81 @@ async def test_filters_are_passed_to_sql_as_json_arrays(mock_prisma: MagicMock):
|
|||
assert call.args[4] == json.dumps(["gpt-5.1", "claude-opus-4-8"])
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_admin_roles_do_not_scope_cache_activity_to_user(mock_prisma: MagicMock):
|
||||
await get_global_activity(
|
||||
start_date="2026-07-01",
|
||||
end_date="2026-07-27",
|
||||
key_aliases=[],
|
||||
models=[],
|
||||
user_api_key_dict=auth_for_role(LitellmUserRoles.PROXY_ADMIN),
|
||||
)
|
||||
|
||||
groups_call = next(call for call in mock_prisma.db.query_raw.call_args_list if "GROUP BY" in call.args[0])
|
||||
assert 'sl."user" = $5' in groups_call.args[0]
|
||||
assert groups_call.args[5] is None
|
||||
assert all(call.args[-1] is None for call in mock_prisma.db.query_raw.call_args_list)
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"internal_role",
|
||||
[LitellmUserRoles.INTERNAL_USER, LitellmUserRoles.INTERNAL_USER_VIEW_ONLY],
|
||||
)
|
||||
@pytest.mark.asyncio
|
||||
async def test_cache_hits_activity_scopes_internal_roles_to_own_spend(
|
||||
mock_prisma: MagicMock,
|
||||
internal_role: LitellmUserRoles,
|
||||
):
|
||||
response = await get_global_activity(
|
||||
start_date="2026-07-01",
|
||||
end_date="2026-07-27",
|
||||
key_aliases=[],
|
||||
models=[],
|
||||
user_api_key_dict=auth_for_role(internal_role, user_id="internal-user"),
|
||||
)
|
||||
|
||||
assert response.groups
|
||||
groups_call = next(call for call in mock_prisma.db.query_raw.call_args_list if "GROUP BY" in call.args[0])
|
||||
error_call = next(call for call in mock_prisma.db.query_raw.call_args_list if call.args[0] == ERROR_BREAKDOWN_SQL)
|
||||
key_alias_call = next(
|
||||
call for call in mock_prisma.db.query_raw.call_args_list if call.args[0] == KEY_ALIAS_OPTIONS_SQL
|
||||
)
|
||||
model_call = next(call for call in mock_prisma.db.query_raw.call_args_list if call.args[0] == MODEL_OPTIONS_SQL)
|
||||
assert 'sl."user" = $5' in groups_call.args[0]
|
||||
assert groups_call.args[5] == "internal-user"
|
||||
assert 'sl."user" = $5' in error_call.args[0]
|
||||
assert error_call.args[5] == "internal-user"
|
||||
assert 'sl."user" = $3' in key_alias_call.args[0]
|
||||
assert key_alias_call.args[3] == "internal-user"
|
||||
assert 'sl."user" = $3' in model_call.args[0]
|
||||
assert model_call.args[3] == "internal-user"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_cache_hits_activity_rejects_internal_role_without_user_id(mock_prisma: MagicMock):
|
||||
with pytest.raises(HTTPException, match="No user_id found") as exc_info:
|
||||
await get_global_activity(
|
||||
start_date="2026-07-01",
|
||||
end_date="2026-07-27",
|
||||
key_aliases=[],
|
||||
models=[],
|
||||
user_api_key_dict=auth_for_role(LitellmUserRoles.INTERNAL_USER, user_id=None),
|
||||
)
|
||||
|
||||
assert exc_info.value.status_code == 400
|
||||
mock_prisma.db.query_raw.assert_not_called()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_rejects_malformed_dates_with_400(mock_prisma: MagicMock):
|
||||
with pytest.raises(HTTPException) as exc_info:
|
||||
await get_global_activity(start_date="07/01/2026", end_date="2026-07-27", key_aliases=[], models=[])
|
||||
await get_global_activity(
|
||||
start_date="07/01/2026",
|
||||
end_date="2026-07-27",
|
||||
key_aliases=[],
|
||||
models=[],
|
||||
user_api_key_dict=auth_for_role(LitellmUserRoles.PROXY_ADMIN),
|
||||
)
|
||||
|
||||
assert exc_info.value.status_code == 400
|
||||
mock_prisma.db.query_raw.assert_not_called()
|
||||
Loading…
Add table
Reference in a new issue