From c41b1418d4d08428469499d307f565e82948334e Mon Sep 17 00:00:00 2001 From: Krrish Dholakia Date: Sat, 30 Dec 2023 16:51:39 +0530 Subject: [PATCH] test(test_router_init.py): fix test router init --- litellm/tests/test_router_init.py | 24 ++++++++++++++---------- 1 file changed, 14 insertions(+), 10 deletions(-) diff --git a/litellm/tests/test_router_init.py b/litellm/tests/test_router_init.py index e36b8319ae8..581506d1391 100644 --- a/litellm/tests/test_router_init.py +++ b/litellm/tests/test_router_init.py @@ -41,14 +41,17 @@ def test_init_clients(): ] router = Router(model_list=model_list) for elem in router.model_list: - assert elem["client"] is not None - assert elem["async_client"] is not None - assert elem["stream_client"] is not None - assert elem["stream_async_client"] is not None + model_id = elem["model_info"]["id"] + assert router.cache.get_cache(f"{model_id}_client") is not None + assert router.cache.get_cache(f"{model_id}_async_client") is not None + assert router.cache.get_cache(f"{model_id}_stream_client") is not None + assert router.cache.get_cache(f"{model_id}_stream_async_client") is not None # check if timeout for stream/non stream clients is set correctly - async_client = elem["async_client"] - stream_async_client = elem["stream_async_client"] + async_client = router.cache.get_cache(f"{model_id}_async_client") + stream_async_client = router.cache.get_cache( + f"{model_id}_stream_async_client" + ) assert async_client.timeout == 0.01 assert stream_async_client.timeout == 0.000_001 @@ -79,10 +82,11 @@ def test_init_clients_basic(): ] router = Router(model_list=model_list) for elem in router.model_list: - assert elem["client"] is not None - assert elem["async_client"] is not None - assert elem["stream_client"] is not None - assert elem["stream_async_client"] is not None + model_id = elem["model_info"]["id"] + assert router.cache.get_cache(f"{model_id}_client") is not None + assert router.cache.get_cache(f"{model_id}_async_client") is not None + assert router.cache.get_cache(f"{model_id}_stream_client") is not None + assert router.cache.get_cache(f"{model_id}_stream_async_client") is not None print("PASSED !") # see if we can init clients without timeout or max retries set