From 35e640076b01588d78485e8f76d5171d67093007 Mon Sep 17 00:00:00 2001 From: Krrish Dholakia Date: Fri, 19 Jul 2024 17:06:49 -0700 Subject: [PATCH 1/3] fix(user_api_key_auth.py): update valid token cache with updated team object cache --- litellm/proxy/_experimental/out/404.html | 1 - .../proxy/_experimental/out/model_hub.html | 1 - .../proxy/_experimental/out/onboarding.html | 1 - litellm/proxy/_new_secret_config.yaml | 7 ++-- litellm/proxy/auth/auth_checks.py | 15 +++++++- litellm/proxy/auth/user_api_key_auth.py | 22 ++++++++++- .../management_endpoints/team_endpoints.py | 25 ++++++++++--- litellm/tests/test_user_api_key_auth.py | 37 +++++++++++++++++++ 8 files changed, 94 insertions(+), 15 deletions(-) delete mode 100644 litellm/proxy/_experimental/out/404.html delete mode 100644 litellm/proxy/_experimental/out/model_hub.html delete mode 100644 litellm/proxy/_experimental/out/onboarding.html diff --git a/litellm/proxy/_experimental/out/404.html b/litellm/proxy/_experimental/out/404.html deleted file mode 100644 index 26066d6b95d..00000000000 --- a/litellm/proxy/_experimental/out/404.html +++ /dev/null @@ -1 +0,0 @@ -404: This page could not be found.LiteLLM Dashboard

404

This page could not be found.

