From ec1dc7c6aae68ef43d65e814ff8d58b2a52d8e64 Mon Sep 17 00:00:00 2001 From: Quang Tran Date: Wed, 8 Jul 2026 23:13:20 +0700 Subject: [PATCH] fix: clean shutdown for StreamingTranscriptionClient --- tests/test_streaming_client.py | 48 +++++++++++++++++++++++++++++++++- whisper_live/client.py | 19 +++++++++++--- 2 files changed, 62 insertions(+), 5 deletions(-) diff --git a/tests/test_streaming_client.py b/tests/test_streaming_client.py index 4c8933a..d986467 100644 --- a/tests/test_streaming_client.py +++ b/tests/test_streaming_client.py @@ -1,4 +1,5 @@ import json +import time import unittest from unittest.mock import patch, MagicMock @@ -153,12 +154,57 @@ class TestConnectLifecycle(StreamingClientTestCase): def test_close_sends_end_of_audio(self): self._server_ready() + self._inner.recording = False # pretend server already closed with patch.object(self._inner, 'send_packet_to_server') as mock_send, \ patch.object(self._inner, 'close_websocket') as mock_close: - self.client.close(drain_seconds=0) + self.client.close() mock_send.assert_called_once_with(Client.END_OF_AUDIO.encode("utf-8")) mock_close.assert_called_once() + def test_close_waits_for_server_then_times_out(self): + self._server_ready() + self.assertTrue(self._inner.recording) # server still "processing" + start = time.time() + with patch.object(self._inner, 'send_packet_to_server'), \ + patch.object(self._inner, 'close_websocket') as mock_close: + self.client.close(timeout=0.2) + self.assertGreaterEqual(time.time() - start, 0.2) + mock_close.assert_called_once() + + def test_close_returns_early_when_server_closes(self): + self._server_ready() + + def close_soon(_msg): + self._inner.recording = False + + with patch.object(self._inner, 'send_packet_to_server', side_effect=close_soon), \ + patch.object(self._inner, 'close_websocket') as mock_close: + start = time.time() + self.client.close(timeout=10.0) + self.assertLess(time.time() - start, 1.0) + mock_close.assert_called_once() + + +class TestErrorHandling(StreamingClientTestCase): + def test_close_frame_not_reported_as_error(self): + """A normal CLOSE control frame (opcode 8) must not fire on_error.""" + self._server_ready() + errors = [] + self.client._client._on_error_hook = errors.append + close_frame = MagicMock() + close_frame.opcode = 8 + self._inner.on_error(self.mock_ws_app, close_frame) + self.assertEqual(errors, []) + self.assertFalse(self._inner.server_error) + + def test_real_error_still_reported(self): + self._server_ready() + errors = [] + self.client._client._on_error_hook = errors.append + self._inner.on_error(self.mock_ws_app, RuntimeError("boom")) + self.assertEqual(len(errors), 1) + self.assertTrue(self._inner.server_error) + if __name__ == '__main__': unittest.main() diff --git a/whisper_live/client.py b/whisper_live/client.py index 88a9677..759e2e7 100644 --- a/whisper_live/client.py +++ b/whisper_live/client.py @@ -933,6 +933,10 @@ class _HookedClient(Client): self._on_session_started() def on_error(self, ws, error): + # websocket-client surfaces the server's CLOSE control frame (opcode 8) + # through on_error during shutdown; a normal close is not an error. + if getattr(error, "opcode", None) == 8: + return if self._on_error_hook: self._on_error_hook(error) super().on_error(ws, error) @@ -1108,18 +1112,25 @@ class StreamingTranscriptionClient: # Alias for ``last_partial``; kept for readability at call sites. last_segment = last_partial - def close(self, drain_seconds: float = 2.0) -> None: - """Signal end-of-stream, wait briefly for final transcripts, then close. + def close(self, timeout: float = 15.0) -> None: + """Signal end-of-stream, wait for the server to finish, then close. + + After ``END_OF_AUDIO`` the server transcribes any buffered audio, sends + the final committed segment, and closes the connection. Waiting for that + server-initiated close keeps the last segment from being dropped. Args: - drain_seconds: Seconds to wait for the server to flush remaining audio. + timeout: Maximum seconds to wait for the server to close before + forcing the connection shut. """ if self._closed: return self._closed = True try: self._client.send_packet_to_server(Client.END_OF_AUDIO.encode("utf-8")) - time.sleep(drain_seconds) + deadline = time.time() + timeout + while self._client.recording and time.time() < deadline: + time.sleep(0.05) finally: self._client.close_websocket()