From 18bce1864afb2ebb08fd4061488e9aeea6015e63 Mon Sep 17 00:00:00 2001 From: Aaron Boxer Date: Fri, 17 Apr 2026 09:30:19 -0400 Subject: [PATCH] Validate server parameters on startup - max_clients must be >= 1 - max_connection_time must be > 0 - batch_max_size must be >= 1 (when batch enabled) - batch_window_ms must be >= 0 (when batch enabled) - Added 5 new tests for parameter validation --- tests/test_server_extended.py | 25 +++++++++++++++++++++++++ whisper_live/server.py | 10 ++++++++++ 2 files changed, 35 insertions(+) diff --git a/tests/test_server_extended.py b/tests/test_server_extended.py index 3172f4e..3ec4aa4 100644 --- a/tests/test_server_extended.py +++ b/tests/test_server_extended.py @@ -215,6 +215,31 @@ class TestTranscriptionServerInit(unittest.TestCase): whisper_tensorrt_path="/nonexistent/path", ) + def test_run_max_clients_zero_raises(self): + server = TranscriptionServer() + with self.assertRaises(ValueError): + server.run(host="localhost", port=9090, max_clients=0) + + def test_run_max_clients_negative_raises(self): + server = TranscriptionServer() + with self.assertRaises(ValueError): + server.run(host="localhost", port=9090, max_clients=-1) + + def test_run_max_connection_time_zero_raises(self): + server = TranscriptionServer() + with self.assertRaises(ValueError): + server.run(host="localhost", port=9090, max_connection_time=0) + + def test_run_batch_max_size_zero_raises(self): + server = TranscriptionServer() + with self.assertRaises(ValueError): + server.run(host="localhost", port=9090, batch_enabled=True, batch_max_size=0) + + def test_run_batch_window_ms_negative_raises(self): + server = TranscriptionServer() + with self.assertRaises(ValueError): + server.run(host="localhost", port=9090, batch_enabled=True, batch_window_ms=-1) + class TestTranscriptionServerGetAudio(unittest.TestCase): def setUp(self): diff --git a/whisper_live/server.py b/whisper_live/server.py index e342119..54da64a 100644 --- a/whisper_live/server.py +++ b/whisper_live/server.py @@ -468,6 +468,16 @@ class TranscriptionServer: """ self.cache_path = cache_path self.raw_pcm_input = raw_pcm_input + + if max_clients < 1: + raise ValueError(f"max_clients must be >= 1, got {max_clients}") + if max_connection_time <= 0: + raise ValueError(f"max_connection_time must be > 0, got {max_connection_time}") + if batch_enabled and batch_max_size < 1: + raise ValueError(f"batch_max_size must be >= 1, got {batch_max_size}") + if batch_enabled and batch_window_ms < 0: + raise ValueError(f"batch_window_ms must be >= 0, got {batch_window_ms}") + self.client_manager = ClientManager(max_clients, max_connection_time) if faster_whisper_custom_model_path is not None and not os.path.exists(faster_whisper_custom_model_path): if "/" not in faster_whisper_custom_model_path: