From 3911d62bbeee55f940ac4294275ada2c5580bf88 Mon Sep 17 00:00:00 2001 From: mateo-berri <277851410+mateo-berri@users.noreply.github.com> Date: Fri, 18 Sep 2026 17:44:28 -0700 Subject: [PATCH] fix(vertex_ai): prune a discarded turn's id once its marker is delivered --- .../audio_transcription/realtime_backend.py | 27 ++++++++++++++----- .../test_vertex_ai_realtime_backend.py | 1 + 2 files changed, 21 insertions(+), 7 deletions(-) diff --git a/litellm/llms/vertex_ai/audio_transcription/realtime_backend.py b/litellm/llms/vertex_ai/audio_transcription/realtime_backend.py index 4c8338c027e..859a883463c 100644 --- a/litellm/llms/vertex_ai/audio_transcription/realtime_backend.py +++ b/litellm/llms/vertex_ai/audio_transcription/realtime_backend.py @@ -87,10 +87,16 @@ class _TurnResult: @dataclass(frozen=True, slots=True) class _TurnDiscarded: - pass + turn: int -_OutboxItem = str | _TurnResult | _StreamFailure | _Closed +@dataclass(frozen=True, slots=True) +class _TurnDiscardedEvent: + turn: int + event: str + + +_OutboxItem = str | _TurnResult | _TurnDiscardedEvent | _StreamFailure | _Closed def open_speech_client(target: SpeechStreamingTarget, access_token: str) -> SpeechStreamingClient: @@ -304,6 +310,9 @@ class SpeechStreamingBackend: raise _normal_closure() case _TurnResult(): return None if item.turn in self._discarded_turns else item.event + case _TurnDiscardedEvent(): + self._discarded_turns -= {item.turn} + return item.event case str(): return item case _: @@ -345,7 +354,10 @@ class SpeechStreamingBackend: self._billed_before += await link.relay(self._outbox, self._billed_before) case _TurnDiscarded(): await self._outbox.put( - VertexSpeechStreamingTurnDiscarded(billed_seconds=self._billed_before).model_dump_json() + _TurnDiscardedEvent( + turn=link.turn, + event=VertexSpeechStreamingTurnDiscarded(billed_seconds=self._billed_before).model_dump_json(), + ) ) case _: assert_never(link) @@ -396,10 +408,11 @@ class SpeechStreamingBackend: await self._link(_TURN_FINISHED_EVENT) async def _discard_turn(self) -> None: - turn: Final = self._turn + streams: Final = self._turn + turn: Final = self._turn_index self._turn = () - self._discarded_turns |= {self._turn_index} + self._discarded_turns |= {turn} self._turn_index += 1 - for stream in turn: + for stream in streams: stream.cancel() - await self._link(_TurnDiscarded()) + await self._link(_TurnDiscarded(turn=turn)) diff --git a/tests/test_litellm/llms/vertex_ai/audio_transcription/test_vertex_ai_realtime_backend.py b/tests/test_litellm/llms/vertex_ai/audio_transcription/test_vertex_ai_realtime_backend.py index 15601c5ca6c..d5e88706e23 100644 --- a/tests/test_litellm/llms/vertex_ai/audio_transcription/test_vertex_ai_realtime_backend.py +++ b/tests/test_litellm/llms/vertex_ai/audio_transcription/test_vertex_ai_realtime_backend.py @@ -260,6 +260,7 @@ async def test_discard_turn_drops_its_queued_results_and_keeps_google_billed_sec await _until(lambda: len(client.streams[0]) == 3) await backend.send(DISCARD_TURN) assert await _recv(backend) == {"kind": "turn_discarded", "billed_seconds": 2.0} + assert backend._discarded_turns == frozenset() await backend.send(b"\x03\x03") fresh = await _recv(backend) assert fresh["results"] == [{"transcript": "fresh", "is_final": True}]