diff --git a/litellm/litellm_core_utils/audio_utils/utils.py b/litellm/litellm_core_utils/audio_utils/utils.py index dab3e48f91a..74cdb2cefa6 100644 --- a/litellm/litellm_core_utils/audio_utils/utils.py +++ b/litellm/litellm_core_utils/audio_utils/utils.py @@ -5,7 +5,7 @@ Utils used for litellm.transcription() and litellm.atranscription() import hashlib import os from dataclasses import dataclass -from typing import Final +from typing import BinaryIO, Final from litellm.types.files import ( AUDIO_FILE_TYPES, @@ -244,7 +244,7 @@ def get_audio_file_content_hash(file_obj: FileTypes) -> str: return hash_object.hexdigest() -def get_audio_file_for_health_check() -> FileTypes: +def get_audio_file_for_health_check() -> BinaryIO: """ Get an audio file for health check diff --git a/litellm/litellm_core_utils/health_check_helpers.py b/litellm/litellm_core_utils/health_check_helpers.py index a0f027cd58f..d77463e9435 100644 --- a/litellm/litellm_core_utils/health_check_helpers.py +++ b/litellm/litellm_core_utils/health_check_helpers.py @@ -193,6 +193,13 @@ class HealthCheckHelpers: from litellm.litellm_core_utils.health_check_utils import _filter_model_params from litellm.realtime_api.main import _realtime_health_check + async def audio_transcription_health_check(): + with get_audio_file_for_health_check() as audio_file: + return await litellm.atranscription( + **_filter_model_params(model_params=model_params), + file=audio_file, + ) + return { "chat": lambda: litellm.acompletion( **model_params, @@ -212,10 +219,7 @@ class HealthCheckHelpers: }, input=prompt or "test", ), - "audio_transcription": lambda: litellm.atranscription( - **_filter_model_params(model_params=model_params), - file=get_audio_file_for_health_check(), - ), + "audio_transcription": audio_transcription_health_check, "image_generation": lambda: litellm.aimage_generation( **_filter_model_params(model_params=model_params), prompt=prompt, diff --git a/tests/unit/litellm_core_utils/test_health_check_helpers.py b/tests/unit/litellm_core_utils/test_health_check_helpers.py index 47c4576f91f..6e7e54c60e6 100644 --- a/tests/unit/litellm_core_utils/test_health_check_helpers.py +++ b/tests/unit/litellm_core_utils/test_health_check_helpers.py @@ -37,6 +37,33 @@ def _distinct_rgb_colors(png: bytes) -> set[bytes]: return {bytes(row[i : i + 3]) for row in rows for i in range(1, row_size, 3)} +@pytest.mark.asyncio +@pytest.mark.parametrize("raises", [False, True]) +async def test_audio_transcription_health_check_closes_file(raises: bool): + async def transcription(**kwargs): + file = kwargs["file"] + assert not file.closed + assert file.read(4) == b"RIFF" + if raises: + raise RuntimeError("transcription failed") + return {"text": "healthy"} + + handlers = HealthCheckHelpers.get_mode_handlers( + model="whisper-1", + custom_llm_provider="openai", + model_params={"model": "openai/whisper-1", "api_key": "sk-test"}, + ) + with patch("litellm.atranscription", side_effect=transcription) as mock_transcription: + if raises: + with pytest.raises(RuntimeError, match="transcription failed"): + await handlers["audio_transcription"]() + else: + assert await handlers["audio_transcription"]() == {"text": "healthy"} + + mock_transcription.assert_called_once() + assert mock_transcription.call_args.kwargs["file"].closed + + @pytest.mark.asyncio async def test_image_edit_health_check_handler_uses_descriptive_prompt_and_multicolor_png(): model_params = {"model": "openai/gpt-image-1", "api_key": "sk-test"}