\ No newline at end of file diff --git a/litellm/proxy/_experimental/out/model_hub.html b/litellm/proxy/_experimental/out/model_hub.html deleted file mode 100644 index dfcd24baab4..00000000000 --- a/litellm/proxy/_experimental/out/model_hub.html +++ /dev/null @@ -1 +0,0 @@ -LiteLLM Dashboard \ No newline at end of file diff --git a/litellm/proxy/_experimental/out/onboarding.html b/litellm/proxy/_experimental/out/onboarding.html deleted file mode 100644 index 36485485ce6..00000000000 --- a/litellm/proxy/_experimental/out/onboarding.html +++ /dev/null @@ -1 +0,0 @@ -LiteLLM Dashboard \ No newline at end of file diff --git a/litellm/proxy/_new_secret_config.yaml b/litellm/proxy/_new_secret_config.yaml index 68ee59d866e..a70bbb7be61 100644 --- a/litellm/proxy/_new_secret_config.yaml +++ b/litellm/proxy/_new_secret_config.yaml @@ -1,5 +1,6 @@ model_list: - - model_name: bad-azure-model + - model_name: azure-chatgpt litellm_params: - model: gpt-4 - request_timeout: 1 + model: azure/chatgpt-v-2 + api_key: os.environ/AZURE_API_KEY + api_base: os.environ/AZURE_API_BASE diff --git a/litellm/proxy/auth/auth_checks.py b/litellm/proxy/auth/auth_checks.py index 655de7964e1..96171f2efb7 100644 --- a/litellm/proxy/auth/auth_checks.py +++ b/litellm/proxy/auth/auth_checks.py @@ -59,7 +59,7 @@ def common_checks( 6. [OPTIONAL] If 'litellm.max_budget' is set (>0), is proxy under budget """ _model = request_body.get("model", None) - if team_object is not None and team_object.blocked == True: + if team_object is not None and team_object.blocked is True: raise Exception( f"Team={team_object.team_id} is blocked. Update via `/team/unblock` if your admin." ) @@ -349,6 +349,15 @@ async def get_user_object( ) +async def _cache_team_object( + team_id: str, + team_table: LiteLLM_TeamTable, + user_api_key_cache: DualCache, +): + key = "team_id:{}".format(team_id) + await user_api_key_cache.async_set_cache(key=key, value=team_table) + + @log_to_opentelemetry async def get_team_object( team_id: str, @@ -386,7 +395,9 @@ async def get_team_object( _response = LiteLLM_TeamTable(**response.dict()) # save the team object to cache - await user_api_key_cache.async_set_cache(key=key, value=_response) + await _cache_team_object( + team_id=team_id, team_table=_response, user_api_key_cache=user_api_key_cache + ) return _response except Exception as e: diff --git a/litellm/proxy/auth/user_api_key_auth.py b/litellm/proxy/auth/user_api_key_auth.py index dcd2cbb804a..0feb13bc850 100644 --- a/litellm/proxy/auth/user_api_key_auth.py +++ b/litellm/proxy/auth/user_api_key_auth.py @@ -453,6 +453,27 @@ async def user_api_key_auth( return valid_token + if ( + valid_token is not None + and isinstance(valid_token, UserAPIKeyAuth) + and valid_token.team_id is not None + and user_api_key_cache.get_cache( + key="team_id:{}".format(valid_token.team_id) + ) + is not None + ): + ## UPDATE TEAM VALUES BASED ON CACHED TEAM OBJECT - allows `/team/update` values to work for cached token + team_obj: LiteLLM_TeamTable = user_api_key_cache.get_cache( + key="team_id:{}".format(valid_token.team_id) + ) + + team_obj_dict = team_obj.__dict__ + + for k, v in team_obj_dict.items(): + field_name = f"team_{k}" + if field_name in valid_token.__fields__: + setattr(valid_token, field_name, v) + try: is_master_key_valid = secrets.compare_digest(api_key, master_key) # type: ignore except Exception as e: @@ -504,7 +525,6 @@ async def user_api_key_auth( raise Exception("No connected db.") ## check for cache hit (In-Memory Cache) - original_api_key = api_key # (Patch: For DynamoDB Backwards Compatibility) _user_role = None if api_key.startswith("sk-"): api_key = hash_token(token=api_key) diff --git a/litellm/proxy/management_endpoints/team_endpoints.py b/litellm/proxy/management_endpoints/team_endpoints.py index 4a5faaf77a1..bb98a02ec3e 100644 --- a/litellm/proxy/management_endpoints/team_endpoints.py +++ b/litellm/proxy/management_endpoints/team_endpoints.py @@ -328,11 +328,13 @@ async def update_team( }' ``` """ + from litellm.proxy.auth.auth_checks import _cache_team_object from litellm.proxy.proxy_server import ( _duration_in_seconds, create_audit_log_for_update, litellm_proxy_admin_name, prisma_client, + user_api_key_cache, ) if prisma_client is None: @@ -361,11 +363,22 @@ async def update_team( # set the budget_reset_at in DB updated_kv["budget_reset_at"] = reset_at - team_row = await prisma_client.update_data( - update_key_values=updated_kv, - data=updated_kv, - table_name="team", - team_id=data.team_id, + team_row: Optional[ + LiteLLM_TeamTable + ] = await prisma_client.db.litellm_teamtable.update( + where={"team_id": data.team_id}, data=updated_kv # type: ignore + ) + + if team_row is None or team_row.team_id is None: + raise HTTPException( + status_code=400, + detail={"error": "Team doesn't exist. Got={}".format(team_row)}, + ) + + await _cache_team_object( + team_id=team_row.team_id, + team_table=team_row, + user_api_key_cache=user_api_key_cache, ) # Enterprise Feature - Audit Logging. Enable with litellm.store_audit_logs = True @@ -392,7 +405,7 @@ async def update_team( ) ) - return team_row + return {"team_id": team_row.team_id, "data": team_row} @router.post( diff --git a/litellm/tests/test_user_api_key_auth.py b/litellm/tests/test_user_api_key_auth.py index 8d3a0af53e8..6460dfa5e3e 100644 --- a/litellm/tests/test_user_api_key_auth.py +++ b/litellm/tests/test_user_api_key_auth.py @@ -44,3 +44,40 @@ def test_check_valid_ip( request = Request(client_ip) assert _check_valid_ip(allowed_ips, request) == expected_result # type: ignore + + +@pytest.mark.asyncio +async def test_check_blocked_team(): + """ + cached valid_token obj has team_blocked = true + + cached team obj has team_blocked = false + + assert team is not blocked + """ + from fastapi import Request + from starlette.datastructures import URL + + from litellm.proxy._types import LiteLLM_TeamTable, UserAPIKeyAuth + from litellm.proxy.auth.user_api_key_auth import user_api_key_auth + from litellm.proxy.proxy_server import hash_token, user_api_key_cache + + _team_id = "1234" + team_obj = LiteLLM_TeamTable(team_id=_team_id, blocked=False) + + user_key = "sk-12345678" + + valid_token = UserAPIKeyAuth( + team_id=_team_id, team_blocked=True, token=hash_token(user_key) + ) + user_api_key_cache.set_cache(key=hash_token(user_key), value=valid_token) + user_api_key_cache.set_cache(key="team_id:{}".format(_team_id), value=team_obj) + + setattr(litellm.proxy.proxy_server, "user_api_key_cache", user_api_key_cache) + setattr(litellm.proxy.proxy_server, "master_key", "sk-1234") + setattr(litellm.proxy.proxy_server, "prisma_client", "hello-world") + + request = Request(scope={"type": "http"}) + request._url = URL(url="/chat/completions") + + await user_api_key_auth(request=request, api_key="Bearer " + user_key) From 99aa311083a3817e65214de43e6be9e9eefd7bfa Mon Sep 17 00:00:00 2001 From: Krrish Dholakia Date: Fri, 19 Jul 2024 17:35:59 -0700 Subject: [PATCH 2/3] fix(user_api_key_auth.py): update team values in token cache if refreshed more recently --- litellm/proxy/_types.py | 4 ++++ litellm/proxy/auth/user_api_key_auth.py | 18 +++++++++++++----- litellm/proxy/utils.py | 4 +++- litellm/tests/test_user_api_key_auth.py | 14 +++++++++++--- 4 files changed, 31 insertions(+), 9 deletions(-) diff --git a/litellm/proxy/_types.py b/litellm/proxy/_types.py index 7acd38e8b3e..e9371c1d8d9 100644 --- a/litellm/proxy/_types.py +++ b/litellm/proxy/_types.py @@ -886,6 +886,7 @@ class LiteLLM_TeamTable(TeamBase): budget_duration: Optional[str] = None budget_reset_at: Optional[datetime] = None model_id: Optional[int] = None + last_refreshed_at: Optional[float] = None model_config = ConfigDict(protected_namespaces=()) @@ -1238,6 +1239,9 @@ class LiteLLM_VerificationTokenView(LiteLLM_VerificationToken): end_user_rpm_limit: Optional[int] = None end_user_max_budget: Optional[float] = None + # Time stamps + last_refreshed_at: Optional[float] = None # last time joint view was pulled from db + class UserAPIKeyAuth( LiteLLM_VerificationTokenView diff --git a/litellm/proxy/auth/user_api_key_auth.py b/litellm/proxy/auth/user_api_key_auth.py index 0feb13bc850..c5549ffcb66 100644 --- a/litellm/proxy/auth/user_api_key_auth.py +++ b/litellm/proxy/auth/user_api_key_auth.py @@ -467,12 +467,17 @@ async def user_api_key_auth( key="team_id:{}".format(valid_token.team_id) ) - team_obj_dict = team_obj.__dict__ + if ( + team_obj.last_refreshed_at is not None + and valid_token.last_refreshed_at is not None + and team_obj.last_refreshed_at > valid_token.last_refreshed_at + ): + team_obj_dict = team_obj.__dict__ - for k, v in team_obj_dict.items(): - field_name = f"team_{k}" - if field_name in valid_token.__fields__: - setattr(valid_token, field_name, v) + for k, v in team_obj_dict.items(): + field_name = f"team_{k}" + if field_name in valid_token.__fields__: + setattr(valid_token, field_name, v) try: is_master_key_valid = secrets.compare_digest(api_key, master_key) # type: ignore @@ -541,10 +546,13 @@ async def user_api_key_auth( parent_otel_span=parent_otel_span, proxy_logging_obj=proxy_logging_obj, ) + if _valid_token is not None: + ## update cached token valid_token = UserAPIKeyAuth( **_valid_token.model_dump(exclude_none=True) ) + verbose_proxy_logger.debug("Token from db: %s", valid_token) elif valid_token is not None and isinstance(valid_token, UserAPIKeyAuth): verbose_proxy_logger.debug("API Key Cache Hit!") diff --git a/litellm/proxy/utils.py b/litellm/proxy/utils.py index f06b6570564..0f87e962abc 100644 --- a/litellm/proxy/utils.py +++ b/litellm/proxy/utils.py @@ -1331,7 +1331,9 @@ class PrismaClient: response["team_models"] = [] if response["team_blocked"] is None: response["team_blocked"] = False - response = LiteLLM_VerificationTokenView(**response) + response = LiteLLM_VerificationTokenView( + **response, last_refreshed_at=time.time() + ) # for prisma we need to cast the expires time to str if response.expires is not None and isinstance( response.expires, datetime diff --git a/litellm/tests/test_user_api_key_auth.py b/litellm/tests/test_user_api_key_auth.py index 6460dfa5e3e..5a8402bdcc6 100644 --- a/litellm/tests/test_user_api_key_auth.py +++ b/litellm/tests/test_user_api_key_auth.py @@ -55,6 +55,9 @@ async def test_check_blocked_team(): assert team is not blocked """ + import asyncio + import time + from fastapi import Request from starlette.datastructures import URL @@ -63,12 +66,17 @@ async def test_check_blocked_team(): from litellm.proxy.proxy_server import hash_token, user_api_key_cache _team_id = "1234" - team_obj = LiteLLM_TeamTable(team_id=_team_id, blocked=False) - user_key = "sk-12345678" valid_token = UserAPIKeyAuth( - team_id=_team_id, team_blocked=True, token=hash_token(user_key) + team_id=_team_id, + team_blocked=True, + token=hash_token(user_key), + last_refreshed_at=time.time(), + ) + await asyncio.sleep(1) + team_obj = LiteLLM_TeamTable( + team_id=_team_id, blocked=False, last_refreshed_at=time.time() ) user_api_key_cache.set_cache(key=hash_token(user_key), value=valid_token) user_api_key_cache.set_cache(key="team_id:{}".format(_team_id), value=team_obj) From ccb8035949d053b711578020772daaeb9f171a78 Mon Sep 17 00:00:00 2001 From: Krrish Dholakia Date: Fri, 19 Jul 2024 18:26:13 -0700 Subject: [PATCH 3/3] fix(batches/main.py): fix linting error --- litellm/batches/main.py | 6 +++--- litellm/types/llms/openai.py | 2 +- 2 files changed, 4 insertions(+), 4 deletions(-) diff --git a/litellm/batches/main.py b/litellm/batches/main.py index af2dc5059d3..79aefa5f518 100644 --- a/litellm/batches/main.py +++ b/litellm/batches/main.py @@ -50,7 +50,7 @@ async def acreate_batch( extra_headers: Optional[Dict[str, str]] = None, extra_body: Optional[Dict[str, str]] = None, **kwargs, -) -> Coroutine[Any, Any, Batch]: +) -> Batch: """ Async: Creates and executes a batch from an uploaded file of request @@ -89,7 +89,7 @@ async def acreate_batch( def create_batch( completion_window: Literal["24h"], - endpoint: Literal["/v1/chat/completions", "/v1/embeddings"], + endpoint: Literal["/v1/chat/completions", "/v1/embeddings", "/v1/completions"], input_file_id: str, custom_llm_provider: Literal["openai"] = "openai", metadata: Optional[Dict[str, str]] = None, @@ -189,7 +189,7 @@ async def aretrieve_batch( extra_headers: Optional[Dict[str, str]] = None, extra_body: Optional[Dict[str, str]] = None, **kwargs, -) -> Coroutine[Any, Any, Batch]: +) -> Batch: """ Async: Retrieves a batch. diff --git a/litellm/types/llms/openai.py b/litellm/types/llms/openai.py index 42f1dac3d6d..294e299dbf4 100644 --- a/litellm/types/llms/openai.py +++ b/litellm/types/llms/openai.py @@ -257,7 +257,7 @@ class CreateBatchRequest(TypedDict, total=False): """ completion_window: Literal["24h"] - endpoint: Literal["/v1/chat/completions", "/v1/embeddings"] + endpoint: Literal["/v1/chat/completions", "/v1/embeddings", "/v1/completions"] input_file_id: str metadata: Optional[Dict[str, str]] extra_headers: Optional[Dict[str, str]]