diff --git a/docs/my-website/docs/proxy/virtual_keys.md b/docs/my-website/docs/proxy/virtual_keys.md index c20bbef1127..e1c89bbc216 100644 --- a/docs/my-website/docs/proxy/virtual_keys.md +++ b/docs/my-website/docs/proxy/virtual_keys.md @@ -73,7 +73,8 @@ curl 'http://0.0.0.0:8000/key/generate' \ "models": ["gpt-3.5-turbo", "gpt-4", "claude-2"], "duration": "20m", "metadata": {"user": "ishaan@berri.ai"}, - "team_id": "core-infra" + "team_id": "core-infra", + "max_budget": 10, }' ``` @@ -88,6 +89,8 @@ Request Params: - `team_id`: *str or null (optional)* Specify team_id for the associated key +- `max_budget`: *float or null (optional)* Specify max budget (in Dollars $) for a given key. If no value is set, the key has no budget + ### Response ```python @@ -282,6 +285,77 @@ Request Params: } ``` +## Set Budgets - Per Key + +Set `max_budget` in (USD $) param in the `key/generate` request. By default the `max_budget` is set to `null` and is not checked for keys + +```shell +curl 'http://0.0.0.0:8000/key/generate' \ +--header 'Authorization: Bearer ' \ +--header 'Content-Type: application/json' \ +--data-raw '{ + "metadata": {"user": "ishaan@berri.ai"}, + "team_id": "core-infra", + "max_budget": 10, +}' +``` + +#### Expected Behaviour +- Costs Per key get auto-populated in `LiteLLM_VerificationToken` Table +- After the key crosses it's `max_budget`, requests fail + +Example Request to `/chat/completions` when key has crossed budget + +```shell +curl --location 'http://0.0.0.0:8000/chat/completions' \ + --header 'Content-Type: application/json' \ + --header 'Authorization: Bearer sk-ULl_IKCVFy2EZRzQB16RUA' \ + --data ' { + "model": "azure-gpt-3.5", + "user": "e09b4da8-ed80-4b05-ac93-e16d9eb56fca", + "messages": [ + { + "role": "user", + "content": "respond in 50 lines" + } + ], +}' +``` + + +Expected Response from `/chat/completions` when key has crossed budget +```shell +{ + "detail":"Authentication Error, ExceededTokenBudget: Current spend for token: 7.2e-05; Max Budget for Token: 2e-07" +} +``` + + +## Set Budgets - Per User + +LiteLLM exposes a `/user/new` endpoint to create budgets for users, that persist across multiple keys. + +This is documented in the swagger (live on your server root endpoint - e.g. `http://0.0.0.0:8000/`). Here's an example request. + +```shell +curl --location 'http://localhost:8000/user/new' \ +--header 'Authorization: Bearer ' \ +--header 'Content-Type: application/json' \ +--data-raw '{"models": ["azure-models"], "max_budget": 0, "user_id": "krrish3@berri.ai"}' +``` +The request is a normal `/key/generate` request body + a `max_budget` field. + +**Sample Response** + +```shell +{ + "key": "sk-YF2OxDbrgd1y2KgwxmEA2w", + "expires": "2023-12-22T09:53:13.861000Z", + "user_id": "krrish3@berri.ai", + "max_budget": 0.0 +} +``` + ## Tracking Spend You can get spend for a key by using the `/key/info` endpoint. @@ -317,32 +391,6 @@ This is automatically updated (in USD) when calls are made to /completions, /cha ``` - -## Set Budgets - -LiteLLM exposes a `/user/new` endpoint to create budgets for users, that persist across multiple keys. - -This is documented in the swagger (live on your server root endpoint - e.g. `http://0.0.0.0:8000/`). Here's an example request. - -```shell -curl --location 'http://localhost:8000/user/new' \ ---header 'Authorization: Bearer ' \ ---header 'Content-Type: application/json' \ ---data-raw '{"models": ["azure-models"], "max_budget": 0, "user_id": "krrish3@berri.ai"}' -``` -The request is a normal `/key/generate` request body + a `max_budget` field. - -**Sample Response** - -```shell -{ - "key": "sk-YF2OxDbrgd1y2KgwxmEA2w", - "expires": "2023-12-22T09:53:13.861000Z", - "user_id": "krrish3@berri.ai", - "max_budget": 0.0 -} -``` - ## Custom Auth You can now override the default api key auth. diff --git a/litellm/proxy/_types.py b/litellm/proxy/_types.py index 72b7273e513..bb56ad6bf1b 100644 --- a/litellm/proxy/_types.py +++ b/litellm/proxy/_types.py @@ -128,6 +128,7 @@ class GenerateKeyRequest(LiteLLMBase): aliases: Optional[dict] = {} config: Optional[dict] = {} spend: Optional[float] = 0 + max_budget: Optional[float] = None user_id: Optional[str] = None team_id: Optional[str] = None max_parallel_requests: Optional[int] = None @@ -145,6 +146,7 @@ class UpdateKeyRequest(LiteLLMBase): aliases: Optional[dict] = None config: Optional[dict] = None spend: Optional[float] = None + max_budget: Optional[float] = None user_id: Optional[str] = None max_parallel_requests: Optional[int] = None metadata: Optional[dict] = None @@ -162,6 +164,7 @@ class UserAPIKeyAuth(LiteLLMBase): # the expected response object for user api aliases: dict = {} config: dict = {} spend: Optional[float] = 0 + max_budget: Optional[float] = None user_id: Optional[str] = None max_parallel_requests: Optional[int] = None duration: str = "1h" @@ -294,6 +297,7 @@ class ConfigYAML(LiteLLMBase): class LiteLLM_VerificationToken(LiteLLMBase): token: str spend: float = 0.0 + max_budget: Optional[float] = None expires: Union[str, None] models: List[str] aliases: Dict[str, str] = {} diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index 9ef9a815818..24caf5b94fd 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -307,6 +307,7 @@ async def user_api_key_auth( # 1. If token can call model # 2. If user_id for this token is in budget # 3. If token is expired + # 4. If token spend is under Budget for the token # Check 1. If token can call model litellm.model_alias_map = valid_token.aliases @@ -407,6 +408,13 @@ async def user_api_key_auth( detail=f"Authentication Error - Expired Key. Key Expiry time {expiry_time} and current time {current_time}", ) + # Check 4. Token Spend is under budget + if valid_token.spend is not None and valid_token.max_budget is not None: + if valid_token.spend > valid_token.max_budget: + raise Exception( + f"ExceededTokenBudget: Current spend for token: {valid_token.spend}; Max Budget for Token: {valid_token.max_budget}" + ) + # Token passed all checks # Add token to cache user_api_key_cache.set_cache(key=api_key, value=valid_token, ttl=60) @@ -669,7 +677,9 @@ async def update_database( if prisma_client is not None: # Fetch the existing cost for the given token existing_spend_obj = await prisma_client.get_data(token=token) - verbose_proxy_logger.debug(f"existing spend: {existing_spend_obj}") + verbose_proxy_logger.debug( + f"_update_key_db: existing spend: {existing_spend_obj}" + ) if existing_spend_obj is None: existing_spend = 0 else: @@ -680,12 +690,19 @@ async def update_database( verbose_proxy_logger.debug(f"new cost: {new_spend}") # Update the cost column for the given token await prisma_client.update_data(token=token, data={"spend": new_spend}) + + valid_token = user_api_key_cache.get_cache(key=token) + if valid_token is not None: + valid_token.spend = new_spend + user_api_key_cache.set_cache(key=token, value=valid_token) elif custom_db_client is not None: # Fetch the existing cost for the given token existing_spend_obj = await custom_db_client.get_data( key=token, table_name="key" ) - verbose_proxy_logger.debug(f"existing spend: {existing_spend_obj}") + verbose_proxy_logger.debug( + f"_update_key_db existing spend: {existing_spend_obj}" + ) if existing_spend_obj is None: existing_spend = 0 else: @@ -699,6 +716,11 @@ async def update_database( key=token, value={"spend": new_spend}, table_name="key" ) + valid_token = user_api_key_cache.get_cache(key=token) + if valid_token is not None: + valid_token.spend = new_spend + user_api_key_cache.set_cache(key=token, value=valid_token) + async def _insert_spend_log_to_db(): # Helper to generate payload to log verbose_proxy_logger.debug("inserting spend log to db") @@ -1128,7 +1150,8 @@ async def generate_key_helper_fn( aliases: dict, config: dict, spend: float, - max_budget: Optional[float] = None, + key_max_budget: Optional[float] = None, # key_max_budget is used to Budget Per key + max_budget: Optional[float] = None, # max_budget is used to Budget Per user token: Optional[str] = None, user_id: Optional[str] = None, team_id: Optional[str] = None, @@ -1201,6 +1224,7 @@ async def generate_key_helper_fn( "aliases": aliases_json, "config": config_json, "spend": spend, + "max_budget": key_max_budget, "user_id": user_id, "team_id": team_id, "max_parallel_requests": max_parallel_requests, @@ -2163,6 +2187,7 @@ async def generate_key_fn( - aliases: Optional[dict] - Any alias mappings, on top of anything in the config.yaml model list. - https://docs.litellm.ai/docs/proxy/virtual_keys#managing-auth---upgradedowngrade-models - config: Optional[dict] - any key-specific configs, overrides config in config.yaml - spend: Optional[int] - Amount spent by key. Default is 0. Will be updated by proxy whenever key is used. https://docs.litellm.ai/docs/proxy/virtual_keys#managing-auth---tracking-spend + - max_budget: Optional[float] - Specify max budget for a given key. - max_parallel_requests: Optional[int] - Rate limit a user based on the number of parallel requests. Raises 429 error, if user's parallel requests > x. - metadata: Optional[dict] - Metadata for key, store information for key. Example metadata = {"team": "core-infra", "app": "app2", "email": "ishaan@berri.ai" } @@ -2182,6 +2207,11 @@ async def generate_key_fn( raise HTTPException(status_code=status.HTTP_403_FORBIDDEN, detail=message) data_json = data.json() # type: ignore + + # if we get max_budget passed to /key/generate, then use it as key_max_budget. Since generate_key_helper_fn is used to make new users + if "max_budget" in data_json: + data_json["key_max_budget"] = data_json.pop("max_budget", None) + response = await generate_key_helper_fn(**data_json) return GenerateKeyResponse( key=response["token"], expires=response["expires"], user_id=response["user_id"] diff --git a/litellm/proxy/schema.prisma b/litellm/proxy/schema.prisma index 1c73a1405a4..931a1581258 100644 --- a/litellm/proxy/schema.prisma +++ b/litellm/proxy/schema.prisma @@ -33,6 +33,7 @@ model LiteLLM_VerificationToken { metadata Json @default("{}") tpm_limit BigInt? rpm_limit BigInt? + max_budget Float? @default(0.0) } model LiteLLM_Config { diff --git a/litellm/tests/test_key_generate_dynamodb.py b/litellm/tests/test_key_generate_dynamodb.py index 2cfa9c95312..a0772f87a37 100644 --- a/litellm/tests/test_key_generate_dynamodb.py +++ b/litellm/tests/test_key_generate_dynamodb.py @@ -25,9 +25,14 @@ sys.path.insert( ) # Adds the parent directory to the system path import pytest, logging, asyncio import litellm, asyncio -from litellm.proxy.proxy_server import new_user, user_api_key_auth, user_update +from litellm.proxy.proxy_server import ( + new_user, + user_api_key_auth, + user_update, + generate_key_fn, +) -from litellm.proxy._types import NewUserRequest, DynamoDBArgs +from litellm.proxy._types import NewUserRequest, DynamoDBArgs, GenerateKeyRequest from litellm.proxy.utils import DBClient from starlette.datastructures import URL @@ -175,7 +180,7 @@ def test_call_with_valid_model(custom_db_client): pytest.fail(f"An exception occurred - {str(e)}") -def test_call_with_key_over_budget(custom_db_client): +def test_call_with_user_over_budget(custom_db_client): # 5. Make a call with a key over budget, expect to fail setattr(litellm.proxy.proxy_server, "custom_db_client", custom_db_client) setattr(litellm.proxy.proxy_server, "master_key", "sk-1234") @@ -245,7 +250,7 @@ def test_call_with_key_over_budget(custom_db_client): print(vars(e)) -def test_call_with_key_over_budget_stream(custom_db_client): +def test_call_with_user_over_budget_stream(custom_db_client): # 6. Make a call with a key over budget, expect to fail setattr(litellm.proxy.proxy_server, "custom_db_client", custom_db_client) setattr(litellm.proxy.proxy_server, "master_key", "sk-1234") @@ -315,3 +320,145 @@ def test_call_with_key_over_budget_stream(custom_db_client): error_detail = e.detail assert "Authentication Error, ExceededBudget:" in error_detail print(vars(e)) + + +def test_call_with_user_key_budget(custom_db_client): + # 7. Make a call with a key over budget, expect to fail + setattr(litellm.proxy.proxy_server, "custom_db_client", custom_db_client) + setattr(litellm.proxy.proxy_server, "master_key", "sk-1234") + from litellm._logging import verbose_proxy_logger + import logging + + verbose_proxy_logger.setLevel(logging.DEBUG) + try: + + async def test(): + request = GenerateKeyRequest(max_budget=0.00001) + key = await generate_key_fn(request) + print(key) + + generated_key = key.key + user_id = key.user_id + bearer_token = "Bearer " + generated_key + + request = Request(scope={"type": "http"}) + request._url = URL(url="/chat/completions") + + # use generated key to auth in + result = await user_api_key_auth(request=request, api_key=bearer_token) + print("result from user auth with new key", result) + + # update spend using track_cost callback, make 2nd request, it should fail + from litellm.proxy.proxy_server import track_cost_callback + from litellm import ModelResponse, Choices, Message, Usage + + resp = ModelResponse( + id="chatcmpl-e41836bb-bb8b-4df2-8e70-8f3e160155ac", + choices=[ + Choices( + finish_reason=None, + index=0, + message=Message( + content=" Sure! Here is a short poem about the sky:\n\nA canvas of blue, a", + role="assistant", + ), + ) + ], + model="gpt-35-turbo", # azure always has model written like this + usage=Usage(prompt_tokens=210, completion_tokens=200, total_tokens=410), + ) + await track_cost_callback( + kwargs={ + "stream": False, + "litellm_params": { + "metadata": { + "user_api_key": generated_key, + "user_api_key_user_id": user_id, + } + }, + }, + completion_response=resp, + ) + + # use generated key to auth in + result = await user_api_key_auth(request=request, api_key=bearer_token) + print("result from user auth with new key", result) + pytest.fail(f"This should have failed!. They key crossed it's budget") + + asyncio.run(test()) + except Exception as e: + error_detail = e.detail + assert "Authentication Error, ExceededTokenBudget:" in error_detail + print(vars(e)) + + +def test_call_with_key_over_budget_stream(custom_db_client): + # 8. Make a call with a key over budget, expect to fail + setattr(litellm.proxy.proxy_server, "custom_db_client", custom_db_client) + setattr(litellm.proxy.proxy_server, "master_key", "sk-1234") + from litellm._logging import verbose_proxy_logger + import logging + + litellm.set_verbose = True + verbose_proxy_logger.setLevel(logging.DEBUG) + try: + + async def test(): + request = GenerateKeyRequest(max_budget=0.00001) + key = await generate_key_fn(request) + print(key) + + generated_key = key.key + user_id = key.user_id + bearer_token = "Bearer " + generated_key + + request = Request(scope={"type": "http"}) + request._url = URL(url="/chat/completions") + + # use generated key to auth in + result = await user_api_key_auth(request=request, api_key=bearer_token) + print("result from user auth with new key", result) + + # update spend using track_cost callback, make 2nd request, it should fail + from litellm.proxy.proxy_server import track_cost_callback + from litellm import ModelResponse, Choices, Message, Usage + + resp = ModelResponse( + id="chatcmpl-e41836bb-bb8b-4df2-8e70-8f3e160155ac", + choices=[ + Choices( + finish_reason=None, + index=0, + message=Message( + content=" Sure! Here is a short poem about the sky:\n\nA canvas of blue, a", + role="assistant", + ), + ) + ], + model="gpt-35-turbo", # azure always has model written like this + usage=Usage(prompt_tokens=210, completion_tokens=200, total_tokens=410), + ) + await track_cost_callback( + kwargs={ + "stream": True, + "complete_streaming_response": resp, + "litellm_params": { + "metadata": { + "user_api_key": generated_key, + "user_api_key_user_id": user_id, + } + }, + }, + completion_response=ModelResponse(), + ) + + # use generated key to auth in + result = await user_api_key_auth(request=request, api_key=bearer_token) + print("result from user auth with new key", result) + pytest.fail(f"This should have failed!. They key crossed it's budget") + + asyncio.run(test()) + except Exception as e: + error_detail = e.detail + assert "Authentication Error, ExceededTokenBudget:" in error_detail + print(vars(e)) diff --git a/litellm/tests/test_key_generate_prisma.py b/litellm/tests/test_key_generate_prisma.py index 5ecfc89d75c..4734e103022 100644 --- a/litellm/tests/test_key_generate_prisma.py +++ b/litellm/tests/test_key_generate_prisma.py @@ -3,13 +3,15 @@ # 2. Make a call with invalid key, expect it to fail # 3. Make a call to a key with invalid model - expect to fail # 4. Make a call to a key with valid model - expect to pass -# 5. Make a call with key over budget, expect to fail -# 6. Make a streaming chat/completions call with key over budget, expect to fail +# 5. Make a call with user over budget, expect to fail +# 6. Make a streaming chat/completions call with user over budget, expect to fail # 7. Make a call with an key that never expires, expect to pass # 8. Make a call with an expired key, expect to fail # 9. Delete a Key # 10. Generate a key, call key/info. Assert info returned is the same as generated key info # 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 # function to call to generate key - async def new_user(data: NewUserRequest): @@ -39,6 +41,7 @@ from litellm.proxy.proxy_server import ( delete_key_fn, info_key_fn, update_key_fn, + generate_key_fn, ) from litellm.proxy.utils import PrismaClient, ProxyLogging from litellm._logging import verbose_proxy_logger @@ -47,6 +50,7 @@ verbose_proxy_logger.setLevel(level=logging.DEBUG) from litellm.proxy._types import ( NewUserRequest, + GenerateKeyRequest, DynamoDBArgs, DeleteKeyRequest, UpdateKeyRequest, @@ -205,7 +209,7 @@ def test_call_with_valid_model(prisma_client): pytest.fail(f"An exception occurred - {str(e)}") -def test_call_with_key_over_budget(prisma_client): +def test_call_with_user_over_budget(prisma_client): # 5. Make a call with a key over budget, expect to fail setattr(litellm.proxy.proxy_server, "prisma_client", prisma_client) setattr(litellm.proxy.proxy_server, "master_key", "sk-1234") @@ -274,7 +278,7 @@ def test_call_with_key_over_budget(prisma_client): print(vars(e)) -def test_call_with_key_over_budget_stream(prisma_client): +def test_call_with_user_over_budget_stream(prisma_client): # 6. Make a call with a key over budget, expect to fail setattr(litellm.proxy.proxy_server, "prisma_client", prisma_client) setattr(litellm.proxy.proxy_server, "master_key", "sk-1234") @@ -651,7 +655,6 @@ def test_key_generate_with_custom_auth(prisma_client): try: async def test(): - await litellm.proxy.proxy_server.prisma_client.connect() try: request = GenerateKeyRequest() key = await generate_key_fn(request) @@ -679,3 +682,146 @@ def test_key_generate_with_custom_auth(prisma_client): print("Got Exception", e) print(e.detail) pytest.fail(f"An exception occurred - {str(e)}") + + +def test_call_with_key_over_budget(prisma_client): + # 12. Make a call with a key over budget, expect to fail + setattr(litellm.proxy.proxy_server, "prisma_client", prisma_client) + setattr(litellm.proxy.proxy_server, "master_key", "sk-1234") + try: + + async def test(): + await litellm.proxy.proxy_server.prisma_client.connect() + request = GenerateKeyRequest(max_budget=0.00001) + key = await generate_key_fn(request) + print(key) + + generated_key = key.key + user_id = key.user_id + bearer_token = "Bearer " + generated_key + + request = Request(scope={"type": "http"}) + request._url = URL(url="/chat/completions") + + # use generated key to auth in + result = await user_api_key_auth(request=request, api_key=bearer_token) + print("result from user auth with new key", result) + + # update spend using track_cost callback, make 2nd request, it should fail + from litellm.proxy.proxy_server import track_cost_callback + from litellm import ModelResponse, Choices, Message, Usage + + resp = ModelResponse( + id="chatcmpl-e41836bb-bb8b-4df2-8e70-8f3e160155ac", + choices=[ + Choices( + finish_reason=None, + index=0, + message=Message( + content=" Sure! Here is a short poem about the sky:\n\nA canvas of blue, a", + role="assistant", + ), + ) + ], + model="gpt-35-turbo", # azure always has model written like this + usage=Usage(prompt_tokens=210, completion_tokens=200, total_tokens=410), + ) + await track_cost_callback( + kwargs={ + "stream": False, + "litellm_params": { + "metadata": { + "user_api_key": generated_key, + "user_api_key_user_id": user_id, + } + }, + }, + completion_response=resp, + start_time=datetime.now(), + end_time=datetime.now(), + ) + + # use generated key to auth in + result = await user_api_key_auth(request=request, api_key=bearer_token) + print("result from user auth with new key", result) + pytest.fail(f"This should have failed!. They key crossed it's budget") + + asyncio.run(test()) + except Exception as e: + error_detail = e.detail + assert "Authentication Error, ExceededTokenBudget:" in error_detail + print(vars(e)) + + +def test_call_with_key_over_budget_stream(prisma_client): + # 14. Make a call with a key over budget, expect to fail + setattr(litellm.proxy.proxy_server, "prisma_client", prisma_client) + setattr(litellm.proxy.proxy_server, "master_key", "sk-1234") + from litellm._logging import verbose_proxy_logger + import logging + + litellm.set_verbose = True + verbose_proxy_logger.setLevel(logging.DEBUG) + try: + + async def test(): + await litellm.proxy.proxy_server.prisma_client.connect() + request = GenerateKeyRequest(max_budget=0.00001) + key = await generate_key_fn(request) + print(key) + + generated_key = key.key + user_id = key.user_id + bearer_token = "Bearer " + generated_key + + request = Request(scope={"type": "http"}) + request._url = URL(url="/chat/completions") + + # use generated key to auth in + result = await user_api_key_auth(request=request, api_key=bearer_token) + print("result from user auth with new key", result) + + # update spend using track_cost callback, make 2nd request, it should fail + from litellm.proxy.proxy_server import track_cost_callback + from litellm import ModelResponse, Choices, Message, Usage + + resp = ModelResponse( + id="chatcmpl-e41836bb-bb8b-4df2-8e70-8f3e160155ac", + choices=[ + Choices( + finish_reason=None, + index=0, + message=Message( + content=" Sure! Here is a short poem about the sky:\n\nA canvas of blue, a", + role="assistant", + ), + ) + ], + model="gpt-35-turbo", # azure always has model written like this + usage=Usage(prompt_tokens=210, completion_tokens=200, total_tokens=410), + ) + await track_cost_callback( + kwargs={ + "stream": True, + "complete_streaming_response": resp, + "litellm_params": { + "metadata": { + "user_api_key": generated_key, + "user_api_key_user_id": user_id, + } + }, + }, + completion_response=ModelResponse(), + start_time=datetime.now(), + end_time=datetime.now(), + ) + + # use generated key to auth in + result = await user_api_key_auth(request=request, api_key=bearer_token) + print("result from user auth with new key", result) + pytest.fail(f"This should have failed!. They key crossed it's budget") + + except Exception as e: + error_detail = e.detail + assert "Authentication Error, ExceededTokenBudget:" in error_detail + print(vars(e)) diff --git a/schema.prisma b/schema.prisma index 07d4d342242..1212b0c661a 100644 --- a/schema.prisma +++ b/schema.prisma @@ -33,6 +33,7 @@ model LiteLLM_VerificationToken { metadata Json @default("{}") tpm_limit BigInt? rpm_limit BigInt? + max_budget Float? @default(0.0) } model LiteLLM_Config {