diff --git a/litellm/proxy/management_endpoints/common_daily_activity.py b/litellm/proxy/management_endpoints/common_daily_activity.py index e13d623c73b..27b0960c823 100644 --- a/litellm/proxy/management_endpoints/common_daily_activity.py +++ b/litellm/proxy/management_endpoints/common_daily_activity.py @@ -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, diff --git a/litellm/proxy/management_endpoints/daily_activity_routes.py b/litellm/proxy/management_endpoints/daily_activity_routes.py index 065f23c3829..fa3c745c88e 100644 --- a/litellm/proxy/management_endpoints/daily_activity_routes.py +++ b/litellm/proxy/management_endpoints/daily_activity_routes.py @@ -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, ), diff --git a/litellm/proxy/management_endpoints/team_endpoints.py b/litellm/proxy/management_endpoints/team_endpoints.py index 88d51729c13..8943f5a7416 100644 --- a/litellm/proxy/management_endpoints/team_endpoints.py +++ b/litellm/proxy/management_endpoints/team_endpoints.py @@ -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 diff --git a/tests/unit/proxy/management_endpoints/test_common_daily_activity.py b/tests/unit/proxy/management_endpoints/test_common_daily_activity.py index 45820adb4ce..8808b73f89d 100644 --- a/tests/unit/proxy/management_endpoints/test_common_daily_activity.py +++ b/tests/unit/proxy/management_endpoints/test_common_daily_activity.py @@ -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() diff --git a/tests/unit/proxy/management_endpoints/test_daily_activity_routes.py b/tests/unit/proxy/management_endpoints/test_daily_activity_routes.py index 076538e5cd0..c3f4fdfef54 100644 --- a/tests/unit/proxy/management_endpoints/test_daily_activity_routes.py +++ b/tests/unit/proxy/management_endpoints/test_daily_activity_routes.py @@ -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() diff --git a/tests/unit/proxy/management_endpoints/test_team_endpoints.py b/tests/unit/proxy/management_endpoints/test_team_endpoints.py index 7e8e588fec6..3e71cc70099 100644 --- a/tests/unit/proxy/management_endpoints/test_team_endpoints.py +++ b/tests/unit/proxy/management_endpoints/test_team_endpoints.py @@ -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"