fix(transcription): allow null request metadata

Signed-off-by: 1fanwang <1fannnw@gmail.com>
This commit is contained in:
1fanwang 2026-09-19 18:08:49 -07:00
parent 61a73c59b0
commit 21cf9054a1
2 changed files with 32 additions and 2 deletions

View file

@ -1207,7 +1207,7 @@ def function_setup(
# Lazy import audio_utils.utils only when needed for transcription calls
audio_utils: Final = _get_cached_audio_utils()
file_checksum: Final = audio_utils.get_audio_file_content_hash(file_obj=_file_obj)
if "metadata" in kwargs:
if kwargs.get("metadata") is not None:
kwargs["metadata"]["file_checksum"] = file_checksum
else:
kwargs["metadata"] = {"file_checksum": file_checksum}

View file

@ -2,6 +2,7 @@ import asyncio
import base64
import contextlib
import contextvars
import hashlib
import io
import json
import logging
@ -12,8 +13,9 @@ from collections.abc import Callable, Iterator, Mapping
from concurrent.futures import Future, ThreadPoolExecutor
from datetime import datetime, timedelta, timezone
from pathlib import PurePath
from typing import Final, cast
from typing import Final, Literal, cast
from unittest.mock import AsyncMock, MagicMock, patch
from uuid import uuid4
import httpx
import pytest
@ -4859,6 +4861,34 @@ async def test_wrapper_aresponses_reads_cache_once_and_replays_from_that_read(
assert standard_logging_object["cache_hit"] is True
@pytest.mark.parametrize("call_type", [CallTypes.transcription, CallTypes.atranscription])
@pytest.mark.parametrize("metadata_mode", ["omitted", "null", "empty", "tagged"])
def test_function_setup_transcription_optional_metadata(
call_type: CallTypes,
metadata_mode: Literal["omitted", "null", "empty", "tagged"],
) -> None:
audio: Final = b"audio contents"
supplied_metadata: Final = {"label": "transcription"} if metadata_mode == "tagged" else {}
expected_metadata: Final = {
**supplied_metadata,
"file_checksum": hashlib.sha256(audio).hexdigest(),
}
metadata_kwargs: Final = (
{} if metadata_mode == "omitted" else {"metadata": None if metadata_mode == "null" else supplied_metadata}
)
_, returned_kwargs = litellm.utils.function_setup(
original_function=call_type.value,
rules_obj=litellm.utils.Rules(),
start_time=datetime.now(timezone.utc),
model="whisper-1",
file=io.BytesIO(audio),
litellm_call_id=str(uuid4()),
**metadata_kwargs,
)
assert returned_kwargs["metadata"] == expected_metadata
def test_function_setup_failure_after_logging_construction_restores_context(monkeypatch):
"""If function_setup() constructs Logging() (which already mutated
trace_id_var/session_id_var in __init__) but then raises before returning,