From 5dfe54d20d05c289503c6fc486f21437293d802d Mon Sep 17 00:00:00 2001 From: Krrish Dholakia Date: Fri, 26 Jan 2024 15:31:37 -0800 Subject: [PATCH] feat(proxy_server.py): save abbreviated key name if allow_user_auth enabled --- litellm/proxy/_types.py | 3 ++ litellm/proxy/proxy_server.py | 5 +++ litellm/proxy/schema.prisma | 2 + litellm/tests/test_key_generate_prisma.py | 47 +++++++++++++++++++++++ schema.prisma | 2 + 5 files changed, 59 insertions(+) diff --git a/litellm/proxy/_types.py b/litellm/proxy/_types.py index 9a5acc44064..0c38504b89f 100644 --- a/litellm/proxy/_types.py +++ b/litellm/proxy/_types.py @@ -140,6 +140,7 @@ class GenerateRequestBase(LiteLLMBase): class GenerateKeyRequest(GenerateRequestBase): + key_alias: Optional[str] = None duration: Optional[str] = "1h" aliases: Optional[dict] = {} config: Optional[dict] = {} @@ -304,6 +305,8 @@ class ConfigYAML(LiteLLMBase): class LiteLLM_VerificationToken(LiteLLMBase): token: str + key_name: Optional[str] = None + key_alias: Optional[str] = None spend: float = 0.0 max_budget: Optional[float] = None expires: Union[str, None] diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index 4a84847e0d9..d390d7a4132 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -243,6 +243,7 @@ async def user_api_key_auth( response = await user_custom_auth(request=request, api_key=api_key) return UserAPIKeyAuth.model_validate(response) ### LITELLM-DEFINED AUTH FUNCTION ### + assert api_key.startswith("sk-") # prevent token hashes from being used if master_key is None: if isinstance(api_key, str): return UserAPIKeyAuth(api_key=api_key) @@ -1239,6 +1240,7 @@ async def generate_key_helper_fn( rpm_limit: Optional[int] = None, query_type: Literal["insert_data", "update_data"] = "insert_data", update_key_values: Optional[dict] = None, + key_alias: Optional[str] = None, ): global prisma_client, custom_db_client @@ -1312,6 +1314,7 @@ async def generate_key_helper_fn( } key_data = { "token": token, + "key_alias": key_alias, "expires": expires, "models": models, "aliases": aliases_json, @@ -1327,6 +1330,8 @@ async def generate_key_helper_fn( "budget_duration": key_budget_duration, "budget_reset_at": key_reset_at, } + if general_settings.get("allow_user_auth", False) == True: + key_data["key_name"] = f"sk-...{token[-4:]}" if prisma_client is not None: ## CREATE USER (If necessary) verbose_proxy_logger.debug(f"prisma_client: Creating User={user_data}") diff --git a/litellm/proxy/schema.prisma b/litellm/proxy/schema.prisma index 2eb6332092a..0b379a23716 100644 --- a/litellm/proxy/schema.prisma +++ b/litellm/proxy/schema.prisma @@ -24,6 +24,8 @@ model LiteLLM_UserTable { // required for token gen model LiteLLM_VerificationToken { token String @unique + key_name String? + key_alias String? spend Float @default(0.0) expires DateTime? models String[] diff --git a/litellm/tests/test_key_generate_prisma.py b/litellm/tests/test_key_generate_prisma.py index 98a056730d6..9cc1a1754d2 100644 --- a/litellm/tests/test_key_generate_prisma.py +++ b/litellm/tests/test_key_generate_prisma.py @@ -12,6 +12,8 @@ # 11. Generate a Key, cal key/info, call key/update, call key/info # 12. Make a call with key over budget, expect to fail # 14. Make a streaming chat/completions call with key over budget, expect to fail +# 15. Generate key, when `allow_user_auth`=False - check if `/key/info` returns key_name=null +# 16. Generate key, when `allow_user_auth`=True - check if `/key/info` returns key_name=sk... # function to call to generate key - async def new_user(data: NewUserRequest): @@ -1140,3 +1142,48 @@ async def test_view_spend_per_key(prisma_client): except Exception as e: print("Got Exception", e) pytest.fail(f"Got exception {e}") + + +@pytest.mark.asyncio() +async def test_key_name_null(prisma_client): + """ + - create key + - get key info + - assert key_name is null + """ + setattr(litellm.proxy.proxy_server, "prisma_client", prisma_client) + setattr(litellm.proxy.proxy_server, "master_key", "sk-1234") + await litellm.proxy.proxy_server.prisma_client.connect() + try: + request = GenerateKeyRequest() + key = await generate_key_fn(request) + generated_key = key.key + result = await info_key_fn(key=generated_key) + print("result from info_key_fn", result) + assert result["info"]["key_name"] is None + except Exception as e: + print("Got Exception", e) + pytest.fail(f"Got exception {e}") + + +@pytest.mark.asyncio() +async def test_key_name_set(prisma_client): + """ + - create key + - get key info + - assert key_name is not null + """ + setattr(litellm.proxy.proxy_server, "prisma_client", prisma_client) + setattr(litellm.proxy.proxy_server, "master_key", "sk-1234") + setattr(litellm.proxy.proxy_server, "general_settings", {"allow_user_auth": True}) + await litellm.proxy.proxy_server.prisma_client.connect() + try: + request = GenerateKeyRequest() + key = await generate_key_fn(request) + generated_key = key.key + result = await info_key_fn(key=generated_key) + print("result from info_key_fn", result) + assert isinstance(result["info"]["key_name"], str) + except Exception as e: + print("Got Exception", e) + pytest.fail(f"Got exception {e}") diff --git a/schema.prisma b/schema.prisma index 0882c650c81..02e4114e5d9 100644 --- a/schema.prisma +++ b/schema.prisma @@ -25,6 +25,8 @@ model LiteLLM_UserTable { // Generate Tokens for Proxy model LiteLLM_VerificationToken { token String @unique + key_name String? + key_alias String? spend Float @default(0.0) expires DateTime? models String[]