Merge pull request #36 from makaveli10/add_client_queue

Add client queue.
This commit is contained in:
Marcus Edel
2023-08-09 11:48:03 -04:00
committed by GitHub
13 changed files with 404 additions and 66 deletions
+5 -3
View File
@@ -190,17 +190,19 @@ async function stopCapture() {
* Listens for messages from the runtime and performs corresponding actions. * Listens for messages from the runtime and performs corresponding actions.
* @param {Object} message - The message received from the runtime. * @param {Object} message - The message received from the runtime.
*/ */
chrome.runtime.onMessage.addListener((message) => { chrome.runtime.onMessage.addListener(async (message) => {
if (message.action === "startCapture") { if (message.action === "startCapture") {
startCapture(message); startCapture(message);
} else if (message.action === "stopCapture") { } else if (message.action === "stopCapture") {
stopCapture(); stopCapture();
} else if (message.action === "updateSelectedLanguage") { } else if (message.action === "updateSelectedLanguage") {
console.log("Selected language");
console.log(message.detectedLanguage);
const detectedLanguage = message.detectedLanguage; const detectedLanguage = message.detectedLanguage;
chrome.runtime.sendMessage({ action: "updateSelectedLanguage", detectedLanguage }); chrome.runtime.sendMessage({ action: "updateSelectedLanguage", detectedLanguage });
chrome.storage.local.set({ selectedLanguage: detectedLanguage }); chrome.storage.local.set({ selectedLanguage: detectedLanguage });
} else if (message.action === "toggleCaptureButtons") {
chrome.runtime.sendMessage({ action: "toggleCaptureButtons", data: false });
chrome.storage.local.set({ capturingState: { isCapturing: false } })
stopCapture();
} }
}); });
+53
View File
@@ -6,6 +6,52 @@ var elem_text = null;
var segments = []; var segments = [];
var text_segments = []; var text_segments = [];
function initPopupElement() {
if (document.getElementById('popupElement')) {
return;
}
const popupContainer = document.createElement('div');
popupContainer.id = 'popupElement';
popupContainer.style.cssText = 'position: fixed; top: 50%; left: 50%; transform: translate(-50%, -50%); background: white; color: black; padding: 16px; border-radius: 10px; box-shadow: 0px 0px 10px rgba(0, 0, 0, 0.5); display: none; text-align: center;';
const popupText = document.createElement('span');
popupText.textContent = 'Default Text';
popupText.className = 'popupText';
popupText.style.fontSize = '24px';
popupContainer.appendChild(popupText);
const buttonContainer = document.createElement('div');
buttonContainer.style.marginTop = '8px';
const closePopupButton = document.createElement('button');
closePopupButton.textContent = 'Close';
closePopupButton.style.backgroundColor = '#65428A';
closePopupButton.style.color = 'white';
closePopupButton.style.border = 'none';
closePopupButton.style.padding = '8px 16px'; // Add padding for better click area
closePopupButton.style.cursor = 'pointer';
closePopupButton.addEventListener('click', async () => {
popupContainer.style.display = 'none';
await browser.runtime.sendMessage({ action: 'toggleCaptureButtons', data: false });
});
buttonContainer.appendChild(closePopupButton);
popupContainer.appendChild(buttonContainer);
document.body.appendChild(popupContainer);
}
function showPopup(customText) {
const popup = document.getElementById('popupElement');
const popupText = popup.querySelector('.popupText');
if (popup && popupText) {
popupText.textContent = customText || 'Default Text'; // Set default text if custom text is not provided
popup.style.display = 'block';
}
}
function init_element() { function init_element() {
if (document.getElementById('transcription')) { if (document.getElementById('transcription')) {
return; return;
@@ -128,11 +174,18 @@ chrome.runtime.onMessage.addListener((request, sender, sendResponse) => {
remove_element(); remove_element();
sendResponse({data: "STOPPED"}); sendResponse({data: "STOPPED"});
return; return;
} else if (type === "showWaitPopup"){
initPopupElement();
showPopup(`Estimated wait time ~ ${Math.round(data)} minutes`);
sendResponse({data: "popup"});
return;
} }
init_element(); init_element();
message = JSON.parse(data); message = JSON.parse(data);
message = message["segments"];
var text = ''; var text = '';
for (var i = 0; i < message.length; i++) { for (var i = 0; i < message.length; i++) {
-1
View File
@@ -10,7 +10,6 @@
</head> </head>
<body> <body>
<script src="ort.min.js"></script>
<script src="options.js"></script> <script src="options.js"></script>
</body> </body>
+31 -2
View File
@@ -66,6 +66,16 @@ function resampleTo16kHZ(audioData, origSampleRate = 44100) {
return resampledData; return resampledData;
} }
function generateUUID() {
let dt = new Date().getTime();
const uuid = 'xxxxxxxx-xxxx-4xxx-yxxx-xxxxxxxxxxxx'.replace(/[xy]/g, function(c) {
const r = (dt + Math.random() * 16) % 16 | 0;
dt = Math.floor(dt / 16);
return (c === 'x' ? r : (r & 0x3 | 0x8)).toString(16);
});
return uuid;
}
/** /**
* Starts recording audio from the captured tab. * Starts recording audio from the captured tab.
@@ -73,6 +83,7 @@ function resampleTo16kHZ(audioData, origSampleRate = 44100) {
*/ */
async function startRecord(option) { async function startRecord(option) {
const stream = await captureTabAudio(); const stream = await captureTabAudio();
const uuid = generateUUID();
if (stream) { if (stream) {
// call when the stream inactive // call when the stream inactive
@@ -88,6 +99,7 @@ async function startRecord(option) {
socket.onopen = function(e) { socket.onopen = function(e) {
socket.send( socket.send(
JSON.stringify({ JSON.stringify({
uid: uuid,
multilingual: option.multilingual, multilingual: option.multilingual,
language: option.language, language: option.language,
task: option.task task: option.task
@@ -96,13 +108,26 @@ async function startRecord(option) {
}; };
socket.onmessage = async (event) => { socket.onmessage = async (event) => {
const data = JSON.parse(event.data);
if (data["uid"] !== uuid)
return;
if (data["status"] === "WAIT"){
await sendMessageToTab(option.currentTabId, {
type: "showWaitPopup",
data: data["message"],
});
chrome.runtime.sendMessage({ action: "toggleCaptureButtons", data: false })
chrome.runtime.sendMessage({ action: "stopCapture" })
return;
}
if (isServerReady === false){ if (isServerReady === false){
isServerReady = true; isServerReady = true;
return; return;
} }
if (language === null) { if (language === null) {
const data = JSON.parse(event.data);
language = data["language"]; language = data["language"];
// send message to popup.js to update dropdown // send message to popup.js to update dropdown
@@ -115,6 +140,11 @@ async function startRecord(option) {
return; return;
} }
if (data["message"] === "DISCONNECT"){
chrome.runtime.sendMessage({ action: "toggleCaptureButtons", data: false })
return;
}
res = await sendMessageToTab(option.currentTabId, { res = await sendMessageToTab(option.currentTabId, {
type: "transcript", type: "transcript",
data: event.data, data: event.data,
@@ -135,7 +165,6 @@ async function startRecord(option) {
audioDataCache.push(inputData); audioDataCache.push(inputData);
// feed inputs and run
socket.send(audioData16kHz); socket.send(audioData16kHz);
}; };
+9
View File
@@ -119,6 +119,8 @@ document.addEventListener("DOMContentLoaded", function () {
startButton.disabled = isCapturing; startButton.disabled = isCapturing;
stopButton.disabled = !isCapturing; stopButton.disabled = !isCapturing;
useServerCheckbox.disabled = isCapturing; useServerCheckbox.disabled = isCapturing;
useMultilingualCheckbox.disabled = isCapturing;
startButton.classList.toggle("disabled", isCapturing); startButton.classList.toggle("disabled", isCapturing);
stopButton.classList.toggle("disabled", !isCapturing); stopButton.classList.toggle("disabled", !isCapturing);
} }
@@ -166,4 +168,11 @@ document.addEventListener("DOMContentLoaded", function () {
} }
}); });
chrome.runtime.onMessage.addListener(async (request, sender, sendResponse) => {
if (request.action === "toggleCaptureButtons") {
toggleCaptureButtons(false);
chrome.storage.local.set({ capturingState: { isCapturing: false } })
}
});
}); });
+49 -5
View File
@@ -1,7 +1,7 @@
browser.runtime.onMessage.addListener(function(request, sender, sendResponse) { browser.runtime.onMessage.addListener(async function(request, sender, sendResponse) {
const { action, data } = request; const { action, data } = request;
if (action === "transcript") { if (action === "transcript") {
browser.tabs.query({ active: true, currentWindow: true }) await browser.tabs.query({ active: true, currentWindow: true })
.then((tabs) => { .then((tabs) => {
const tabId = tabs[0].id; const tabId = tabs[0].id;
browser.tabs.sendMessage(tabId, { action: "show_transcript", data }); browser.tabs.sendMessage(tabId, { action: "show_transcript", data });
@@ -12,9 +12,53 @@ browser.runtime.onMessage.addListener(function(request, sender, sendResponse) {
} }
if (action === "updateSelectedLanguage") { if (action === "updateSelectedLanguage") {
const detectedLanguage = data; const detectedLanguage = data;
if (detectedLanguage) { try {
browser.runtime.sendMessage({ action: "updateSelectedLanguage", detectedLanguage }); await browser.storage.local.set({ selectedLanguage: detectedLanguage });
browser.storage.local.set({ selectedLanguage: detectedLanguage }); browser.tabs.query({ active: true, currentWindow: true }).then((tabs) => {
const tabId = tabs[0].id;
browser.tabs.sendMessage(tabId, { action: "updateSelectedLanguage", detectedLanguage });
});
} catch (error) {
console.error("Error updateSelectedLanguage:", error);
}
}
if (action === "toggleCaptureButtons") {
try {
await browser.storage.local.set({ capturingState: { isCapturing: false } });
browser.tabs.query({ active: true, currentWindow: true }).then((tabs) => {
const tabId = tabs[0].id;
browser.tabs.sendMessage(tabId, { action: "toggleCaptureButtons", data: false });
});
} catch (error) {
console.error("Error updating capturing state:", error);
}
try{
await browser.tabs.query({ active: true, currentWindow: true })
.then((tabs) => {
const tabId = tabs[0].id;
browser.tabs.sendMessage(tabId, { action: "stopCapture", data });
})
.catch((error) => {
console.error("Error retrieving active tab:", error);
});
} catch (error) {
console.error(error);
}
}
if (action === "showPopup") {
try{
await browser.tabs.query({ active: true, currentWindow: true })
.then((tabs) => {
const tabId = tabs[0].id;
browser.tabs.sendMessage(tabId, { action: "showWaitPopup", data });
})
.catch((error) => {
console.error(error);
});
} catch (error) {
console.error(error);
} }
} }
}); });
+99 -12
View File
@@ -5,6 +5,30 @@ let audioContext = null;
let scriptProcessor = null; let scriptProcessor = null;
let language = null; let language = null;
let isPaused = false;
const mediaElements = document.querySelectorAll('video, audio');
mediaElements.forEach((mediaElement) => {
mediaElement.addEventListener('play', handlePlaybackStateChange);
mediaElement.addEventListener('pause', handlePlaybackStateChange);
});
function handlePlaybackStateChange(event) {
isPaused = event.target.paused;
}
function generateUUID() {
let dt = new Date().getTime();
const uuid = 'xxxxxxxx-xxxx-4xxx-yxxx-xxxxxxxxxxxx'.replace(/[xy]/g, function(c) {
const r = (dt + Math.random() * 16) % 16 | 0;
dt = Math.floor(dt / 16);
return (c === 'x' ? r : (r & 0x3 | 0x8)).toString(16);
});
return uuid;
}
/** /**
* Resamples the audio data to a target sample rate of 16kHz. * Resamples the audio data to a target sample rate of 16kHz.
* @param {Array|ArrayBuffer|TypedArray} audioData - The input audio data. * @param {Array|ArrayBuffer|TypedArray} audioData - The input audio data.
@@ -45,9 +69,12 @@ function startRecording(data) {
if (language === null && !data.useMultilingual) { if (language === null && !data.useMultilingual) {
language = 'en'; language = 'en';
} }
const uuid = generateUUID();
socket.onopen = function(e) { socket.onopen = function(e) {
socket.send( socket.send(
JSON.stringify({ JSON.stringify({
uid: uuid,
multilingual: data.useMultilingual, multilingual: data.useMultilingual,
language: data.language, language: data.language,
task: data.task task: data.task
@@ -56,25 +83,33 @@ function startRecording(data) {
}; };
let isServerReady = false; let isServerReady = false;
socket.onmessage = (event) => { socket.onmessage = async (event) => {
if (!isServerReady){ const data = JSON.parse(event.data);
if (data["uid"] !== uuid)
return;
if (data["status"] === "WAIT"){
await browser.runtime.sendMessage({ action: "showPopup", data: data["message"] })
return;
}
if (!isServerReady && data["message"] === "SERVER_READY"){
isServerReady = true; isServerReady = true;
return; return;
} }
if (language === null ){ if (language === null ){
const data = JSON.parse(event.data);
language = data["language"]; language = data["language"];
await browser.runtime.sendMessage({ action: "updateSelectedLanguage", data: language })
browser.runtime.sendMessage({ action: "updateSelectedLanguage", data: language })
.catch(function(error) {
console.error("Error sending message:", error);
});
return return
} }
const data = event.data;; if (data["message"] === "DISCONNECT"){
browser.runtime.sendMessage({ action: "transcript", data }) await browser.runtime.sendMessage({ action: "toggleCaptureButtons", data: false })
return
}
await browser.runtime.sendMessage({ action: "transcript", data: event.data })
.catch(function(error) { .catch(function(error) {
console.error("Error sending message:", error); console.error("Error sending message:", error);
}); });
@@ -90,14 +125,13 @@ function startRecording(data) {
recorder = audioContext.createScriptProcessor(4096, 1, 1); recorder = audioContext.createScriptProcessor(4096, 1, 1);
recorder.onaudioprocess = async (event) => { recorder.onaudioprocess = async (event) => {
if (!audioContext || !isCapturing || !isServerReady) return; if (!audioContext || !isCapturing || !isServerReady || isPaused) return;
const inputData = event.inputBuffer.getChannelData(0); const inputData = event.inputBuffer.getChannelData(0);
const audioData16kHz = resampleTo16kHZ(inputData, audioContext.sampleRate); const audioData16kHz = resampleTo16kHZ(inputData, audioContext.sampleRate);
audioDataCache.push(inputData); audioDataCache.push(inputData);
// feed inputs and run
socket.send(audioData16kHz); socket.send(audioData16kHz);
}; };
@@ -113,6 +147,52 @@ var elem_text = null;
var segments = []; var segments = [];
var text_segments = []; var text_segments = [];
function initPopupElement() {
if (document.getElementById('popupElement')) {
return;
}
const popupContainer = document.createElement('div');
popupContainer.id = 'popupElement';
popupContainer.style.cssText = 'position: fixed; top: 50%; left: 50%; transform: translate(-50%, -50%); background: white; color: black; padding: 16px; border-radius: 10px; box-shadow: 0px 0px 10px rgba(0, 0, 0, 0.5); display: none; text-align: center;';
const popupText = document.createElement('span');
popupText.textContent = 'Default Text';
popupText.className = 'popupText';
popupText.style.fontSize = '24px';
popupContainer.appendChild(popupText);
const buttonContainer = document.createElement('div');
buttonContainer.style.marginTop = '8px';
const closePopupButton = document.createElement('button');
closePopupButton.textContent = 'Close';
closePopupButton.style.backgroundColor = '#65428A';
closePopupButton.style.color = 'white';
closePopupButton.style.border = 'none';
closePopupButton.style.padding = '8px 16px'; // Add padding for better click area
closePopupButton.style.cursor = 'pointer';
closePopupButton.addEventListener('click', async () => {
popupContainer.style.display = 'none';
await browser.runtime.sendMessage({ action: 'toggleCaptureButtons', data: false });
});
buttonContainer.appendChild(closePopupButton);
popupContainer.appendChild(buttonContainer);
document.body.appendChild(popupContainer);
}
function showPopup(customText) {
const popup = document.getElementById('popupElement');
const popupText = popup.querySelector('.popupText');
if (popup && popupText) {
popupText.textContent = customText || 'Default Text'; // Set default text if custom text is not provided
popup.style.display = 'block';
}
}
function init_element() { function init_element() {
if (document.getElementById('transcription')) { if (document.getElementById('transcription')) {
return; return;
@@ -250,10 +330,17 @@ browser.runtime.onMessage.addListener((request, sender, sendResponse) => {
remove_element(); remove_element();
} else if (action === "showWaitPopup") {
initPopupElement();
showPopup(`Estimated wait time ~ ${Math.round(data)} minutes`);
} else if (action === "show_transcript"){ } else if (action === "show_transcript"){
if (!isCapturing) return; if (!isCapturing) return;
init_element(); init_element();
message = JSON.parse(data); message = JSON.parse(data);
message = message["segments"];
var text = ''; var text = '';
for (var i = 0; i < message.length; i++) { for (var i = 0; i < message.length; i++) {
+2
View File
@@ -15,6 +15,8 @@
<input type="checkbox" id="useServerCheckbox"> <input type="checkbox" id="useServerCheckbox">
<label for="useServerCheckbox">Use Collabora Whisper-Live Server</label> <label for="useServerCheckbox">Use Collabora Whisper-Live Server</label>
</div> </div>
<textarea id="waitTextBox" style="display: none;"></textarea>
<div class="checkbox-container"> <div class="checkbox-container">
<input type="checkbox" id="useMultilingualCheckbox"> <input type="checkbox" id="useMultilingualCheckbox">
<label for="useMultilingualCheckbox">Use Multilingual Model</label> <label for="useMultilingualCheckbox">Use Multilingual Model</label>
+14 -2
View File
@@ -113,7 +113,9 @@ document.addEventListener("DOMContentLoaded", function() {
function toggleCaptureButtons(isCapturing) { function toggleCaptureButtons(isCapturing) {
startButton.disabled = isCapturing; startButton.disabled = isCapturing;
stopButton.disabled = !isCapturing; stopButton.disabled = !isCapturing;
useServerCheckbox.disabled = isCapturing; // Disable checkbox if capturing useServerCheckbox.disabled = isCapturing;
useMultilingualCheckbox.disabled = isCapturing;
startButton.classList.toggle("disabled", isCapturing); startButton.classList.toggle("disabled", isCapturing);
stopButton.classList.toggle("disabled", !isCapturing); stopButton.classList.toggle("disabled", !isCapturing);
} }
@@ -152,7 +154,7 @@ document.addEventListener("DOMContentLoaded", function() {
browser.runtime.onMessage.addListener((request, sender, sendResponse) => { browser.runtime.onMessage.addListener((request, sender, sendResponse) => {
if (request.action === "updateSelectedLanguage") { if (request.action === "updateSelectedLanguage") {
const detectedLanguage = request.detectedLanguage; const detectedLanguage = request.data;
if (detectedLanguage) { if (detectedLanguage) {
languageDropdown.value = detectedLanguage; languageDropdown.value = detectedLanguage;
@@ -161,4 +163,14 @@ document.addEventListener("DOMContentLoaded", function() {
} }
} }
}); });
browser.runtime.onMessage.addListener((request, sender, sendResponse) => {
if (request.action === "toggleCaptureButtons") {
toggleCaptureButtons(false);
browser.storage.local.set({ capturingState: { isCapturing: false } })
.catch(function(error) {
console.error("Error storing capturing state:", error);
});
}
});
}); });
+1 -2
View File
@@ -1,6 +1,5 @@
from whisper_live.server import TranscriptionServer from whisper_live.server import TranscriptionServer
if __name__ == "__main__": if __name__ == "__main__":
server = TranscriptionServer() server = TranscriptionServer()
server.run("0.0.0.0", 9090) server.run("0.0.0.0")
+1 -1
View File
@@ -9,7 +9,7 @@ README = (HERE / "README.md").read_text()
# This call to setup() does all the work # This call to setup() does all the work
setup(name="whisper-live", setup(name="whisper-live",
version="0.0.5", version="0.0.6",
description="A nearly-live implementation of OpenAI's Whisper.", description="A nearly-live implementation of OpenAI's Whisper.",
long_description=README, long_description=README,
long_description_content_type="text/markdown", long_description_content_type="text/markdown",
+30 -6
View File
@@ -10,6 +10,7 @@ import threading
import textwrap import textwrap
import json import json
import websocket import websocket
import uuid
def resample(file: str, sr: int = 16000): def resample(file: str, sr: int = 16000):
@@ -47,10 +48,12 @@ class Client:
CHANNELS = 1 CHANNELS = 1
RATE = 16000 RATE = 16000
RECORD_SECONDS = 60000 RECORD_SECONDS = 60000
START_RECORDING = False RECORDING = False
multilingual = False multilingual = False
language = None language = None
task = "transcribe" task = "transcribe"
uid = str(uuid.uuid4())
WAITING = False
def __init__(self, host=None, port=None, is_multilingual=False, lang=None, translate=False): def __init__(self, host=None, port=None, is_multilingual=False, lang=None, translate=False):
Client.multilingual = is_multilingual Client.multilingual = is_multilingual
@@ -90,16 +93,32 @@ class Client:
@staticmethod @staticmethod
def on_message(ws, message): def on_message(ws, message):
message = json.loads(message) message = json.loads(message)
if message == "SERVER_READY": if message.get('uid')!=Client.uid:
Client.START_RECORDING = True print("[ERROR]: invalid client uid")
return return
if isinstance(message, dict): if "status" in message.keys() and message["status"] == "WAIT":
Client.WAITING = True
print(f"[INFO]:Server is full. Estimated wait time {round(message['message'])} minutes.")
if "message" in message.keys() and message["message"] == "DISCONNECT":
print("[INFO]: Server overtime disconnected.")
Client.RECORDING = False
if "message" in message.keys() and message["message"] == "SERVER_READY":
Client.RECORDING = True
return
if "language" in message.keys():
Client.language = message.get("language") Client.language = message.get("language")
lang_prob = message.get("language_prob") lang_prob = message.get("language_prob")
print(f"[INFO]: Server detected language {Client.language} with probability {lang_prob}") print(f"[INFO]: Server detected language {Client.language} with probability {lang_prob}")
return return
if "segments" not in message.keys():
return
message = message["segments"]
text = [] text = []
if len(message): if len(message):
for seg in message: for seg in message:
@@ -134,6 +153,7 @@ class Client:
print("[INFO]: Opened connection") print("[INFO]: Opened connection")
ws.send(json.dumps({ ws.send(json.dumps({
'uid': Client.uid,
'multilingual': Client.multilingual, 'multilingual': Client.multilingual,
'language': Client.language, 'language': Client.language,
'task': Client.task 'task': Client.task
@@ -162,7 +182,7 @@ class Client:
output=True, output=True,
frames_per_buffer=self.CHUNK) frames_per_buffer=self.CHUNK)
try: try:
while True: while Client.RECORDING:
data = self.wf.readframes(self.CHUNK) data = self.wf.readframes(self.CHUNK)
if data==b'': break if data==b'': break
@@ -210,6 +230,7 @@ class Client:
os.makedirs("chunks", exist_ok=True) os.makedirs("chunks", exist_ok=True)
try: try:
for _ in range(0, int(self.RATE / self.CHUNK * self.RECORD_SECONDS)): for _ in range(0, int(self.RATE / self.CHUNK * self.RECORD_SECONDS)):
if not Client.RECORDING: break
data = self.stream.read(self.CHUNK) data = self.stream.read(self.CHUNK)
self.frames += data self.frames += data
@@ -264,7 +285,10 @@ class TranscriptionClient:
def __call__(self, audio=None): def __call__(self, audio=None):
print("[INFO]: Waiting for server ready ...") print("[INFO]: Waiting for server ready ...")
while not Client.START_RECORDING: while not Client.RECORDING:
if Client.WAITING:
self.client.close_websocket()
return
pass pass
print("[INFO]: Server Ready!") print("[INFO]: Server Ready!")
if audio is not None: if audio is not None:
+108 -30
View File
@@ -2,11 +2,12 @@ import websockets
import pickle, struct, time, pyaudio import pickle, struct, time, pyaudio
import threading import threading
import os, json import os, json
import base64
import wave import wave
import textwrap import textwrap
import logging import logging
logging.basicConfig(level = logging.INFO) # logging.basicConfig(level = logging.INFO)
from collections import deque from collections import deque
from dataclasses import dataclass from dataclasses import dataclass
@@ -14,6 +15,7 @@ from websockets.sync.server import serve
import torch import torch
import numpy as np import numpy as np
import time
from whisper_live.transcriber import WhisperModel from whisper_live.transcriber import WhisperModel
@@ -24,9 +26,30 @@ class TranscriptionServer:
Attributes: Attributes:
clients (dict): A dictionary to store connected clients. clients (dict): A dictionary to store connected clients.
""" """
RATE = 16000
def __init__(self): def __init__(self):
# voice activity detection model
self.vad_model, _ = torch.hub.load(repo_or_dir='snakers4/silero-vad',
model='silero_vad',
force_reload=True,
onnx=True
)
self.vad_threshold = 0.4
self.clients = {} self.clients = {}
self.websockets = {}
self.clients_start_time = {}
self.max_clients = 4
self.max_connection_time = 600 # in seconds
def get_wait_time(self):
wait_time = None
for k,v in self.clients_start_time.items():
current_client_time_remaining = self.max_connection_time - (time.time() - v)
if wait_time is None:
wait_time = current_client_time_remaining
elif current_client_time_remaining < wait_time:
wait_time = current_client_time_remaining
return wait_time/60
def recv_audio(self, websocket): def recv_audio(self, websocket):
""" """
@@ -35,30 +58,73 @@ class TranscriptionServer:
Args: Args:
websocket (WebSocket): The WebSocket connection for the client. websocket (WebSocket): The WebSocket connection for the client.
""" """
logging.info("New client connected")
options = websocket.recv() options = websocket.recv()
options = json.loads(options) options = json.loads(options)
if len(self.clients) >= self.max_clients:
logging.warning("Client Queue Full. Asking client to wait ...")
wait_time = self.get_wait_time()
response = {
"uid" : options["uid"],
"status": "WAIT",
"message": wait_time,
}
websocket.send(json.dumps(response))
websocket.close()
del websocket
return
client = ServeClient( client = ServeClient(
websocket, websocket,
multilingual=options["multilingual"], multilingual=options["multilingual"],
language=options["language"], language=options["language"],
task=options["task"], task=options["task"],
client_uid=options["uid"]
) )
self.clients[websocket] = client self.clients[websocket] = client
self.clients_start_time[websocket] = time.time()
while True: while True:
try: try:
frame_data = websocket.recv() frame_data = websocket.recv()
frame_np = np.frombuffer(frame_data, np.float32) frame_np = np.frombuffer(frame_data, dtype=np.float32)
try:
speech_prob = self.vad_model(torch.from_numpy(frame_np.copy()), self.RATE).item()
if speech_prob < self.vad_threshold:
continue
except Exception as e:
logging.error(e)
return
self.clients[websocket].add_frames(frame_np) self.clients[websocket].add_frames(frame_np)
elapsed_time = time.time() - self.clients_start_time[websocket]
if elapsed_time >= self.max_connection_time:
self.clients[websocket].disconnect()
logging.warning(f"{self.clients[websocket]} Client disconnected due to overtime.")
self.clients[websocket].cleanup()
self.clients.pop(websocket)
self.clients_start_time.pop(websocket)
websocket.close()
del websocket
break
except Exception as e: except Exception as e:
logging.error(e)
self.clients[websocket].cleanup() self.clients[websocket].cleanup()
self.clients.pop(websocket) self.clients.pop(websocket)
self.clients_start_time.pop(websocket)
logging.info("Connection Closed.") logging.info("Connection Closed.")
logging.info(self.clients)
del websocket
break break
def run(self, host, port): def run(self, host, port=9090):
""" """
Run the transcription server. Run the transcription server.
@@ -73,8 +139,10 @@ class TranscriptionServer:
class ServeClient: class ServeClient:
RATE = 16000 RATE = 16000
SERVER_READY = "SERVER_READY" SERVER_READY = "SERVER_READY"
DISCONNECT = "DISCONNECT"
def __init__(self, websocket, task="transcribe", device=None, multilingual=False, language=None): def __init__(self, websocket, task="transcribe", device=None, multilingual=False, language=None, client_uid=None):
self.client_uid = client_uid
self.data = b"" self.data = b""
self.frames = b"" self.frames = b""
self.language = language if multilingual else "en" self.language = language if multilingual else "en"
@@ -87,14 +155,6 @@ class ServeClient:
local_files_only=False, local_files_only=False,
) )
# voice activity detection model
self.vad_model, _ = torch.hub.load(repo_or_dir='snakers4/silero-vad',
model='silero_vad',
force_reload=True,
onnx=True
)
self.vad_threshold = 0.4
self.timestamp_offset = 0.0 self.timestamp_offset = 0.0
self.frames_np = None self.frames_np = None
self.frames_offset = 0.0 self.frames_offset = 0.0
@@ -117,7 +177,14 @@ class ServeClient:
self.websocket = websocket self.websocket = websocket
self.trans_thread = threading.Thread(target=self.speech_to_text) self.trans_thread = threading.Thread(target=self.speech_to_text)
self.trans_thread.start() self.trans_thread.start()
self.websocket.send(json.dumps(self.SERVER_READY)) self.websocket.send(
json.dumps(
{
"uid": self.client_uid,
"message": self.SERVER_READY
}
)
)
def fill_output(self, output): def fill_output(self, output):
""" """
@@ -142,15 +209,6 @@ class ServeClient:
return wrapped return wrapped
def add_frames(self, frame_np): def add_frames(self, frame_np):
try:
speech_prob = self.vad_model(torch.from_numpy(frame_np.copy()), self.RATE).item()
if speech_prob < self.vad_threshold:
return
except Exception as e:
logging.error(e)
return
if self.frames_np is not None and self.frames_np.shape[0] > 45*self.RATE: if self.frames_np is not None and self.frames_np.shape[0] > 45*self.RATE:
self.frames_offset += 30.0 self.frames_offset += 30.0
self.frames_np = self.frames_np[int(30*self.RATE):] self.frames_np = self.frames_np[int(30*self.RATE):]
@@ -179,7 +237,8 @@ class ServeClient:
task=self.task task=self.task
) )
logging.info(f"Detected language {self.language} with probability {lang_prob}") logging.info(f"Detected language {self.language} with probability {lang_prob}")
self.websocket.send(json.dumps({"language": self.language, "language_prob": lang_prob})) self.websocket.send(json.dumps(
{"uid": self.client_uid, "language": self.language, "language_prob": lang_prob}))
while True: while True:
if self.exit: if self.exit:
@@ -227,9 +286,14 @@ class ServeClient:
segments = segments + [last_segment] segments = segments + [last_segment]
try: try:
self.websocket.send(json.dumps(segments)) self.websocket.send(
json.dumps({
"uid": self.client_uid,
"segments": segments
})
)
except Exception as e: except Exception as e:
logging.info(f"[ERROR]: {e}") logging.error(f"[ERROR]: {e}")
else: else:
# show previous output if there is pause i.e. no output from whisper # show previous output if there is pause i.e. no output from whisper
segments = [] segments = []
@@ -246,11 +310,16 @@ class ServeClient:
self.text.append('') self.text.append('')
try: try:
self.websocket.send(json.dumps(segments)) self.websocket.send(
json.dumps({
"uid": self.client_uid,
"segments": segments
})
)
except Exception as e: except Exception as e:
logging.info(f"[INFO]: {e}") logging.error(f"[ERROR]: {e}")
except Exception as e: except Exception as e:
logging.info(f"[INFO]: {e}") logging.error(f"[ERROR]: {e}")
time.sleep(0.01) time.sleep(0.01)
def update_segments(self, segments, duration): def update_segments(self, segments, duration):
@@ -321,8 +390,17 @@ class ServeClient:
return last_segment return last_segment
def disconnect(self):
self.websocket.send(
json.dumps(
{
"uid": self.client_uid,
"message": self.DISCONNECT
}
)
)
def cleanup(self): def cleanup(self):
logging.info("Cleaning up.") logging.info("Cleaning up.")
self.exit = True self.exit = True
self.transcriber.destroy() self.transcriber.destroy()