mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-09 22:31:41 +00:00
test: add unit tests
This commit is contained in:
parent
eec8e1a746
commit
35c6f6f83f
2 changed files with 49 additions and 11 deletions
|
|
@ -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:]}"
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue