feat(proxy): add expires filter to GET /key/list (#32953)

* feat(proxy): add expires filter to GET /key/list

Add an opt-in expires query param to GET /key/list so callers can fetch
only expired or only active keys without paginating every page and
filtering client-side. 'expired' matches keys whose expires is in the
past (NULL expires excluded); 'active' matches keys that never expire or
expire in the future. Omitting the param preserves existing behavior for
every caller. An unrecognized value returns HTTP 400 rather than silently
returning all keys.

The filter is pushed to the database via the existing Prisma where
builder so callers avoid pulling the full key table into application
memory.

Resolves LIT-3387

* refactor(proxy): declare VALID_EXPIRES_FILTER_VALUES before its first use
This commit is contained in:
ryan-crabbe-berri 2026-07-11 16:48:25 -07:00 • committed by GitHub
parent 03970d7c3a
commit cb24864a7f
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
3 changed files with 288 additions and 0 deletions

View file

@ -5141,6 +5141,9 @@ async def get_member_team_ids(
return _get_member_team_ids_from_objects(user_api_key_dict, team_objects)
VALID_EXPIRES_FILTER_VALUES = frozenset({"active", "expired"})
@router.get(
"/key/list",
tags=["key management"],
@ -5180,6 +5183,10 @@ async def list_keys(
False,
description="If true (proxy admins only), match user_id/key_alias as case-insensitive substrings instead of exact values. Defaults to false: /key/list matched these exactly before substring search was added, and an exact user_id/key_alias filter must never return another user's keys.",
),
expires: str | None = Query(
None,
description="Filter keys by expiration. 'expired' returns keys whose expires is in the past; 'active' returns keys that never expire or expire in the future. Omit to return keys regardless of expiration.",
),
) -> KeyListResponseObject:
"""
List all keys for a given user / team / organization.
@ -5215,6 +5222,12 @@ async def list_keys(
detail={"error": "Invalid status value. Currently only 'deleted' is supported."},
)
if isinstance(expires, str) and expires not in VALID_EXPIRES_FILTER_VALUES:
raise HTTPException(
status_code=400,
detail={"error": "Invalid expires value. Supported: 'active', 'expired'."},
)
complete_user_info = await validate_key_list_check(
user_api_key_dict=user_api_key_dict,
user_id=user_id,
@ -5295,6 +5308,7 @@ async def list_keys(
access_group_id=access_group_id,
agent_id=agent_id,
use_substring_matching=use_substring_matching,
expires_filter=expires if isinstance(expires, str) else None,
)
verbose_proxy_logger.debug("Successfully prepared response")
@ -5502,6 +5516,12 @@ def _validate_sort_params(sort_by: Optional[str], sort_order: str) -> Optional[D
return order_by
def _build_expires_where_clause(expires_filter: str, now: datetime) -> dict[str, Any]:
if expires_filter == "expired":
return {"AND": [{"expires": {"not": None}}, {"expires": {"lt": now}}]}
return {"OR": [{"expires": None}, {"expires": {"gte": now}}]}
def _build_key_filter_conditions(
user_id: Optional[str],
team_id: Optional[str],
@ -5516,6 +5536,7 @@ def _build_key_filter_conditions(
access_group_id: Optional[str] = None,
agent_id: Optional[str] = None,
use_substring_matching: bool = False,
expires_filter: str | None = None,
) -> Dict[str, Union[str, Dict[str, Any], List[Dict[str, Any]]]]:
"""Build filter conditions for key listing.
@ -5623,6 +5644,8 @@ def _build_key_filter_conditions(
where = {"AND": [where, {"access_group_ids": {"hasSome": [access_group_id]}}]}
if agent_id and isinstance(agent_id, str):
where = {"AND": [where, {"agent_id": agent_id}]}
if expires_filter is not None and expires_filter in VALID_EXPIRES_FILTER_VALUES:
where = {"AND": [where, _build_expires_where_clause(expires_filter, datetime.now(timezone.utc))]}
verbose_proxy_logger.debug(f"Filter conditions: {where}")
return where
@ -5652,6 +5675,7 @@ async def _list_key_helper(
access_group_id: Optional[str] = None,
agent_id: Optional[str] = None,
use_substring_matching: bool = False,
expires_filter: str | None = None,
) -> KeyListResponseObject:
"""
Helper function to list keys
@ -5689,6 +5713,7 @@ async def _list_key_helper(
access_group_id=access_group_id,
agent_id=agent_id,
use_substring_matching=use_substring_matching,
expires_filter=expires_filter,
)
# Calculate skip for pagination

View file

@ -14530,3 +14530,264 @@ def test_generate_key_helper_fn_accepts_per_tag_rate_limits():
request_type="user",
tag_rpm_limit={"cell-1": 5},
)
def _find_expires_clauses(node):
"""Recursively collect every value keyed 'expires' anywhere in a Prisma where dict."""
found = []
if isinstance(node, dict):
for key, value in node.items():
if key == "expires":
found.append(value)
else:
found.extend(_find_expires_clauses(value))
elif isinstance(node, list):
for item in node:
found.extend(_find_expires_clauses(item))
return found
def test_build_expires_where_clause_expired_shape():
"""'expired' must exclude never-expiring (NULL) keys and match expires < now."""
from datetime import datetime, timezone
from litellm.proxy.management_endpoints.key_management_endpoints import (
_build_expires_where_clause,
)
now = datetime(2026, 1, 2, 3, 4, 5, tzinfo=timezone.utc)
assert _build_expires_where_clause("expired", now) == {
"AND": [{"expires": {"not": None}}, {"expires": {"lt": now}}]
}
def test_build_expires_where_clause_active_shape():
"""'active' must include never-expiring (NULL) keys and match expires >= now."""
from datetime import datetime, timezone
from litellm.proxy.management_endpoints.key_management_endpoints import (
_build_expires_where_clause,
)
now = datetime(2026, 1, 2, 3, 4, 5, tzinfo=timezone.utc)
assert _build_expires_where_clause("active", now) == {
"OR": [{"expires": None}, {"expires": {"gte": now}}]
}
def test_build_key_filter_conditions_expired_applies_lt_clause():
"""expires_filter='expired' ANDs in a not-NULL + lt(now) constraint."""
from litellm.proxy.management_endpoints.key_management_endpoints import (
_build_key_filter_conditions,
)
where = _build_key_filter_conditions(
user_id="u1",
team_id=None,
organization_id=None,
key_alias=None,
key_hash=None,
exclude_team_id=None,
admin_team_ids=None,
expires_filter="expired",
)
clauses = _find_expires_clauses(where)
assert {"not": None} in clauses
lt_clauses = [c for c in clauses if isinstance(c, dict) and "lt" in c]
assert len(lt_clauses) == 1
assert "gte" not in str(clauses)
def test_build_key_filter_conditions_active_applies_gte_and_null():
"""expires_filter='active' ANDs in a NULL-or-gte(now) constraint."""
from litellm.proxy.management_endpoints.key_management_endpoints import (
_build_key_filter_conditions,
)
where = _build_key_filter_conditions(
user_id="u1",
team_id=None,
organization_id=None,
key_alias=None,
key_hash=None,
exclude_team_id=None,
admin_team_ids=None,
expires_filter="active",
)
clauses = _find_expires_clauses(where)
assert None in clauses
gte_clauses = [c for c in clauses if isinstance(c, dict) and "gte" in c]
assert len(gte_clauses) == 1
assert "lt" not in str(clauses)
def test_build_key_filter_conditions_no_expires_filter_omits_clause():
"""Default (no expires_filter) must not add any expires constraint — preserves existing callers."""
from litellm.proxy.management_endpoints.key_management_endpoints import (
_build_key_filter_conditions,
)
where = _build_key_filter_conditions(
user_id="u1",
team_id=None,
organization_id=None,
key_alias=None,
key_hash=None,
exclude_team_id=None,
admin_team_ids=None,
)
assert _find_expires_clauses(where) == []
def test_build_key_filter_conditions_invalid_expires_filter_omits_clause():
"""An unrecognized expires_filter value is ignored, not applied blindly."""
from litellm.proxy.management_endpoints.key_management_endpoints import (
_build_key_filter_conditions,
)
where = _build_key_filter_conditions(
user_id="u1",
team_id=None,
organization_id=None,
key_alias=None,
key_hash=None,
exclude_team_id=None,
admin_team_ids=None,
expires_filter="garbage",
)
assert _find_expires_clauses(where) == []
def test_build_key_filter_conditions_expires_now_is_call_time_utc():
"""The lt(now) boundary is computed at call time as a tz-aware UTC datetime."""
from datetime import datetime, timezone
from litellm.proxy.management_endpoints.key_management_endpoints import (
_build_key_filter_conditions,
)
before = datetime.now(timezone.utc)
where = _build_key_filter_conditions(
user_id="u1",
team_id=None,
organization_id=None,
key_alias=None,
key_hash=None,
exclude_team_id=None,
admin_team_ids=None,
expires_filter="expired",
)
after = datetime.now(timezone.utc)
lt_values = [c["lt"] for c in _find_expires_clauses(where) if isinstance(c, dict) and "lt" in c]
assert len(lt_values) == 1
now_value = lt_values[0]
assert now_value.tzinfo is not None
assert now_value.utcoffset().total_seconds() == 0
assert before <= now_value <= after
@pytest.mark.asyncio
async def test_list_keys_rejects_invalid_expires():
"""A typo'd expires value must 400, never silently fall back to returning all keys."""
from unittest.mock import Mock, patch
mock_prisma_client = AsyncMock()
mock_user_api_key_dict = UserAPIKeyAuth(user_role=LitellmUserRoles.PROXY_ADMIN)
with patch("litellm.proxy.proxy_server.prisma_client", mock_prisma_client):
with pytest.raises(ProxyException) as exc_info:
await list_keys(
request=Mock(),
user_api_key_dict=mock_user_api_key_dict,
status=None,
expires="expred",
)
assert exc_info.value.code == "400"
assert "Invalid expires value" in str(exc_info.value.message)
@pytest.mark.asyncio
@pytest.mark.parametrize(
"expires_value, expected_forward",
[("expired", "expired"), ("active", "active"), (None, None)],
)
async def test_list_keys_forwards_expires_filter(expires_value, expected_forward):
"""list_keys forwards a valid/None expires value verbatim to _list_key_helper as expires_filter."""
from unittest.mock import Mock, patch
mock_prisma_client = AsyncMock()
mock_user_api_key_dict = UserAPIKeyAuth(user_role=LitellmUserRoles.PROXY_ADMIN)
mock_user_info = LiteLLM_UserTable(
user_id="admin-user",
user_email="admin@example.com",
teams=[],
organization_memberships=[],
)
mock_helper = AsyncMock(
return_value={"keys": [], "total_count": 0, "current_page": 1, "total_pages": 0}
)
with (
patch("litellm.proxy.proxy_server.prisma_client", mock_prisma_client),
patch(
"litellm.proxy.management_endpoints.key_management_endpoints.validate_key_list_check",
return_value=mock_user_info,
),
patch(
"litellm.proxy.management_endpoints.key_management_endpoints._list_key_helper",
mock_helper,
),
):
await list_keys(
request=Mock(),
user_api_key_dict=mock_user_api_key_dict,
status=None,
expires=expires_value,
)
mock_helper.assert_called_once()
assert mock_helper.call_args.kwargs["expires_filter"] == expected_forward
@pytest.mark.asyncio
async def test_list_keys_without_expires_param_forwards_none():
"""Existing callers that never pass `expires` must not 400 and must forward expires_filter=None."""
from unittest.mock import Mock, patch
mock_prisma_client = AsyncMock()
mock_user_api_key_dict = UserAPIKeyAuth(user_role=LitellmUserRoles.PROXY_ADMIN)
mock_user_info = LiteLLM_UserTable(
user_id="admin-user",
user_email="admin@example.com",
teams=[],
organization_memberships=[],
)
mock_helper = AsyncMock(
return_value={"keys": [], "total_count": 0, "current_page": 1, "total_pages": 0}
)
with (
patch("litellm.proxy.proxy_server.prisma_client", mock_prisma_client),
patch(
"litellm.proxy.management_endpoints.key_management_endpoints.validate_key_list_check",
return_value=mock_user_info,
),
patch(
"litellm.proxy.management_endpoints.key_management_endpoints._list_key_helper",
mock_helper,
),
):
await list_keys(
request=Mock(),
user_api_key_dict=mock_user_api_key_dict,
status=None,
)
mock_helper.assert_called_once()
assert mock_helper.call_args.kwargs["expires_filter"] is None

View file

@ -42428,6 +42428,8 @@ export interface operations {
agent_id?: string | null;
/** @description If true (proxy admins only), match user_id/key_alias as case-insensitive substrings instead of exact values. Defaults to false: /key/list matched these exactly before substring search was added, and an exact user_id/key_alias filter must never return another user's keys. */
substring_matching?: boolean;
/** @description Filter keys by expiration. 'expired' returns keys whose expires is in the past; 'active' returns keys that never expire or expire in the future. Omit to return keys regardless of expiration. */
expires?: string | null;
};
header?: never;
path?: never;