mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-07 02:59:05 +00:00
Merge pull request #4519 from BerriAI/litellm_re_use_openai_azure_clients_whisper
[Fix+Test] /audio/transcriptions - use initialized OpenAI / Azure OpenAI clients
This commit is contained in:
commit
90a0db5618
2 changed files with 60 additions and 1 deletions
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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}")
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue