"""MeetVault local live speech engine. Protocol -------- WebSocket: ws://127.0.0.1:8765/ws/live 1. Client sends a JSON config message. 2. Client streams signed 16-bit little-endian mono PCM at 16 kHz. 3. Server returns JSON events: ready, partial, final, health, error. Real mode combines faster-whisper with diart speaker diarization. Mock mode is available so the desktop UI can be exercised without ML dependencies/models. """ from __future__ import annotations import argparse import asyncio import json import os import queue import threading import time import uuid from dataclasses import dataclass from typing import Optional import numpy as np from fastapi import FastAPI, WebSocket, WebSocketDisconnect from fastapi.middleware.cors import CORSMiddleware import uvicorn SAMPLE_RATE = 16000 @dataclass class EngineConfig: model: str = "small" compute: str = "auto" language: str = "auto" diarization: bool = True sample_rate: int = SAMPLE_RATE class SpeakerTimeline: """Thread-safe speaker timeline built from diart streaming predictions.""" def __init__(self) -> None: self._segments: list[tuple[float, float, str]] = [] self._label_map: dict[str, str] = {} self._lock = threading.Lock() def _stable_label(self, raw: str) -> str: with self._lock: if raw not in self._label_map: self._label_map[raw] = f"speaker_{len(self._label_map) + 1}" return self._label_map[raw] def update_from_annotation(self, annotation) -> None: fresh: list[tuple[float, float, str]] = [] try: for segment, _, label in annotation.itertracks(yield_label=True): fresh.append((float(segment.start), float(segment.end), self._stable_label(str(label)))) except Exception: return if not fresh: return latest_end = max(end for _, end, _ in fresh) with self._lock: # Keep a rolling two-minute history, then merge the latest prediction. self._segments = [s for s in self._segments if s[1] >= latest_end - 120] self._segments.extend(fresh) def resolve(self, start: float, end: float) -> str: with self._lock: candidates = list(self._segments) best_label = "speaker_1" best_overlap = 0.0 for s, e, label in candidates: overlap = max(0.0, min(end, e) - max(start, s)) if overlap > best_overlap: best_overlap = overlap best_label = label return best_label class QueueAudioSource: """diart AudioSource fed by PCM pushed from the WebSocket thread.""" def __init__(self, sample_rate: int): from diart.sources import AudioSource class _Source(AudioSource): def __init__(self, sr: int): super().__init__("meetvault-live", sr) self.q: queue.Queue[Optional[np.ndarray]] = queue.Queue() self.cursor = 0 @property def is_regular(self): return False @property def duration(self): return None def push(self, samples: np.ndarray): self.q.put(samples.astype(np.float32, copy=False)) def close(self): self.q.put(None) def read(self): from pyannote.core import SlidingWindow, SlidingWindowFeature while True: item = self.q.get() if item is None: self.stream.on_completed() return start = self.cursor / self.sample_rate window = SlidingWindow(start=start, duration=1 / self.sample_rate, step=1 / self.sample_rate) feature = SlidingWindowFeature(item.reshape(-1, 1), window) self.cursor += len(item) self.stream.on_next(feature) self.source = _Source(sample_rate) def push(self, samples: np.ndarray) -> None: self.source.push(samples) def close(self) -> None: self.source.close() class DiartWorker: def __init__(self, enabled: bool, sample_rate: int = SAMPLE_RATE): self.enabled = enabled self.timeline = SpeakerTimeline() self.audio_source: Optional[QueueAudioSource] = None self.thread: Optional[threading.Thread] = None self.error: Optional[str] = None if enabled: self._start(sample_rate) def _start(self, sample_rate: int) -> None: try: from diart import SpeakerDiarization from diart.inference import StreamingInference pipeline = SpeakerDiarization() expected_rate = int(getattr(pipeline.config, "sample_rate", sample_rate)) if expected_rate != sample_rate: raise RuntimeError(f"diart expects {expected_rate} Hz but MeetVault sends {sample_rate} Hz") self.audio_source = QueueAudioSource(sample_rate) inference = StreamingInference(pipeline, self.audio_source.source, do_plot=False) def hook(annotation_and_waveform): try: self.timeline.update_from_annotation(annotation_and_waveform[0]) except Exception: pass inference.attach_hooks(hook) self.thread = threading.Thread(target=inference, name="meetvault-diart", daemon=True) self.thread.start() except Exception as exc: self.error = str(exc) self.enabled = False def push(self, samples: np.ndarray) -> None: if self.enabled and self.audio_source: self.audio_source.push(samples) def close(self) -> None: if self.audio_source: self.audio_source.close() def resolve(self, start: float, end: float) -> str: return self.timeline.resolve(start, end) if self.enabled else "speaker_1" class LiveSpeechSession: def __init__(self, config: EngineConfig): from faster_whisper import WhisperModel self.config = config device = config.compute if device == "auto": device = "cpu" compute_type = "float16" if device == "cuda" else "int8" try: self.model = WhisperModel(config.model, device=device, compute_type=compute_type) except Exception: # Useful fallback when CUDA was selected but is not available. self.model = WhisperModel(config.model, device="cpu", compute_type="int8") self.diarizer = DiartWorker(config.diarization, config.sample_rate) self.samples = np.empty(0, dtype=np.float32) self.lock = threading.Lock() self.emitted_until = 0.0 self.last_process_samples = 0 self.closed = False def push_pcm(self, raw: bytes) -> None: pcm = np.frombuffer(raw, dtype=" max_samples: dropped = self.samples.size - max_samples self.samples = self.samples[dropped:] self.emitted_until = max(0.0, self.emitted_until - dropped / self.config.sample_rate) def total_seconds(self) -> float: with self.lock: return self.samples.size / self.config.sample_rate def should_process(self) -> bool: with self.lock: new_samples = self.samples.size - self.last_process_samples return new_samples >= int(self.config.sample_rate * 1.4) def transcribe_once(self) -> tuple[list[dict], Optional[str]]: with self.lock: total = self.samples.size if total < self.config.sample_rate: return [], None window_size = min(total, self.config.sample_rate * 10) audio = self.samples[-window_size:].copy() self.last_process_samples = total window_start = max(0.0, total / self.config.sample_rate - len(audio) / self.config.sample_rate) segments_iter, _ = self.model.transcribe( audio, language=None if self.config.language in ("", "auto") else self.config.language, beam_size=1, best_of=1, vad_filter=True, word_timestamps=True, condition_on_previous_text=False, temperature=0.0, ) segments = list(segments_iter) results: list[dict] = [] partial: Optional[str] = None safe_cutoff = total / self.config.sample_rate - 0.65 for seg in segments: text = seg.text.strip() if not text: continue abs_start = window_start + float(seg.start) abs_end = window_start + float(seg.end) if abs_end <= self.emitted_until + 0.08: continue if abs_end > safe_cutoff: partial = text continue speaker_id = self.diarizer.resolve(abs_start, abs_end) results.append({ "id": str(uuid.uuid4()), "speakerId": speaker_id, "start": abs_start, "end": abs_end, "text": text, "final": True, }) self.emitted_until = max(self.emitted_until, abs_end) return results, partial def close(self) -> None: self.closed = True self.diarizer.close() class MockSession: def __init__(self, config: EngineConfig): self.config = config self.sample_count = 0 self.index = 0 self.lines = [ ("speaker_1", "We should finalize the upload flow before Friday."), ("speaker_2", "I can take the local recording and video playback work."), ("speaker_1", "Great. Keep the backend in processing state for this demo."), ("speaker_3", "The overlay should keep updating while another application is open."), ] def push_pcm(self, raw: bytes) -> Optional[dict]: self.sample_count += len(raw) // 2 threshold = self.config.sample_rate * 2.4 if self.sample_count < threshold: return None self.sample_count = 0 speaker, text = self.lines[self.index % len(self.lines)] self.index += 1 end = self.index * 2.4 return { "id": str(uuid.uuid4()), "speakerId": speaker, "start": end - 2.1, "end": end, "text": text, "final": True, } def close(self) -> None: pass app = FastAPI(title="MeetVault Local Speech Engine") app.add_middleware(CORSMiddleware, allow_origins=["*"], allow_methods=["*"], allow_headers=["*"]) MOCK_ONLY = False @app.get("/health") async def health(): return {"ok": True, "mode": "mock" if MOCK_ONLY else "real-capable", "sample_rate": SAMPLE_RATE} @app.websocket("/ws/live") async def live(websocket: WebSocket): await websocket.accept() session = None process_task = None try: first = await websocket.receive_text() payload = json.loads(first) config = EngineConfig( model=str(payload.get("model", "small")), compute=str(payload.get("compute", "auto")), language=str(payload.get("language", "auto")), diarization=bool(payload.get("diarization", True)), sample_rate=int(payload.get("sample_rate", SAMPLE_RATE)), ) if config.sample_rate != SAMPLE_RATE: await websocket.send_json({"type": "error", "message": "MeetVault speech engine currently requires 16 kHz mono PCM."}) return if MOCK_ONLY: session = MockSession(config) await websocket.send_json({"type": "ready", "engine": "mock-whisper", "diarization": "mock-diart"}) else: try: session = await asyncio.to_thread(LiveSpeechSession, config) except Exception as exc: await websocket.send_json({ "type": "error", "message": ( f"Could not initialize Whisper/diart: {exc}. " "Use --mock for UI testing. Real diarization may require Hugging Face authentication and accepted pyannote model terms." ), }) return diarization_state = "diart ready" if session.diarizer.enabled else f"diart unavailable: {session.diarizer.error}" await websocket.send_json({"type": "ready", "engine": f"Whisper {config.model}", "diarization": diarization_state}) async def processor(): while not session.closed: await asyncio.sleep(0.35) if not session.should_process(): continue try: finals, partial = await asyncio.to_thread(session.transcribe_once) for segment in finals: await websocket.send_json({"type": "final", "segment": segment}) if partial: await websocket.send_json({"type": "partial", "text": partial}) except Exception as exc: await websocket.send_json({"type": "error", "message": f"Live transcription error: {exc}"}) await asyncio.sleep(1) process_task = asyncio.create_task(processor()) while True: message = await websocket.receive() if message.get("bytes") is not None: raw = message["bytes"] if isinstance(session, MockSession): segment = session.push_pcm(raw) if segment: await websocket.send_json({"type": "final", "segment": segment}) else: session.push_pcm(raw) elif message.get("text"): command = json.loads(message["text"]) if command.get("type") == "stop": break except WebSocketDisconnect: pass except Exception as exc: try: await websocket.send_json({"type": "error", "message": str(exc)}) except Exception: pass finally: if session: session.close() if process_task: process_task.cancel() def main() -> None: parser = argparse.ArgumentParser() parser.add_argument("--host", default="127.0.0.1") parser.add_argument("--port", type=int, default=8765) parser.add_argument("--mock", action="store_true", help="Run without Whisper/diart models") args = parser.parse_args() global MOCK_ONLY MOCK_ONLY = args.mock uvicorn.run(app, host=args.host, port=args.port, log_level="info") if __name__ == "__main__": main()