diff --git a/litellm/main.py b/litellm/main.py index 55ac01935eb..d6819b5ec01 100644 --- a/litellm/main.py +++ b/litellm/main.py @@ -48,6 +48,7 @@ from litellm import ( # type: ignore get_litellm_params, get_optional_params, ) +from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj from litellm.utils import ( CustomStreamWrapper, Usage, @@ -4262,7 +4263,7 @@ def transcription( api_base: Optional[str] = None, api_version: Optional[str] = None, max_retries: Optional[int] = None, - litellm_logging_obj=None, + litellm_logging_obj: Optional[LiteLLMLoggingObj] = None, custom_llm_provider=None, **kwargs, ): @@ -4277,6 +4278,18 @@ def transcription( proxy_server_request = kwargs.get("proxy_server_request", None) model_info = kwargs.get("model_info", None) metadata = kwargs.get("metadata", {}) + client: Optional[ + Union[ + openai.AsyncOpenAI, + openai.OpenAI, + openai.AzureOpenAI, + openai.AsyncAzureOpenAI, + ] + ] = kwargs.pop("client", None) + + if litellm_logging_obj: + litellm_logging_obj.model_call_details["client"] = str(client) + if max_retries is None: max_retries = openai.DEFAULT_MAX_RETRIES @@ -4316,6 +4329,7 @@ def transcription( optional_params=optional_params, model_response=model_response, atranscription=atranscription, + client=client, timeout=timeout, logging_obj=litellm_logging_obj, api_base=api_base, @@ -4349,6 +4363,7 @@ def transcription( optional_params=optional_params, model_response=model_response, atranscription=atranscription, + client=client, timeout=timeout, logging_obj=litellm_logging_obj, max_retries=max_retries, diff --git a/tests/test_whisper.py b/tests/test_whisper.py index 1debbbc1db3..09819f796c8 100644 --- a/tests/test_whisper.py +++ b/tests/test_whisper.py @@ -8,6 +8,9 @@ from openai import AsyncOpenAI import sys, os, dotenv from typing import Optional from dotenv import load_dotenv +from litellm.integrations.custom_logger import CustomLogger +import litellm +import logging # Get the current directory of the file being run pwd = os.path.dirname(os.path.realpath(__file__)) @@ -84,9 +87,32 @@ async def test_transcription_async_openai(): assert isinstance(transcript.text, str) +# This file includes the custom callbacks for LiteLLM Proxy +# Once defined, these can be passed in proxy_config.yaml +class MyCustomHandler(CustomLogger): + def __init__(self): + self.openai_client = None + + async def async_log_success_event(self, kwargs, response_obj, start_time, end_time): + try: + # init logging config + print("logging a transcript kwargs: ", kwargs) + print("openai client=", kwargs.get("client")) + self.openai_client = kwargs.get("client") + + except: + pass + + +proxy_handler_instance = MyCustomHandler() + + +# Set litellm.callbacks = [proxy_handler_instance] on the proxy +# need to set litellm.callbacks = [proxy_handler_instance] # on the proxy @pytest.mark.asyncio async def test_transcription_on_router(): litellm.set_verbose = True + litellm.callbacks = [proxy_handler_instance] print("\n Testing async transcription on router\n") try: model_list = [ @@ -108,11 +134,29 @@ async def test_transcription_on_router(): ] router = Router(model_list=model_list) + + router_level_clients = [] + for deployment in router.model_list: + _deployment_openai_client = router._get_client( + deployment=deployment, + kwargs={"model": "whisper-1"}, + client_type="async", + ) + + router_level_clients.append(str(_deployment_openai_client)) + response = await router.atranscription( model="whisper", file=audio_file, ) print(response) + + # PROD Test + # Ensure we ONLY use OpenAI/Azure client initialized on the router level + await asyncio.sleep(5) + print("OpenAI Client used= ", proxy_handler_instance.openai_client) + print("all router level clients= ", router_level_clients) + assert proxy_handler_instance.openai_client in router_level_clients except Exception as e: traceback.print_exc() pytest.fail(f"Error occurred: {e}")