diff --git a/whisper_live/client.py b/whisper_live/client.py index 070845d..633c123 100644 --- a/whisper_live/client.py +++ b/whisper_live/client.py @@ -50,7 +50,7 @@ class Client: INSTANCES = {} def __init__( - self, host=None, port=None, is_multilingual=False, lang=None, translate=False + self, host=None, port=None, is_multilingual=False, lang=None, translate=False, model_size="small" ): """ Initializes a Client instance for audio recording and streaming to a server. @@ -80,7 +80,9 @@ class Client: self.last_response_recieved = None self.disconnect_if_no_response_for = 15 self.multilingual = is_multilingual - self.language = lang if is_multilingual else "en" + self.language = lang + self.model_size = model_size + self.server_error = False if translate: self.task = "translate" @@ -140,11 +142,16 @@ class Client: print("[ERROR]: invalid client uid") return - if "status" in message.keys() and message["status"] == "WAIT": - self.waiting = True - print( - f"[INFO]:Server is full. Estimated wait time {round(message['message'])} minutes." - ) + if "status" in message.keys(): + if message["status"] == "WAIT": + self.waiting = True + print( + f"[INFO]:Server is full. Estimated wait time {round(message['message'])} minutes." + ) + elif message["status"] == "ERROR": + print(f"Message from Server: {message['message']}") + self.server_error = True + return if "message" in message.keys() and message["message"] == "DISCONNECT": print("[INFO]: Server overtime disconnected.") @@ -213,6 +220,7 @@ class Client: "multilingual": self.multilingual, "language": self.language, "task": self.task, + "model_size": self.model_size, } ) ) @@ -497,8 +505,8 @@ class TranscriptionClient: transcription_client() ``` """ - def __init__(self, host, port, is_multilingual=False, lang=None, translate=False): - self.client = Client(host, port, is_multilingual, lang, translate) + def __init__(self, host, port, is_multilingual=False, lang=None, translate=False, model_size="small"): + self.client = Client(host, port, is_multilingual, lang, translate, model_size) def __call__(self, audio=None, hls_url=None): """ @@ -514,10 +522,10 @@ class TranscriptionClient: """ print("[INFO]: Waiting for server ready ...") while not self.client.recording: - if self.client.waiting: + if self.client.waiting or self.client.server_error: self.client.close_websocket() return - pass + print("[INFO]: Server Ready!") if hls_url is not None: self.client.process_hls_stream(hls_url)