#!/usr/bin/env python3 """Small OpenAI-compatible Whisper service for the dedicated Jetson.""" from __future__ import annotations import cgi import base64 import hashlib import json import os import re import socket import struct import tempfile import threading import time import wave from http.server import BaseHTTPRequestHandler, ThreadingHTTPServer from pathlib import Path import torch import whisper HOST = os.getenv("HERMES_STT_HOST", "0.0.0.0") PORT = int(os.getenv("HERMES_STT_PORT", "9000")) MODEL_NAME = os.getenv("HERMES_STT_MODEL", "large-v3-turbo") ROLLING_MODEL_NAME = os.getenv("HERMES_STT_ROLLING_MODEL", "tiny") CACHE_DIR = Path(os.getenv("HERMES_STT_CACHE", "/cache/whisper")) # A short, neutral vocabulary/style prime for the accurate final decode only. # Whisper treats ``initial_prompt`` as soft bias, not a transcript, so a handful # of well-known accented place names across EN/ES/RU teaches correct # capitalization and diacritics on the FIRST pass (no "say it twice") without # memorizing any single user's phrases. It is never applied to the fast rolling # previews, which must stay clean and cheap. DEFAULT_INITIAL_PROMPT = ( "Proper nouns and names keep their capitalization and accents, for " "example Córdoba, Cancún, Málaga, Moscú, Москва, New York, and names " "like Amy, Claude, and Hermes. Technical and government acronyms are " "written as uppercase letters, for example CUI, FOUO, DoD, NIST, CMMC, " "FIPS, RMF, POA&M, ATO, and SBU." ) INITIAL_PROMPT = os.getenv("HERMES_STT_INITIAL_PROMPT", DEFAULT_INITIAL_PROMPT).strip() def _parse_temperature_fallback(raw: str) -> tuple[float, ...]: """Parse a comma-separated Whisper temperature fallback ladder.""" values: list[float] = [] for token in raw.split(","): token = token.strip() if not token: continue try: values.append(float(token)) except ValueError: continue return tuple(values) if values else (0.0,) # The accurate committed transcript retries a low-confidence greedy pass at a # few higher temperatures; the fast rolling preview never does. FINAL_BEAM_SIZE = max(1, int(os.getenv("HERMES_STT_FINAL_BEAM_SIZE", "2"))) FINAL_TEMPERATURE = _parse_temperature_fallback( os.getenv("HERMES_STT_FINAL_TEMPERATURE", "0.0,0.2,0.4") ) MAX_AUDIO_BYTES = 30 * 1024 * 1024 STREAM_SAMPLE_RATE = 16_000 STREAM_SAMPLE_WIDTH = 2 STREAM_CHANNELS = 1 MAX_STREAM_SECONDS = int(os.getenv("HERMES_STT_STREAM_MAX_SECONDS", "90")) MAX_STREAM_AUDIO_BYTES = min( MAX_AUDIO_BYTES, STREAM_SAMPLE_RATE * STREAM_SAMPLE_WIDTH * MAX_STREAM_SECONDS, ) MAX_WS_MESSAGE_BYTES = 256 * 1024 STREAM_IDLE_SECONDS = int(os.getenv("HERMES_STT_STREAM_IDLE_SECONDS", "120")) VAD_END_SILENCE_MS = int(os.getenv("HERMES_STT_VAD_END_SILENCE_MS", "650")) VAD_TAIL_MS = int(os.getenv("HERMES_STT_VAD_TAIL_MS", "220")) ROLLING_MIN_AUDIO_MS = max( 300, min(2_000, int(os.getenv("HERMES_STT_ROLLING_MIN_AUDIO_MS", "1800"))), ) ROLLING_INTERVAL_MS = max( 400, min(5_000, int(os.getenv("HERMES_STT_ROLLING_INTERVAL_MS", "750"))), ) ROLLING_WINDOW_MS = max( 3_000, min(30_000, int(os.getenv("HERMES_STT_ROLLING_WINDOW_MS", "12000"))), ) ROLLING_MAX_RESULT_LAG_MS = max( 500, min(5_000, int(os.getenv("HERMES_STT_ROLLING_MAX_LAG_MS", "2000"))), ) _WS_GUID = "258EAFA5-E914-47DA-95CA-C5AB0DC85B11" class InferenceGate: """Give authoritative transcripts priority over disposable partial work.""" def __init__(self) -> None: self._condition = threading.Condition() self._busy = False self._final_waiters = 0 self._last_final_wait_ms = 0.0 self._last_final_decode_ms = 0.0 self._last_background_decode_ms = 0.0 self._last_rolling_decode_ms = 0.0 self._last_speculative_decode_ms = 0.0 self._background_skips = 0 def run_final(self, callback): """Wait fairly enough that newly arriving partials cannot steal the GPU.""" wait_started = time.monotonic() with self._condition: self._final_waiters += 1 try: while self._busy: self._condition.wait() self._busy = True finally: self._final_waiters -= 1 self._last_final_wait_ms = (time.monotonic() - wait_started) * 1000 decode_started = time.monotonic() try: return callback() finally: with self._condition: self._last_final_decode_ms = (time.monotonic() - decode_started) * 1000 self._busy = False self._condition.notify_all() def try_background(self, callback, *, kind: str): """Run one useful predecode only when no final is active or waiting.""" if kind not in {"rolling", "speculative"}: raise ValueError("unknown background inference kind") with self._condition: if self._busy or self._final_waiters: self._background_skips += 1 return None self._busy = True decode_started = time.monotonic() try: return callback() finally: with self._condition: elapsed_ms = (time.monotonic() - decode_started) * 1000 self._last_background_decode_ms = elapsed_ms if kind == "rolling": self._last_rolling_decode_ms = elapsed_ms else: self._last_speculative_decode_ms = elapsed_ms self._busy = False self._condition.notify_all() def snapshot(self) -> dict: """Expose bounded timing telemetry for live latency verification.""" with self._condition: return { "busy": self._busy, "final_waiters": self._final_waiters, "last_final_wait_ms": round(self._last_final_wait_ms, 1), "last_final_decode_ms": round(self._last_final_decode_ms, 1), "last_background_decode_ms": round(self._last_background_decode_ms, 1), "last_rolling_decode_ms": round(self._last_rolling_decode_ms, 1), "last_speculative_decode_ms": round( self._last_speculative_decode_ms, 1 ), "background_skips": self._background_skips, } INFERENCE_GATE = InferenceGate() class StreamTelemetry: """Retain bounded evidence about the most recently committed voice turn.""" def __init__(self) -> None: self._lock = threading.Lock() self._speculative_cache_hits = 0 self._direct_final_decodes = 0 self._last_commit_cache_hit = False self._last_commit_pending_wait_ms = 0.0 self._last_commit_total_ms = 0.0 self._last_commit_audio_ms = 0.0 def record_commit( self, *, cache_hit: bool, pending_wait_ms: float, total_ms: float, audio_bytes: int, ) -> None: """Record whether commit reused its exact speculative EOS snapshot.""" with self._lock: if cache_hit: self._speculative_cache_hits += 1 else: self._direct_final_decodes += 1 self._last_commit_cache_hit = cache_hit self._last_commit_pending_wait_ms = pending_wait_ms self._last_commit_total_ms = total_ms self._last_commit_audio_ms = ( audio_bytes * 1000 / (STREAM_SAMPLE_RATE * STREAM_SAMPLE_WIDTH) ) def snapshot(self) -> dict: """Return the small aggregate used by health checks and release tests.""" with self._lock: return { "speculative_cache_hits": self._speculative_cache_hits, "direct_final_decodes": self._direct_final_decodes, "last_commit_cache_hit": self._last_commit_cache_hit, "last_commit_pending_wait_ms": round( self._last_commit_pending_wait_ms, 1 ), "last_commit_total_ms": round(self._last_commit_total_ms, 1), "last_commit_audio_ms": round(self._last_commit_audio_ms, 1), } STREAM_TELEMETRY = StreamTelemetry() STARTUP_WARMUP_MS: float | None = None STARTUP_ROLLING_WARMUP_MS: float | None = None def _repetitive_token(token: str) -> bool: """Identify long periodic Whisper hallucinations caused by steady noise.""" letters = "".join(re.findall(r"[a-z]+", token.lower())) if len(letters) < 10: return False for period in range(1, 5): pattern = letters[:period] matches = sum( character == pattern[index % period] for index, character in enumerate(letters) ) if matches / len(letters) >= 0.86: return True return ( max(letters.count(character) for character in set(letters)) / len(letters) >= 0.78 ) def _clean_transcript(result: dict) -> str: """Drop noise-only segments and repetitive tokens while retaining speech.""" segments = result.get("segments") if not isinstance(segments, list): segments = [{"text": result.get("text") or ""}] kept: list[str] = [] for segment in segments: if not isinstance(segment, dict): continue text = str(segment.get("text") or "").strip() if not text: continue no_speech = float(segment.get("no_speech_prob") or 0.0) average_logprob = float(segment.get("avg_logprob") or 0.0) if no_speech >= 0.55 and average_logprob <= -0.55: continue words = [word for word in text.split() if not _repetitive_token(word)] if words: kept.append(" ".join(words)) return " ".join(kept).strip() def _detected_language(result: object) -> str: """Return the bare ISO-639 code Whisper decoded with, or nothing at all. ``whisper.transcribe`` reports the language it auto-detected (or the one it was told to use) as a plain lowercase token such as ``en``/``ru``/``yue``. Anything that is not that exact shape is dropped rather than guessed at, so a surprising model result can never become a downstream voice selector. """ if not isinstance(result, dict): return "" value = result.get("language") if not isinstance(value, str): return "" code = value.strip().lower() if not 2 <= len(code) <= 3 or not code.isascii() or not code.isalpha(): return "" return code def _transcription_payload(result: dict) -> dict: """Build the transcription contract: text plus the model's own language.""" return { "text": _clean_transcript(result), "model": MODEL_NAME, "language": _detected_language(result), } # Shared hallucination guards; the accuracy/latency split lives in the two # wrappers below, not here. _COMMON_DECODE_PARAMS = { "task": "transcribe", "condition_on_previous_text": False, "compression_ratio_threshold": 2.0, "logprob_threshold": -0.8, "no_speech_threshold": 0.5, "verbose": False, } def _decode(model: object, path: str, language: str, *, params: dict) -> dict: """Decode one path with an explicit quality profile (final vs. rolling).""" result = model.transcribe( path, language=None if language in {"", "auto"} else language, fp16=torch.cuda.is_available(), **_COMMON_DECODE_PARAMS, **params, ) return _transcription_payload(result) def _final_decode_params() -> dict: """Accurate profile: beam search, a small temperature fallback, and priming. Used for the authoritative POST decode, the committed streaming turn, and the speculative EOS predecode that the commit reuses verbatim (so the frozen cache stays byte-identical to a direct final decode). """ params = {"beam_size": FINAL_BEAM_SIZE, "temperature": FINAL_TEMPERATURE} if INITIAL_PROMPT: params["initial_prompt"] = INITIAL_PROMPT return params # Fast, clean profile for disposable rolling previews: greedy, no fallback ladder # and deliberately no ``initial_prompt`` so live partials stay cheap and snappy. _ROLLING_DECODE_PARAMS = {"beam_size": 1, "temperature": 0} def _decode_final(model: object, path: str, language: str) -> dict: """Decode with the accurate, primed profile for the committed transcript.""" return _decode(model, path, language, params=_final_decode_params()) def _decode_rolling(model: object, path: str, language: str) -> dict: """Decode with the fast greedy profile for a disposable live preview.""" return _decode(model, path, language, params=_ROLLING_DECODE_PARAMS) def _transcribe_path(model: object, path: str, language: str) -> dict: """Run the accurate full-utterance decode behind the authoritative gate.""" return INFERENCE_GATE.run_final(lambda: _decode_final(model, path, language)) def _write_pcm_wav(pcm: bytes) -> str: """Write an aligned browser PCM snapshot to a temporary Whisper input.""" with tempfile.NamedTemporaryFile( prefix="atlas-stt-stream-", suffix=".wav", delete=False ) as temp: temp_path = temp.name with wave.open(temp_path, "wb") as wav_file: wav_file.setnchannels(STREAM_CHANNELS) wav_file.setsampwidth(STREAM_SAMPLE_WIDTH) wav_file.setframerate(STREAM_SAMPLE_RATE) wav_file.writeframes(pcm) return temp_path def _transcribe_pcm(model: object, pcm: bytes, language: str) -> dict: """Decode one complete 16-kHz mono PCM snapshot through the canonical path.""" temp_path = "" try: temp_path = _write_pcm_wav(pcm) return _transcribe_path(model, temp_path, language) finally: if temp_path: try: os.unlink(temp_path) except OSError: pass def _warm_model(model: object, decode=_decode_rolling) -> float: """Pay Whisper's lazy CUDA initialization cost before readiness. Each model is warmed with the exact profile it will serve so the first real request never inherits beam-search or fallback CUDA kernel compilation. """ temp_path = _write_pcm_wav(bytes(_pcm_bytes_for_ms(1_000))) started = time.monotonic() try: decode(model, temp_path, "auto") torch.cuda.synchronize() finally: try: os.unlink(temp_path) except OSError: pass return (time.monotonic() - started) * 1000 def _transcribe_pcm_rolling(model: object, pcm: bytes, language: str) -> dict | None: """Decode a rolling window only when the single GPU worker is immediately free. Partial work is disposable. It must never form an inference backlog ahead of a final utterance or another user's request, so a busy model means this speech epoch is skipped rather than retried. """ temp_path = "" try: temp_path = _write_pcm_wav(pcm) return INFERENCE_GATE.try_background( lambda: _decode_rolling(model, temp_path, language), kind="rolling", ) finally: if temp_path: try: os.unlink(temp_path) except OSError: pass def _transcribe_pcm_speculative( model: object, pcm: bytes, language: str ) -> dict | None: """Predecode one EOS snapshot only when the model is currently idle. This runs on the accurate final model with the final profile, so a commit that reuses the frozen EOS cache gets the exact same high-accuracy result a direct final decode would have produced. """ temp_path = "" try: temp_path = _write_pcm_wav(pcm) return INFERENCE_GATE.try_background( lambda: _decode_final(model, temp_path, language), kind="speculative", ) finally: if temp_path: try: os.unlink(temp_path) except OSError: pass def _pcm_bytes_for_ms(milliseconds: int) -> int: """Convert milliseconds to an aligned byte count for the stream format.""" return ( STREAM_SAMPLE_RATE * STREAM_SAMPLE_WIDTH * milliseconds // 1000 // STREAM_SAMPLE_WIDTH * STREAM_SAMPLE_WIDTH ) def _token_identity(token: str) -> str: """Normalize punctuation/case drift before comparing consecutive decodes.""" return re.sub(r"[^\w']+", "", token.casefold(), flags=re.UNICODE) def _stable_token_prefix(previous: str, current: str) -> str: """Return only tokens repeated in the same order by consecutive decodes.""" old_tokens = previous.split() new_tokens = current.split() agreed = 0 for old, new in zip(old_tokens, new_tokens): if not _token_identity(old) or _token_identity(old) != _token_identity(new): break agreed += 1 return " ".join(new_tokens[:agreed]) class WebSocketError(Exception): """Protocol error carrying an RFC 6455 close code safe to expose.""" def __init__(self, message: str, close_code: int = 1002) -> None: super().__init__(message) self.close_code = close_code class WebSocketConnection: """Minimal server-side RFC 6455 framing for the private streaming endpoint.""" def __init__(self, handler: BaseHTTPRequestHandler) -> None: self.handler = handler self._send_lock = threading.Lock() self._closed = False @staticmethod def accept_value(key: str) -> str: """Return the RFC 6455 handshake response for a client nonce.""" digest = hashlib.sha1((key + _WS_GUID).encode("ascii")).digest() return base64.b64encode(digest).decode("ascii") def _read_exact(self, size: int) -> bytes: data = bytearray() while len(data) < size: chunk = self.handler.rfile.read(size - len(data)) if not chunk: raise EOFError("websocket peer disconnected") data.extend(chunk) return bytes(data) def receive(self) -> tuple[int, bytes]: """Read and reassemble one client message, answering ping inline.""" opcode = 0 payload = bytearray() while True: header = self._read_exact(2) first, second = header final = bool(first & 0x80) reserved = first & 0x70 frame_opcode = first & 0x0F masked = bool(second & 0x80) length = second & 0x7F if reserved or not masked: raise WebSocketError("invalid websocket frame") if length == 126: length = struct.unpack("!H", self._read_exact(2))[0] elif length == 127: length = struct.unpack("!Q", self._read_exact(8))[0] if frame_opcode >= 0x8 and (not final or length > 125): raise WebSocketError("invalid websocket control frame") if ( length > MAX_WS_MESSAGE_BYTES or len(payload) + length > MAX_WS_MESSAGE_BYTES ): raise WebSocketError("websocket message is too large", 1009) mask = self._read_exact(4) frame = bytearray(self._read_exact(length)) for index in range(length): frame[index] ^= mask[index % 4] if frame_opcode == 0x8: raise EOFError("websocket peer closed") if frame_opcode == 0x9: self.send(bytes(frame), opcode=0xA) continue if frame_opcode == 0xA: continue if frame_opcode in {0x1, 0x2}: if opcode: raise WebSocketError("nested websocket message") opcode = frame_opcode elif frame_opcode != 0x0 or not opcode: raise WebSocketError("unexpected websocket continuation") payload.extend(frame) if final: return opcode, bytes(payload) def send(self, payload: bytes, opcode: int) -> None: """Send one unmasked server frame while serializing background writers.""" with self._send_lock: if self._closed: return length = len(payload) header = bytearray([0x80 | opcode]) if length < 126: header.append(length) elif length <= 0xFFFF: header.extend((126,)) header.extend(struct.pack("!H", length)) else: header.extend((127,)) header.extend(struct.pack("!Q", length)) self.handler.wfile.write(bytes(header) + payload) self.handler.wfile.flush() def send_json(self, payload: dict) -> None: """Send a compact UTF-8 JSON message.""" self.send(json.dumps(payload, separators=(",", ":")).encode("utf-8"), 0x1) def close(self, code: int = 1000, reason: str = "") -> None: """Close once; callers may safely race disconnect and worker completion.""" reason_bytes = reason.encode("utf-8")[:123] with self._send_lock: if self._closed: return try: payload = struct.pack("!H", code) + reason_bytes self.handler.wfile.write(bytes([0x88, len(payload)]) + payload) self.handler.wfile.flush() except (BrokenPipeError, ConnectionError, OSError): pass self._closed = True def _pcm_rms(pcm: bytes) -> float: """Return normalized RMS for aligned signed-16-bit little-endian PCM.""" sample_count = len(pcm) // 2 if not sample_count: return 0.0 samples = struct.unpack(f"<{sample_count}h", pcm[: sample_count * 2]) square_mean = sum(sample * sample for sample in samples) / sample_count return (square_mean**0.5) / 32768.0 class StreamingTranscription: """Own one turn's bounded PCM, VAD state and speculative Whisper work.""" def __init__( self, connection: WebSocketConnection, model: object, rolling_model: object | None = None, ) -> None: self.connection = connection self.model = model self.rolling_model = rolling_model if rolling_model is not None else model self.turn_id = "" self.language = "auto" self._pcm = bytearray() self._lock = threading.Condition() self._started = False self._closed = False self._committing = False self._heard_speech = False self._at_eos = False self._silence_bytes = 0 self._last_speech_byte = 0 self._noise_floor = 0.004 self._client_active = False self._epoch = 0 self._pending_key = "" self._cached_key = "" self._cached_payload: dict | None = None self._eos_pcm = b"" self._eos_key = "" self._rolling_inflight = False self._rolling_last_started_at = 0.0 self._rolling_speech_bytes = 0 self._rolling_last_speech_bytes = 0 self._rolling_previous = "" self._rolling_previous_window_start = 0 self._rolling_revision = 0 self._rolling_attempted = False def _response(self, message_type: str, **values: object) -> dict: return {"type": message_type, "turn_id": self.turn_id, **values} def error(self, message: str, *, fatal: bool = False) -> None: """Report a turn-scoped error without leaking implementation details.""" self.connection.send_json(self._response("error", error=message, fatal=fatal)) def start(self, message: dict) -> None: """Validate and initialize the only accepted browser audio contract.""" turn_id = str(message.get("turn_id") or "") language = str(message.get("language") or "auto").strip().lower() if self._started: raise WebSocketError("stream is already started", 1008) if not re.fullmatch(r"[A-Za-z0-9_.:-]{1,128}", turn_id): raise WebSocketError("invalid turn_id", 1008) if message.get("format") != "pcm_s16le": raise WebSocketError("pcm_s16le audio is required", 1003) try: sample_rate = int(message.get("sample_rate") or 0) except (TypeError, ValueError) as exc: raise WebSocketError("invalid sample rate", 1003) from exc if sample_rate != STREAM_SAMPLE_RATE: raise WebSocketError("16000 Hz audio is required", 1003) if language != "auto" and not re.fullmatch(r"[a-z]{2,3}", language): raise WebSocketError("invalid language", 1008) self.turn_id = turn_id self.language = language self._started = True def _snapshot_locked(self) -> tuple[bytes, str]: """Keep the full utterance and discard only bounded post-speech silence. Browser ``resume`` is authoritative: once sent, every following PCM frame is preserved until the next speculative EOS marker, even when a quiet syllable falls below the server's conservative energy threshold. """ if self._at_eos and self._eos_pcm and self._eos_key: return self._eos_pcm, self._eos_key cutoff = len(self._pcm) if self._heard_speech and self._last_speech_byte: tail_bytes = STREAM_SAMPLE_RATE * STREAM_SAMPLE_WIDTH * VAD_TAIL_MS // 1000 cutoff = min(cutoff, self._last_speech_byte + tail_bytes) pcm = bytes(self._pcm[:cutoff]) key = f"{self._epoch}:{hashlib.sha256(pcm).hexdigest()}" return pcm, key def _clear_eos_snapshot_locked(self) -> None: """Invalidate frozen EOS audio when authoritative speech resumes.""" self._eos_pcm = b"" self._eos_key = "" def _rolling_snapshot_locked(self) -> tuple[bytes, int]: """Return a bounded tail window and its absolute start in milliseconds.""" window_bytes = _pcm_bytes_for_ms(ROLLING_WINDOW_MS) start = max(0, len(self._pcm) - window_bytes) start -= start % STREAM_SAMPLE_WIDTH pcm = bytes(self._pcm[start:]) start_ms = start * 1000 // (STREAM_SAMPLE_RATE * STREAM_SAMPLE_WIDTH) return pcm, start_ms def _reset_rolling_history_locked(self, *, rearm: bool) -> None: """Discard provisional agreement after the client's speech epoch changes.""" self._rolling_previous = "" self._rolling_previous_window_start = 0 # Only speech in the new epoch may satisfy the rolling threshold. self._rolling_speech_bytes = 0 self._rolling_last_speech_bytes = 0 self._rolling_last_started_at = 0.0 if rearm: self._rolling_attempted = False def _begin_rolling_locked(self, now: float) -> tuple[bytes, int, int, int] | None: """Reserve one due rolling decode without ever queuing a second one.""" total_bytes = len(self._pcm) minimum_bytes = _pcm_bytes_for_ms(ROLLING_MIN_AUDIO_MS) interval_bytes = _pcm_bytes_for_ms(ROLLING_INTERVAL_MS) required_new_bytes = ( interval_bytes if self._rolling_last_speech_bytes else minimum_bytes ) if ( self._rolling_inflight or self._rolling_attempted or self._rolling_speech_bytes < minimum_bytes or self._rolling_speech_bytes - self._rolling_last_speech_bytes < required_new_bytes or ( self._rolling_last_started_at and (now - self._rolling_last_started_at) * 1000 < ROLLING_INTERVAL_MS ) ): return None pcm, window_start_ms = self._rolling_snapshot_locked() # Reserve the epoch before spawning. A busy gate, decode error or stale # result must not create a retry loop that can collide with the final. self._rolling_attempted = True self._rolling_inflight = True self._rolling_last_started_at = now self._rolling_last_speech_bytes = self._rolling_speech_bytes return pcm, window_start_ms, total_bytes, self._epoch def _start_rolling( self, pcm: bytes, window_start_ms: int, audio_bytes: int, epoch: int, ) -> None: """Run one disposable partial decode and publish only a current result.""" def worker() -> None: payload: dict | None = None try: payload = _transcribe_pcm_rolling( self.rolling_model, pcm, self.language, ) except Exception as exc: # Partial inference is best-effort. The canonical commit path # remains authoritative and reports its own actionable errors. print(f"[stt] rolling transcription skipped: {exc}", flush=True) message: dict | None = None with self._lock: self._rolling_inflight = False lag_bytes = max(0, len(self._pcm) - audio_bytes) max_lag_bytes = _pcm_bytes_for_ms(ROLLING_MAX_RESULT_LAG_MS) current = ( payload is not None and not self._closed and not self._committing and not self._at_eos and self._epoch == epoch and lag_bytes <= max_lag_bytes ) if current: transcript = str(payload.get("text") or "") stable = "" if self._rolling_previous_window_start == window_start_ms: stable = _stable_token_prefix( self._rolling_previous, transcript, ) self._rolling_previous = transcript self._rolling_previous_window_start = window_start_ms self._rolling_revision += 1 message = self._response( "partial", transcript=transcript, stable_transcript=stable, language=payload.get("language") or "", speculative=True, rolling=True, revision=self._rolling_revision, epoch=epoch, window_start_ms=window_start_ms, ) self._lock.notify_all() if message is not None: try: self.connection.send_json(message) except (BrokenPipeError, ConnectionError, OSError): self.cancel() threading.Thread( target=worker, name=f"stt-rolling-{self.turn_id}", daemon=True, ).start() def append(self, pcm: bytes) -> None: """Append PCM while overlapping rolling understanding with active speech.""" if not self._started: raise WebSocketError("start must precede audio", 1008) if not pcm or len(pcm) % STREAM_SAMPLE_WIDTH: raise WebSocketError("unaligned PCM audio", 1003) auto_speculate = False rolling: tuple[bytes, int, int, int] | None = None with self._lock: if self._closed or self._committing: raise WebSocketError("turn is no longer accepting audio", 1008) if len(self._pcm) + len(pcm) > MAX_STREAM_AUDIO_BYTES: raise WebSocketError("stream audio is too long", 1009) self._pcm.extend(pcm) rms = _pcm_rms(pcm) threshold = max(0.012, self._noise_floor * 2.5 + 0.004) speech = self._client_active or rms >= threshold if not self._heard_speech and not speech: self._noise_floor = self._noise_floor * 0.96 + rms * 0.04 if speech: if self._at_eos: self._epoch += 1 self._cached_key = "" self._cached_payload = None self._clear_eos_snapshot_locked() self._reset_rolling_history_locked(rearm=True) self._heard_speech = True self._at_eos = False self._silence_bytes = 0 self._last_speech_byte = len(self._pcm) self._rolling_speech_bytes += len(pcm) rolling = self._begin_rolling_locked(time.monotonic()) elif self._heard_speech: self._silence_bytes += len(pcm) required = ( STREAM_SAMPLE_RATE * STREAM_SAMPLE_WIDTH * VAD_END_SILENCE_MS // 1000 ) if self._silence_bytes >= required and not self._at_eos: self._at_eos = True self._reset_rolling_history_locked(rearm=False) auto_speculate = True if auto_speculate: self.speculate() elif rolling is not None: self._start_rolling(*rolling) def resume(self) -> None: """Invalidate an EOS snapshot when client-side VAD hears resumed speech.""" with self._lock: if not self._started or self._closed: return self._at_eos = False self._silence_bytes = 0 self._epoch += 1 self._client_active = True self._cached_key = "" self._cached_payload = None self._clear_eos_snapshot_locked() self._reset_rolling_history_locked(rearm=True) def speculate(self) -> None: """Decode a stable snapshot once; only an unchanged turn may consume it.""" with self._lock: if not self._started or self._closed or self._committing or not self._pcm: return self._client_active = False self._at_eos = True pcm, key = self._snapshot_locked() # The browser keeps forwarding PCM until recorder shutdown. Freeze # the exact EOS snapshot so trailing silence cannot invalidate a # completed Whisper-small speculative result. A resume message or # newly detected speech clears this snapshot before more audio is # accepted into the authoritative turn. self._eos_pcm = pcm self._eos_key = key if not pcm or key in {self._pending_key, self._cached_key}: return self._pending_key = key def worker() -> None: try: payload = _transcribe_pcm_speculative(self.model, pcm, self.language) except Exception as exc: print(f"[stt] speculative transcription failed: {exc}", flush=True) with self._lock: if self._pending_key == key: self._pending_key = "" self._lock.notify_all() if not self._closed and not self._committing: self.error("speculative transcription failed") return if payload is None: with self._lock: if self._pending_key == key: self._pending_key = "" self._lock.notify_all() return should_send = False with self._lock: if self._pending_key == key: self._pending_key = "" current_pcm, current_key = self._snapshot_locked() del current_pcm if not self._closed and current_key == key: self._cached_key = key self._cached_payload = payload should_send = not self._committing self._lock.notify_all() if should_send: try: self.connection.send_json( self._response( "partial", transcript=payload["text"], language=payload["language"], speculative=True, ) ) except (BrokenPipeError, ConnectionError, OSError): self.cancel() threading.Thread( target=worker, name=f"stt-speculate-{self.turn_id}", daemon=True, ).start() def commit(self) -> None: """Emit one final result from the complete stable utterance snapshot.""" commit_started = time.monotonic() pending_wait_ms = 0.0 cache_hit = False with self._lock: if not self._started or self._closed or self._committing: raise WebSocketError("turn cannot be committed", 1008) if not self._pcm: raise WebSocketError("turn has no audio", 1008) self._committing = True pcm, key = self._snapshot_locked() deadline = time.monotonic() + STREAM_IDLE_SECONDS pending_wait_started = time.monotonic() while self._pending_key == key and not self._cached_payload: remaining = deadline - time.monotonic() if remaining <= 0: break self._lock.wait(timeout=remaining) pending_wait_ms = (time.monotonic() - pending_wait_started) * 1000 payload = self._cached_payload if self._cached_key == key else None cache_hit = payload is not None try: if payload is None: payload = _transcribe_pcm(self.model, pcm, self.language) self.connection.send_json( self._response( "final", transcript=payload["text"], language=payload["language"], model=payload["model"], ) ) STREAM_TELEMETRY.record_commit( cache_hit=cache_hit, pending_wait_ms=pending_wait_ms, total_ms=(time.monotonic() - commit_started) * 1000, audio_bytes=len(pcm), ) except Exception as exc: print(f"[stt] streaming transcription failed: {exc}", flush=True) self.error("transcription failed", fatal=True) finally: with self._lock: self._closed = True self._lock.notify_all() def cancel(self) -> None: """Invalidate queued work and prevent background writers after disconnect.""" with self._lock: self._closed = True self._cached_payload = None self._cached_key = "" self._clear_eos_snapshot_locked() self._lock.notify_all() def _json(handler: BaseHTTPRequestHandler, status: int, payload: dict) -> None: body = json.dumps(payload).encode("utf-8") handler.send_response(status) handler.send_header("Content-Type", "application/json") handler.send_header("Content-Length", str(len(body))) handler.send_header("Cache-Control", "no-store") handler.end_headers() handler.wfile.write(body) class SpeechHandler(BaseHTTPRequestHandler): """Serve health and transcription without exposing a general runtime.""" server_version = "AtlasWhisper/1" def log_message(self, message: str, *args: object) -> None: print(f"[stt] {self.address_string()} {message % args}", flush=True) def do_GET(self) -> None: path = self.path.split("?", 1)[0] if path == "/v1/audio/transcriptions/stream": self._serve_stream() return if path != "/health": _json(self, 404, {"error": "not found"}) return _json( self, 200, { "ok": True, "model": MODEL_NAME, "rolling_model": ROLLING_MODEL_NAME, "device": "cuda" if torch.cuda.is_available() else "cpu", "inference": INFERENCE_GATE.snapshot(), "stream_commits": STREAM_TELEMETRY.snapshot(), "startup": { "warmed": ( STARTUP_WARMUP_MS is not None and STARTUP_ROLLING_WARMUP_MS is not None ), "warmup_ms": round(STARTUP_WARMUP_MS or 0.0, 1), "rolling_warmup_ms": round( STARTUP_ROLLING_WARMUP_MS or 0.0, 1, ), }, "streaming": { "enabled": True, "path": "/v1/audio/transcriptions/stream", "format": "pcm_s16le", "sample_rate": STREAM_SAMPLE_RATE, "max_seconds": MAX_STREAM_SECONDS, "server_vad": True, "speculative": True, "rolling": { "enabled": True, "minimum_audio_ms": ROLLING_MIN_AUDIO_MS, "interval_ms": ROLLING_INTERVAL_MS, "window_ms": ROLLING_WINDOW_MS, "stable_prefix": True, }, }, }, ) def _serve_stream(self) -> None: """Upgrade one request and process a single, isolated voice turn.""" upgrade = self.headers.get("Upgrade", "").strip().lower() connection_tokens = { token.strip().lower() for token in self.headers.get("Connection", "").split(",") } key = self.headers.get("Sec-WebSocket-Key", "").strip() version = self.headers.get("Sec-WebSocket-Version", "").strip() if upgrade != "websocket" or "upgrade" not in connection_tokens or not key: _json(self, 426, {"error": "websocket upgrade is required"}) return if version != "13": self.send_response(426) self.send_header("Sec-WebSocket-Version", "13") self.send_header("Content-Length", "0") self.end_headers() return try: decoded_key = base64.b64decode(key, validate=True) except (ValueError, TypeError): decoded_key = b"" if len(decoded_key) != 16: _json(self, 400, {"error": "invalid websocket key"}) return self.send_response(101, "Switching Protocols") self.send_header("Upgrade", "websocket") self.send_header("Connection", "Upgrade") self.send_header("Sec-WebSocket-Accept", WebSocketConnection.accept_value(key)) self.end_headers() self.wfile.flush() self.close_connection = True self.connection.settimeout(STREAM_IDLE_SECONDS) websocket = WebSocketConnection(self) session = StreamingTranscription( websocket, self.server.model, # type: ignore[attr-defined] self.server.rolling_model, # type: ignore[attr-defined] ) close_code = 1000 close_reason = "" try: while True: opcode, payload = websocket.receive() if opcode == 0x2: session.append(payload) continue try: message = json.loads(payload.decode("utf-8")) except (UnicodeDecodeError, json.JSONDecodeError): raise WebSocketError("control message must be JSON", 1007) if not isinstance(message, dict): raise WebSocketError("control message must be an object", 1007) message_type = str(message.get("type") or "") if message_type == "start": session.start(message) continue if message.get("turn_id") != session.turn_id: raise WebSocketError("turn_id does not match", 1008) if message_type == "speculate": session.speculate() elif message_type == "resume": session.resume() elif message_type == "commit": session.commit() break elif message_type == "cancel": break else: raise WebSocketError("unknown stream control", 1008) except WebSocketError as exc: close_code = exc.close_code close_reason = str(exc) try: session.error(str(exc), fatal=True) except (BrokenPipeError, ConnectionError, OSError): pass except (EOFError, BrokenPipeError, ConnectionError, socket.timeout, OSError): close_reason = "peer disconnected" finally: session.cancel() websocket.close(close_code, close_reason) def do_POST(self) -> None: if self.path != "/v1/audio/transcriptions": _json(self, 404, {"error": "not found"}) return content_length = int(self.headers.get("Content-Length", "0") or "0") if content_length <= 0 or content_length > MAX_AUDIO_BYTES: _json(self, 413, {"error": "audio payload is missing or too large"}) return content_type = self.headers.get("Content-Type", "") if not content_type.lower().startswith("multipart/form-data"): _json(self, 400, {"error": "multipart/form-data is required"}) return form = cgi.FieldStorage( fp=self.rfile, headers=self.headers, environ={ "REQUEST_METHOD": "POST", "CONTENT_TYPE": content_type, "CONTENT_LENGTH": str(content_length), }, ) audio = form["file"] if "file" in form else None if audio is None or not getattr(audio, "file", None): _json(self, 400, {"error": "file is required"}) return suffix = Path(getattr(audio, "filename", "audio.wav") or "audio.wav").suffix suffix = ( suffix if suffix in {".wav", ".webm", ".ogg", ".mp3", ".m4a"} else ".wav" ) language = str(form.getfirst("language", "auto") or "auto").strip().lower() temp_path = "" try: with tempfile.NamedTemporaryFile( prefix="atlas-stt-", suffix=suffix, delete=False ) as temp: temp_path = temp.name while True: chunk = audio.file.read(1024 * 1024) if not chunk: break temp.write(chunk) payload = _transcribe_path( self.server.model, # type: ignore[attr-defined] temp_path, language, ) _json(self, 200, payload) except Exception as exc: print(f"[stt] transcription failed: {exc}", flush=True) _json(self, 500, {"error": "transcription failed"}) finally: if temp_path: try: os.unlink(temp_path) except OSError: pass def main() -> None: """Warm Whisper once, then serve concurrent clients through one GPU lock.""" global STARTUP_ROLLING_WARMUP_MS, STARTUP_WARMUP_MS CACHE_DIR.mkdir(parents=True, exist_ok=True) if not torch.cuda.is_available(): raise RuntimeError("CUDA is required for the Atlas Whisper service") print(f"[stt] loading Whisper {MODEL_NAME} for final CUDA", flush=True) model = whisper.load_model(MODEL_NAME, device="cuda", download_root=str(CACHE_DIR)) print(f"[stt] loading Whisper {ROLLING_MODEL_NAME} for rolling CUDA", flush=True) rolling_model = whisper.load_model( ROLLING_MODEL_NAME, device="cuda", download_root=str(CACHE_DIR), ) STARTUP_ROLLING_WARMUP_MS = _warm_model(rolling_model) print( "[stt] rolling CUDA warm-up completed in " f"{STARTUP_ROLLING_WARMUP_MS:.1f}ms", flush=True, ) # Warm the authoritative model last, with its own beam-search profile, so # the first real final never inherits lazy CUDA work from either model's # startup path. STARTUP_WARMUP_MS = _warm_model(model, _decode_final) print( f"[stt] final CUDA warm-up completed in {STARTUP_WARMUP_MS:.1f}ms", flush=True ) server = ThreadingHTTPServer((HOST, PORT), SpeechHandler) server.model = model # type: ignore[attr-defined] server.rolling_model = rolling_model # type: ignore[attr-defined] print(f"[stt] ready on {HOST}:{PORT}", flush=True) server.serve_forever(poll_interval=0.25) if __name__ == "__main__": main()