mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-10 03:28:53 +00:00
fix(transcription): synthesize srt/vtt output for adapters without native subtitle formats
Extract the Soniox SRT/VTT cue grouping and rendering into a shared litellm_core_utils/audio_utils/subtitle_utils module, have Gemini transcription request word timestamps whenever response_format is srt or vtt, and let the http handler rewrite the response text into the synthesized subtitle document (dropping the internally requested words array) for any provider config that opts in via supports_subtitle_synthesis
This commit is contained in:
parent
493bca667b
commit
5b80fb0fc0
9 changed files with 478 additions and 146 deletions
189
litellm/litellm_core_utils/audio_utils/subtitle_utils.py
Normal file
189
litellm/litellm_core_utils/audio_utils/subtitle_utils.py
Normal file
|
|
@ -0,0 +1,189 @@
|
|||
"""Provider-agnostic SRT/WebVTT subtitle synthesis from timestamped transcription tokens."""
|
||||
|
||||
from collections.abc import Sequence
|
||||
from dataclasses import dataclass
|
||||
from functools import reduce
|
||||
from typing import Final
|
||||
|
||||
from pydantic import BaseModel, ConfigDict, TypeAdapter, ValidationError
|
||||
|
||||
CUE_MAX_TOKENS: Final = 15
|
||||
CUE_MAX_DURATION_MS: Final = 5000
|
||||
|
||||
SRT_RESPONSE_FORMAT: Final = "srt"
|
||||
VTT_RESPONSE_FORMAT: Final = "vtt"
|
||||
SUBTITLE_RESPONSE_FORMATS: Final = frozenset((SRT_RESPONSE_FORMAT, VTT_RESPONSE_FORMAT))
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class SubtitleToken:
|
||||
text: str
|
||||
start_ms: int | None = None
|
||||
end_ms: int | None = None
|
||||
speaker: str | int | None = None
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class SubtitleCue:
|
||||
start_ms: int
|
||||
end_ms: int
|
||||
text: str
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class _CueAccumulator:
|
||||
cues: tuple[SubtitleCue, ...] = ()
|
||||
texts: tuple[str, ...] = ()
|
||||
start_ms: int | None = None
|
||||
end_ms: int | None = None
|
||||
speaker: str | int | None = None
|
||||
|
||||
|
||||
def _completed_cue(accumulator: _CueAccumulator) -> tuple[SubtitleCue, ...]:
|
||||
if not accumulator.texts or accumulator.start_ms is None:
|
||||
return ()
|
||||
text: Final = "".join(accumulator.texts).strip()
|
||||
if not text:
|
||||
return ()
|
||||
end_ms: Final = accumulator.end_ms if accumulator.end_ms is not None else accumulator.start_ms
|
||||
return (SubtitleCue(start_ms=accumulator.start_ms, end_ms=end_ms, text=text),)
|
||||
|
||||
|
||||
def _cue_break_reached(accumulator: _CueAccumulator, token: SubtitleToken) -> bool:
|
||||
if len(accumulator.texts) >= CUE_MAX_TOKENS:
|
||||
return True
|
||||
return (
|
||||
accumulator.start_ms is not None
|
||||
and token.start_ms is not None
|
||||
and token.start_ms - accumulator.start_ms >= CUE_MAX_DURATION_MS
|
||||
)
|
||||
|
||||
|
||||
def _absorb_token(accumulator: _CueAccumulator, token: SubtitleToken) -> _CueAccumulator:
|
||||
if token.start_ms is None and accumulator.start_ms is None:
|
||||
return accumulator
|
||||
if token.speaker is not None and token.speaker != accumulator.speaker:
|
||||
return _CueAccumulator(
|
||||
cues=accumulator.cues + _completed_cue(accumulator),
|
||||
texts=(token.text,),
|
||||
start_ms=token.start_ms,
|
||||
end_ms=token.end_ms,
|
||||
speaker=token.speaker,
|
||||
)
|
||||
if _cue_break_reached(accumulator, token):
|
||||
return _CueAccumulator(
|
||||
cues=accumulator.cues + _completed_cue(accumulator),
|
||||
texts=(token.text,),
|
||||
start_ms=token.start_ms,
|
||||
end_ms=token.end_ms,
|
||||
speaker=accumulator.speaker,
|
||||
)
|
||||
return _CueAccumulator(
|
||||
cues=accumulator.cues,
|
||||
texts=(*accumulator.texts, token.text),
|
||||
start_ms=accumulator.start_ms if accumulator.start_ms is not None else token.start_ms,
|
||||
end_ms=token.end_ms if token.end_ms is not None else accumulator.end_ms,
|
||||
speaker=accumulator.speaker,
|
||||
)
|
||||
|
||||
|
||||
def group_subtitle_tokens_into_cues(tokens: Sequence[SubtitleToken]) -> tuple[SubtitleCue, ...]:
|
||||
accumulator: Final = reduce(_absorb_token, tokens, _CueAccumulator())
|
||||
return accumulator.cues + _completed_cue(accumulator)
|
||||
|
||||
|
||||
def _format_timestamp(total_ms: int, millis_separator: str) -> str:
|
||||
clamped: Final = max(total_ms, 0)
|
||||
hours, hour_remainder = divmod(clamped, 3_600_000)
|
||||
minutes, minute_remainder = divmod(hour_remainder, 60_000)
|
||||
seconds, millis = divmod(minute_remainder, 1_000)
|
||||
return f"{hours:02d}:{minutes:02d}:{seconds:02d}{millis_separator}{millis:03d}"
|
||||
|
||||
|
||||
def _render_srt(cues: Sequence[SubtitleCue]) -> str:
|
||||
lines: Final = tuple(
|
||||
line
|
||||
for index, cue in enumerate(cues, start=1)
|
||||
for line in (
|
||||
str(index),
|
||||
f"{_format_timestamp(cue.start_ms, ',')} --> {_format_timestamp(cue.end_ms, ',')}",
|
||||
cue.text,
|
||||
"",
|
||||
)
|
||||
)
|
||||
return "\n".join(lines)
|
||||
|
||||
|
||||
def _render_vtt(cues: Sequence[SubtitleCue]) -> str:
|
||||
cue_lines: Final = tuple(
|
||||
line
|
||||
for cue in cues
|
||||
for line in (
|
||||
f"{_format_timestamp(cue.start_ms, '.')} --> {_format_timestamp(cue.end_ms, '.')}",
|
||||
cue.text,
|
||||
"",
|
||||
)
|
||||
)
|
||||
return "\n".join(("WEBVTT", "", *cue_lines))
|
||||
|
||||
|
||||
def render_subtitle_tokens_as_srt(tokens: Sequence[SubtitleToken]) -> str:
|
||||
"""Render tokens as an SRT document; empty string when no token has timestamp data."""
|
||||
cues: Final = group_subtitle_tokens_into_cues(tokens)
|
||||
if not cues:
|
||||
return ""
|
||||
return _render_srt(cues)
|
||||
|
||||
|
||||
def render_subtitle_tokens_as_vtt(tokens: Sequence[SubtitleToken]) -> str:
|
||||
"""Render tokens as a WebVTT document; the WEBVTT header is emitted even without cues."""
|
||||
return _render_vtt(group_subtitle_tokens_into_cues(tokens))
|
||||
|
||||
|
||||
class TranscriptionWordTiming(BaseModel):
|
||||
model_config = ConfigDict(frozen=True, extra="ignore")
|
||||
|
||||
word: str = ""
|
||||
start: float | None = None
|
||||
end: float | None = None
|
||||
speaker: str | None = None
|
||||
|
||||
|
||||
_WORD_TIMINGS_ADAPTER: Final = TypeAdapter(tuple[TranscriptionWordTiming, ...])
|
||||
|
||||
|
||||
def _seconds_to_ms(seconds: float | None) -> int | None:
|
||||
if seconds is None:
|
||||
return None
|
||||
return round(seconds * 1000)
|
||||
|
||||
|
||||
def _word_to_subtitle_token(word: TranscriptionWordTiming) -> SubtitleToken:
|
||||
return SubtitleToken(
|
||||
text=f"{word.word} ",
|
||||
start_ms=_seconds_to_ms(word.start),
|
||||
end_ms=_seconds_to_ms(word.end),
|
||||
speaker=word.speaker,
|
||||
)
|
||||
|
||||
|
||||
def _parse_word_timings(words: object) -> tuple[TranscriptionWordTiming, ...]:
|
||||
try:
|
||||
return _WORD_TIMINGS_ADAPTER.validate_python(words)
|
||||
except ValidationError:
|
||||
return ()
|
||||
|
||||
|
||||
def synthesize_subtitle_document(words: object, response_format: str) -> str | None:
|
||||
"""
|
||||
Build an SRT/VTT document from OpenAI verbose_json-style word dicts
|
||||
(word/start/end in float seconds, optional speaker). Returns None when the
|
||||
format is not a subtitle format or the words carry no usable timestamps.
|
||||
"""
|
||||
if response_format not in SUBTITLE_RESPONSE_FORMATS:
|
||||
return None
|
||||
tokens: Final = tuple(_word_to_subtitle_token(word) for word in _parse_word_timings(words))
|
||||
cues: Final = group_subtitle_tokens_into_cues(tokens)
|
||||
if not cues:
|
||||
return None
|
||||
return _render_srt(cues) if response_format == SRT_RESPONSE_FORMAT else _render_vtt(cues)
|
||||
|
|
@ -40,6 +40,16 @@ class BaseAudioTranscriptionConfig(BaseConfig, ABC):
|
|||
def get_supported_openai_params(self, model: str) -> list[OpenAIAudioTranscriptionOptionalParams]:
|
||||
pass
|
||||
|
||||
@property
|
||||
def supports_subtitle_synthesis(self) -> bool:
|
||||
"""
|
||||
Opt-in for providers without a native srt/vtt response body: when True
|
||||
and the user asked for response_format srt/vtt, the http handler
|
||||
synthesizes the subtitle document from the word timestamps the
|
||||
provider's TranscriptionResponse carries in `words`.
|
||||
"""
|
||||
return False
|
||||
|
||||
def get_complete_url(
|
||||
self,
|
||||
api_base: str | None,
|
||||
|
|
|
|||
|
|
@ -25,6 +25,7 @@ from litellm.litellm_core_utils.agentic_loop_settings import (
|
|||
validated_max_agentic_loops,
|
||||
)
|
||||
from litellm.litellm_core_utils.asyncify import run_async_function
|
||||
from litellm.litellm_core_utils.audio_utils.subtitle_utils import synthesize_subtitle_document
|
||||
from litellm.litellm_core_utils.llm_request_utils import serialize_multipart_form_fields
|
||||
from litellm.litellm_core_utils.realtime_errors import realtime_error_event, websocket_close_reason
|
||||
from litellm.litellm_core_utils.realtime_streaming import RealTimeStreaming
|
||||
|
|
@ -1296,9 +1297,23 @@ class BaseLLMHTTPHandler:
|
|||
api_key: str | None,
|
||||
) -> TranscriptionResponse:
|
||||
"""Shared logic for transforming audio transcription responses."""
|
||||
return provider_config.transform_audio_transcription_response(
|
||||
transformed: Final = provider_config.transform_audio_transcription_response(
|
||||
raw_response=response,
|
||||
)
|
||||
if not provider_config.supports_subtitle_synthesis:
|
||||
return transformed
|
||||
requested_format: Final = optional_params.get("response_format")
|
||||
if not isinstance(requested_format, str):
|
||||
return transformed
|
||||
document: Final = synthesize_subtitle_document(
|
||||
words=transformed.get("words"),
|
||||
response_format=requested_format,
|
||||
)
|
||||
if document is None:
|
||||
return transformed
|
||||
transformed.text = document
|
||||
delattr(transformed, "words")
|
||||
return transformed
|
||||
|
||||
def audio_transcriptions(
|
||||
self,
|
||||
|
|
|
|||
|
|
@ -4,6 +4,7 @@ from typing import Final
|
|||
|
||||
from httpx import Headers, Response
|
||||
|
||||
from litellm.litellm_core_utils.audio_utils.subtitle_utils import SUBTITLE_RESPONSE_FORMATS
|
||||
from litellm.litellm_core_utils.audio_utils.utils import (
|
||||
normalize_transcription_language_to_bcp47,
|
||||
process_audio_file,
|
||||
|
|
@ -48,6 +49,10 @@ class GeminiAudioTranscriptionConfig(BaseAudioTranscriptionConfig):
|
|||
) -> list[OpenAIAudioTranscriptionOptionalParams]: # mutable-ok: BaseAudioTranscriptionConfig signature
|
||||
return ["language", "response_format", "timestamp_granularities"] # mutable-ok: base contract returns a list
|
||||
|
||||
@property
|
||||
def supports_subtitle_synthesis(self) -> bool:
|
||||
return True
|
||||
|
||||
def map_openai_params(
|
||||
self,
|
||||
non_default_params: Mapping[str, object],
|
||||
|
|
@ -215,16 +220,17 @@ def _language_config(language: object) -> GeminiTranscriptionConfig:
|
|||
return language_config
|
||||
|
||||
|
||||
def _timestamp_config(timestamp_granularities: object) -> GeminiTranscriptionConfig:
|
||||
if isinstance(timestamp_granularities, list) and "word" in timestamp_granularities:
|
||||
return _WORD_TIMESTAMP_CONFIG
|
||||
return _EMPTY_TRANSCRIPTION_CONFIG
|
||||
def _timestamp_config(timestamp_granularities: object, response_format: object) -> GeminiTranscriptionConfig:
|
||||
wants_word_timestamps: Final = (
|
||||
isinstance(timestamp_granularities, list) and "word" in timestamp_granularities
|
||||
) or response_format in SUBTITLE_RESPONSE_FORMATS
|
||||
return _WORD_TIMESTAMP_CONFIG if wants_word_timestamps else _EMPTY_TRANSCRIPTION_CONFIG
|
||||
|
||||
|
||||
def _build_transcription_config(optional_params: Mapping[str, object]) -> GeminiTranscriptionConfig:
|
||||
transcription_config: Final[GeminiTranscriptionConfig] = {
|
||||
**_language_config(optional_params.get("language")),
|
||||
**_timestamp_config(optional_params.get("timestamp_granularities")),
|
||||
**_timestamp_config(optional_params.get("timestamp_granularities"), optional_params.get("response_format")),
|
||||
}
|
||||
return transcription_config
|
||||
|
||||
|
|
|
|||
|
|
@ -4,6 +4,11 @@ Shared utilities for the Soniox provider (https://soniox.com).
|
|||
|
||||
from typing import Any, Final
|
||||
|
||||
from litellm.litellm_core_utils.audio_utils.subtitle_utils import (
|
||||
SubtitleToken,
|
||||
render_subtitle_tokens_as_srt,
|
||||
render_subtitle_tokens_as_vtt,
|
||||
)
|
||||
from litellm.llms.base_llm.chat.transformation import BaseLLMException
|
||||
|
||||
# Soniox API base URL.
|
||||
|
|
@ -109,121 +114,13 @@ def render_soniox_tokens(tokens: list[dict[str, Any]]) -> str:
|
|||
return "".join(text_parts)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# SRT / VTT subtitle rendering
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
# Maximum number of tokens to group into a single subtitle cue.
|
||||
_CUE_MAX_TOKENS: Final[int] = 15
|
||||
|
||||
# Maximum duration (in ms) for a single cue before forcing a break.
|
||||
_CUE_MAX_DURATION_MS: Final[int] = 5000
|
||||
|
||||
|
||||
def _format_timestamp_srt(ms: int) -> str:
|
||||
"""Format milliseconds as SRT timestamp: HH:MM:SS,mmm"""
|
||||
ms = max(ms, 0)
|
||||
hours: Final = ms // 3_600_000
|
||||
ms %= 3_600_000
|
||||
minutes: Final = ms // 60_000
|
||||
ms %= 60_000
|
||||
seconds: Final = ms // 1_000
|
||||
millis: Final = ms % 1_000
|
||||
return f"{hours:02d}:{minutes:02d}:{seconds:02d},{millis:03d}"
|
||||
|
||||
|
||||
def _format_timestamp_vtt(ms: int) -> str:
|
||||
"""Format milliseconds as VTT timestamp: HH:MM:SS.mmm"""
|
||||
ms = max(ms, 0)
|
||||
hours: Final = ms // 3_600_000
|
||||
ms %= 3_600_000
|
||||
minutes: Final = ms // 60_000
|
||||
ms %= 60_000
|
||||
seconds: Final = ms // 1_000
|
||||
millis: Final = ms % 1_000
|
||||
return f"{hours:02d}:{minutes:02d}:{seconds:02d}.{millis:03d}"
|
||||
|
||||
|
||||
def _group_tokens_into_cues(
|
||||
tokens: list[dict[str, Any]],
|
||||
) -> list[dict[str, Any]]:
|
||||
"""
|
||||
Group Soniox tokens into subtitle cues.
|
||||
|
||||
Each cue has:
|
||||
- start_ms: int
|
||||
- end_ms: int
|
||||
- text: str
|
||||
|
||||
Grouping heuristics:
|
||||
- A new cue starts when token count exceeds _CUE_MAX_TOKENS.
|
||||
- A new cue starts when duration exceeds _CUE_MAX_DURATION_MS.
|
||||
- A new cue starts when the speaker changes (if diarization is on).
|
||||
- Tokens without timestamps are appended to the current cue.
|
||||
"""
|
||||
cues: Final[list[dict[str, Any]]] = []
|
||||
current_tokens: list[str] = []
|
||||
current_start: int | None = None
|
||||
current_end: int | None = None
|
||||
current_speaker: Any | None = None
|
||||
|
||||
def _flush() -> None:
|
||||
if current_tokens and current_start is not None:
|
||||
text: Final = "".join(current_tokens).strip()
|
||||
if text:
|
||||
cues.append(
|
||||
{
|
||||
"start_ms": current_start,
|
||||
"end_ms": (current_end if current_end is not None else current_start),
|
||||
"text": text,
|
||||
}
|
||||
)
|
||||
|
||||
for token in tokens:
|
||||
start_ms = token.get("start_ms")
|
||||
end_ms = token.get("end_ms")
|
||||
text = token.get("text", "")
|
||||
speaker = token.get("speaker")
|
||||
|
||||
# Skip tokens with no timestamp data entirely if we have no cue started
|
||||
if start_ms is None and current_start is None:
|
||||
continue
|
||||
|
||||
# Speaker change forces a new cue
|
||||
if speaker is not None and speaker != current_speaker:
|
||||
_flush()
|
||||
current_tokens = []
|
||||
current_start = start_ms
|
||||
current_end = end_ms
|
||||
current_speaker = speaker
|
||||
current_tokens.append(text)
|
||||
continue
|
||||
|
||||
# Duration or token count exceeded -> flush
|
||||
should_break = False
|
||||
if (
|
||||
len(current_tokens) >= _CUE_MAX_TOKENS
|
||||
or current_start is not None
|
||||
and start_ms is not None
|
||||
and (start_ms - current_start) >= _CUE_MAX_DURATION_MS
|
||||
):
|
||||
should_break = True
|
||||
|
||||
if should_break:
|
||||
_flush()
|
||||
current_tokens = []
|
||||
current_start = start_ms
|
||||
current_end = end_ms
|
||||
current_tokens.append(text)
|
||||
else:
|
||||
if current_start is None:
|
||||
current_start = start_ms
|
||||
if end_ms is not None:
|
||||
current_end = end_ms
|
||||
current_tokens.append(text)
|
||||
|
||||
_flush()
|
||||
return cues
|
||||
def _soniox_token_to_subtitle_token(token: dict[str, Any]) -> SubtitleToken:
|
||||
return SubtitleToken(
|
||||
text=token.get("text", ""),
|
||||
start_ms=token.get("start_ms"),
|
||||
end_ms=token.get("end_ms"),
|
||||
speaker=token.get("speaker"),
|
||||
)
|
||||
|
||||
|
||||
def render_soniox_tokens_as_srt(tokens: list[dict[str, Any]]) -> str:
|
||||
|
|
@ -232,20 +129,7 @@ def render_soniox_tokens_as_srt(tokens: list[dict[str, Any]]) -> str:
|
|||
|
||||
Returns an empty string if no tokens have timestamp data.
|
||||
"""
|
||||
cues: Final = _group_tokens_into_cues(tokens)
|
||||
if not cues:
|
||||
return ""
|
||||
|
||||
lines: Final[list[str]] = []
|
||||
for idx, cue in enumerate(cues, start=1):
|
||||
start = _format_timestamp_srt(cue["start_ms"])
|
||||
end = _format_timestamp_srt(cue["end_ms"])
|
||||
lines.append(str(idx))
|
||||
lines.append(f"{start} --> {end}")
|
||||
lines.append(cue["text"])
|
||||
lines.append("") # blank line between cues
|
||||
|
||||
return "\n".join(lines)
|
||||
return render_subtitle_tokens_as_srt(tuple(_soniox_token_to_subtitle_token(token) for token in tokens))
|
||||
|
||||
|
||||
def render_soniox_tokens_as_vtt(tokens: list[dict[str, Any]]) -> str:
|
||||
|
|
@ -254,14 +138,4 @@ def render_soniox_tokens_as_vtt(tokens: list[dict[str, Any]]) -> str:
|
|||
|
||||
Returns the VTT header even if no cues are present.
|
||||
"""
|
||||
cues: Final = _group_tokens_into_cues(tokens)
|
||||
|
||||
lines: Final[list[str]] = ["WEBVTT", ""]
|
||||
for cue in cues:
|
||||
start = _format_timestamp_vtt(cue["start_ms"])
|
||||
end = _format_timestamp_vtt(cue["end_ms"])
|
||||
lines.append(f"{start} --> {end}")
|
||||
lines.append(cue["text"])
|
||||
lines.append("") # blank line between cues
|
||||
|
||||
return "\n".join(lines)
|
||||
return render_subtitle_tokens_as_vtt(tuple(_soniox_token_to_subtitle_token(token) for token in tokens))
|
||||
|
|
|
|||
|
|
@ -0,0 +1,134 @@
|
|||
from litellm.litellm_core_utils.audio_utils.subtitle_utils import (
|
||||
SubtitleToken,
|
||||
render_subtitle_tokens_as_srt,
|
||||
render_subtitle_tokens_as_vtt,
|
||||
synthesize_subtitle_document,
|
||||
)
|
||||
|
||||
|
||||
class TestRenderSubtitleTokensAsSrt:
|
||||
def test_single_cue_full_document(self):
|
||||
tokens = (
|
||||
SubtitleToken(text="Hello ", start_ms=0, end_ms=500),
|
||||
SubtitleToken(text="world.", start_ms=500, end_ms=1000),
|
||||
)
|
||||
assert render_subtitle_tokens_as_srt(tokens) == "1\n00:00:00,000 --> 00:00:01,000\nHello world.\n"
|
||||
|
||||
def test_speaker_change_starts_a_new_cue(self):
|
||||
tokens = (
|
||||
SubtitleToken(text="Hi.", start_ms=0, end_ms=1000, speaker="spk:0"),
|
||||
SubtitleToken(text="Hey.", start_ms=1500, end_ms=2500, speaker="spk:1"),
|
||||
)
|
||||
assert render_subtitle_tokens_as_srt(tokens) == (
|
||||
"1\n00:00:00,000 --> 00:00:01,000\nHi.\n\n2\n00:00:01,500 --> 00:00:02,500\nHey.\n"
|
||||
)
|
||||
|
||||
def test_token_cap_starts_a_new_cue_after_15_tokens(self):
|
||||
tokens = tuple(
|
||||
SubtitleToken(text=f"{index} ", start_ms=index * 100, end_ms=index * 100 + 100) for index in range(16)
|
||||
)
|
||||
assert render_subtitle_tokens_as_srt(tokens) == (
|
||||
"1\n00:00:00,000 --> 00:00:01,500\n0 1 2 3 4 5 6 7 8 9 10 11 12 13 14\n"
|
||||
"\n2\n00:00:01,500 --> 00:00:01,600\n15\n"
|
||||
)
|
||||
|
||||
def test_duration_cap_starts_a_new_cue_at_5000ms(self):
|
||||
tokens = (
|
||||
SubtitleToken(text="Alpha ", start_ms=0, end_ms=400),
|
||||
SubtitleToken(text="beta ", start_ms=2000, end_ms=2400),
|
||||
SubtitleToken(text="gamma.", start_ms=5000, end_ms=5400),
|
||||
)
|
||||
assert render_subtitle_tokens_as_srt(tokens) == (
|
||||
"1\n00:00:00,000 --> 00:00:02,400\nAlpha beta\n\n2\n00:00:05,000 --> 00:00:05,400\ngamma.\n"
|
||||
)
|
||||
|
||||
def test_timestampless_token_joins_the_current_cue(self):
|
||||
tokens = (
|
||||
SubtitleToken(text="Hello ", start_ms=0, end_ms=500),
|
||||
SubtitleToken(text="there "),
|
||||
SubtitleToken(text="world.", start_ms=900, end_ms=1300),
|
||||
)
|
||||
assert render_subtitle_tokens_as_srt(tokens) == "1\n00:00:00,000 --> 00:00:01,300\nHello there world.\n"
|
||||
|
||||
def test_only_timestampless_tokens_renders_empty(self):
|
||||
assert render_subtitle_tokens_as_srt((SubtitleToken(text="no timestamps"),)) == ""
|
||||
|
||||
def test_empty_tokens_render_empty(self):
|
||||
assert render_subtitle_tokens_as_srt(()) == ""
|
||||
|
||||
def test_timestamps_past_one_hour(self):
|
||||
tokens = (SubtitleToken(text="Late.", start_ms=3_661_001, end_ms=3_662_002),)
|
||||
assert render_subtitle_tokens_as_srt(tokens) == "1\n01:01:01,001 --> 01:01:02,002\nLate.\n"
|
||||
|
||||
def test_negative_timestamps_clamp_to_zero(self):
|
||||
tokens = (SubtitleToken(text="Early.", start_ms=-100, end_ms=-50),)
|
||||
assert render_subtitle_tokens_as_srt(tokens) == "1\n00:00:00,000 --> 00:00:00,000\nEarly.\n"
|
||||
|
||||
def test_missing_end_falls_back_to_cue_start(self):
|
||||
tokens = (SubtitleToken(text="Open.", start_ms=1200),)
|
||||
assert render_subtitle_tokens_as_srt(tokens) == "1\n00:00:01,200 --> 00:00:01,200\nOpen.\n"
|
||||
|
||||
|
||||
class TestRenderSubtitleTokensAsVtt:
|
||||
def test_single_cue_full_document(self):
|
||||
tokens = (
|
||||
SubtitleToken(text="Hello ", start_ms=0, end_ms=500),
|
||||
SubtitleToken(text="world.", start_ms=500, end_ms=1000),
|
||||
)
|
||||
assert render_subtitle_tokens_as_vtt(tokens) == "WEBVTT\n\n00:00:00.000 --> 00:00:01.000\nHello world.\n"
|
||||
|
||||
def test_empty_tokens_render_header_only(self):
|
||||
assert render_subtitle_tokens_as_vtt(()) == "WEBVTT\n"
|
||||
|
||||
def test_timestamps_past_one_hour_use_dot_separator(self):
|
||||
tokens = (SubtitleToken(text="Late.", start_ms=3_661_001, end_ms=3_662_002),)
|
||||
assert render_subtitle_tokens_as_vtt(tokens) == "WEBVTT\n\n01:01:01.001 --> 01:01:02.002\nLate.\n"
|
||||
|
||||
def test_speaker_change_starts_a_new_cue(self):
|
||||
tokens = (
|
||||
SubtitleToken(text="Hi.", start_ms=0, end_ms=1000, speaker=1),
|
||||
SubtitleToken(text="Hey.", start_ms=1500, end_ms=2500, speaker=2),
|
||||
)
|
||||
assert render_subtitle_tokens_as_vtt(tokens) == (
|
||||
"WEBVTT\n\n00:00:00.000 --> 00:00:01.000\nHi.\n\n00:00:01.500 --> 00:00:02.500\nHey.\n"
|
||||
)
|
||||
|
||||
|
||||
class TestSynthesizeSubtitleDocument:
|
||||
WORDS = [
|
||||
{"word": "Four", "start": 0.4, "end": 0.7, "speaker": "spk:0"},
|
||||
{"word": "score", "start": 0.7, "end": 1.1, "speaker": "spk:0"},
|
||||
]
|
||||
|
||||
def test_srt_from_words_converts_seconds_to_milliseconds(self):
|
||||
assert synthesize_subtitle_document(self.WORDS, "srt") == "1\n00:00:00,400 --> 00:00:01,100\nFour score\n"
|
||||
|
||||
def test_vtt_from_words_converts_seconds_to_milliseconds(self):
|
||||
assert synthesize_subtitle_document(self.WORDS, "vtt") == (
|
||||
"WEBVTT\n\n00:00:00.400 --> 00:00:01.100\nFour score\n"
|
||||
)
|
||||
|
||||
def test_speaker_change_splits_cues(self):
|
||||
words = [
|
||||
{"word": "Hi", "start": 0.0, "end": 0.5, "speaker": "spk:0"},
|
||||
{"word": "Hey", "start": 0.6, "end": 1.0, "speaker": "spk:1"},
|
||||
]
|
||||
assert synthesize_subtitle_document(words, "srt") == (
|
||||
"1\n00:00:00,000 --> 00:00:00,500\nHi\n\n2\n00:00:00,600 --> 00:00:01,000\nHey\n"
|
||||
)
|
||||
|
||||
def test_non_subtitle_format_returns_none(self):
|
||||
assert synthesize_subtitle_document(self.WORDS, "verbose_json") is None
|
||||
assert synthesize_subtitle_document(self.WORDS, "json") is None
|
||||
|
||||
def test_missing_words_returns_none(self):
|
||||
assert synthesize_subtitle_document(None, "srt") is None
|
||||
assert synthesize_subtitle_document([], "srt") is None
|
||||
|
||||
def test_words_without_timestamps_return_none(self):
|
||||
assert synthesize_subtitle_document([{"word": "Hello"}], "srt") is None
|
||||
assert synthesize_subtitle_document([{"word": "Hello"}], "vtt") is None
|
||||
|
||||
def test_malformed_words_return_none(self):
|
||||
assert synthesize_subtitle_document("not words", "srt") is None
|
||||
assert synthesize_subtitle_document([{"word": "ok", "start": "not-a-number"}], "srt") is None
|
||||
|
|
@ -1901,6 +1901,35 @@ async def test_async_audio_transcriptions_sends_dict_data_as_json_body():
|
|||
assert response.text == "transcribed"
|
||||
|
||||
|
||||
class _WordTimestampAudioTranscriptionConfig(_JSONBodyAudioTranscriptionConfig):
|
||||
def transform_audio_transcription_response(self, raw_response):
|
||||
payload = raw_response.json()
|
||||
response = TranscriptionResponse(text=payload["text"])
|
||||
response["words"] = payload["words"]
|
||||
return response
|
||||
|
||||
|
||||
def test_transform_audio_transcription_response_without_subtitle_opt_in_keeps_text_and_words():
|
||||
words = [
|
||||
{"word": "hello", "start": 0.0, "end": 0.5},
|
||||
{"word": "world", "start": 0.5, "end": 1.0},
|
||||
]
|
||||
raw_response = httpx.Response(200, json={"text": "hello world", "words": words})
|
||||
|
||||
response = BaseLLMHTTPHandler()._transform_audio_transcription_response(
|
||||
provider_config=_WordTimestampAudioTranscriptionConfig(),
|
||||
model="test-model",
|
||||
response=raw_response,
|
||||
model_response=TranscriptionResponse(),
|
||||
logging_obj=Mock(),
|
||||
optional_params={"response_format": "srt"},
|
||||
api_key=None,
|
||||
)
|
||||
|
||||
assert response.text == "hello world"
|
||||
assert response["words"] == words
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_async_retrieve_file_content_raises_on_http_error():
|
||||
"""
|
||||
|
|
|
|||
|
|
@ -169,6 +169,33 @@ class TestTransformRequest:
|
|||
}
|
||||
}
|
||||
|
||||
@pytest.mark.parametrize("response_format", ["srt", "vtt"])
|
||||
def test_subtitle_response_format_requests_word_timestamps(self, config, response_format):
|
||||
request_data = config.transform_audio_transcription_request(
|
||||
model="gemini-3.5-transcribe",
|
||||
audio_file=("sample.wav", AUDIO_BYTES, "audio/wav"),
|
||||
optional_params={"response_format": response_format},
|
||||
litellm_params={},
|
||||
)
|
||||
transcription_config = request_data.data["generation_config"]["transcription_config"]
|
||||
assert json.loads(json.dumps(transcription_config)) == {
|
||||
"mode": {
|
||||
"type": "verbatim",
|
||||
"timestamp_granularities": ["word"],
|
||||
"diarization_mode": "speaker",
|
||||
}
|
||||
}
|
||||
|
||||
@pytest.mark.parametrize("response_format", ["json", "text", "verbose_json"])
|
||||
def test_non_subtitle_response_format_sends_no_mode(self, config, response_format):
|
||||
request_data = config.transform_audio_transcription_request(
|
||||
model="gemini-3.5-transcribe",
|
||||
audio_file=("sample.wav", AUDIO_BYTES, "audio/wav"),
|
||||
optional_params={"response_format": response_format},
|
||||
litellm_params={},
|
||||
)
|
||||
assert "generation_config" not in request_data.data
|
||||
|
||||
def test_segment_granularity_sends_no_mode(self, config):
|
||||
request_data = config.transform_audio_transcription_request(
|
||||
model="gemini-3.5-transcribe",
|
||||
|
|
@ -214,6 +241,54 @@ class TestTransformResponse:
|
|||
assert response.get("duration") is None
|
||||
|
||||
|
||||
class TestSubtitleSynthesisThroughHandler:
|
||||
def _transform(self, config, response_format):
|
||||
from unittest.mock import Mock
|
||||
|
||||
from litellm.llms.custom_httpx.llm_http_handler import BaseLLMHTTPHandler
|
||||
from litellm.types.utils import TranscriptionResponse
|
||||
|
||||
return BaseLLMHTTPHandler()._transform_audio_transcription_response(
|
||||
provider_config=config,
|
||||
model="gemini-3.5-transcribe",
|
||||
response=make_response(COMPLETED_RESPONSE),
|
||||
model_response=TranscriptionResponse(),
|
||||
logging_obj=Mock(),
|
||||
optional_params={"response_format": response_format},
|
||||
api_key=None,
|
||||
)
|
||||
|
||||
def test_supports_subtitle_synthesis(self, config):
|
||||
assert config.supports_subtitle_synthesis is True
|
||||
|
||||
def test_srt_synthesizes_subtitle_document_and_drops_words(self, config):
|
||||
response = self._transform(config, "srt")
|
||||
assert response.text == (
|
||||
"1\n00:00:00,100 --> 00:00:00,400\nHello\n\n2\n00:00:00,500 --> 00:00:00,900\nworld.\n"
|
||||
)
|
||||
assert "words" not in response
|
||||
assert response["task"] == "transcribe"
|
||||
assert response["duration"] == 0.9
|
||||
assert response.usage.total_tokens == 200
|
||||
|
||||
def test_vtt_synthesizes_subtitle_document_and_drops_words(self, config):
|
||||
response = self._transform(config, "vtt")
|
||||
assert response.text == (
|
||||
"WEBVTT\n\n00:00:00.100 --> 00:00:00.400\nHello\n\n00:00:00.500 --> 00:00:00.900\nworld.\n"
|
||||
)
|
||||
assert "words" not in response
|
||||
assert response.usage.total_tokens == 200
|
||||
|
||||
@pytest.mark.parametrize("response_format", ["json", "verbose_json"])
|
||||
def test_non_subtitle_formats_keep_plain_text_and_words(self, config, response_format):
|
||||
response = self._transform(config, response_format)
|
||||
assert response.text == "Hello world."
|
||||
assert response["words"] == [
|
||||
{"word": "Hello", "start": 0.1, "end": 0.4, "speaker": "spk:0"},
|
||||
{"word": "world.", "start": 0.5, "end": 0.9, "speaker": "spk:1"},
|
||||
]
|
||||
|
||||
|
||||
class TestCostRegression:
|
||||
@pytest.fixture
|
||||
def local_cost_map(self, monkeypatch):
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue