251 lines
11 KiB
Python
251 lines
11 KiB
Python
"""Account-scoped voice text for reply readers; no audio or network on the DB thread."""
|
|
from __future__ import annotations
|
|
|
|
import hashlib
|
|
import os
|
|
from pathlib import Path
|
|
import re
|
|
import sqlite3
|
|
import queue
|
|
import threading
|
|
import time
|
|
|
|
VOICE_TYPES = {4, 16}
|
|
_UUID = re.compile(r"(?<![0-9a-f])[0-9a-f]{8}-[0-9a-f]{4}-[0-9a-f]{4}-[0-9a-f]{4}-[0-9a-f]{12}(?![0-9a-f])", re.I)
|
|
_FILE = re.compile(r"(?<![\w.-])[a-z0-9_-]{16,100}\.(?:silk|amr|wav|mp3|m4a|ogg|pcm)(?!\w)", re.I)
|
|
|
|
|
|
def is_voice(message):
|
|
try:
|
|
return int(message.get('content_type') or 0) in VOICE_TYPES
|
|
except (ValueError, TypeError):
|
|
return False
|
|
|
|
|
|
def clean_transcript(value):
|
|
if not isinstance(value, str):
|
|
return ''
|
|
value = ''.join(c for c in value if c.isprintable() or c in '\n\t').strip()
|
|
if not value or len(value) > 32000 or value in {'[语音]', '语音', '未识别出文字', '转文字失败'}:
|
|
return ''
|
|
return value
|
|
|
|
|
|
def _references(raw):
|
|
if isinstance(raw, bytes):
|
|
raw = raw.decode('utf-8', errors='ignore')
|
|
raw = str(raw or '')
|
|
# References only: never fetch a URL, accept a caller-supplied path or guess a recent file.
|
|
return sorted(set(m.group(0).lower() for pattern in (_UUID, _FILE) for m in pattern.finditer(raw)))[:32]
|
|
|
|
|
|
def cached_voice_fields(connection, account, conv_id, server_id, raw):
|
|
fields = {'voice_transcribed': False, 'voice_status': 'missing',
|
|
'content': '[语音待转文字]', 'voice_refs': _references(raw)}
|
|
sid = str(server_id or '')
|
|
if connection is None or not sid.isdecimal() or not 0 < int(sid) < 2 ** 63:
|
|
return fields
|
|
try:
|
|
# message_id in this table is the SERVER id, not message_table.message_id/rowid.
|
|
# Reject a corrupt collision across conversations in the same account.
|
|
targets = connection.execute('SELECT DISTINCT conversation_id FROM message_table WHERE server_id=? LIMIT 2', (sid,)).fetchall()
|
|
if len(targets) != 1 or str(targets[0][0]) != conv_id:
|
|
return fields
|
|
row = connection.execute('SELECT text, voice_id FROM msg_voice2text WHERE message_id=?', (int(sid),)).fetchone()
|
|
except sqlite3.Error:
|
|
return fields
|
|
if row:
|
|
fields['voice_refs'] = sorted(set(fields['voice_refs'] + _references(row[1])))[:32]
|
|
transcript = clean_transcript(row[0])
|
|
if transcript:
|
|
fields.update(content='[语音转文字] ' + transcript, voice_transcribed=True,
|
|
voice_status='ready', voice_transcription_source='wecom_cache',
|
|
voice_binding=f'{account}\0{conv_id}\0{sid}')
|
|
return fields
|
|
|
|
|
|
def verified_transcript(message):
|
|
return bool(is_voice(message) and message.get('voice_transcribed') is True
|
|
and message.get('voice_status') == 'ready'
|
|
and message.get('voice_binding') == f"{message.get('account')}\0{message.get('conv_id')}\0{message.get('server_id')}"
|
|
and str(message.get('content') or '').startswith('[语音转文字] ')
|
|
and clean_transcript(str(message.get('content') or '')[8:]))
|
|
|
|
|
|
def unanswered_messages(context):
|
|
from contact_reply_guard import system_contact_status
|
|
batch = []
|
|
for message in reversed((context or {}).get('messages') or []):
|
|
if message.get('is_self') or system_contact_status(message) == 'restored':
|
|
break
|
|
batch.append(message)
|
|
return list(reversed(batch))
|
|
|
|
|
|
def voice_batch_status(context):
|
|
"""Caller delays the model for pending voices and never loses mixed text+voice turns."""
|
|
voices = [m for m in unanswered_messages(context) if is_voice(m)]
|
|
if not voices:
|
|
return {'status': 'none'}
|
|
ids = '\0'.join(str(m.get('dedup_key') or '') for m in voices)
|
|
identity = hashlib.sha256(ids.encode('utf-8')).hexdigest()
|
|
if any(m.get('account') != context.get('account') or m.get('conv_id') != context.get('conv_id') for m in voices):
|
|
return {'status': 'error', 'key': identity, 'reason': '语音消息账号或会话未能核验'}
|
|
missing = [m for m in voices if not verified_transcript(m)]
|
|
if not missing:
|
|
return {'status': 'ready', 'key': identity}
|
|
errors = [m for m in missing if m.get('voice_status') == 'error']
|
|
if errors:
|
|
return {'status': 'error', 'key': identity, 'reason': '语音自动转文字失败(' + str(errors[0].get('voice_error_code') or 'transcription_failed') + '),请人工核对'}
|
|
return {'status': 'pending', 'key': identity,
|
|
'reason': '语音正在本机转文字' if any(m.get('voice_status') == 'pending' for m in missing) else '等待企业微信语音文件或转写缓存'}
|
|
|
|
|
|
class VoiceMessageReader:
|
|
"""Resolve only exact media references inside one account's Voice directory."""
|
|
def __init__(self, db_base, service_factory=None, *, background_index=True):
|
|
self.base = Path(db_base).resolve()
|
|
self.services = {}
|
|
self.indexes = {}
|
|
self.service_factory = service_factory
|
|
self._background_index = background_index
|
|
self._index_lock = threading.RLock()
|
|
self._indexing = set()
|
|
self._index_jobs = queue.Queue(maxsize=8)
|
|
self._index_thread = None
|
|
self._closed = threading.Event()
|
|
|
|
def _root(self, account):
|
|
if not str(account).isdecimal():
|
|
return None
|
|
expected = self.base / str(account) / 'Cache' / 'Voice'
|
|
root = expected.resolve()
|
|
# Reject junctions pointing to another account or outside the source tree.
|
|
if os.path.normcase(str(root)) != os.path.normcase(str(expected)):
|
|
return None
|
|
return root
|
|
|
|
def _index(self, account, root):
|
|
cached = self.indexes.get(account)
|
|
if cached and time.monotonic() - cached[0] < 10:
|
|
return cached[1]
|
|
if not self._background_index:
|
|
result = self._build_index(root)
|
|
self.indexes[account] = (time.monotonic(), result)
|
|
return result
|
|
# Keep a known mapping during background refresh. Each request still
|
|
# revalidates its exact file signature in the ASR service. Never regress
|
|
# ready text to missing just because a slow model call crossed the TTL.
|
|
previous = cached[1] if cached else {}
|
|
with self._index_lock:
|
|
if self._closed.is_set():
|
|
return {}
|
|
if account in self._indexing:
|
|
return previous
|
|
try:
|
|
self._index_jobs.put_nowait((account, root))
|
|
except queue.Full:
|
|
return previous
|
|
self._indexing.add(account)
|
|
if self._index_thread is None:
|
|
self._index_thread = threading.Thread(target=self._run_indexes, daemon=True, name='voice-cache-index')
|
|
self._index_thread.start()
|
|
return previous
|
|
|
|
def _run_indexes(self):
|
|
while not self._closed.is_set():
|
|
try:
|
|
account, root = self._index_jobs.get(timeout=0.25)
|
|
except queue.Empty:
|
|
continue
|
|
try:
|
|
result = self._build_index(root)
|
|
except Exception:
|
|
result = {}
|
|
with self._index_lock:
|
|
if not self._closed.is_set():
|
|
self.indexes[account] = (time.monotonic(), result)
|
|
self._indexing.discard(account)
|
|
self._index_jobs.task_done()
|
|
|
|
def _build_index(self, root):
|
|
index, seen_dirs = {}, set()
|
|
visited = 0
|
|
for folder, dirs, files in os.walk(root, followlinks=False):
|
|
resolved = Path(folder).resolve()
|
|
if self._closed.is_set() or len(seen_dirs) >= 5000:
|
|
return {}
|
|
if resolved in seen_dirs or not resolved.is_relative_to(root):
|
|
dirs[:] = []
|
|
continue
|
|
seen_dirs.add(resolved)
|
|
dirs[:] = [d for d in dirs if (Path(folder) / d).resolve().is_relative_to(root)
|
|
and (Path(folder) / d).resolve() not in seen_dirs]
|
|
for name in files:
|
|
visited += 1
|
|
if visited > 20000:
|
|
# A partial index could hide a duplicate; fail closed, not select the first match.
|
|
return {}
|
|
path = Path(folder) / name
|
|
if not path.resolve().is_relative_to(root):
|
|
continue
|
|
tokens = {name.lower()}
|
|
stem = path.stem.lower()
|
|
if _UUID.fullmatch(stem):
|
|
tokens.add(stem)
|
|
for token in tokens:
|
|
index.setdefault(token, set()).add(str(path))
|
|
return index
|
|
|
|
def prepare(self, messages, *, include_history=False):
|
|
# Reply generation keeps the unanswered-only guard; the read-only
|
|
# message library explicitly opts into resolving every visible voice.
|
|
candidates = list(messages or []) if include_history else unanswered_messages({'messages': messages})
|
|
for message in candidates:
|
|
if not is_voice(message) or verified_transcript(message):
|
|
continue
|
|
account, conv_id, sid = (str(message.get(k) or '') for k in ('account', 'conv_id', 'server_id'))
|
|
if not sid.isdecimal() or int(sid) <= 0:
|
|
continue
|
|
root = self._root(account)
|
|
if root is None or not root.is_dir():
|
|
continue
|
|
index = self._index(account, root)
|
|
matches = set()
|
|
for ref in message.get('voice_refs') or []:
|
|
matches.update(index.get(ref, set()))
|
|
if len(matches) != 1:
|
|
continue
|
|
try:
|
|
service = self.services.get(account)
|
|
if service is None:
|
|
factory = self.service_factory
|
|
if factory is None:
|
|
from voice_transcription import VoiceTranscriptionService
|
|
factory = VoiceTranscriptionService
|
|
service = self.services[account] = factory(allowed_roots=(root,))
|
|
result = service.request(account, conv_id, sid, next(iter(matches)))
|
|
status = result.get('status') or 'error'
|
|
message['voice_status'] = status
|
|
if status == 'ready':
|
|
text = clean_transcript(result.get('text'))
|
|
if not text:
|
|
raise ValueError('empty_transcript')
|
|
message.update(content='[语音转文字] ' + text, voice_transcribed=True,
|
|
voice_transcription_source='local_asr',
|
|
voice_binding=f'{account}\0{conv_id}\0{sid}')
|
|
elif status == 'error':
|
|
code = str(result.get('error_code') or 'transcription_failed')
|
|
message['voice_error_code'] = code
|
|
if code in {'audio_missing', 'audio_changed', 'queue_full', 'recognition_failed'}:
|
|
# Worker owns backoff; reply task owns the overall bounded wait deadline.
|
|
message['voice_status'] = 'pending'
|
|
except Exception:
|
|
message.update(voice_status='error', voice_error_code='local_transcription_unavailable')
|
|
|
|
def close(self):
|
|
self._closed.set()
|
|
for service in self.services.values():
|
|
service.close()
|
|
self.services.clear()
|