From f2904cbb4e7829cb35b69ebbf98442e87385e09f Mon Sep 17 00:00:00 2001 From: Krrish Dholakia Date: Thu, 12 Dec 2024 16:07:06 -0800 Subject: [PATCH] fix(proxy_server.py): fix _delete_deployment to handle base case where db_model list is empty don't delete all router models b/c of empty list Fixes https://github.com/BerriAI/litellm/issues/7196 --- litellm/proxy/proxy_server.py | 30 ++++++++++------ tests/local_testing/test_config.py | 56 ++++++++++++++++++++++++++++++ 2 files changed, 76 insertions(+), 10 deletions(-) diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index cb540c5f05e..0f0d73905b5 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -2088,7 +2088,10 @@ class ProxyConfig: """ global user_config_file_path, llm_router combined_id_list = [] - if llm_router is None: + + ## BASE CASES ## + # if llm_router is None or db_models is empty, return 0 + if llm_router is None or len(db_models) == 0: return 0 ## DB MODELS ## @@ -2418,6 +2421,19 @@ class ProxyConfig: return config + async def _get_models_from_db(self, prisma_client: PrismaClient) -> list: + try: + new_models = await prisma_client.db.litellm_proxymodeltable.find_many() + except Exception as e: + verbose_proxy_logger.exception( + "litellm.proxy_server.py::add_deployment() - Error getting new models from DB - {}".format( + str(e) + ) + ) + new_models = [] + + return new_models + async def add_deployment( self, prisma_client: PrismaClient, @@ -2435,15 +2451,9 @@ class ProxyConfig: raise ValueError( f"Master key is not initialized or formatted. master_key={master_key}" ) - try: - new_models = await prisma_client.db.litellm_proxymodeltable.find_many() - except Exception as e: - verbose_proxy_logger.exception( - "litellm.proxy_server.py::add_deployment() - Error getting new models from DB - {}".format( - str(e) - ) - ) - new_models = [] + + new_models = await self._get_models_from_db(prisma_client=prisma_client) + # update llm router await self._update_llm_router( new_models=new_models, proxy_logging_obj=proxy_logging_obj diff --git a/tests/local_testing/test_config.py b/tests/local_testing/test_config.py index a63816e8e22..213f5095eaa 100644 --- a/tests/local_testing/test_config.py +++ b/tests/local_testing/test_config.py @@ -175,6 +175,62 @@ async def test_add_existing_deployment(): assert init_len_list == len(llm_router.model_list) +@pytest.mark.asyncio +async def test_db_error_new_model_check(): + """ + - if error in db, don't delete existing models + + Relevant issue: https://github.com/BerriAI/litellm/blob/ddfe687b13e9f31db2fb2322887804e3d01dd467/litellm/proxy/proxy_server.py#L2461 + """ + import base64 + + litellm_params = LiteLLM_Params( + model="gpt-3.5-turbo", + api_key=os.getenv("AZURE_API_KEY"), + api_base=os.getenv("AZURE_API_BASE"), + api_version=os.getenv("AZURE_API_VERSION"), + ) + deployment = Deployment(model_name="gpt-3.5-turbo", litellm_params=litellm_params) + deployment_2 = Deployment( + model_name="gpt-3.5-turbo-2", litellm_params=litellm_params + ) + + llm_router = litellm.Router( + model_list=[ + deployment.to_json(exclude_none=True), + deployment_2.to_json(exclude_none=True), + ] + ) + + init_len_list = len(llm_router.model_list) + print(f"llm_router: {llm_router}") + master_key = "sk-1234" + setattr(litellm.proxy.proxy_server, "llm_router", llm_router) + setattr(litellm.proxy.proxy_server, "master_key", master_key) + pc = ProxyConfig() + + encrypted_litellm_params = litellm_params.dict(exclude_none=True) + + for k, v in encrypted_litellm_params.items(): + if isinstance(v, str): + encrypted_value = encrypt_value(v, master_key) + encrypted_litellm_params[k] = base64.b64encode(encrypted_value).decode( + "utf-8" + ) + db_model = DBModel( + model_id=deployment.model_info.id, + model_name="gpt-3.5-turbo", + litellm_params=encrypted_litellm_params, + model_info={"id": deployment.model_info.id}, + ) + + db_models = [] + deleted_deployments = await pc._delete_deployment(db_models=db_models) + assert deleted_deployments == 0 + + assert init_len_list == len(llm_router.model_list) + + litellm_params = LiteLLM_Params( model="azure/chatgpt-v-2", api_key=os.getenv("AZURE_API_KEY"),