From c2833e693e18e373f3cfb3c3b0f57b05e2868073 Mon Sep 17 00:00:00 2001 From: "Jugal D. Bhatt" <55304795+jugaldb@users.noreply.github.com> Date: Fri, 25 Jul 2025 22:39:28 +0530 Subject: [PATCH] clean and verify key before inserting (#12840) * clean and verify key * change checking logic * Add unit test --- .../key_management_endpoints.py | 25 +++++++++++----- litellm/proxy/utils.py | 18 +++++++++++ tests/test_litellm/test_utils.py | 30 +++++++++++++++++++ 3 files changed, 65 insertions(+), 8 deletions(-) diff --git a/litellm/proxy/management_endpoints/key_management_endpoints.py b/litellm/proxy/management_endpoints/key_management_endpoints.py index 4565ea8ada3..70b8002ebb6 100644 --- a/litellm/proxy/management_endpoints/key_management_endpoints.py +++ b/litellm/proxy/management_endpoints/key_management_endpoints.py @@ -59,6 +59,7 @@ from litellm.proxy.utils import ( _hash_token_if_needed, handle_exception_on_proxy, jsonify_object, + is_valid_api_key, ) from litellm.router import Router from litellm.secret_managers.main import get_secret @@ -2684,10 +2685,14 @@ async def block_key( if prisma_client is None: raise Exception("{}".format(CommonProxyErrors.db_not_connected_error.value)) - if data.key.startswith("sk-"): - hashed_token = hash_token(token=data.key) - else: - hashed_token = data.key + if not is_valid_api_key(data.key): + raise ProxyException( + message="Invalid key format.", + type=ProxyErrorTypes.bad_request_error, + param="key", + code=status.HTTP_400_BAD_REQUEST, + ) + hashed_token = hash_token(token=data.key) if litellm.store_audit_logs is True: # make an audit log for key update @@ -2791,10 +2796,14 @@ async def unblock_key( if prisma_client is None: raise Exception("{}".format(CommonProxyErrors.db_not_connected_error.value)) - if data.key.startswith("sk-"): - hashed_token = hash_token(token=data.key) - else: - hashed_token = data.key + if not is_valid_api_key(data.key): + raise ProxyException( + message="Invalid key format.", + type=ProxyErrorTypes.bad_request_error, + param="key", + code=status.HTTP_400_BAD_REQUEST, + ) + hashed_token = hash_token(token=data.key) if litellm.store_audit_logs is True: # make an audit log for key update diff --git a/litellm/proxy/utils.py b/litellm/proxy/utils.py index 008b151afce..655b50cfc20 100644 --- a/litellm/proxy/utils.py +++ b/litellm/proxy/utils.py @@ -3184,3 +3184,21 @@ def get_prisma_client_or_throw(message: str): detail={"error": message}, ) return prisma_client + + +def is_valid_api_key(key: str) -> bool: + """ + Validates API key format: + - sk- keys: must match ^sk-[A-Za-z0-9_-]+$ + - hashed keys: must match ^[a-fA-F0-9]{64}$ + - Length between 20 and 100 characters + """ + import re + if not isinstance(key, str): + return False + if 3 <= len(key) <= 100: + if re.match(r"^sk-[A-Za-z0-9_-]+$", key): + return True + if re.match(r"^[a-fA-F0-9]{64}$", key): + return True + return False diff --git a/tests/test_litellm/test_utils.py b/tests/test_litellm/test_utils.py index 320ae754416..6b136f03c87 100644 --- a/tests/test_litellm/test_utils.py +++ b/tests/test_litellm/test_utils.py @@ -23,6 +23,7 @@ from litellm.utils import ( get_llm_provider, get_optional_params_image_gen, ) +from litellm.proxy.utils import is_valid_api_key # Adds the parent directory to the system path @@ -2155,6 +2156,35 @@ def test_image_response_utils(): image_response = ImageResponse(**result) +def test_is_valid_api_key(): + import hashlib + # Valid sk- keys + assert is_valid_api_key("sk-abc123") + assert is_valid_api_key("sk-ABC_123-xyz") + # Valid hashed key (64 hex chars) + assert is_valid_api_key("a" * 64) + assert is_valid_api_key("0123456789abcdef" * 4) # 16*4 = 64 + # Real SHA-256 hash + real_hash = hashlib.sha256(b"my_secret_key").hexdigest() + assert len(real_hash) == 64 + assert is_valid_api_key(real_hash) + # Invalid: too short + assert not is_valid_api_key("sk-") + assert not is_valid_api_key("") + # Invalid: too long + assert not is_valid_api_key("sk-" + "a" * 200) + # Invalid: wrong prefix + assert not is_valid_api_key("pk-abc123") + # Invalid: wrong chars in sk- key + assert not is_valid_api_key("sk-abc$%#@!") + # Invalid: not a string + assert not is_valid_api_key(None) + assert not is_valid_api_key(12345) + # Invalid: wrong length for hash + assert not is_valid_api_key("a" * 63) + assert not is_valid_api_key("a" * 65) + + if __name__ == "__main__": # Allow running this test file directly for debugging pytest.main([__file__, "-v"])