diff --git a/litellm/proxy/management_endpoints/key_management_endpoints.py b/litellm/proxy/management_endpoints/key_management_endpoints.py index ab78bd1d977..9215060dfed 100644 --- a/litellm/proxy/management_endpoints/key_management_endpoints.py +++ b/litellm/proxy/management_endpoints/key_management_endpoints.py @@ -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:]}" diff --git a/tests/test_litellm/proxy/management_endpoints/test_key_management_endpoints.py b/tests/test_litellm/proxy/management_endpoints/test_key_management_endpoints.py index ddecb94ab64..54909999038 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_key_management_endpoints.py +++ b/tests/test_litellm/proxy/management_endpoints/test_key_management_endpoints.py @@ -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)