mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-09 03:18:44 +00:00
clean and verify key before inserting (#12840)
* clean and verify key * change checking logic * Add unit test
This commit is contained in:
parent
b4da29c83e
commit
c2833e693e
3 changed files with 65 additions and 8 deletions
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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"])
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue