375 lines
17 KiB
Python
375 lines
17 KiB
Python
"""Bounded, offline voice recognition. Never discovers accounts, downloads or sends.
|
|
|
|
Call request() repeatedly with an exact account-scoped audio path. It performs
|
|
metadata checks only; a single worker reads, hashes, decodes and recognizes audio.
|
|
Returned state is pending/ready/error. No transcripts are persisted or logged.
|
|
"""
|
|
from __future__ import annotations
|
|
|
|
from collections import OrderedDict
|
|
import hashlib
|
|
import io
|
|
import json
|
|
import math
|
|
import os
|
|
from pathlib import Path
|
|
import queue
|
|
import re
|
|
import stat
|
|
import threading
|
|
import time
|
|
|
|
MODEL_REVISION = "5ff83050b029cf4730c8c89ce0c835b2ea15c6f9"
|
|
MODEL_SHA256 = "d01c3014881c9c6f3133c182f3d2887eb6ca1c789a7538c5c007196857a0a6a9"
|
|
MODEL_FILES = ("config.json", "model.bin", "tokenizer.json", "vocabulary.txt")
|
|
MAX_AUDIO_BYTES = 8 * 1024 * 1024
|
|
MAX_DURATION_SECONDS = 120
|
|
SAMPLE_RATE = 16000
|
|
MAX_TEXT_CHARS = 12000
|
|
|
|
ERROR_MESSAGES = {
|
|
"invalid_identity": "语音消息身份不完整",
|
|
"path_not_allowed": "语音文件不在当前账号允许的缓存目录内",
|
|
"audio_missing": "语音文件尚未缓存到本机",
|
|
"audio_changed": "语音文件正在变化,稍后重试",
|
|
"audio_too_large": "语音文件超过本地识别大小限制",
|
|
"audio_too_long": "语音超过本地识别时长限制",
|
|
"audio_format": "语音格式不可识别或文件不完整",
|
|
"audio_empty": "语音中没有可识别的声音",
|
|
"model_missing": "本地语音模型缺失,请修复安装",
|
|
"model_invalid": "本地语音模型校验失败,请修复安装",
|
|
"dependency_missing": "本地语音识别组件缺失,请修复安装",
|
|
"recognition_failed": "本地语音识别失败,稍后重试",
|
|
"low_confidence": "语音识别结果不可靠,需要人工确认",
|
|
"queue_full": "语音识别队列已满,稍后重试",
|
|
"stopped": "语音识别已停止",
|
|
}
|
|
|
|
|
|
class VoiceTranscriptionError(RuntimeError):
|
|
def __init__(self, code):
|
|
self.code = code if code in ERROR_MESSAGES else "recognition_failed"
|
|
super().__init__(ERROR_MESSAGES[self.code])
|
|
|
|
|
|
def _error(code, *, retry_at=0.0):
|
|
return {"status": "error", "text": "", "error_code": code,
|
|
"message": ERROR_MESSAGES[code], "retry_at": retry_at, "audio_sha256": ""}
|
|
|
|
|
|
def _identity(account, conv_id, server_id):
|
|
if not all(isinstance(value, str) for value in (account, conv_id, server_id)):
|
|
raise VoiceTranscriptionError("invalid_identity")
|
|
if not re.fullmatch(r"[1-9][0-9]{0,23}", account) or not re.fullmatch(r"[1-9][0-9]{0,23}", server_id):
|
|
raise VoiceTranscriptionError("invalid_identity")
|
|
if re.fullmatch(r"M:[1-9][0-9]{0,23}", conv_id):
|
|
if conv_id[2:] != account:
|
|
return account, conv_id, server_id
|
|
elif re.fullmatch(r"S:[1-9][0-9]{0,23}_[1-9][0-9]{0,23}", conv_id):
|
|
parts = conv_id[2:].split("_")
|
|
if parts.count(account) == 1:
|
|
return account, conv_id, server_id
|
|
raise VoiceTranscriptionError("invalid_identity")
|
|
|
|
|
|
def _path_and_signature(audio_path, allowed_roots):
|
|
try:
|
|
path = Path(audio_path)
|
|
if not path.is_absolute() or str(path).startswith(("\\\\", "//")) or any(":" in part for part in path.parts[1:]):
|
|
raise VoiceTranscriptionError("path_not_allowed")
|
|
resolved = path.resolve(strict=True)
|
|
if not allowed_roots or not any(resolved.is_relative_to(root) for root in allowed_roots):
|
|
raise VoiceTranscriptionError("path_not_allowed")
|
|
info = resolved.stat()
|
|
if not stat.S_ISREG(info.st_mode):
|
|
raise VoiceTranscriptionError("path_not_allowed")
|
|
if info.st_size <= 0:
|
|
raise VoiceTranscriptionError("audio_empty")
|
|
if info.st_size > MAX_AUDIO_BYTES:
|
|
raise VoiceTranscriptionError("audio_too_large")
|
|
return resolved, (info.st_dev, info.st_ino, info.st_size, info.st_mtime_ns, info.st_ctime_ns)
|
|
except VoiceTranscriptionError:
|
|
raise
|
|
except (OSError, TypeError, ValueError):
|
|
raise VoiceTranscriptionError("audio_missing") from None
|
|
|
|
|
|
def _read_audio(path, roots, expected=None):
|
|
path, before = _path_and_signature(path, roots)
|
|
if expected is not None and before != expected:
|
|
raise VoiceTranscriptionError("audio_changed")
|
|
try:
|
|
with path.open("rb") as handle:
|
|
info = os.fstat(handle.fileno())
|
|
# Windows Path.stat and fstat disagree on ctime semantics on 3.12.
|
|
if (info.st_dev, info.st_ino, info.st_size, info.st_mtime_ns) != before[:4]:
|
|
raise VoiceTranscriptionError("audio_changed")
|
|
raw = handle.read(MAX_AUDIO_BYTES + 1)
|
|
after_handle = os.fstat(handle.fileno())
|
|
if (after_handle.st_dev, after_handle.st_ino, after_handle.st_size, after_handle.st_mtime_ns) != before[:4]:
|
|
raise VoiceTranscriptionError("audio_changed")
|
|
_, after = _path_and_signature(path, roots)
|
|
if before != after or len(raw) != before[2]:
|
|
raise VoiceTranscriptionError("audio_changed")
|
|
return raw, hashlib.sha256(raw).hexdigest()
|
|
except VoiceTranscriptionError:
|
|
raise
|
|
except OSError:
|
|
raise VoiceTranscriptionError("audio_missing") from None
|
|
|
|
|
|
class _BoundedPCM(io.BytesIO):
|
|
def write(self, data):
|
|
if self.tell() + len(data) > MAX_DURATION_SECONDS * SAMPLE_RATE * 2:
|
|
raise VoiceTranscriptionError("audio_too_long")
|
|
return super().write(data)
|
|
|
|
|
|
def _decode_audio(raw):
|
|
"""Decode only known in-memory formats; PyAV never sees a path or URL."""
|
|
try:
|
|
import numpy as np
|
|
if raw.startswith((b"#!SILK_V3", b"\x02#!SILK_V3")):
|
|
import pysilk
|
|
output = _BoundedPCM()
|
|
pysilk.decode(io.BytesIO(raw), output, SAMPLE_RATE)
|
|
pcm = output.getvalue()
|
|
if not pcm or len(pcm) % 2:
|
|
raise VoiceTranscriptionError("audio_format")
|
|
audio = np.frombuffer(pcm, dtype="<i2").astype(np.float32) / 32768.0
|
|
else:
|
|
if raw.startswith(b"RIFF") and raw[8:12] == b"WAVE":
|
|
kind = "wav"
|
|
elif raw.startswith((b"#!AMR\n", b"#!AMR-WB\n")):
|
|
kind = "amr"
|
|
else:
|
|
raise VoiceTranscriptionError("audio_format")
|
|
import av
|
|
chunks, samples = [], 0
|
|
with av.open(io.BytesIO(raw), mode="r", format=kind) as container:
|
|
if len(container.streams.audio) != 1:
|
|
raise VoiceTranscriptionError("audio_format")
|
|
stream = container.streams.audio[0]
|
|
channels = stream.codec_context.channels
|
|
rate = stream.codec_context.sample_rate
|
|
if channels not in (1, 2) or not 8000 <= rate <= 96000:
|
|
raise VoiceTranscriptionError("audio_format")
|
|
if stream.duration is not None and stream.time_base is not None and float(stream.duration * stream.time_base) > MAX_DURATION_SECONDS:
|
|
raise VoiceTranscriptionError("audio_too_long")
|
|
resampler = av.AudioResampler(format="s16", layout="mono", rate=SAMPLE_RATE)
|
|
def append(frames):
|
|
nonlocal samples
|
|
for frame in frames:
|
|
samples += frame.samples
|
|
if samples > MAX_DURATION_SECONDS * SAMPLE_RATE:
|
|
raise VoiceTranscriptionError("audio_too_long")
|
|
chunks.append(frame.to_ndarray().reshape(-1))
|
|
for frame in container.decode(audio=0):
|
|
append(resampler.resample(frame))
|
|
append(resampler.resample(None))
|
|
audio = np.concatenate(chunks).astype(np.float32) / 32768.0 if chunks else np.zeros(0, dtype=np.float32)
|
|
if len(audio) > MAX_DURATION_SECONDS * SAMPLE_RATE:
|
|
raise VoiceTranscriptionError("audio_too_long")
|
|
if len(audio) < SAMPLE_RATE // 10 or not np.isfinite(audio).all() or float(np.max(np.abs(audio), initial=0)) < 0.0001:
|
|
raise VoiceTranscriptionError("audio_empty")
|
|
return audio
|
|
except VoiceTranscriptionError:
|
|
raise
|
|
except ImportError:
|
|
raise VoiceTranscriptionError("dependency_missing") from None
|
|
except Exception:
|
|
raise VoiceTranscriptionError("audio_format") from None
|
|
|
|
|
|
def default_model_dir():
|
|
from runtime_paths import resource_path
|
|
return resource_path("assets", "asr", "faster-whisper-base")
|
|
|
|
|
|
def _verified_model_directory(model_dir):
|
|
path = Path(model_dir).resolve()
|
|
try:
|
|
manifest = json.loads((path / "manifest.json").read_text(encoding="utf-8"))
|
|
if manifest.get("revision") != MODEL_REVISION:
|
|
raise VoiceTranscriptionError("model_invalid")
|
|
for name in MODEL_FILES:
|
|
file = path / name
|
|
info = manifest["files"][name]
|
|
if not file.is_file():
|
|
raise VoiceTranscriptionError("model_missing")
|
|
if file.stat().st_size != info["size"]:
|
|
raise VoiceTranscriptionError("model_invalid")
|
|
digest = hashlib.sha256()
|
|
with file.open("rb") as stream:
|
|
while chunk := stream.read(1024 * 1024):
|
|
digest.update(chunk)
|
|
if digest.hexdigest() != info["sha256"] or (name == "model.bin" and digest.hexdigest() != MODEL_SHA256):
|
|
raise VoiceTranscriptionError("model_invalid")
|
|
return path
|
|
except VoiceTranscriptionError:
|
|
raise
|
|
except FileNotFoundError:
|
|
raise VoiceTranscriptionError("model_missing") from None
|
|
except (OSError, ValueError, TypeError, KeyError):
|
|
raise VoiceTranscriptionError("model_invalid") from None
|
|
|
|
|
|
class OfflineRecognizer:
|
|
def __init__(self, model_dir=None):
|
|
self.model_dir = model_dir
|
|
self._model = None
|
|
|
|
def __call__(self, audio):
|
|
if self._model is None:
|
|
path = _verified_model_directory(self.model_dir or default_model_dir())
|
|
try:
|
|
from faster_whisper import WhisperModel
|
|
# tokenizer.json is mandatory above: its absence would make the
|
|
# library fetch an online tokenizer despite local model loading.
|
|
self._model = WhisperModel(str(path), device="cpu", compute_type="int8",
|
|
cpu_threads=2, num_workers=1, local_files_only=True, use_auth_token=False)
|
|
except ImportError:
|
|
raise VoiceTranscriptionError("dependency_missing") from None
|
|
except Exception:
|
|
raise VoiceTranscriptionError("model_invalid") from None
|
|
segments, _ = self._model.transcribe(audio, language="zh", beam_size=5,
|
|
temperature=0.0, condition_on_previous_text=False, vad_filter=True,
|
|
log_progress=False, compression_ratio_threshold=2.4, log_prob_threshold=-1.0,
|
|
no_speech_threshold=0.6, max_new_tokens=256)
|
|
parts = []
|
|
for segment in segments:
|
|
if (not math.isfinite(float(segment.avg_logprob)) or segment.avg_logprob < -1.0
|
|
or segment.no_speech_prob > 0.6 or segment.compression_ratio > 2.4):
|
|
raise VoiceTranscriptionError("low_confidence")
|
|
parts.append(segment.text)
|
|
return "".join(parts)
|
|
|
|
|
|
def _recognize(raw, digest, recognizer):
|
|
audio = _decode_audio(raw)
|
|
try:
|
|
text = recognizer(audio)
|
|
if not isinstance(text, str) or not text.strip():
|
|
raise VoiceTranscriptionError("audio_empty")
|
|
text = text.strip()
|
|
if len(text) > MAX_TEXT_CHARS or "\0" in text or not any(char.isalnum() for char in text):
|
|
raise VoiceTranscriptionError("low_confidence")
|
|
return {"status": "ready", "text": text, "error_code": "", "message": "语音已转成文字",
|
|
"retry_at": 0.0, "audio_sha256": digest, "duration_seconds": round(len(audio) / SAMPLE_RATE, 3),
|
|
"source": "local_asr"}
|
|
except VoiceTranscriptionError:
|
|
raise
|
|
except Exception:
|
|
raise VoiceTranscriptionError("recognition_failed") from None
|
|
|
|
|
|
def transcribe_file(audio_path, *, model_dir=None, allowed_roots=(), recognizer=None):
|
|
"""Blocking worker/testing entry; no state writes, network or model downloads."""
|
|
roots = tuple(Path(root).resolve() for root in allowed_roots)
|
|
raw, digest = _read_audio(audio_path, roots)
|
|
return _recognize(raw, digest, recognizer or OfflineRecognizer(model_dir))
|
|
|
|
|
|
class VoiceTranscriptionService:
|
|
"""One daemon worker, bounded waiting queue/cache, account/message isolation."""
|
|
def __init__(self, model_dir=None, allowed_roots=(), max_pending=32, retry_seconds=30.0,
|
|
recognizer=None, max_cache=128):
|
|
self._roots = tuple(Path(root).resolve() for root in allowed_roots)
|
|
self._recognizer = recognizer or OfflineRecognizer(model_dir)
|
|
self._max_pending = max(1, min(128, int(max_pending)))
|
|
self._max_cache = max(self._max_pending, min(512, int(max_cache)))
|
|
self._retry_seconds = max(1.0, float(retry_seconds))
|
|
self._jobs = queue.Queue(maxsize=self._max_pending)
|
|
self._states = OrderedDict()
|
|
self._results = OrderedDict()
|
|
self._lock = threading.RLock()
|
|
self._stop = threading.Event()
|
|
self._worker_thread = None
|
|
|
|
def request(self, account, conv_id, server_id, audio_path):
|
|
"""Metadata-only nonblocking enqueue/poll; text appears only when ready."""
|
|
if self._stop.is_set():
|
|
return _error("stopped")
|
|
try:
|
|
identity = _identity(account, conv_id, server_id)
|
|
path, signature = _path_and_signature(audio_path, self._roots)
|
|
except VoiceTranscriptionError as exc:
|
|
return _error(exc.code, retry_at=time.time() + self._retry_seconds)
|
|
token = (str(path), signature)
|
|
with self._lock:
|
|
if self._stop.is_set():
|
|
return _error("stopped")
|
|
state = self._states.get(identity)
|
|
if state is not None and state["token"] == token:
|
|
self._states.move_to_end(identity)
|
|
result = state["result"]
|
|
if result["status"] != "error" or time.time() < result["retry_at"]:
|
|
return dict(result)
|
|
if sum(row["result"]["status"] == "pending" for row in self._states.values()) >= self._max_pending:
|
|
return _error("queue_full", retry_at=time.time() + self._retry_seconds)
|
|
attempts = state["attempts"] + 1 if state is not None and state["token"] == token else 1
|
|
result = {"status": "pending", "text": "", "error_code": "", "message": "正在本机识别语音",
|
|
"retry_at": 0.0, "audio_sha256": ""}
|
|
state = {"token": token, "result": result, "attempts": attempts}
|
|
self._states[identity] = state
|
|
self._states.move_to_end(identity)
|
|
while len(self._states) > self._max_cache:
|
|
old = next((key for key, row in self._states.items() if row["result"]["status"] != "pending"), None)
|
|
if old is None:
|
|
break
|
|
self._states.pop(old)
|
|
try:
|
|
self._jobs.put_nowait((identity, state, path, signature))
|
|
except queue.Full:
|
|
state["result"] = _error("queue_full", retry_at=time.time() + self._retry_seconds)
|
|
return dict(state["result"])
|
|
if self._worker_thread is None:
|
|
self._worker_thread = threading.Thread(target=self._run, daemon=True, name="offline-voice-recognition")
|
|
self._worker_thread.start()
|
|
return dict(result)
|
|
|
|
def _run(self):
|
|
while not self._stop.is_set():
|
|
try:
|
|
identity, state, path, signature = self._jobs.get(timeout=0.25)
|
|
except queue.Empty:
|
|
continue
|
|
try:
|
|
with self._lock:
|
|
if self._states.get(identity) is not state or self._stop.is_set():
|
|
continue
|
|
raw, digest = _read_audio(path, self._roots, signature)
|
|
key = (*identity, digest)
|
|
with self._lock:
|
|
result = self._results.get(key)
|
|
if result is None:
|
|
result = _recognize(raw, digest, self._recognizer)
|
|
# A changed source must never receive the earlier file's text.
|
|
if _path_and_signature(path, self._roots)[1] != signature:
|
|
raise VoiceTranscriptionError("audio_changed")
|
|
with self._lock:
|
|
if self._stop.is_set():
|
|
continue
|
|
self._results[key] = dict(result)
|
|
self._results.move_to_end(key)
|
|
while len(self._results) > self._max_cache:
|
|
self._results.popitem(last=False)
|
|
except VoiceTranscriptionError as exc:
|
|
delay = min(300.0, self._retry_seconds * 2 ** min(state["attempts"] - 1, 4))
|
|
result = _error(exc.code, retry_at=time.time() + delay)
|
|
except Exception:
|
|
result = _error("recognition_failed", retry_at=time.time() + self._retry_seconds)
|
|
finally:
|
|
self._jobs.task_done()
|
|
with self._lock:
|
|
if self._states.get(identity) is state and not self._stop.is_set():
|
|
state["result"] = dict(result)
|
|
|
|
def close(self):
|
|
"""Stop accepting work immediately; discard pending/output state."""
|
|
self._stop.set()
|
|
with self._lock:
|
|
self._states.clear()
|
|
self._results.clear()
|
|
|