fix(proxy): stop hashing raw sk- values in list searches

The search= param on /key/list, /audit, and /spend/logs/ui, plus key_hash= on /key/list, now compare the pasted value verbatim. Only a copied key ID (the hash) matches, so a raw virtual key never needs to travel in a GET query string

Claude-Session: https://claude.ai/code/session_01Q5sbiogJzPcCRmYSbaHxZf
This commit is contained in:
ryan-crabbe-berri 2026-09-03 15:54:42 -07:00
parent e504477a69
commit 9baa19c7d1
9 changed files with 43 additions and 141 deletions

View file

@ -18,7 +18,6 @@ from litellm_enterprise.types.proxy.audit_logging_endpoints import (
from litellm.proxy._types import CommonProxyErrors, UserAPIKeyAuth
from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
from litellm.proxy.utils import _hash_token_if_needed
from litellm.repositories.prisma_protocols import TableActions
from litellm.repositories.table_repositories import AuditLogRepository
@ -50,14 +49,13 @@ def _build_json_field_or_condition(json_key: str, value: str) -> dict[str, objec
def _build_search_condition(search: str) -> dict[str, object]:
"""Match any id column; a raw sk- key is hashed for the two columns that store key hashes."""
hashed: Final = _hash_token_if_needed(search)
"""Match a row whose id, changed_by, object_id, or changed_by_api_key equals the search value."""
return {
"OR": (
{"id": search},
{"changed_by": search},
{"object_id": hashed},
{"changed_by_api_key": hashed},
{"object_id": search},
{"changed_by_api_key": search},
)
}
@ -99,10 +97,7 @@ async def get_audit_logs(
),
search: str | None = Query(
None,
description=(
"Match a row whose id, object_id, changed_by, or changed_by_api_key equals this value "
"(a raw sk- virtual key is hashed first)"
),
description="Match a row whose id, object_id, changed_by, or changed_by_api_key equals this value",
),
# Sorting parameters
sort_by: str | None = Query(
@ -159,7 +154,7 @@ async def get_audit_logs(
{sort_by: sort_order} if sort_by and isinstance(sort_by, str) else {"updated_at": sort_order}
)
audit_log_table: Final[TableActions["prisma_models.LiteLLM_AuditLog"]] = AuditLogRepository(prisma_client).table
audit_log_table: Final[TableActions[prisma_models.LiteLLM_AuditLog]] = AuditLogRepository(prisma_client).table
# Get paginated results
audit_logs: Final = await audit_log_table.find_many(
@ -221,7 +216,7 @@ async def get_audit_log_by_id(
detail={"message": CommonProxyErrors.db_not_connected_error.value},
)
audit_log_table: Final[TableActions["prisma_models.LiteLLM_AuditLog"]] = AuditLogRepository(prisma_client).table
audit_log_table: Final[TableActions[prisma_models.LiteLLM_AuditLog]] = AuditLogRepository(prisma_client).table
# Get the audit log by ID
audit_log: Final = await audit_log_table.find_unique(where={"id": id})

View file

@ -5803,7 +5803,7 @@ async def list_keys(
),
search: str | None = Query(
None,
description="Combined search: matches keys whose token (key hash) equals the value, hashing a raw sk- key first, OR whose key_alias contains it (case-insensitive).",
description="Combined search: matches keys whose token (key hash) equals the value OR whose key_alias contains it (case-insensitive).",
),
return_full_object: bool = Query(False, description="Return full key object"),
include_team_keys: bool = Query(False, description="Include all keys for teams that user is an admin of."),
@ -5867,17 +5867,13 @@ async def list_keys(
detail={"error": "Invalid expires value. Supported: 'active', 'expired'."},
)
hashed_key_hash: Final[str | None] = (
_hash_token_if_needed(token=key_hash) if isinstance(key_hash, str) else None
)
complete_user_info: Final = await validate_key_list_check(
user_api_key_dict=user_api_key_dict,
user_id=user_id,
team_id=team_id,
organization_id=organization_id,
key_alias=key_alias,
key_hash=hashed_key_hash,
key_hash=key_hash,
prisma_client=prisma_client,
)
@ -5937,7 +5933,7 @@ async def list_keys(
user_id=user_id,
team_id=team_id,
key_alias=key_alias,
key_hash=hashed_key_hash,
key_hash=key_hash,
return_full_object=return_full_object,
organization_id=organization_id,
admin_team_ids=admin_team_ids,
@ -6175,7 +6171,7 @@ def _build_expires_where_clause(expires_filter: str, now: datetime) -> dict[str,
def _build_key_search_where(search: str) -> KeySearchWhere:
search_where: Final[KeySearchWhere] = {
"OR": (
{"token": _hash_token_if_needed(token=search)},
{"token": search},
{"key_alias": {"contains": search, "mode": "insensitive"}},
)
}

View file

@ -2240,20 +2240,18 @@ def _build_spend_log_search_condition(
end_date: datetime,
next_param_index: int,
) -> _SpendLogSearchCondition:
"""request_id (indexed) matches across all time; the unindexed id columns only inside the window (sk- keys hashed)."""
"""request_id (indexed) matches across all time; the unindexed id columns only inside the window."""
raw: Final = f"${next_param_index}"
hashed: Final = f"${next_param_index + 1}"
window_start: Final = f"${next_param_index + 2}"
window_end: Final = f"${next_param_index + 3}"
window_start: Final = f"${next_param_index + 1}"
window_end: Final = f"${next_param_index + 2}"
sql: Final = (
f"(request_id = {raw} OR ("
f"\"startTime\" >= ({window_start}::timestamptz AT TIME ZONE 'UTC') "
f"AND \"startTime\" <= ({window_end}::timestamptz AT TIME ZONE 'UTC') "
f'AND (api_key = {hashed} OR team_id = {raw} OR "user" = {raw} OR end_user = {raw} '
f'AND (api_key = {raw} OR team_id = {raw} OR "user" = {raw} OR end_user = {raw} '
f"OR session_id = {raw} OR model_id = {raw})))"
)
hashed_search: Final = hash_token(token=search) if search.startswith("sk-") else search
return _SpendLogSearchCondition(sql=sql, params=(search, hashed_search, start_date, end_date))
return _SpendLogSearchCondition(sql=sql, params=(search, start_date, end_date))
@router.get(
@ -2359,7 +2357,7 @@ async def ui_view_spend_logs(
search: str | None = fastapi.Query(
default=None,
description=(
"Match a log whose request_id, api_key (a raw sk- key is hashed first), team_id, user, end_user, "
"Match a log whose request_id, api_key (hash), team_id, user, end_user, "
"session_id, or model_id equals this value. request_id matches across all time; the other columns "
"match inside start_date/end_date, which stay required"
),

View file

@ -16,7 +16,7 @@ class KeyAliasContainsWhere(TypedDict):
class KeySearchWhere(TypedDict):
"""Prisma filter behind `/key/list?search=`: exact token (sk- keys hashed) or alias substring, case-insensitive."""
"""Prisma filter behind `/key/list?search=`: exact token or case-insensitive alias substring."""
OR: ReadOnly[tuple[KeyTokenWhere, KeyAliasContainsWhere]]

View file

@ -1,4 +1,3 @@
import hashlib
from datetime import datetime, timedelta
from typing import Final
from unittest.mock import AsyncMock, MagicMock, patch
@ -172,24 +171,6 @@ def test_search_matches_any_id_column_alongside_the_other_filters(mock_prisma_cl
}
def test_search_hashes_a_raw_virtual_key_for_the_hashed_columns(mock_prisma_client):
where: Final = _list_audit_logs_where(mock_prisma_client, "search=sk-raw")
hashed: Final = hashlib.sha256(b"sk-raw").hexdigest()
assert where == {
"AND": (
{
"OR": (
{"id": "sk-raw"},
{"changed_by": "sk-raw"},
{"object_id": hashed},
{"changed_by_api_key": hashed},
)
},
)
}
def test_an_empty_search_leaves_the_where_clause_unchanged(mock_prisma_client):
where: Final = _list_audit_logs_where(mock_prisma_client, "action=create&search=")

View file

@ -6351,33 +6351,15 @@ def _search_clause(search: str, token: str) -> dict:
return {"OR": [{"token": token}, {"key_alias": {"contains": search, "mode": "insensitive"}}]}
def test_build_key_filter_conditions_search_hashes_raw_key_and_ors_alias_contains():
def test_build_key_filter_conditions_search_ors_token_and_alias_contains():
"""
LIT-4741: `search` matches a key by its alias (case-insensitive contains) OR by
its ID. A pasted raw sk- key is hashed to its token first; an already-hashed
value is used verbatim.
its ID (the token column), with the pasted value used verbatim.
"""
from litellm.proxy._types import hash_token
from litellm.proxy.management_endpoints.key_management_endpoints import (
_build_key_filter_conditions,
)
raw_where = json.loads(
json.dumps(
_build_key_filter_conditions(
user_id=None,
team_id=None,
organization_id=None,
key_alias=None,
key_hash=None,
exclude_team_id=None,
admin_team_ids=None,
search="sk-raw",
)
)
)
assert _search_clause("sk-raw", hash_token("sk-raw")) in raw_where["AND"], f"raw search not ANDed: {raw_where}"
hashed_where = json.loads(
json.dumps(
_build_key_filter_conditions(
@ -6402,7 +6384,6 @@ def test_build_key_filter_conditions_search_narrows_team_admin_visibility():
LIT-4741, same class as LIT-3243: `search` must be a top-level AND so it
narrows a team admin's admin-team branch instead of being bypassed by it.
"""
from litellm.proxy._types import hash_token
from litellm.proxy.management_endpoints.key_management_endpoints import (
_build_key_filter_conditions,
)
@ -6419,21 +6400,19 @@ def test_build_key_filter_conditions_search_narrows_team_admin_visibility():
admin_team_ids=["team-a"],
member_team_ids=["team-a"],
include_created_by_keys=False,
search="sk-member",
search="member-key-id",
)
)
)
assert where.get("AND"), f"expected top-level AND, got: {where}"
assert _search_clause("sk-member", hash_token("sk-member")) in where["AND"], f"search not ANDed: {where}"
assert _search_clause("member-key-id", "member-key-id") in where["AND"], f"search not ANDed: {where}"
assert json.dumps({"team_id": {"in": ["team-a"]}}) in json.dumps(where)
@pytest.mark.asyncio
async def test_list_key_helper_applies_search_to_prisma_where():
"""LIT-4741: `search` given to _list_key_helper must reach the Prisma where clause."""
from litellm.proxy._types import hash_token
mock_prisma_client = AsyncMock()
mock_find_many = AsyncMock(return_value=[])
mock_prisma_client.db.litellm_verificationtoken.find_many = mock_find_many
@ -6448,11 +6427,11 @@ async def test_list_key_helper_applies_search_to_prisma_where():
organization_id=None,
key_alias=None,
key_hash=None,
search="sk-raw",
search="key-id-123",
)
where = json.loads(json.dumps(mock_find_many.call_args.kwargs["where"]))
assert _search_clause("sk-raw", hash_token("sk-raw")) in where["AND"], f"search not in Prisma where: {where}"
assert _search_clause("key-id-123", "key-id-123") in where["AND"], f"search not in Prisma where: {where}"
@pytest.mark.asyncio
@ -14978,47 +14957,13 @@ async def test_list_keys_non_admin_cannot_opt_into_substring():
assert kwargs["user_id"] == "alice"
@pytest.mark.asyncio
async def test_list_keys_hashes_raw_key_hash_before_validation():
"""LIT-4741: a raw sk- key pasted as key_hash is hashed before the ownership
check and the query, so a non-admin filtering by their own raw key gets the
row instead of the 'Key Hash not found.' 403."""
from litellm.proxy._types import hash_token
user = UserAPIKeyAuth(user_role=LitellmUserRoles.INTERNAL_USER, user_id="alice")
validate = AsyncMock(
return_value=LiteLLM_UserTable(
user_id="alice", user_email="alice@example.com", teams=[], organization_memberships=[]
)
)
helper = AsyncMock(return_value={"keys": [], "total_count": 0, "current_page": 1, "total_pages": 0})
with (
patch("litellm.proxy.proxy_server.prisma_client", AsyncMock()),
patch(
"litellm.proxy.management_endpoints.key_management_endpoints.validate_key_list_check",
validate,
),
patch("litellm.proxy.management_endpoints.key_management_endpoints._list_key_helper", helper),
):
await list_keys(
request=MagicMock(),
user_api_key_dict=user,
status=None,
user_id=None,
key_hash="sk-raw",
)
assert validate.call_args.kwargs["key_hash"] == hash_token("sk-raw")
assert helper.call_args.kwargs["key_hash"] == hash_token("sk-raw")
@pytest.mark.asyncio
async def test_list_keys_search_is_honored_for_non_admin():
"""LIT-4741: unlike substring_matching, `search` is not admin-gated. A non-admin's
search reaches the helper while their own-user scoping stays in place."""
user = UserAPIKeyAuth(user_role=LitellmUserRoles.INTERNAL_USER, user_id="alice")
kwargs = await _list_keys_capture_helper_kwargs(user, user_id=None, search="sk-raw")
assert kwargs["search"] == "sk-raw"
kwargs = await _list_keys_capture_helper_kwargs(user, user_id=None, search="key-id-123")
assert kwargs["search"] == "key-id-123"
assert kwargs["user_id"] == "alice"

View file

@ -61,7 +61,7 @@ def _filter_logs_by_date_range(logs, where):
_SEARCH_CLAUSE_RE = re.compile(
r'\(request_id = \$(\d+) OR \("startTime" >= \(\$(\d+)::timestamptz AT TIME ZONE \'UTC\'\) '
r'AND "startTime" <= \(\$(\d+)::timestamptz AT TIME ZONE \'UTC\'\) '
r'AND \(api_key = \$(\d+) OR team_id = \$\1 OR "user" = \$\1 OR end_user = \$\1 '
r'AND \(api_key = \$\1 OR team_id = \$\1 OR "user" = \$\1 OR end_user = \$\1 '
r"OR session_id = \$\1 OR model_id = \$\1\)\)\)"
)
@ -72,9 +72,8 @@ def _matches_spend_log_search(log, search):
return True
if not _filter_logs_by_date_range([log], {"startTime": {"gte": search["gte"], "lte": search["lte"]}}):
return False
if log.get("api_key") == search["api_key"]:
return True
return any(log.get(col) == search["value"] for col in ("team_id", "user", "end_user", "session_id", "model_id"))
columns = ("api_key", "team_id", "user", "end_user", "session_id", "model_id")
return any(log.get(col) == search["value"] for col in columns)
def _reconstruct_ui_where_from_sql(sql_query, params):
@ -98,10 +97,9 @@ def _reconstruct_ui_where_from_sql(sql_query, params):
search_clause = _SEARCH_CLAUSE_RE.search(clause.group(1))
if search_clause:
raw_index, start_index, end_index, hashed_index = (int(g) for g in search_clause.groups())
raw_index, start_index, end_index = (int(g) for g in search_clause.groups())
where["search"] = {
"value": params[raw_index - 1],
"api_key": params[hashed_index - 1],
"gte": _iso(params[start_index - 1]),
"lte": _iso(params[end_index - 1]),
}
@ -2384,31 +2382,20 @@ async def test_ui_view_spend_logs_request_id_owner_scoped_by_id_only(
def test_build_spend_log_search_condition_windows_every_branch_except_request_id():
"""LIT-4741: request_id matches across all time; the six other id columns only inside the window,
and a raw sk- key is hashed for the api_key branch alone."""
all comparing the pasted value verbatim."""
start = datetime.datetime(2026, 8, 1, tzinfo=timezone.utc)
end = datetime.datetime(2026, 8, 2, tzinfo=timezone.utc)
condition = spend_management_endpoints._build_spend_log_search_condition(
search="sk-raw-key", start_date=start, end_date=end, next_param_index=3
search="key-hash-7", start_date=start, end_date=end, next_param_index=3
)
assert condition.sql == (
"(request_id = $3 OR (\"startTime\" >= ($5::timestamptz AT TIME ZONE 'UTC') "
"AND \"startTime\" <= ($6::timestamptz AT TIME ZONE 'UTC') "
'AND (api_key = $4 OR team_id = $3 OR "user" = $3 OR end_user = $3 OR session_id = $3 OR model_id = $3)))'
"(request_id = $3 OR (\"startTime\" >= ($4::timestamptz AT TIME ZONE 'UTC') "
"AND \"startTime\" <= ($5::timestamptz AT TIME ZONE 'UTC') "
'AND (api_key = $3 OR team_id = $3 OR "user" = $3 OR end_user = $3 OR session_id = $3 OR model_id = $3)))'
)
assert condition.params == ("sk-raw-key", hashlib.sha256(b"sk-raw-key").hexdigest(), start, end)
def test_build_spend_log_search_condition_leaves_non_key_values_unhashed():
start = datetime.datetime(2026, 8, 1, tzinfo=timezone.utc)
end = datetime.datetime(2026, 8, 2, tzinfo=timezone.utc)
condition = spend_management_endpoints._build_spend_log_search_condition(
search="sess-42", start_date=start, end_date=end, next_param_index=1
)
assert condition.params == ("sess-42", "sess-42", start, end)
assert condition.params == ("key-hash-7", start, end)
def _search_fixture_logs(today):
@ -2427,7 +2414,7 @@ def _search_fixture_logs(today):
return [
{**base, "request_id": "req-session", "session_id": "sess-42", "startTime": recent},
{**base, "request_id": "req-session-old", "session_id": "sess-42", "startTime": old},
{**base, "request_id": "req-key", "api_key": hashlib.sha256(b"sk-raw-key").hexdigest(), "startTime": recent},
{**base, "request_id": "req-key", "api_key": "hashed-7", "startTime": recent},
{**base, "request_id": "req-team", "team_id": "team-7", "startTime": recent},
{**base, "request_id": "req-user", "user": "user-7", "startTime": recent},
{**base, "request_id": "req-end-user", "end_user": "cust-7", "startTime": recent},
@ -2461,7 +2448,7 @@ def _five_day_window(today):
[
("req-session-old", {"req-session-old"}),
("sess-42", {"req-session"}),
("sk-raw-key", {"req-key"}),
("hashed-7", {"req-key"}),
("team-7", {"req-team"}),
("user-7", {"req-user"}),
("cust-7", {"req-end-user"}),

View file

@ -525,14 +525,14 @@ describe("useKeys", () => {
json: async () => mockKeysResponse,
});
const { result } = renderHook(() => useKeys(1, 10, { search: "sk-pasted-key" }), { wrapper });
const { result } = renderHook(() => useKeys(1, 10, { search: "pasted-key-id" }), { wrapper });
await waitFor(() => {
expect(result.current.isLoading).toBe(false);
});
const callUrl = new URL(mockFetch.mock.calls[0][0], "http://localhost");
expect(callUrl.searchParams.get("search")).toBe("sk-pasted-key");
expect(callUrl.searchParams.get("search")).toBe("pasted-key-id");
expect(callUrl.searchParams.has("key_alias")).toBe(false);
expect(callUrl.searchParams.has("key_hash")).toBe(false);
});

View file

@ -40904,7 +40904,7 @@ export interface operations {
object_team_id?: string | null;
/** @description Filter by token (key hash) present in before_value or updated_values JSON (PostgreSQL only) */
object_key_hash?: string | null;
/** @description Match a row whose id, object_id, changed_by, or changed_by_api_key equals this value (a raw sk- virtual key is hashed first) */
/** @description Match a row whose id, object_id, changed_by, or changed_by_api_key equals this value */
search?: string | null;
/** @description Column to sort by (e.g. 'updated_at', 'action', 'table_name') */
sort_by?: string | null;
@ -49611,7 +49611,7 @@ export interface operations {
key_hash?: string | null;
/** @description Filter keys by key alias. Exact match by default; set substring_matching=true (admin only) for case-insensitive substring matching. */
key_alias?: string | null;
/** @description Combined search: matches keys whose token (key hash) equals the value, hashing a raw sk- key first, OR whose key_alias contains it (case-insensitive). */
/** @description Combined search: matches keys whose token (key hash) equals the value OR whose key_alias contains it (case-insensitive). */
search?: string | null;
/** @description Return full key object */
return_full_object?: boolean;
@ -56871,7 +56871,7 @@ export interface operations {
group_by_session?: boolean;
/** @description Keyset cursor '<last_activity>|<api_key>|<session_key>' from a previous group_by_session page. UI route only, honored when sorting by startTime */
session_cursor?: string | null;
/** @description Match a log whose request_id, api_key (a raw sk- key is hashed first), team_id, user, end_user, session_id, or model_id equals this value. request_id matches across all time; the other columns match inside start_date/end_date, which stay required */
/** @description Match a log whose request_id, api_key (hash), team_id, user, end_user, session_id, or model_id equals this value. request_id matches across all time; the other columns match inside start_date/end_date, which stay required */
search?: string | null;
};
header?: never;
@ -56989,7 +56989,7 @@ export interface operations {
group_by_session?: boolean;
/** @description Keyset cursor '<last_activity>|<api_key>|<session_key>' from a previous group_by_session page. UI route only, honored when sorting by startTime */
session_cursor?: string | null;
/** @description Match a log whose request_id, api_key (a raw sk- key is hashed first), team_id, user, end_user, session_id, or model_id equals this value. request_id matches across all time; the other columns match inside start_date/end_date, which stay required */
/** @description Match a log whose request_id, api_key (hash), team_id, user, end_user, session_id, or model_id equals this value. request_id matches across all time; the other columns match inside start_date/end_date, which stay required */
search?: string | null;
};
header?: never;