test: add unit tests

This commit is contained in:
Krrish Dholakia 2025-06-26 15:56:10 -07:00
parent eec8e1a746
commit 35c6f6f83f
2 changed files with 49 additions and 11 deletions

View file

@ -1835,6 +1835,21 @@ async def _rotate_master_key(
)
def get_new_token(data: Optional[RegenerateKeyRequest]) -> str:
if data and data.new_key is not None:
new_token = data.new_key
if not data.new_key.startswith("sk-"):
raise HTTPException(
status_code=status.HTTP_400_BAD_REQUEST,
detail={
"error": "New key must start with 'sk-'. This is to distinguish a key hash (used by litellm for logging / internal logic) from the actual key."
},
)
else:
new_token = f"sk-{secrets.token_urlsafe(LENGTH_OF_LITELLM_GENERATED_KEY)}"
return new_token
@router.post(
"/key/{key:path}/regenerate",
tags=["key management"],
@ -1987,17 +2002,7 @@ async def regenerate_key_fn(
verbose_proxy_logger.debug("key_in_db: %s", _key_in_db)
if data and data.new_key is not None:
new_token = data.new_key
if not data.new_key.startswith("sk-"):
raise HTTPException(
status_code=status.HTTP_400_BAD_REQUEST,
detail={
"error": "New key must start with 'sk-'. This is to distinguish a key hash (used by litellm for logging / internal logic) from the actual key."
},
)
else:
new_token = f"sk-{secrets.token_urlsafe(LENGTH_OF_LITELLM_GENERATED_KEY)}"
new_token = get_new_token(data=data)
new_token_hash = hash_token(new_token)
new_token_key_name = f"sk-...{new_token[-4:]}"

View file

@ -459,3 +459,36 @@ async def test_key_update_object_permissions_missing_permission_record(monkeypat
# Verify upsert was called to create new record
mock_prisma_client.db.litellm_objectpermissiontable.upsert.assert_called_once()
def test_get_new_token_with_valid_key():
"""Test get_new_token function when provided with a valid key that starts with 'sk-'"""
from litellm.proxy._types import RegenerateKeyRequest
from litellm.proxy.management_endpoints.key_management_endpoints import (
get_new_token,
)
# Test with valid new_key
data = RegenerateKeyRequest(new_key="sk-test123456789")
result = get_new_token(data)
assert result == "sk-test123456789"
def test_get_new_token_with_invalid_key():
"""Test get_new_token function when provided with an invalid key that doesn't start with 'sk-'"""
from fastapi import HTTPException
from litellm.proxy._types import RegenerateKeyRequest
from litellm.proxy.management_endpoints.key_management_endpoints import (
get_new_token,
)
# Test with invalid new_key (doesn't start with 'sk-')
data = RegenerateKeyRequest(new_key="invalid-key-123")
with pytest.raises(HTTPException) as exc_info:
get_new_token(data)
assert exc_info.value.status_code == 400
assert "New key must start with 'sk-'" in str(exc_info.value.detail)