From b42ced9816d95c3aa5e9032df75fb37dcd6904f3 Mon Sep 17 00:00:00 2001 From: makaveli10 Date: Fri, 16 Feb 2024 20:31:38 +0530 Subject: [PATCH] fix: tests for end of speech message while mocking pyaudio --- tests/test_client.py | 1 + tests/test_server.py | 4 ++-- whisper_live/client.py | 11 ++++++++--- whisper_live/server.py | 3 +++ 4 files changed, 14 insertions(+), 5 deletions(-) diff --git a/tests/test_client.py b/tests/test_client.py index 468b5a1..56d5dbc 100644 --- a/tests/test_client.py +++ b/tests/test_client.py @@ -46,6 +46,7 @@ class TestClientCallbacks(BaseTestCase): "language": self.client.language, "task": self.client.task, "model": self.client.model, + "use_vad": True }) self.client.on_open(self.mock_ws_app) self.mock_ws_app.send.assert_called_with(expected_message) diff --git a/tests/test_server.py b/tests/test_server.py index cd14bb3..e5d630f 100644 --- a/tests/test_server.py +++ b/tests/test_server.py @@ -69,7 +69,7 @@ class TestServerConnection(unittest.TestCase): class TestServerInferenceAccuracy(unittest.TestCase): @classmethod def setUpClass(cls): - cls.server_process = subprocess.Popen(["python", "run_server.py"]) # Adjust the command as needed + cls.server_process = subprocess.Popen(["python", "run_server.py"]) time.sleep(2) @classmethod @@ -134,4 +134,4 @@ class TestExceptionHandling(unittest.TestCase): for message in log.output: print(message) print() - self.assertTrue(any("Unexpected error: Unexpected error" in message for message in log.output)) + self.assertTrue(any("Unexpected error" in message for message in log.output)) diff --git a/whisper_live/client.py b/whisper_live/client.py index 104824c..6d32257 100644 --- a/whisper_live/client.py +++ b/whisper_live/client.py @@ -58,6 +58,7 @@ class Client: self.server_error = False self.srt_file_path = srt_file_path self.use_vad = use_vad + self.last_recieved_segment = None if translate: self.task = "translate" @@ -123,6 +124,10 @@ class Client: (not self.transcript or float(seg['start']) >= float(self.transcript[-1]['end']))): self.transcript.append(seg) + # update last received segment and last valild responsne time + if self.last_recieved_segment is None or self.last_recieved_segment != segments[-1]["text"]: + self.last_response_recieved = time.time() + self.last_recieved_segment = segments[-1]["text"] # Truncate to last 3 entries for brevity. text = text[-3:] @@ -142,7 +147,6 @@ class Client: message (str): The received message from the server. """ - self.last_response_recieved = time.time() message = json.loads(message) if self.uid != message.get("uid"): @@ -158,6 +162,7 @@ class Client: self.recording = False if "message" in message.keys() and message["message"] == "SERVER_READY": + self.last_response_recieved = time.time() self.recording = True self.server_backend = message["backend"] print(f"[INFO]: Server Running with backend {self.server_backend}") @@ -275,11 +280,11 @@ class Client: self.stream.write(data) wavfile.close() - self.send_packet_to_server(Client.END_OF_AUDIO.encode('utf-8')) # Ensure it's sent as bytes + assert self.last_response_recieved while time.time() - self.last_response_recieved < self.disconnect_if_no_response_for: continue - + self.send_packet_to_server(Client.END_OF_AUDIO.encode('utf-8')) # Ensure it's sent as bytes if self.server_backend == "faster_whisper": self.write_srt_file(self.srt_file_path) self.stream.close() diff --git a/whisper_live/server.py b/whisper_live/server.py index 353097e..7bebe39 100644 --- a/whisper_live/server.py +++ b/whisper_live/server.py @@ -207,6 +207,9 @@ class TranscriptionServer: except json.JSONDecodeError: logging.error("Failed to decode JSON from client") return False + except ConnectionClosed: + logging.info("Connection closed by client") + return False except Exception as e: logging.error(f"Error during new connection initialization: {str(e)}") return False