From c9bfb7f0abb488b439f80cdae963ca61751a8587 Mon Sep 17 00:00:00 2001 From: mateo-berri <277851410+mateo-berri@users.noreply.github.com> Date: Tue, 18 Aug 2026 15:37:57 -0700 Subject: [PATCH] fix(guardrails): cap the date window accepted by /guardrails/usage endpoints --- litellm/proxy/guardrails/usage_endpoints.py | 50 ++++++++++---- ruff-strict-budget.json | 2 +- .../proxy/guardrails/test_usage_endpoints.py | 69 +++++++++++++++++++ 3 files changed, 107 insertions(+), 14 deletions(-) diff --git a/litellm/proxy/guardrails/usage_endpoints.py b/litellm/proxy/guardrails/usage_endpoints.py index ca89c7587ba..494a15d8948 100644 --- a/litellm/proxy/guardrails/usage_endpoints.py +++ b/litellm/proxy/guardrails/usage_endpoints.py @@ -5,7 +5,7 @@ GET /guardrails/usage/overview, /guardrails/usage/detail/:id, /guardrails/usage/ import json from collections.abc import Callable, Iterable, Mapping, Sequence -from datetime import datetime, timedelta, timezone +from datetime import date, datetime, timedelta, timezone from itertools import groupby from types import MappingProxyType from typing import TYPE_CHECKING, Any, Final, Literal, overload @@ -48,6 +48,35 @@ router: Final = APIRouter() _EMPTY_UNITS: Final[Mapping[str, int]] = MappingProxyType({}) +_USAGE_MAX_RANGE_DAYS: Final = 366 + + +def _resolve_usage_window(start_date: str | None, end_date: str | None) -> tuple[str, str]: + from fastapi import HTTPException, status + + now: Final = datetime.now(timezone.utc) + end: Final = end_date or now.strftime("%Y-%m-%d") + start: Final = start_date or (now - timedelta(days=7)).strftime("%Y-%m-%d") + try: + parsed: Final = (date.fromisoformat(start), date.fromisoformat(end)) + except ValueError: + raise HTTPException( + status_code=status.HTTP_400_BAD_REQUEST, + detail="start_date and end_date must be in YYYY-MM-DD format", + ) + start_obj, end_obj = parsed + if end_obj < start_obj: + raise HTTPException( + status_code=status.HTTP_400_BAD_REQUEST, + detail="start_date must be on or before end_date", + ) + if end_obj - start_obj > timedelta(days=_USAGE_MAX_RANGE_DAYS): + raise HTTPException( + status_code=status.HTTP_400_BAD_REQUEST, + detail=f"Date range too large; maximum is {_USAGE_MAX_RANGE_DAYS} days", + ) + return start, end + def _guardrails_table( prisma_client: "PrismaClient", @@ -457,9 +486,7 @@ async def guardrails_usage_overview( rows=[], chart=[], totalRequests=0, totalBlocked=0, passRate=100.0, totalUsageUnits=_EMPTY_UNITS ) - now: Final = datetime.now(timezone.utc) - end: Final = end_date or now.strftime("%Y-%m-%d") - start: Final = start_date or (now - timedelta(days=7)).strftime("%Y-%m-%d") + start, end = _resolve_usage_window(start_date, end_date) from litellm.proxy.guardrails.guardrail_registry import IN_MEMORY_GUARDRAIL_HANDLER @@ -477,7 +504,7 @@ async def guardrails_usage_overview( ) # Previous period for trend - start_prev: Final = (datetime.strptime(start, "%Y-%m-%d") - timedelta(days=7)).strftime("%Y-%m-%d") + start_prev: Final = (date.fromisoformat(start) - timedelta(days=7)).isoformat() metrics_prev: Sequence[prisma_models.LiteLLM_DailyGuardrailMetrics] = await _find_daily_guardrail_metrics( prisma_client, where={"date": {"gte": start_prev, "lt": start}} ) @@ -531,9 +558,7 @@ async def guardrails_usage_detail( raise HTTPException(status_code=500, detail="Prisma client not initialized") - now: Final = datetime.now(timezone.utc) - end: Final = end_date or now.strftime("%Y-%m-%d") - start: Final = start_date or (now - timedelta(days=7)).strftime("%Y-%m-%d") + start, end = _resolve_usage_window(start_date, end_date) from litellm.proxy.guardrails.guardrail_registry import IN_MEMORY_GUARDRAIL_HANDLER @@ -556,11 +581,12 @@ async def guardrails_usage_detail( "date": {"gte": start, "lte": end}, }, ) + start_prev: Final = (date.fromisoformat(start) - timedelta(days=7)).isoformat() metrics_prev: Final[Sequence[prisma_models.LiteLLM_DailyGuardrailMetrics]] = await _find_daily_guardrail_metrics( prisma_client, where={ "guardrail_id": {"in": metric_ids}, - "date": {"lt": start}, + "date": {"gte": start_prev, "lt": start}, }, ) units_where: Final[prisma_types.LiteLLM_DailyGuardrailUsageUnitsWhereInput] = { @@ -838,9 +864,7 @@ async def policies_usage_overview( rows=[], chart=[], totalRequests=0, totalBlocked=0, passRate=100.0, totalUsageUnits=_EMPTY_UNITS ) - now: Final = datetime.now(timezone.utc) - end: Final = end_date or now.strftime("%Y-%m-%d") - start: Final = start_date or (now - timedelta(days=7)).strftime("%Y-%m-%d") + start, end = _resolve_usage_window(start_date, end_date) try: policies: Final = await _policies_table(prisma_client).find_many() @@ -851,7 +875,7 @@ async def policies_usage_overview( prisma_client, where={ "date": { - "gte": (datetime.strptime(start, "%Y-%m-%d") - timedelta(days=7)).strftime("%Y-%m-%d"), + "gte": (date.fromisoformat(start) - timedelta(days=7)).isoformat(), "lt": start, } }, diff --git a/ruff-strict-budget.json b/ruff-strict-budget.json index 4a6ff6af902..7ed9635c9cd 100644 --- a/ruff-strict-budget.json +++ b/ruff-strict-budget.json @@ -96,7 +96,7 @@ "limit": 10 }, "DTZ007": { - "limit": 19 + "limit": 17 }, "DTZ011": { "limit": 3 diff --git a/tests/test_litellm/proxy/guardrails/test_usage_endpoints.py b/tests/test_litellm/proxy/guardrails/test_usage_endpoints.py index f63c08a2c39..d87e607a865 100644 --- a/tests/test_litellm/proxy/guardrails/test_usage_endpoints.py +++ b/tests/test_litellm/proxy/guardrails/test_usage_endpoints.py @@ -27,6 +27,7 @@ from litellm.proxy.guardrails.usage_endpoints import ( guardrails_usage_detail, guardrails_usage_logs, guardrails_usage_overview, + policies_usage_overview, ) from litellm.types.guardrails import Guardrail, LitellmParams @@ -356,3 +357,71 @@ async def test_logs_resolves_config_guardrail_logical_name(): ) where = prisma.db.litellm_spendlogguardrailindex.find_many.call_args.kwargs["where"] assert where["guardrail_id"] == {"in": ["yaml-uuid", "yaml-pii"]} + + +# ---- date window cap (LIT-5762) --------------------------------------------- + + +@pytest.mark.asyncio +async def test_overview_rejects_range_over_max_days(): + prisma = _prisma() + handler = _config_handler() + p1, p2 = _patches(prisma, handler) + with p1, p2, pytest.raises(HTTPException) as exc: + await guardrails_usage_overview(start_date="2020-01-01", end_date=END, user_api_key_dict=ADMIN) + assert exc.value.status_code == 400 + assert "366" in str(exc.value.detail) + + +@pytest.mark.asyncio +async def test_overview_accepts_range_at_exactly_max_days(): + prisma = _prisma() + handler = _config_handler() + p1, p2 = _patches(prisma, handler) + with p1, p2: + resp = await guardrails_usage_overview(start_date="2025-04-26", end_date="2026-04-27", user_api_key_dict=ADMIN) + assert resp.totalRequests == 0 + + +@pytest.mark.asyncio +async def test_overview_rejects_malformed_dates(): + prisma = _prisma() + handler = _config_handler() + p1, p2 = _patches(prisma, handler) + with p1, p2, pytest.raises(HTTPException) as exc: + await guardrails_usage_overview(start_date="not-a-date", end_date=END, user_api_key_dict=ADMIN) + assert exc.value.status_code == 400 + + +@pytest.mark.asyncio +async def test_detail_rejects_reversed_dates(): + prisma = _prisma(find_unique=_db_row()) + handler = _config_handler() + p1, p2 = _patches(prisma, handler) + with p1, p2, pytest.raises(HTTPException) as exc: + await guardrails_usage_detail(guardrail_id="db-1", start_date=END, end_date=START, user_api_key_dict=ADMIN) + assert exc.value.status_code == 400 + + +@pytest.mark.asyncio +async def test_policies_overview_rejects_range_over_max_days(): + prisma = _prisma() + handler = _config_handler() + p1, p2 = _patches(prisma, handler) + with p1, p2, pytest.raises(HTTPException) as exc: + await policies_usage_overview(start_date="2020-01-01", end_date=END, user_api_key_dict=ADMIN) + assert exc.value.status_code == 400 + + +@pytest.mark.asyncio +async def test_detail_prev_trend_query_is_bounded(): + """Regression: the trend query scanned every metrics row before start_date.""" + prisma = _prisma(find_unique=_db_row()) + handler = _config_handler() + p1, p2 = _patches(prisma, handler) + with p1, p2: + await guardrails_usage_detail(guardrail_id="db-1", start_date=START, end_date=END, user_api_key_dict=ADMIN) + wheres = [c.kwargs["where"] for c in prisma.db.litellm_dailyguardrailmetrics.find_many.await_args_list] + prev_wheres = [w for w in wheres if "lt" in w.get("date", {})] + assert prev_wheres + assert all("gte" in w["date"] for w in prev_wheres)