Hirely-Backend / backend /whispercode /socket_server.py
NaikPranav11's picture
Initial clean deployment
1207440
Raw
History Blame Contribute Delete
4.58 kB
import asyncio
import websockets
import numpy as np
from .whisper_vad import RealtimeWhisperVAD
import json
import logging
import time
import aiohttp
logger = logging.getLogger(__name__)
clients = set()
# Simple transcript assembly utilities
class TranscriptBuffer:
def __init__(self, silence_threshold=1.0):
self.words = []
self.last_word_time = None
self.silence_threshold = silence_threshold
def add_word(self, word_obj):
self.words.append(word_obj.word)
self.last_word_time = time.time()
def should_flush(self):
if not self.words:
return False
return (time.time() - (self.last_word_time or 0)) > self.silence_threshold
def flush(self):
txt = " ".join(self.words).strip()
self.words = []
self.last_word_time = None
return txt
async def post_transcript_to_rag(transcript: str):
url = "http://localhost:5001/evaluate"
payload = {"transcript": transcript}
try:
async with aiohttp.ClientSession() as session:
async with session.post(url, json=payload, timeout=10) as resp:
text = await resp.text()
logger.info(f"Posted transcript to RAG; response status={resp.status}")
return resp.status, text
except Exception as e:
logger.error(f"Failed to post transcript to RAG: {e}")
return None, str(e)
# Callback function for RealtimeWhisperVAD to call per word
transcript_buffer = TranscriptBuffer(silence_threshold=1.0)
def word_callback(word):
# This callback runs in a thread; schedule async broadcast and flush checks
try:
# Add to transcript buffer
transcript_buffer.add_word(word)
# Broadcast word to connected clients asynchronously
msg = json.dumps({"word": word.word, "start": word.start, "end": word.end})
asyncio.run_coroutine_threadsafe(broadcast_message(msg), asyncio.get_event_loop())
except Exception as e:
logger.error(f"Error in word_callback: {e}")
async def broadcast_message(message: str):
dead = []
for client in clients:
try:
await client.send(message)
except Exception:
dead.append(client)
for d in dead:
clients.discard(d)
# Configuration for RealtimeWhisperVAD
config = {
"whisper_model": "tiny",
"sample_rate": 16000,
"chunk_duration": 0.9,
"vad_threshold": 0.5,
"min_speech_duration": 0.5,
"max_speech_duration": 5,
"silence_duration": 0.5,
"transcription_callback": word_callback,
}
# Initialize RealtimeWhisperVAD
transcriber = RealtimeWhisperVAD(**config)
async def handler(websocket, path):
logger.info(f"Client connected from {websocket.remote_address}")
clients.add(websocket)
try:
# If STT backend is unavailable, inform client and close
if not getattr(transcriber, 'use_faster_whisper', False):
await websocket.send(json.dumps({"error": "STT backend unavailable"}))
await websocket.close()
return
# Start the transcriber when a client connects
if not transcriber.running:
transcriber.start()
# Polling loop: receive audio chunks and periodically check buffer flush
async for message in websocket:
try:
audio_float32 = np.frombuffer(message, dtype=np.float32)
transcriber.add_audio_chunk(audio_float32)
except Exception as e:
logger.error(f"Error processing audio chunk: {e}")
# If buffer indicates silence, flush transcript and POST to RAG
if transcript_buffer.should_flush():
transcript = transcript_buffer.flush()
if transcript:
status, resp_text = await post_transcript_to_rag(transcript)
# Send evaluation result back to clients
payload = {"event": "evaluation_result", "status": status, "resp": resp_text}
await broadcast_message(json.dumps(payload))
except websockets.exceptions.ConnectionClosed:
logger.info(f"Client disconnected from {websocket.remote_address}")
finally:
clients.discard(websocket)
async def main():
# Use 0.0.0.0 to be reachable; keep localhost if you want local-only
async with websockets.serve(handler, "0.0.0.0", 5002):
logger.info("WebSocket server started on ws://0.0.0.0:5002")
await asyncio.Future() # run forever
if __name__ == "__main__":
asyncio.run(main())