mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-06 02:48:13 +00:00
Merge a82dc87519 into b781d157d7
This commit is contained in:
commit
f630aec973
3 changed files with 37 additions and 6 deletions
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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"}
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue