mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-03 02:22:24 +00:00
fix(proxy): reject non-canonical daily activity dates (#44143)
The daily spend tables store date as text, so a request date that strptime accepts but that is not spelled YYYY-MM-DD (2026-9-24, 2026-09-4, full-width digits) was compared as raw text against canonical rows and matched nothing, and the export route copied it into Content-Disposition, which fails latin-1 encoding and returned 500. A shared parse_canonical_date_range now rejects any spelling whose round trip differs from the input, so every bounded daily activity route (user, team, tag, organization, customer, agent: aggregated, aggregated/keys, search, model top keys, export, cache leakage) and the paginated get_daily_activity path answer 400 before touching the repository, and the export filename is built from the validated dates. Co-authored-by: yassin <yassin@berri.ai> Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
parent
c208306f60
commit
5ccb1a143b
6 changed files with 201 additions and 28 deletions
|
|
@ -2,7 +2,7 @@ import asyncio
|
|||
from collections.abc import Awaitable, Callable, Mapping, Sequence
|
||||
from collections.abc import Set as AbstractSet
|
||||
from dataclasses import dataclass, replace
|
||||
from datetime import datetime, timedelta
|
||||
from datetime import date, datetime, timedelta
|
||||
from types import MappingProxyType
|
||||
from typing import Final, Literal, NoReturn, Protocol
|
||||
|
||||
|
|
@ -66,6 +66,33 @@ def raise_public(error: ScopeDenied | InvalidDateRange) -> NoReturn:
|
|||
assert_never(error)
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class CanonicalDateRange:
|
||||
start: date
|
||||
end: date
|
||||
|
||||
|
||||
def parse_canonical_date(value: str) -> date | None:
|
||||
"""The daily spend tables store ``date`` as text and compare it against the raw request
|
||||
string, so only the exact ``YYYY-MM-DD`` spelling can match a row. Spellings the parser
|
||||
would normalise (``2026-9-24``, ``20260924``, full-width digits) are rejected instead."""
|
||||
try:
|
||||
parsed: Final = date.fromisoformat(value)
|
||||
except ValueError:
|
||||
return None
|
||||
return parsed if parsed.isoformat() == value else None
|
||||
|
||||
|
||||
def parse_canonical_date_range(start_date: str | None, end_date: str | None) -> CanonicalDateRange | InvalidDateRange:
|
||||
if start_date is None or end_date is None:
|
||||
return InvalidDateRange(reason="Please provide start_date and end_date")
|
||||
start: Final = parse_canonical_date(start_date)
|
||||
end: Final = parse_canonical_date(end_date)
|
||||
if start is None or end is None:
|
||||
return InvalidDateRange(reason="start_date and end_date must be valid YYYY-MM-DD dates")
|
||||
return CanonicalDateRange(start=start, end=end)
|
||||
|
||||
|
||||
class DailySpendRecord(Protocol):
|
||||
@property
|
||||
def date(self) -> str: ...
|
||||
|
|
@ -877,10 +904,9 @@ async def get_daily_activity(
|
|||
) -> SpendAnalyticsPaginatedResponse:
|
||||
if prisma_client is None:
|
||||
raise HTTPException(status_code=500, detail={"error": CommonProxyErrors.db_not_connected_error.value})
|
||||
if start_date is None or end_date is None:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_400_BAD_REQUEST, detail={"error": "Please provide start_date and end_date"}
|
||||
)
|
||||
date_range: Final = parse_canonical_date_range(start_date, end_date)
|
||||
if isinstance(date_range, InvalidDateRange):
|
||||
raise_public(date_range)
|
||||
try:
|
||||
scope: Final = daily_activity_scope(
|
||||
table_name,
|
||||
|
|
@ -888,8 +914,8 @@ async def get_daily_activity(
|
|||
entity_id,
|
||||
exclude_entity_ids,
|
||||
api_key,
|
||||
start_date,
|
||||
end_date,
|
||||
date_range.start.isoformat(),
|
||||
date_range.end.isoformat(),
|
||||
model,
|
||||
timezone_offset_minutes,
|
||||
include_current_utc_day,
|
||||
|
|
|
|||
|
|
@ -3,7 +3,7 @@ import io
|
|||
import json
|
||||
from collections.abc import AsyncIterator, Mapping, Sequence
|
||||
from dataclasses import asdict, fields, replace
|
||||
from datetime import datetime
|
||||
from datetime import date, datetime
|
||||
from typing import Annotated, Final, Literal
|
||||
|
||||
from fastapi import APIRouter, Depends, HTTPException, Query
|
||||
|
|
@ -19,6 +19,7 @@ from litellm.proxy.management_endpoints.common_daily_activity import (
|
|||
ScopeDenied,
|
||||
daily_activity_repository,
|
||||
get_daily_activity_aggregated,
|
||||
parse_canonical_date_range,
|
||||
raise_public,
|
||||
spend_logs_window,
|
||||
)
|
||||
|
|
@ -69,9 +70,8 @@ def get_daily_activity_repository() -> DailyActivityRepository:
|
|||
|
||||
def _date_range_error(query: EntityQuery, *, user_aggregated: bool) -> InvalidDateRange | None:
|
||||
if user_aggregated:
|
||||
if query.start_date is None or query.end_date is None:
|
||||
return InvalidDateRange(reason="Please provide start_date and end_date")
|
||||
return None
|
||||
date_range: Final = parse_canonical_date_range(query.start_date, query.end_date)
|
||||
return date_range if isinstance(date_range, InvalidDateRange) else None
|
||||
|
||||
range_error: Final[str | None] = aggregated_date_range_error(query.start_date, query.end_date)
|
||||
return None if range_error is None else InvalidDateRange(reason=range_error)
|
||||
|
|
@ -164,19 +164,19 @@ async def _key_activity_rows(
|
|||
|
||||
def _export_filename(
|
||||
entity: str,
|
||||
start_date: str,
|
||||
end_date: str,
|
||||
start_date: date,
|
||||
end_date: date,
|
||||
export_type: ExportType,
|
||||
file_format: Literal["csv", "json"],
|
||||
) -> str:
|
||||
extension: Final[str] = "csv" if file_format == "csv" else "json"
|
||||
return f"{entity}-usage-{start_date}-{end_date}-{export_type.value}.{extension}"
|
||||
return f"{entity}-usage-{start_date.isoformat()}-{end_date.isoformat()}-{export_type.value}.{extension}"
|
||||
|
||||
|
||||
def _content_disposition(
|
||||
entity: str,
|
||||
start_date: str,
|
||||
end_date: str,
|
||||
start_date: date,
|
||||
end_date: date,
|
||||
export_type: ExportType,
|
||||
file_format: Literal["csv", "json"],
|
||||
) -> str:
|
||||
|
|
@ -479,8 +479,8 @@ def _register_export_route(router: APIRouter, resolver: EntityScopeResolver, pre
|
|||
"Cache-Control": "no-store",
|
||||
"Content-Disposition": _content_disposition(
|
||||
resolver.entity,
|
||||
resolved.scope.start_date,
|
||||
resolved.scope.end_date,
|
||||
date.fromisoformat(resolved.scope.start_date),
|
||||
date.fromisoformat(resolved.scope.end_date),
|
||||
export_type,
|
||||
file_format,
|
||||
),
|
||||
|
|
|
|||
|
|
@ -124,6 +124,10 @@ from litellm.proxy.hooks.model_max_budget_limiter import (
|
|||
)
|
||||
from litellm.proxy.management.teams.access import TEAM_OR_ORG_ADMIN, TeamRole, is_team_admin, team_access_denied
|
||||
from litellm.proxy.management.teams.dependencies import get_team_access
|
||||
from litellm.proxy.management_endpoints.common_daily_activity import (
|
||||
InvalidDateRange,
|
||||
parse_canonical_date_range,
|
||||
)
|
||||
from litellm.proxy.management_endpoints.common_utils import (
|
||||
_check_disable_global_guardrails_caller_permission,
|
||||
_check_passthrough_routes_caller_permission,
|
||||
|
|
@ -6691,16 +6695,12 @@ _MAX_AGGREGATED_RANGE_DAYS: Final = 400
|
|||
def aggregated_date_range_error(start_date: str | None, end_date: str | None) -> str | None:
|
||||
"""The aggregated endpoint has no pagination to bound its work, so malformed
|
||||
dates and ranges wider than the UI ever requests are rejected before querying."""
|
||||
if start_date is None or end_date is None:
|
||||
return "Please provide start_date and end_date"
|
||||
try:
|
||||
parsed_start: Final = datetime.strptime(start_date, "%Y-%m-%d").replace(tzinfo=timezone.utc)
|
||||
parsed_end: Final = datetime.strptime(end_date, "%Y-%m-%d").replace(tzinfo=timezone.utc)
|
||||
except ValueError:
|
||||
return "start_date and end_date must be valid YYYY-MM-DD dates"
|
||||
if parsed_end < parsed_start:
|
||||
date_range: Final = parse_canonical_date_range(start_date, end_date)
|
||||
if isinstance(date_range, InvalidDateRange):
|
||||
return date_range.reason
|
||||
if date_range.end < date_range.start:
|
||||
return "end_date must be on or after start_date"
|
||||
if (parsed_end - parsed_start).days > _MAX_AGGREGATED_RANGE_DAYS:
|
||||
if (date_range.end - date_range.start).days > _MAX_AGGREGATED_RANGE_DAYS:
|
||||
return f"Date range must be at most {_MAX_AGGREGATED_RANGE_DAYS} days"
|
||||
return None
|
||||
|
||||
|
|
|
|||
|
|
@ -1,5 +1,5 @@
|
|||
from collections.abc import Mapping, Sequence
|
||||
from datetime import datetime
|
||||
from datetime import date, datetime
|
||||
from types import SimpleNamespace
|
||||
from typing import Final
|
||||
from unittest.mock import AsyncMock, MagicMock
|
||||
|
|
@ -10,6 +10,7 @@ from fastapi import HTTPException
|
|||
import litellm.proxy.management_endpoints.common_daily_activity as common_daily_activity_module
|
||||
from litellm.constants import USAGE_TOP_API_KEYS_DEFAULT
|
||||
from litellm.proxy.management_endpoints.common_daily_activity import (
|
||||
CanonicalDateRange,
|
||||
InvalidDateRange,
|
||||
_is_user_agent_tag,
|
||||
_ProxyDailyActivityReads,
|
||||
|
|
@ -19,6 +20,8 @@ from litellm.proxy.management_endpoints.common_daily_activity import (
|
|||
daily_activity_scope,
|
||||
get_api_key_metadata,
|
||||
get_daily_activity,
|
||||
parse_canonical_date,
|
||||
parse_canonical_date_range,
|
||||
raise_public,
|
||||
update_metrics,
|
||||
)
|
||||
|
|
@ -2456,3 +2459,56 @@ def test_raise_public_maps_invalid_date_range_to_400() -> None:
|
|||
raise_public(InvalidDateRange(reason="Date range must be at most 400 days"))
|
||||
assert excinfo.value.status_code == 400
|
||||
assert excinfo.value.detail == {"error": "Date range must be at most 400 days"}
|
||||
|
||||
|
||||
@pytest.mark.parametrize("value", ("2026-9-24", "2026-09-24", "2026-09-4", "2026-02-30", "20260924", ""))
|
||||
def test_parse_canonical_date_rejects_spellings_that_do_not_round_trip(value: str) -> None:
|
||||
assert parse_canonical_date(value) is None
|
||||
|
||||
|
||||
def test_parse_canonical_date_accepts_the_exact_yyyy_mm_dd_spelling() -> None:
|
||||
assert parse_canonical_date("2026-09-24") == date(2026, 9, 24)
|
||||
assert parse_canonical_date("0001-01-01") == date(1, 1, 1)
|
||||
|
||||
|
||||
def test_parse_canonical_date_range_reports_missing_then_malformed_dates() -> None:
|
||||
assert parse_canonical_date_range(None, "2026-09-24") == InvalidDateRange(
|
||||
reason="Please provide start_date and end_date"
|
||||
)
|
||||
assert parse_canonical_date_range("2026-09-24", "2026-9-26") == InvalidDateRange(
|
||||
reason="start_date and end_date must be valid YYYY-MM-DD dates"
|
||||
)
|
||||
assert parse_canonical_date_range("2026-09-24", "2026-09-26") == CanonicalDateRange(
|
||||
start=date(2026, 9, 24), end=date(2026, 9, 26)
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("start_date", ("2026-9-24", "2026-09-24", "2026-09-4"))
|
||||
async def test_get_daily_activity_rejects_non_canonical_dates_before_querying(start_date: str) -> None:
|
||||
mock_prisma = MagicMock()
|
||||
mock_prisma.db = MagicMock()
|
||||
mock_table = MagicMock()
|
||||
mock_table.count = AsyncMock(return_value=0)
|
||||
mock_table.find_many = AsyncMock(return_value=[])
|
||||
mock_prisma.db.litellm_dailyteamspend = mock_table
|
||||
|
||||
with pytest.raises(HTTPException) as error:
|
||||
await get_daily_activity(
|
||||
prisma_client=mock_prisma,
|
||||
table_name="litellm_dailyteamspend",
|
||||
entity_id_field="team_id",
|
||||
entity_id="team-a",
|
||||
entity_metadata_field=None,
|
||||
start_date=start_date,
|
||||
end_date="2026-09-26",
|
||||
model=None,
|
||||
api_key=None,
|
||||
page=1,
|
||||
page_size=10,
|
||||
)
|
||||
|
||||
assert error.value.status_code == 400
|
||||
assert error.value.detail == {"error": "start_date and end_date must be valid YYYY-MM-DD dates"}
|
||||
mock_table.count.assert_not_awaited()
|
||||
mock_table.find_many.assert_not_awaited()
|
||||
|
|
|
|||
|
|
@ -1210,6 +1210,9 @@ def test_user_aggregate_keeps_current_day_query_semantics(
|
|||
("0000-01-01", "9999-12-31", "valid YYYY-MM-DD"),
|
||||
("2024-06-01", "2024-01-01", "on or after"),
|
||||
("not-a-date", "2024-01-31", "valid YYYY-MM-DD"),
|
||||
("2026-9-24", "2026-09-26", "valid YYYY-MM-DD"),
|
||||
("2026-09-24", "2026-09-26", "valid YYYY-MM-DD"),
|
||||
("2026-09-01", "2026-09-4", "valid YYYY-MM-DD"),
|
||||
(None, "2024-01-31", "start_date and end_date"),
|
||||
),
|
||||
)
|
||||
|
|
@ -1259,3 +1262,71 @@ def test_user_key_page_rejects_bad_date_ranges(
|
|||
assert response.status_code == 400, response.text
|
||||
assert message in str(response.json()["detail"]), response.text
|
||||
repository.key_page.assert_not_awaited()
|
||||
|
||||
|
||||
_NON_CANONICAL_DATE_RANGES: Final[tuple[tuple[str, str], ...]] = (
|
||||
("2026-9-24", "2026-09-26"),
|
||||
("2026-09-24", "2026-09-26"),
|
||||
("2026-09-01", "2026-09-4"),
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.parametrize(("start_date", "end_date"), _NON_CANONICAL_DATE_RANGES)
|
||||
def test_user_aggregate_rejects_non_canonical_dates(
|
||||
daily_activity_client: tuple[TestClient, _FakeRepository], start_date: str, end_date: str
|
||||
) -> None:
|
||||
client, repository = daily_activity_client
|
||||
response: Final = client.get(
|
||||
"/user/daily/activity/aggregated",
|
||||
params={"start_date": start_date, "end_date": end_date, "user_id": "user-a"},
|
||||
)
|
||||
assert response.status_code == 400, response.text
|
||||
assert response.json()["detail"] == {"error": "start_date and end_date must be valid YYYY-MM-DD dates"}
|
||||
repository.aggregated.assert_not_awaited()
|
||||
|
||||
|
||||
def test_user_aggregate_still_accepts_ranges_wider_than_the_team_limit(
|
||||
daily_activity_client: tuple[TestClient, _FakeRepository],
|
||||
) -> None:
|
||||
client, repository = daily_activity_client
|
||||
response: Final = client.get(
|
||||
"/user/daily/activity/aggregated",
|
||||
params={"start_date": "2020-01-01", "end_date": "2026-12-31", "user_id": "user-a"},
|
||||
)
|
||||
assert response.status_code == 200, response.text
|
||||
repository.aggregated.assert_awaited_once()
|
||||
|
||||
|
||||
@pytest.mark.parametrize(("start_date", "end_date"), _NON_CANONICAL_DATE_RANGES)
|
||||
@pytest.mark.parametrize(("prefix", "query_name", "entity_id"), _ENTITY_CASES)
|
||||
def test_export_routes_reject_non_canonical_dates_before_querying(
|
||||
daily_activity_client: tuple[TestClient, _FakeRepository],
|
||||
prefix: str,
|
||||
query_name: str,
|
||||
entity_id: str,
|
||||
start_date: str,
|
||||
end_date: str,
|
||||
) -> None:
|
||||
client, repository = daily_activity_client
|
||||
repository.export_rows_error = AssertionError("export must not query the repository")
|
||||
response: Final = client.get(
|
||||
f"{prefix}/daily/activity/export",
|
||||
params={query_name: entity_id, "start_date": start_date, "end_date": end_date, "export_type": "daily"},
|
||||
)
|
||||
assert response.status_code == 400, response.text
|
||||
assert response.json()["detail"] == {"error": "start_date and end_date must be valid YYYY-MM-DD dates"}
|
||||
assert "content-disposition" not in response.headers
|
||||
|
||||
|
||||
def test_export_content_disposition_is_ascii_and_built_from_canonical_dates(
|
||||
daily_activity_client: tuple[TestClient, _FakeRepository],
|
||||
) -> None:
|
||||
client, _ = daily_activity_client
|
||||
response: Final = client.get(
|
||||
"/team/daily/activity/export",
|
||||
params={**_entity_params("team_ids", "team-a"), "export_type": ExportType.DAILY.value},
|
||||
)
|
||||
assert response.status_code == 200, response.text
|
||||
disposition: Final = response.headers["content-disposition"]
|
||||
assert disposition == 'attachment; filename="team-usage-2025-01-01-2025-01-02-daily.csv"'
|
||||
assert disposition.isascii()
|
||||
|
|
|
|||
|
|
@ -53,6 +53,7 @@ from litellm.proxy.management_endpoints.team_endpoints import (
|
|||
_update_model_table,
|
||||
_validate_and_populate_member_user_info,
|
||||
_validate_team_member_reset_spend_value,
|
||||
aggregated_date_range_error,
|
||||
delete_team,
|
||||
list_available_teams,
|
||||
reset_team_member_budget_fn,
|
||||
|
|
@ -16819,3 +16820,22 @@ def test_list_team_v2_answers_503_no_db_connection_when_the_callers_user_read_hi
|
|||
|
||||
assert response.status_code == 503, response.text
|
||||
assert response.json() == _DB_OUTAGE_503_BODY
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
("start_date", "end_date"),
|
||||
(
|
||||
("2026-9-24", "2026-09-26"),
|
||||
("2026-09-24", "2026-09-26"),
|
||||
("2026-09-01", "2026-09-4"),
|
||||
("2026-02-30", "2026-09-26"),
|
||||
),
|
||||
)
|
||||
def test_aggregated_date_range_error_rejects_non_canonical_dates(start_date: str, end_date: str) -> None:
|
||||
assert aggregated_date_range_error(start_date, end_date) == "start_date and end_date must be valid YYYY-MM-DD dates"
|
||||
|
||||
|
||||
def test_aggregated_date_range_error_accepts_canonical_dates_and_keeps_range_checks() -> None:
|
||||
assert aggregated_date_range_error("2026-09-24", "2026-09-26") is None
|
||||
assert aggregated_date_range_error("2026-09-26", "2026-09-24") == "end_date must be on or after start_date"
|
||||
assert aggregated_date_range_error("2020-01-01", "2026-12-31") == "Date range must be at most 400 days"
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue