Files
kefu/wechat_rpa/voice_messages.py
T
2026-09-21 10:34:06 +08:00

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()