fix: clean shutdown for StreamingTranscriptionClient
This commit is contained in:
@@ -1,4 +1,5 @@
|
|||||||
import json
|
import json
|
||||||
|
import time
|
||||||
import unittest
|
import unittest
|
||||||
from unittest.mock import patch, MagicMock
|
from unittest.mock import patch, MagicMock
|
||||||
|
|
||||||
@@ -153,12 +154,57 @@ class TestConnectLifecycle(StreamingClientTestCase):
|
|||||||
|
|
||||||
def test_close_sends_end_of_audio(self):
|
def test_close_sends_end_of_audio(self):
|
||||||
self._server_ready()
|
self._server_ready()
|
||||||
|
self._inner.recording = False # pretend server already closed
|
||||||
with patch.object(self._inner, 'send_packet_to_server') as mock_send, \
|
with patch.object(self._inner, 'send_packet_to_server') as mock_send, \
|
||||||
patch.object(self._inner, 'close_websocket') as mock_close:
|
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_send.assert_called_once_with(Client.END_OF_AUDIO.encode("utf-8"))
|
||||||
mock_close.assert_called_once()
|
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__':
|
if __name__ == '__main__':
|
||||||
unittest.main()
|
unittest.main()
|
||||||
|
|||||||
+15
-4
@@ -933,6 +933,10 @@ class _HookedClient(Client):
|
|||||||
self._on_session_started()
|
self._on_session_started()
|
||||||
|
|
||||||
def on_error(self, ws, error):
|
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:
|
if self._on_error_hook:
|
||||||
self._on_error_hook(error)
|
self._on_error_hook(error)
|
||||||
super().on_error(ws, error)
|
super().on_error(ws, error)
|
||||||
@@ -1108,18 +1112,25 @@ class StreamingTranscriptionClient:
|
|||||||
# Alias for ``last_partial``; kept for readability at call sites.
|
# Alias for ``last_partial``; kept for readability at call sites.
|
||||||
last_segment = last_partial
|
last_segment = last_partial
|
||||||
|
|
||||||
def close(self, drain_seconds: float = 2.0) -> None:
|
def close(self, timeout: float = 15.0) -> None:
|
||||||
"""Signal end-of-stream, wait briefly for final transcripts, then close.
|
"""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:
|
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:
|
if self._closed:
|
||||||
return
|
return
|
||||||
self._closed = True
|
self._closed = True
|
||||||
try:
|
try:
|
||||||
self._client.send_packet_to_server(Client.END_OF_AUDIO.encode("utf-8"))
|
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:
|
finally:
|
||||||
self._client.close_websocket()
|
self._client.close_websocket()
|
||||||
|
|
||||||
|
|||||||
Reference in New Issue
Block a user