# -*- coding: utf-8 -*- """Qwen3-ASR websocket streaming service.""" from __future__ import annotations import asyncio import io import json import logging import re import time import unicodedata from difflib import SequenceMatcher from dataclasses import dataclass, field from enum import IntEnum from typing import Any, Dict, List, Optional import numpy as np import soundfile as sf from fastapi import WebSocket, WebSocketDisconnect from app.core.config import settings from app.core.exceptions import create_error_response from app.core.executor import run_sync from app.core.text_cleanup import deduplicate_asr_text from app.services.asr.model_selection import validate_realtime_model_id from app.services.asr.qwen3_engine import Qwen3ASREngine, Qwen3StreamingState from app.services.asr.engines.global_models import get_global_vad_model, get_vad_inference_lock from app.services.realtime_speaker_clusterer import get_realtime_speaker_clusterer from app.services.asr.runtime import RuntimeEngineLease, get_runtime_router from app.services.speaker_registry import get_speaker_registry_service from app.utils.text_processing import normalize_asr_text logger = logging.getLogger(__name__) def _remove_repeated_chars(text: str, max_run: int) -> str: out: List[str] = [] last: Optional[str] = None run = 0 for ch in text: if ch == last: run += 1 else: last = ch run = 1 if run <= max_run: out.append(ch) return "".join(out) def _normalize_transcript_chunk_text(text: str) -> str: collapsed = " ".join(str(text or "").split()).strip() if not collapsed: return "" return _remove_repeated_chars(collapsed, 8) def _is_opening_punctuation(ch: str) -> bool: return ch in '([{"\'“‘' def _is_closing_punctuation(ch: str) -> bool: return ch in '.,!?:;)]}"\'”’。,!?:;、' def _append_with_spacing(left: str, right: str) -> str: if not left: return right if not right: return left prev = left[-1] nxt = right[0] needs_space = ( not prev.isspace() and not nxt.isspace() and not _is_closing_punctuation(nxt) and not _is_opening_punctuation(prev) ) return f"{left} {right}" if needs_space else f"{left}{right}" def _normalize_token(token: str) -> str: return "".join( ch.lower() for ch in token if ch.isalnum() or ch in {"'", "-"} ) def _token_views(text: str) -> List[tuple[str, int]]: views: List[tuple[str, int]] = [] start: Optional[int] = None for idx, ch in enumerate(text): if ch.isspace(): if start is not None: token = text[start:idx] normalized = _normalize_token(token) if normalized: views.append((normalized, start)) start = None continue if start is None: start = idx if start is not None: token = text[start:] normalized = _normalize_token(token) if normalized: views.append((normalized, start)) return views def _dedupe_overlap_word_boundary( previous: str, current: str, min_overlap: int = 3, max_overlap: int = 24, ) -> int: prev_tokens = _token_views(previous) curr_tokens = _token_views(current) if not prev_tokens or not curr_tokens: return 0 upper = min(max_overlap, len(prev_tokens), len(curr_tokens)) lower = max(min_overlap, 1) if upper < lower: return 0 for size in range(upper, lower - 1, -1): left = prev_tokens[-size:] right = curr_tokens[:size] if all(a == b for (a, _), (b, _) in zip(left, right)): if size == len(curr_tokens): return len(current) return curr_tokens[size][1] return 0 def _char_count_to_index(text: str, char_count: int) -> int: if char_count <= 0: return 0 count = 0 for idx, ch in enumerate(text): count += 1 if count == char_count: return idx + 1 return len(text) def _dedupe_overlap_char_boundary( previous: str, current: str, min_chars: int = 6, max_chars: int = 80, ) -> int: prev_chars = list(previous) curr_chars = list(current) if not prev_chars or not curr_chars: return 0 upper = min(max_chars, len(prev_chars), len(curr_chars)) lower = max(min_chars, 1) if upper < lower: return 0 for size in range(upper, lower - 1, -1): left = "".join(ch.lower() for ch in prev_chars[-size:]) right = "".join(ch.lower() for ch in curr_chars[:size]) if left == right: return _char_count_to_index(current, size) return 0 class RealtimeTranscriptAssembler: def __init__(self) -> None: self._merged = "" def push(self, text: str) -> None: cleaned = _normalize_transcript_chunk_text(text) if not cleaned: return if not self._merged: self._merged = cleaned return delta_start = _dedupe_overlap_word_boundary(self._merged, cleaned) if delta_start == 0: delta_start = _dedupe_overlap_char_boundary(self._merged, cleaned) if delta_start >= len(cleaned): return delta = cleaned[delta_start:].lstrip() if not delta: return self._merged = _append_with_spacing(self._merged, delta) def text(self) -> str: return self._merged.strip() def _convert_audio( audio_bytes: bytes, fmt: str, sample_rate: int, ) -> Optional[np.ndarray]: try: if fmt == "pcm": audio = np.frombuffer(audio_bytes, dtype=np.int16).astype(np.float32) / 32768.0 elif fmt == "wav": try: audio, sr = sf.read(io.BytesIO(audio_bytes)) audio = np.asarray(audio, dtype=np.float32) if sr != sample_rate: logger.warning("WAV sample rate mismatch: payload=%s actual=%s", sample_rate, sr) sample_rate = int(sr) except Exception: if len(audio_bytes) > 44: audio_bytes = audio_bytes[44:] audio = np.frombuffer(audio_bytes, dtype=np.int16).astype(np.float32) / 32768.0 else: raise ValueError(f"Unsupported audio format: {fmt}") if getattr(audio, "ndim", 1) > 1: audio = np.mean(audio, axis=1) if sample_rate != 16000: import scipy.signal num = int(len(audio) * 16000 / sample_rate) audio = scipy.signal.resample(audio, num) if isinstance(audio, tuple): audio = audio[0] return np.asarray(audio, dtype=np.float32) except Exception as exc: logger.error("Audio conversion failed: %s", exc) return None class ConnectionState(IntEnum): READY = 1 STARTED = 2 STREAMING = 3 @dataclass class ConnectionContext: state: ConnectionState = ConnectionState.READY params: Dict[str, Any] = field(default_factory=dict) engine_lease: Optional[RuntimeEngineLease] = None engine: Optional[Qwen3ASREngine] = None pre_roll_audio: np.ndarray = field(default_factory=lambda: np.array([], dtype=np.float32)) # segment_audio_buffer 保留“当前断句以来”的完整音频,用于断句时提取声纹。 segment_audio_buffer: np.ndarray = field(default_factory=lambda: np.array([], dtype=np.float32)) # stream_window_buffer 保留最近窗口,供 realtime partial 在 native stream 失效时兜底重转写。 stream_window_buffer: np.ndarray = field(default_factory=lambda: np.array([], dtype=np.float32)) # realtime_stream_state 只负责低延迟 partial,不参与最终 segment 定稿。 realtime_stream_state: Optional[Qwen3StreamingState] = None sentence_active: bool = False silence_samples: int = 0 total_samples: int = 0 confirmed_segments: List[Dict[str, Any]] = field(default_factory=list) speaker_records: List[Dict[str, Any]] = field(default_factory=list) speaker_display_map: Dict[int, str] = field(default_factory=dict) speaker_display_profiles: Dict[str, np.ndarray] = field(default_factory=dict) speaker_display_counts: Dict[str, int] = field(default_factory=dict) next_speaker_display_id: int = 1 timeline_cursor_ms: int = 0 segment_index: int = 0 last_partial_text: str = "" last_partial_language: str = "" last_partial_chunk_id: int = 0 last_partial_raw_text: str = "" last_partial_display_text: str = "" stable_partial_prefix: str = "" segment_observed_text: str = "" segment_observed_language: str = "" best_partial_text: str = "" best_partial_language: str = "" last_partial_decode_samples: int = 0 partial_stable_rounds: int = 0 send_lock: asyncio.Lock = field(default_factory=asyncio.Lock) background_tasks: set[asyncio.Task[Any]] = field(default_factory=set) speaker_job_queue: asyncio.Queue[Any] = field(default_factory=asyncio.Queue) speaker_worker_task: Optional[asyncio.Task[Any]] = None speaker_worker_socket_id: Optional[int] = None MAX_BUFFER = 960000 DEFAULT_MAX_PARTIAL_TEXT_CHARS = 1200 @dataclass class SessionEntry: session_id: str ctx: ConnectionContext attached: bool = False detached_at: Optional[float] = None closed: bool = False class Qwen3ASRService: _MAX_SPEAKER_HISTORY_RECORDS = 64 _MAX_SPEAKER_MATCH_CONTEXT = 24 _MAX_SPEAKER_FINAL_RECLUSTER = 48 _MAX_SPEAKER_REALTIME_RECLUSTER = 12 _MIN_REALTIME_RECLUSTER_SEGMENTS = 5 _MIN_HARD_SEGMENT_SEC = 12.0 _SEGMENT_EVENT_TEXT_TAIL_CHARS = 2000 _SEGMENT_EVENT_TAIL_COUNT = 8 _LIGHTWEIGHT_RECENT_SENTENCE_COUNT = 3 def __init__(self) -> None: self._sessions: Dict[str, SessionEntry] = {} self._session_lock = asyncio.Lock() async def _cleanup_expired_sessions(self) -> None: ttl_sec = max(int(getattr(settings, "REALTIME_SESSION_RESUME_TTL_SEC", 120) or 0), 0) if ttl_sec <= 0: expired_ids = [session_id for session_id, entry in self._sessions.items() if not entry.attached] else: now = time.monotonic() expired_ids = [ session_id for session_id, entry in self._sessions.items() if not entry.attached and ( entry.closed or entry.detached_at is None or now - entry.detached_at > ttl_sec ) ] for session_id in expired_ids: entry = self._sessions.pop(session_id, None) if entry is not None: await self._release_session_resources(entry.ctx) async def _release_session_resources(self, ctx: ConnectionContext) -> None: if ctx.speaker_worker_task is not None and not ctx.speaker_worker_task.done(): await ctx.speaker_job_queue.put(None) if ctx.background_tasks: await asyncio.gather(*list(ctx.background_tasks), return_exceptions=True) if ctx.engine_lease is not None: await ctx.engine_lease.close() ctx.engine_lease = None ctx.engine = None ctx.realtime_stream_state = None async def _detach_session( self, session_id: str, *, keep_for_resume: bool, ) -> None: async with self._session_lock: entry = self._sessions.get(session_id) if entry is None: return entry.attached = False entry.detached_at = time.monotonic() entry.closed = not keep_for_resume self._stop_speaker_worker(entry.ctx) await self._cleanup_expired_sessions() async def _acquire_session( self, session_id: str, ) -> tuple[ConnectionContext, bool]: async with self._session_lock: await self._cleanup_expired_sessions() entry = self._sessions.get(session_id) if entry is not None and not entry.attached and not entry.closed: entry.attached = True entry.detached_at = None return entry.ctx, True ctx = ConnectionContext() self._sessions[session_id] = SessionEntry( session_id=session_id, ctx=ctx, attached=True, detached_at=None, closed=False, ) return ctx, False @staticmethod def _tencent_speaker_id(value: Any) -> int: if value is None: return -1 if isinstance(value, bool): return int(value) if isinstance(value, int): return value text = str(value).strip() if not text: return -1 if re.fullmatch(r"-?\d+", text): return int(text) match = re.fullmatch(r"Speaker(\d+)", text) if match: return max(int(match.group(1)) - 1, 0) return -1 def _build_tencent_sentence( self, payload: Dict[str, Any], *, sentence_type: int, sentence_text: Optional[str] = None, ) -> Dict[str, Any]: speaker_id = self._tencent_speaker_id(payload.get("speaker_id")) if payload.get("speaker_pending"): speaker_id = -1 sentence = { "sentence_id": int(payload.get("index", 0)), "sentence_type": int(sentence_type), "speaker_id": speaker_id, "start_time": int(payload.get("start_ms", 0)), "end_time": int(payload.get("end_ms", 0)), "sentence": str(sentence_text if sentence_text is not None else payload.get("text", "")), } # Tencent-compatible core fields first; project-specific fields remain minimal. for key in ( "speaker_name", "user_id", ): if key in payload: sentence[key] = payload.get(key) return sentence async def _send_json_safe( self, websocket: WebSocket, ctx: ConnectionContext, payload: Dict[str, Any], ) -> None: async with ctx.send_lock: await websocket.send_json(payload) def _track_background_task( self, ctx: ConnectionContext, task: asyncio.Task[Any], ) -> None: ctx.background_tasks.add(task) task.add_done_callback(ctx.background_tasks.discard) def _ensure_speaker_worker( self, websocket: WebSocket, ctx: ConnectionContext, task_id: str, ) -> None: task = ctx.speaker_worker_task socket_id = id(websocket) if ( task is not None and not task.done() and ctx.speaker_worker_socket_id == socket_id ): return worker = asyncio.create_task( self._speaker_worker_loop( websocket, ctx, task_id, ) ) ctx.speaker_worker_task = worker ctx.speaker_worker_socket_id = socket_id self._track_background_task(ctx, worker) def _stop_speaker_worker(self, ctx: ConnectionContext) -> None: task = ctx.speaker_worker_task if task is not None and not task.done(): task.cancel() ctx.speaker_worker_task = None ctx.speaker_worker_socket_id = None async def _speaker_worker_loop( self, websocket: WebSocket, ctx: ConnectionContext, task_id: str, ) -> None: while True: job = await ctx.speaker_job_queue.get() try: if job is None: return await self._resolve_and_emit_segment_speaker( websocket, ctx, task_id, segment_index=int(job["segment_index"]), audio=np.asarray(job["audio"], dtype=np.float32), reason=str(job["reason"]), segment_start_ms=int(job["segment_start_ms"]), segment_end_ms=int(job["segment_end_ms"]), ) finally: ctx.speaker_job_queue.task_done() async def _emit_speaker_update( self, websocket: WebSocket, ctx: ConnectionContext, task_id: str, segment_index: int, ) -> None: segment = next( (item for item in ctx.confirmed_segments if int(item.get("index", -1)) == int(segment_index)), None, ) if segment is None: return sentence_payload = self._build_tencent_sentence(segment, sentence_type=1) await self._send_json_safe( websocket, ctx, { "type": "sentences", "code": 0, "voice_id": task_id, "final": 0, "result": { "slice_type": 2, "index": int(segment_index), "voice_text_str": str(segment.get("text") or ""), }, "sentences": [sentence_payload], }, ) @staticmethod def _public_speaker_view(payload: Dict[str, Any]) -> tuple[Any, Any, Any]: return ( payload.get("speaker_id"), payload.get("speaker_name"), payload.get("user_id"), ) def _create_display_speaker_id( self, ctx: ConnectionContext, ) -> str: display_id = f"Speaker{ctx.next_speaker_display_id:02d}" ctx.next_speaker_display_id += 1 return display_id def _match_display_speaker_id( self, ctx: ConnectionContext, *, cluster_index: Optional[int], embedding: Optional[np.ndarray], ) -> Optional[str]: if embedding is None: return None normalized = np.asarray(embedding, dtype=np.float32) norm = float(np.linalg.norm(normalized)) if norm <= 0: return None normalized = normalized / norm preferred_display_id: Optional[str] = None if cluster_index is not None: mapped = ctx.speaker_display_map.get(int(cluster_index)) if mapped: profile = ctx.speaker_display_profiles.get(mapped) if profile is not None: preferred_score = float(np.dot(normalized, profile)) if preferred_score >= settings.REALTIME_SPEAKER_CONFIRM_THRESHOLD: preferred_display_id = mapped # Do not aggressively reuse a generic display id across unrelated # segments when there is no stable cluster mapping yet; that easily # collapses multiple real speakers into Speaker01. if preferred_display_id is None: preferred_display_id = self._create_display_speaker_id(ctx) count = int(ctx.speaker_display_counts.get(preferred_display_id, 0)) previous = ctx.speaker_display_profiles.get(preferred_display_id) if previous is None or count <= 0: updated = normalized count = 1 else: updated = previous * float(count) + normalized updated_norm = float(np.linalg.norm(updated)) if updated_norm > 0: updated = updated / updated_norm count += 1 ctx.speaker_display_profiles[preferred_display_id] = np.asarray(updated, dtype=np.float32) ctx.speaker_display_counts[preferred_display_id] = count if cluster_index is not None: ctx.speaker_display_map[int(cluster_index)] = preferred_display_id return preferred_display_id @staticmethod def _safe_int(value: Any, default: Optional[int] = None) -> Optional[int]: try: if value is None: return default return int(value) except (TypeError, ValueError): return default def _public_payload(self, payload: Optional[Dict[str, Any]]) -> Optional[Dict[str, Any]]: if payload is None: return None return { key: value for key, value in payload.items() if not str(key).startswith("_") } def _is_stable_speaker_info( self, speaker_info: Optional[Dict[str, Any]], ) -> bool: if not speaker_info: return False if speaker_info.get("registry_speaker_id") or speaker_info.get("user_id"): return True strategy = str(speaker_info.get("speaker_strategy") or "") if strategy == "registry_match": return True if strategy in {"recent_named_attach", "recent_named_inherit"}: confidence = float(speaker_info.get("speaker_confidence") or 0.0) return confidence >= 0.7 if strategy == "embedding_match" and speaker_info.get("_matched_existing"): confidence = float(speaker_info.get("speaker_confidence") or 0.0) return confidence >= 0.62 if strategy == "embedding_match": confidence = float(speaker_info.get("speaker_confidence") or 0.0) duration_ms = float(speaker_info.get("_segment_duration_ms") or 0.0) if duration_ms >= 8000: return confidence >= 0.62 if strategy != "embedding_match": return False confidence = float(speaker_info.get("speaker_confidence") or 0.0) threshold = max( float(getattr(settings, "REALTIME_SPEAKER_CONFIRM_THRESHOLD", 0.62) or 0.62), 0.68, ) return confidence >= threshold def _record_segment_speaker( self, ctx: ConnectionContext, *, segment_index: int, segment_start_ms: int, segment_end_ms: int, speaker_info: Optional[Dict[str, Any]], ) -> None: if not speaker_info: return embedding = speaker_info.get("_embedding") if embedding is None: return ctx.speaker_records.append( { "index": segment_index, "start_ms": int(segment_start_ms), "end_ms": int(segment_end_ms), "embedding": np.asarray(embedding, dtype=np.float32), "embeddings": [ np.asarray(item, dtype=np.float32) for item in (speaker_info.get("_chunk_embeddings") or [embedding]) ], "chunks": [ { "start_ms": int(item.get("start_ms", segment_start_ms)), "end_ms": int(item.get("end_ms", segment_end_ms)), "embedding": np.asarray(item["embedding"], dtype=np.float32), } for item in (speaker_info.get("_chunks") or []) if item.get("embedding") is not None ], "speaker_id": speaker_info.get("speaker_id"), "cluster_index": speaker_info.get("_cluster_index"), "speaker_name": speaker_info.get("speaker_name"), "registry_speaker_id": speaker_info.get("registry_speaker_id"), "user_id": speaker_info.get("user_id"), } ) if len(ctx.speaker_records) > self._MAX_SPEAKER_HISTORY_RECORDS: ctx.speaker_records = ctx.speaker_records[-self._MAX_SPEAKER_HISTORY_RECORDS:] def _inherit_recent_anonymous_speaker( self, ctx: ConnectionContext, speaker_info: Dict[str, Any], ) -> Optional[Dict[str, Any]]: if speaker_info.get("registry_speaker_id") or speaker_info.get("user_id"): return None if speaker_info.get("_matched_existing"): return None duration_ms = int(speaker_info.get("_segment_duration_ms") or 0) if duration_ms < 8000: return None if not ctx.speaker_records: return None last_record = ctx.speaker_records[-1] last_speaker_id = str(last_record.get("speaker_id") or "").strip() if not re.fullmatch(r"Speaker\d+", last_speaker_id): return None inherited = dict(speaker_info) inherited["speaker_id"] = last_speaker_id inherited["speaker_name"] = str(last_record.get("speaker_name") or last_speaker_id) inherited["_cluster_index"] = last_record.get("cluster_index") inherited["_matched_existing"] = True inherited["speaker_confidence"] = max(float(inherited.get("speaker_confidence") or 0.0), 0.66) inherited["speaker_strategy"] = "recent_inherit" return inherited def _inherit_recent_named_speaker( self, ctx: ConnectionContext, speaker_info: Dict[str, Any], ) -> Optional[Dict[str, Any]]: if speaker_info.get("registry_speaker_id") or speaker_info.get("user_id"): return speaker_info if speaker_info.get("_matched_existing"): return None duration_ms = int(speaker_info.get("_segment_duration_ms") or 0) if duration_ms < 4500: return None if not ctx.speaker_records: return None last_record = ctx.speaker_records[-1] if not last_record.get("registry_speaker_id") and not last_record.get("user_id"): return None inherited = dict(speaker_info) inherited["speaker_id"] = last_record.get("registry_speaker_id") or last_record.get("speaker_id") inherited["speaker_name"] = last_record.get("speaker_name") or inherited["speaker_id"] inherited["user_id"] = last_record.get("user_id") inherited["registry_speaker_id"] = last_record.get("registry_speaker_id") inherited["_cluster_index"] = last_record.get("cluster_index") inherited["_matched_existing"] = True inherited["speaker_confidence"] = max(float(inherited.get("speaker_confidence") or 0.0), 0.72) inherited["speaker_strategy"] = "recent_named_inherit" return inherited def _speaker_match_records( self, ctx: ConnectionContext, ) -> List[Dict[str, Any]]: if len(ctx.speaker_records) <= self._MAX_SPEAKER_MATCH_CONTEXT: return ctx.speaker_records return ctx.speaker_records[-self._MAX_SPEAKER_MATCH_CONTEXT:] def _recluster_confirmed_segments( self, ctx: ConnectionContext, *, max_records: Optional[int] = None, only_stable_updates: bool = False, only_pending_updates: bool = False, ) -> List[Dict[str, Any]]: if len(ctx.speaker_records) < 2: return [] record_limit = max_records or self._MAX_SPEAKER_FINAL_RECLUSTER records = ( ctx.speaker_records if len(ctx.speaker_records) <= record_limit else ctx.speaker_records[-record_limit:] ) clusters, assignments = self._cluster_speaker_records(records) if not clusters: return [] cluster_labels: dict[int, Dict[str, Any]] = {} unknown_display_index = 1 for idx, cluster in enumerate(clusters): named_records = [ item for item in cluster["record_refs"] if item.get("registry_speaker_id") or item.get("user_id") or ( item.get("speaker_name") and not re.fullmatch(r"Speaker\d+", str(item.get("speaker_name")).strip()) ) ] if named_records: preferred = max( named_records, key=lambda item: ( 1 if item.get("registry_speaker_id") else 0, 1 if item.get("user_id") else 0, 1 if item.get("speaker_name") else 0, ), ) speaker_id = preferred.get("registry_speaker_id") or preferred.get("speaker_id") speaker_name = preferred.get("speaker_name") or speaker_id or f"Speaker{unknown_display_index:02d}" cluster_labels[idx] = { "speaker_id": speaker_id or f"Speaker{unknown_display_index:02d}", "speaker_name": speaker_name, "user_id": preferred.get("user_id"), "registry_speaker_id": preferred.get("registry_speaker_id"), "speaker_strategy": "final_recluster", "speaker_confidence": 1.0, } else: speaker_id = f"Speaker{unknown_display_index:02d}" cluster_labels[idx] = { "speaker_id": speaker_id, "speaker_name": speaker_id, "user_id": None, "speaker_strategy": "final_recluster", "speaker_confidence": 1.0, } unknown_display_index += 1 updates: List[Dict[str, Any]] = [] for record, cluster_idx in zip(records, assignments): segment_index = int(record["index"]) segment = next( (item for item in ctx.confirmed_segments if int(item.get("index", -1)) == segment_index), None, ) if segment is None: continue if only_pending_updates and not segment.get("speaker_pending"): continue new_info = cluster_labels[cluster_idx] if only_stable_updates: has_named_identity = bool(new_info.get("registry_speaker_id") or new_info.get("user_id")) generic_name = str(new_info.get("speaker_name") or "").strip() if not has_named_identity and re.fullmatch(r"Speaker\d+", generic_name): if not segment.get("speaker_pending"): continue changed = any(segment.get(key) != value for key, value in new_info.items()) if not changed and not segment.get("speaker_pending"): continue segment.update(new_info) segment.pop("speaker_pending", None) updates.append( { "segment_index": segment_index, "segment": self._public_payload(segment), } ) return updates def _maybe_recluster_recent_segments( self, ctx: ConnectionContext, ) -> List[Dict[str, Any]]: if len(ctx.speaker_records) < self._MIN_REALTIME_RECLUSTER_SEGMENTS: return [] return self._recluster_confirmed_segments( ctx, max_records=self._MAX_SPEAKER_REALTIME_RECLUSTER, only_stable_updates=True, only_pending_updates=True, ) def _build_sv_chunks(self, audio: np.ndarray) -> List[np.ndarray]: return get_realtime_speaker_clusterer().build_sv_chunks(audio) async def _extract_chunk_embeddings( self, audio: np.ndarray, ) -> List[np.ndarray]: return await get_realtime_speaker_clusterer().extract_chunk_embeddings(audio) def _cluster_speaker_records( self, records: List[Dict[str, Any]], ) -> tuple[List[Dict[str, Any]], List[int]]: return get_realtime_speaker_clusterer().cluster_records(records) async def _ensure_engine(self, ctx: ConnectionContext) -> Qwen3ASREngine: if ctx.engine is not None: return ctx.engine runtime_router = get_runtime_router() model = validate_realtime_model_id("qwen3-asr") logger.info("Using Qwen3-ASR model: %s", model) ctx.engine_lease = await runtime_router.acquire_engine(model) engine = ctx.engine_lease.engine if not isinstance(engine, Qwen3ASREngine): raise RuntimeError("Current model is not Qwen3-ASR") if not engine.supports_realtime: raise RuntimeError( f"Current device {engine.device} does not support Qwen3-ASR realtime streaming; " "only CUDA vLLM and CPU Rust paths are supported" ) ctx.engine = engine return engine def _has_voice(self, audio: np.ndarray) -> bool: if audio.size == 0: return False rms = float(np.sqrt(np.mean(audio**2))) peak = float(np.max(np.abs(audio))) return rms >= 0.0045 or (rms >= 0.0025 and peak >= 0.08) def _get_dynamic_silence_threshold_samples(self, ctx: ConnectionContext) -> int: silence_ms = max(int(ctx.params.get("silence_duration_ms", 800) or 800), 100) return int(silence_ms * 16) def _normalize_output_text( self, text: str, ctx: ConnectionContext, *, enable_itn: Optional[bool] = None, ) -> str: raw = str(text or "").strip() if "" in raw: _prefix, raw = raw.split("", 1) raw = re.sub(r"^\s*language\s+[A-Za-z][A-Za-z\s-]*\s*", "", raw, flags=re.IGNORECASE) if enable_itn is None: enable_itn = bool(ctx.params.get("enable_inverse_text_normalization", True)) return normalize_asr_text( raw, enable_itn=enable_itn, ) def _normalize_output_language(self, language: Optional[str], ctx: ConnectionContext) -> str: raw = str(language or "").strip() if not raw or raw.lower() == "none": return "" raw = re.sub(r"^[^\w]+", "", raw) match = re.search(r"language\s+([A-Za-z][A-Za-z\s-]*)$", raw, re.IGNORECASE) if match: return match.group(1).strip() return raw def _infer_text_language(self, text: str) -> str: has_cjk = any("\u4e00" <= ch <= "\u9fff" for ch in text) has_thai = any("\u0e00" <= ch <= "\u0e7f" for ch in text) has_latin = any("LATIN" in unicodedata.name(ch, "") for ch in text if ch.isalpha()) if has_thai: return "Thai" if has_cjk: return "Chinese" if has_latin: return "English" return "" def _get_session_dominant_language(self, ctx: ConnectionContext) -> str: counts: Dict[str, int] = {} for segment in ctx.confirmed_segments: language = str(segment.get("language") or "").strip() if not language: continue counts[language] = counts.get(language, 0) + 1 if not counts: return "" language, count = max(counts.items(), key=lambda item: item[1]) return language if count >= 2 else "" def _pre_roll_samples(self, ctx: ConnectionContext) -> int: pre_roll_ms = max(int(ctx.params.get("pre_roll_ms", 240) or 0), 0) return int(pre_roll_ms * 16) def _min_partial_samples(self, ctx: ConnectionContext) -> int: return int( max( float( ctx.params.get( "min_partial_sec", settings.REALTIME_MIN_PARTIAL_SEC, ) or settings.REALTIME_MIN_PARTIAL_SEC ), 0.25, ) * 16000 ) def _partial_emit_interval_samples(self, ctx: ConnectionContext) -> int: """ Control how often we surface partial text to the client. When native Qwen streaming is healthy we can emit much more frequently, because decoding is already happening incrementally in memory. If native streaming falls back to window retranscription, keep the cadence a little slower to avoid excessive temp-file retranscribes. """ interval_sec = max( float( getattr( settings, "REALTIME_PARTIAL_EMIT_INTERVAL_SEC", 0.25, ) or 0.25 ), 0.12, ) if ctx.params.get("enable_native_partial_stream") is False: interval_sec = max(interval_sec, max(self._stream_chunk_size_sec(ctx), 0.9)) elif ctx.realtime_stream_state is None: interval_sec = max(interval_sec, max(min(self._stream_chunk_size_sec(ctx), 0.8), 0.55)) return int(interval_sec * 16000) def _stream_window_samples(self) -> int: window_sec = max( float(getattr(settings, "REALTIME_STREAM_WINDOW_SEC", 8.0) or 8.0), 1.0, ) return int(window_sec * 16000) def _partial_window_samples(self, ctx: ConnectionContext) -> int: window_sec = max( float( ctx.params.get( "partial_window_sec", getattr(settings, "REALTIME_PARTIAL_WINDOW_SEC", 8.0), ) or getattr(settings, "REALTIME_PARTIAL_WINDOW_SEC", 8.0) ), 1.0, ) return int(window_sec * 16000) def _max_sentence_count(self, ctx: ConnectionContext) -> int: return max(int(ctx.params.get("max_sentence_count", 8) or 8), 1) def _max_partial_text_chars(self, ctx: ConnectionContext) -> int: return max( int( ctx.params.get( "max_partial_text_chars", ctx.DEFAULT_MAX_PARTIAL_TEXT_CHARS, ) or ctx.DEFAULT_MAX_PARTIAL_TEXT_CHARS ), 200, ) def _max_segment_samples(self, ctx: ConnectionContext) -> int: configured_sec = float( ctx.params.get( "hard_limit_sec", ctx.params.get("max_segment_sec", settings.REALTIME_FINAL_SEGMENT_HARD_LIMIT_SEC), ) or 0.0 ) if configured_sec <= 0: return 0 return int(max(configured_sec, self._MIN_HARD_SEGMENT_SEC, 1.0) * 16000) def _soft_segment_limit_samples(self, ctx: ConnectionContext) -> int: limit_sec = float( ctx.params.get( "soft_limit_sec", getattr(settings, "REALTIME_FINAL_SEGMENT_SOFT_LIMIT_SEC", 8.0), ) or 0.0 ) if limit_sec <= 0: return 0 return int(max(limit_sec, 1.0) * 16000) def _clip_partial_text( self, text: str, ctx: ConnectionContext, ) -> str: limit = self._max_partial_text_chars(ctx) if len(text) <= limit: return text return text[-limit:] def _enable_realtime_vad_split(self, ctx: ConnectionContext) -> bool: return bool(ctx.params.get("enable_realtime_vad_split", False)) def _enable_realtime_longform(self, ctx: ConnectionContext) -> bool: return bool(ctx.params.get("enable_realtime_longform", False)) def _enable_realtime_refine(self, ctx: ConnectionContext) -> bool: return False @staticmethod def _sentence_count(text: str) -> int: if not text: return 0 return len(re.findall(r"[。!??!]", text)) @staticmethod def _ends_with_sentence_punctuation(text: str) -> bool: stripped = str(text or "").rstrip() return bool(stripped) and stripped[-1] in "。!??!" def _last_sentence_body_len(self, text: str) -> int: stripped = str(text or "").rstrip() if not stripped: return 0 if self._ends_with_sentence_punctuation(stripped): stripped = stripped[:-1].rstrip() last_boundary = max(stripped.rfind(ch) for ch in "。!??!") body = stripped[last_boundary + 1:] if last_boundary >= 0 else stripped return sum(1 for ch in body if ch.isalnum() or "\u4e00" <= ch <= "\u9fff") def _should_commit_complete_sentence(self, ctx: ConnectionContext, text: str) -> bool: candidate = self._sanitize_candidate_text(text) if not candidate or not self._ends_with_sentence_punctuation(candidate): return False duration_sec = float(ctx.segment_audio_buffer.size) / 16000.0 min_sec = max( float( ctx.params.get( "complete_sentence_commit_sec", max(float(ctx.params.get("force_stable_segment_sec", 8.0) or 8.0), 12.0), ) or 12.0 ), 6.0, ) if duration_sec < min_sec: return False min_chars = max( int( ctx.params.get( "complete_sentence_commit_min_chars", max(int(ctx.params.get("force_stable_min_chars", 24) or 24), 48), ) or 48 ), 12, ) if len(candidate) < min_chars: return False # Avoid committing on tiny unstable tails such as "可。" or "啊。". return self._last_sentence_body_len(candidate) >= 6 def _stream_chunk_size_sec(self, ctx: ConnectionContext) -> float: return max( float( ctx.params.get( "chunk_size_sec", settings.REALTIME_STREAM_CHUNK_SEC, ) or settings.REALTIME_STREAM_CHUNK_SEC ), 0.4, ) def _stream_unfixed_chunk_num(self, ctx: ConnectionContext) -> int: return max( int( ctx.params.get( "unfixed_chunk_num", settings.REALTIME_STREAM_MAX_PENDING_CHUNKS, ) or settings.REALTIME_STREAM_MAX_PENDING_CHUNKS ), 1, ) def _stream_unfixed_token_num(self, ctx: ConnectionContext) -> int: return max(int(ctx.params.get("unfixed_token_num", 5) or 5), 1) def _should_use_native_partial_stream( self, ctx: ConnectionContext, engine: Qwen3ASREngine, ) -> bool: _ = engine configured = ctx.params.get("enable_native_partial_stream") if configured is not None: return bool(configured) # 默认使用整段周期重转写;仅在客户端明确要求时启用底层 native partial stream。 return False async def _init_realtime_stream_state( self, ctx: ConnectionContext, engine: Qwen3ASREngine, ) -> None: if not self._should_use_native_partial_stream(ctx, engine): return if ctx.realtime_stream_state is not None: return try: # Native stream 只用来拿低延迟 partial,最终句子仍走整段重转写定稿。 ctx.realtime_stream_state = await run_sync( engine.init_streaming_state, ctx.params.get("context", ""), ctx.params.get("language"), chunk_size_sec=self._stream_chunk_size_sec(ctx), unfixed_chunk_num=self._stream_unfixed_chunk_num(ctx), unfixed_token_num=self._stream_unfixed_token_num(ctx), max_new_tokens=48, ) except Exception as exc: logger.debug("Init realtime partial stream failed: %s", exc) ctx.realtime_stream_state = None async def _push_realtime_stream_audio( self, ctx: ConnectionContext, engine: Qwen3ASREngine, audio: np.ndarray, ) -> None: if audio.size == 0: return if not self._should_use_native_partial_stream(ctx, engine): return if ctx.realtime_stream_state is None: await self._init_realtime_stream_state(ctx, engine) if ctx.realtime_stream_state is None: return try: ctx.realtime_stream_state = await run_sync( engine.streaming_transcribe, audio, ctx.realtime_stream_state, ) except Exception as exc: logger.debug("Realtime partial stream push failed: %s", exc) ctx.realtime_stream_state = None async def _finish_realtime_stream_text( self, ctx: ConnectionContext, engine: Qwen3ASREngine, ) -> tuple[str, str]: if not self._should_use_native_partial_stream(ctx, engine): return "", "" if ctx.realtime_stream_state is None: return "", "" try: ctx.realtime_stream_state = await run_sync( engine.finish_streaming_transcribe, ctx.realtime_stream_state, ) text = self._normalize_output_text( str(ctx.realtime_stream_state.last_text or ""), ctx, enable_itn=False, ) language = self._normalize_output_language( getattr(ctx.realtime_stream_state, "last_language", ""), ctx, ) return text, language except Exception as exc: logger.debug("Finalize realtime partial stream failed: %s", exc) return "", "" finally: ctx.realtime_stream_state = None async def _decode_turn_partial_text( self, engine: Qwen3ASREngine, ctx: ConnectionContext, ) -> tuple[str, str]: if self._should_use_native_partial_stream(ctx, engine) and ctx.realtime_stream_state is not None: stream_text = self._normalize_output_text( str(ctx.realtime_stream_state.last_text or ""), ctx, enable_itn=False, ) stream_language = self._normalize_output_language( getattr(ctx.realtime_stream_state, "last_language", ""), ctx, ) if stream_text.strip(): return stream_text, stream_language # 流式文本为空时回退到当前音频窗口重转写。 partial_audio = np.asarray( ctx.stream_window_buffer if ctx.stream_window_buffer.size > 0 else ctx.segment_audio_buffer, dtype=np.float32, ) fallback = await self._transcribe_audio_text( engine, ctx, partial_audio, ) return fallback, "" def _append_pre_roll(self, ctx: ConnectionContext, audio: np.ndarray) -> None: if audio.size == 0: return ctx.pre_roll_audio = np.concatenate([ctx.pre_roll_audio, audio]) max_samples = self._pre_roll_samples(ctx) if max_samples <= 0: ctx.pre_roll_audio = np.array([], dtype=np.float32) elif ctx.pre_roll_audio.size > max_samples: ctx.pre_roll_audio = ctx.pre_roll_audio[-max_samples:] def _start_turn(self, ctx: ConnectionContext, audio: np.ndarray) -> None: parts: List[np.ndarray] = [] if ctx.pre_roll_audio.size > 0: parts.append(np.asarray(ctx.pre_roll_audio, dtype=np.float32)) parts.append(np.asarray(audio, dtype=np.float32)) ctx.segment_audio_buffer = ( np.concatenate(parts) if len(parts) > 1 else np.asarray(parts[0], dtype=np.float32) ) ctx.pre_roll_audio = np.array([], dtype=np.float32) window_samples = self._stream_window_samples() ctx.stream_window_buffer = np.asarray( ctx.segment_audio_buffer[-window_samples:], dtype=np.float32, ) partial_window_samples = self._partial_window_samples(ctx) if ctx.stream_window_buffer.size > partial_window_samples: ctx.stream_window_buffer = ctx.stream_window_buffer[-partial_window_samples:] ctx.sentence_active = True ctx.silence_samples = 0 ctx.total_samples = int(ctx.segment_audio_buffer.size) ctx.last_partial_decode_samples = 0 self._reset_partial_state(ctx) def _append_turn_audio( self, ctx: ConnectionContext, audio: np.ndarray, *, has_voice: bool, ) -> None: ctx.segment_audio_buffer = np.concatenate([ctx.segment_audio_buffer, audio]) ctx.stream_window_buffer = np.concatenate([ctx.stream_window_buffer, audio]) window_samples = self._partial_window_samples(ctx) if ctx.stream_window_buffer.size > window_samples: ctx.stream_window_buffer = ctx.stream_window_buffer[-window_samples:] ctx.total_samples = int(ctx.segment_audio_buffer.size) if has_voice: ctx.silence_samples = 0 else: ctx.silence_samples += int(audio.size) def _should_decode_turn_partial(self, ctx: ConnectionContext) -> bool: if not ctx.sentence_active: return False current_samples = int(ctx.segment_audio_buffer.size) if current_samples < self._min_partial_samples(ctx): return False return (current_samples - ctx.last_partial_decode_samples) >= self._partial_emit_interval_samples(ctx) def _reset_partial_state(self, ctx: ConnectionContext) -> None: ctx.last_partial_text = "" ctx.last_partial_language = "" ctx.last_partial_chunk_id = 0 ctx.last_partial_raw_text = "" ctx.last_partial_display_text = "" ctx.stable_partial_prefix = "" ctx.segment_observed_text = "" ctx.segment_observed_language = "" ctx.best_partial_text = "" ctx.best_partial_language = "" ctx.last_partial_decode_samples = 0 ctx.partial_stable_rounds = 0 ctx.realtime_stream_state = None def _update_partial_stability( self, ctx: ConnectionContext, text: str, ) -> int: current = text.strip() previous = ctx.last_partial_text.strip() if current and previous and current == previous: ctx.partial_stable_rounds += 1 else: ctx.partial_stable_rounds = 0 return ctx.partial_stable_rounds def _prefer_segment_text( self, primary: str, fallback: str, ) -> str: primary_text = primary.strip() fallback_text = fallback.strip() if not primary_text: return fallback_text if not fallback_text: return primary_text if primary_text == fallback_text: return primary_text if len(fallback_text) > len(primary_text) and ( primary_text in fallback_text or self._common_prefix_len(primary_text, fallback_text) >= min(len(primary_text), 8) ): return fallback_text if len(primary_text) > len(fallback_text) and ( fallback_text in primary_text or self._common_prefix_len(primary_text, fallback_text) >= min(len(fallback_text), 8) ): return primary_text return fallback_text if len(fallback_text) > len(primary_text) else primary_text def _merge_segment_text( self, stable: str, candidate: str, ) -> str: stable_text = stable.strip() candidate_text = candidate.strip() if not stable_text: return candidate_text if not candidate_text: return stable_text if stable_text == candidate_text: return stable_text if stable_text in candidate_text: return candidate_text if candidate_text in stable_text: return stable_text match = SequenceMatcher( None, stable_text, candidate_text, autojunk=False, ).find_longest_match(0, len(stable_text), 0, len(candidate_text)) if match.size >= min(6, len(stable_text), len(candidate_text)): return stable_text[:match.a] + candidate_text[match.b:] return self._prefer_segment_text(stable_text, candidate_text) def _accumulate_segment_text( self, previous: str, current: str, ) -> str: previous_text = previous.strip() current_text = current.strip() if not previous_text: return current_text if not current_text: return previous_text if previous_text == current_text: return previous_text if current_text.startswith(previous_text) or previous_text in current_text: return current_text if previous_text.startswith(current_text): return previous_text common_len = self._common_prefix_len(previous_text, current_text) if common_len >= min(len(previous_text), len(current_text), 8): return current_text if len(current_text) >= len(previous_text) else previous_text return current_text if len(current_text) >= len(previous_text) else previous_text def _sequence_similarity(self, left: str, right: str) -> float: left_text = left.strip() right_text = right.strip() if not left_text or not right_text: return 0.0 if left_text == right_text: return 1.0 return float( SequenceMatcher( None, left_text, right_text, autojunk=False, ).ratio() ) @staticmethod def _overlap_char_views(text: str) -> List[tuple[str, int]]: views: List[tuple[str, int]] = [] for idx, ch in enumerate(str(text or "")): if ch.isalnum() or "\u4e00" <= ch <= "\u9fff": views.append((ch.lower(), idx)) return views def _previous_segment_overlap_cut_index( self, previous: str, current: str, ) -> int: previous_text = str(previous or "").strip() current_text = str(current or "").strip() if not previous_text or not current_text: return 0 exact_cut = _dedupe_overlap_word_boundary(previous_text, current_text) if exact_cut == 0: exact_cut = _dedupe_overlap_char_boundary(previous_text, current_text) if exact_cut > 0: return exact_cut prev_views = self._overlap_char_views(previous_text) curr_views = self._overlap_char_views(current_text) if len(prev_views) < 12 or len(curr_views) < 12: return 0 lookback = min(140, len(prev_views)) lookahead = min(180, len(curr_views)) prev_tail = prev_views[-lookback:] curr_prefix = curr_views[:lookahead] prev_key = "".join(ch for ch, _ in prev_tail) curr_key = "".join(ch for ch, _ in curr_prefix) match = SequenceMatcher(None, prev_key, curr_key, autojunk=False).find_longest_match( 0, len(prev_key), 0, len(curr_key), ) if match.size < 12: return 0 prev_suffix_gap = len(prev_key) - (match.a + match.size) current_prefix_gap = match.b if prev_suffix_gap > 4 or current_prefix_gap > 8: return 0 matched_current_chars = match.b + match.size if matched_current_chars >= len(curr_views): return len(current_text) return curr_views[matched_current_chars][1] def _trim_previous_segment_overlap( self, ctx: ConnectionContext, text: str, ) -> str: current = str(text or "").strip() if not current or not ctx.confirmed_segments: return current previous = str(ctx.confirmed_segments[-1].get("text") or "").strip() cut_index = self._previous_segment_overlap_cut_index(previous, current) if cut_index <= 0: return current trimmed = current[cut_index:].lstrip(" \t\r\n,,。.!!??;;::、") if trimmed != current: logger.debug( "Trimmed realtime segment overlap: removed=%s remaining=%s", cut_index, len(trimmed), ) return trimmed def _is_unstable_expansion( self, stable: str, candidate: str, ) -> bool: stable_text = stable.strip() candidate_text = candidate.strip() if len(stable_text) < 6 or len(candidate_text) <= len(stable_text): return False common_prefix = self._common_prefix_len(stable_text, candidate_text) similarity = self._sequence_similarity(stable_text, candidate_text) grows_too_fast = len(candidate_text) >= len(stable_text) + max(20, len(stable_text) // 2) loses_context = common_prefix < min(6, len(stable_text) // 2) return grows_too_fast and loses_context and similarity < 0.45 def _is_degenerate_repetition(self, text: str) -> bool: stripped = text.strip() if len(stripped) < 16: return False units = re.findall(r"[\u4e00-\u9fff]+|[A-Za-z0-9]+|[^\w\s]", stripped) lexical_units = [unit for unit in units if re.search(r"[\u4e00-\u9fffA-Za-z0-9]", unit)] if len(lexical_units) < 6: return False counts: Dict[str, int] = {} for unit in lexical_units: counts[unit] = counts.get(unit, 0) + 1 dominant = max(counts.values()) if counts else 0 unique = len(counts) if dominant >= 6 and unique <= 2: return True if dominant / max(len(lexical_units), 1) >= 0.7 and len(lexical_units) >= 10: return True return False def _trim_repetitive_suffix(self, text: str) -> str: trimmed = text.strip() if len(trimmed) < 12: return trimmed repeated_char = re.search(r"(.)\1{7,}$", trimmed) if repeated_char: start = repeated_char.start() trimmed = (trimmed[:start] + repeated_char.group(1) * 2).strip() for unit_len in range(1, 7): pattern = re.compile(rf"(.{{{unit_len}}})\1{{4,}}$") match = pattern.search(trimmed) if match: start = match.start() unit = match.group(1) trimmed = (trimmed[:start] + unit * 2).strip() break return trimmed def _sanitize_candidate_text(self, text: str) -> str: candidate = text.strip() if not candidate: return "" candidate = self._trim_repetitive_suffix(candidate) candidate = deduplicate_asr_text(candidate) if self._is_degenerate_repetition(candidate): return "" return candidate def _update_segment_observed_text( self, ctx: ConnectionContext, text: str, language: str, ) -> str: candidate = self._sanitize_candidate_text(text) if not candidate: return ctx.segment_observed_text if self._is_unstable_expansion(ctx.segment_observed_text, candidate): return ctx.segment_observed_text observed = self._accumulate_segment_text(ctx.segment_observed_text, candidate) ctx.segment_observed_text = observed if language: ctx.segment_observed_language = language return observed def _update_best_partial( self, ctx: ConnectionContext, text: str, language: str, ) -> None: candidate = self._sanitize_candidate_text(text) if not candidate: return if self._is_unstable_expansion(ctx.best_partial_text, candidate): return chosen = self._prefer_segment_text(ctx.best_partial_text, candidate) if chosen != ctx.best_partial_text: ctx.best_partial_text = chosen ctx.best_partial_language = language or ctx.best_partial_language def _common_prefix_len(self, left: str, right: str) -> int: limit = min(len(left), len(right)) idx = 0 while idx < limit and left[idx] == right[idx]: idx += 1 return idx def _stable_prefix_cutoff(self, text: str, max_commit_len: int) -> int: if max_commit_len <= 0: return 0 candidate = text[:max_commit_len] if not candidate: return 0 for idx in range(len(candidate) - 1, -1, -1): ch = candidate[idx] if ch.isspace() or unicodedata.category(ch).startswith("P"): return idx + 1 return len(candidate) def _stabilize_partial_text(self, ctx: ConnectionContext, raw_text: str) -> str: text = self._sanitize_candidate_text(raw_text) if not text: ctx.last_partial_raw_text = "" ctx.stable_partial_prefix = "" return "" previous_raw = ctx.last_partial_raw_text stable_prefix = ctx.stable_partial_prefix previous_emitted = ctx.last_partial_text.strip() if previous_raw: common_len = self._common_prefix_len(previous_raw, text) keep_tail_chars = max(int(settings.REALTIME_STREAM_STABLE_TAIL_CHARS), 2) max_commit_len = max(0, common_len - keep_tail_chars) stable_cutoff = self._stable_prefix_cutoff(text, max_commit_len) if stable_cutoff - len(stable_prefix) >= max( int(settings.REALTIME_STREAM_STABLE_MIN_GROW_CHARS), 1, ): stable_prefix = text[:stable_cutoff] if stable_prefix and not text.startswith(stable_prefix): stable_common = self._common_prefix_len(stable_prefix, text) tolerated_divergence = max( int(settings.REALTIME_STREAM_DIVERGENCE_TOLERANCE_CHARS), 1, ) if len(stable_prefix) - stable_common >= tolerated_divergence and previous_emitted: return previous_emitted stable_prefix = text[:stable_common] stable_prefix = text[: self._stable_prefix_cutoff(text, len(stable_prefix))] ctx.stable_partial_prefix = stable_prefix ctx.last_partial_raw_text = text unstable_suffix = text[len(stable_prefix):] return f"{stable_prefix}{unstable_suffix}".strip() def _should_emit_partial( self, ctx: ConnectionContext, text: str, language: str, ) -> bool: stripped = self._sanitize_candidate_text(text) if not stripped: return False if len(stripped) <= 1: return False if self._is_unstable_expansion(ctx.last_partial_display_text, stripped): return False normalized_language = self._normalize_output_language(language, ctx) if not normalized_language: normalized_language = self._infer_text_language(stripped) if ( stripped == ctx.last_partial_display_text and normalized_language == ctx.last_partial_language ): return False dominant_language = self._get_session_dominant_language(ctx) duration_sec = float(ctx.segment_audio_buffer.size) / 16000.0 if len(stripped) <= 6 and self._is_suspicious_segment_text( stripped, language=normalized_language, dominant_language=dominant_language, duration_sec=max(duration_sec, 0.1), explicit_language=ctx.params.get("language"), ): return False return True def _partial_display_text(self, ctx: ConnectionContext, text: str) -> str: stripped = self._sanitize_candidate_text(text) if not stripped: return "" holdback_chars = max( int( getattr( settings, "REALTIME_PARTIAL_HOLDBACK_CHARS", 0, ) ), 0, ) holdback_chars = int( max( int(ctx.params.get("partial_holdback_chars", holdback_chars) or holdback_chars), 0, ) ) if holdback_chars <= 0: return stripped if ctx.partial_stable_rounds >= 1: return stripped if self._sentence_count(stripped) >= 1 and stripped[-1:] in "。!?.!?": return stripped if len(stripped) <= max(holdback_chars + 2, 8): return stripped stable_prefix = stripped[:-holdback_chars].rstrip() return stable_prefix or stripped def _is_suspicious_segment_text( self, text: str, *, language: str, dominant_language: str, duration_sec: float, explicit_language: Optional[str], ) -> bool: stripped = text.strip() if not stripped: return True chars = [ch for ch in stripped if not ch.isspace()] if not chars: return True digit_count = sum(ch.isdigit() for ch in chars) alpha_count = sum(ch.isalpha() for ch in chars) cjk_count = sum("\u4e00" <= ch <= "\u9fff" for ch in chars) thai_count = sum("\u0e00" <= ch <= "\u0e7f" for ch in chars) punct_count = sum(unicodedata.category(ch).startswith("P") for ch in chars) total = len(chars) digit_ratio = digit_count / total punct_ratio = punct_count / total latin_alpha_count = max(alpha_count - thai_count, 0) script_families = sum( 1 for present in ( cjk_count > 0, thai_count > 0, latin_alpha_count > 0, digit_count > 0, ) if present ) if digit_count >= 3 and digit_ratio >= 0.3: return True if total <= 20 and script_families >= 3: return True if total <= 8 and punct_ratio >= 0.5: return True # Allow short pure-Latin utterances to pass during bilingual switching. if ( not explicit_language and language == "English" and latin_alpha_count >= max(total - punct_count - digit_count, 1) and total >= 3 ): return False if ( not explicit_language and dominant_language and language and language != dominant_language and duration_sec <= 6.0 ): return True return False def _is_valid_committed_segment_text(self, text: str, duration_sec: float) -> bool: stripped = self._sanitize_candidate_text(text) if len(stripped) >= 3: return True lexical_count = sum( 1 for ch in stripped if ch.isalnum() or "\u4e00" <= ch <= "\u9fff" ) if lexical_count <= 0: return False # Short acknowledgements such as "嗯。", "好。", "啊?" should still # close the visible segment after silence; speaker attribution may # attach to nearby speech, but the ASR turn itself is valid. return duration_sec >= 0.25 async def _estimate_voiced_duration_ms(self, audio: np.ndarray) -> int: if audio.size == 0: return 0 vad_segments = await self._run_vad_segments(audio) if vad_segments is None: return int(len(audio) / 16) return sum( max(0, int(seg[1]) - int(seg[0])) for seg in vad_segments if len(seg) >= 2 ) async def _run_vad_segments(self, audio: np.ndarray) -> Optional[list[list[int]]]: if audio.size == 0: return [] temp_path: Optional[str] = None try: temp_path = get_speaker_registry_service().save_audio_array_to_temp( audio, sample_rate=16000, ) def _run_vad() -> list[list[int]]: vad_model = get_global_vad_model(settings.DEVICE) with get_vad_inference_lock(): result = vad_model.generate(input=temp_path, cache={}) return result[0].get("value", []) if result else [] return await run_sync(_run_vad) except Exception as exc: logger.debug("Realtime voiced-duration VAD failed: %s", exc) return None finally: get_speaker_registry_service().cleanup_file(temp_path) async def _split_silence_audio_by_vad( self, audio: np.ndarray, ) -> tuple[np.ndarray, np.ndarray]: vad_segments = await self._run_vad_segments(audio) if not vad_segments: return audio, np.array([], dtype=np.float32) valid_segments = [ (int(seg[0]), int(seg[1])) for seg in vad_segments if len(seg) >= 2 and int(seg[1]) > int(seg[0]) ] if not valid_segments: return audio, np.array([], dtype=np.float32) finalized_end_ms = valid_segments[-1][1] finalized_end_sample = min(audio.size, int(finalized_end_ms * 16)) finalized_audio = np.asarray(audio[:finalized_end_sample], dtype=np.float32) return finalized_audio, np.array([], dtype=np.float32) async def _find_completed_segment_split_sample( self, audio: np.ndarray, ) -> Optional[int]: if audio.size < int(3.0 * 16000): return None vad_segments = await self._run_vad_segments(audio) if not vad_segments: return None valid_segments = [ (int(seg[0]), int(seg[1])) for seg in vad_segments if len(seg) >= 2 and int(seg[1]) > int(seg[0]) ] if len(valid_segments) < 2: return None total_ms = int(audio.size / 16) finalize_silence_ms = int(max(settings.REALTIME_VAD_FINALIZE_SILENCE_SEC, 0.2) * 1000) last_start_ms, last_end_ms = valid_segments[-1] if total_ms - last_end_ms >= finalize_silence_ms: return min(audio.size, last_end_ms * 16) # 接近 max_duration 时,优先把倒数第二个已完成语音段刷出去,保留最后一段继续等上下文。 hard_limit_sec = max( float(getattr(settings, "REALTIME_FINAL_SEGMENT_HARD_LIMIT_SEC", 12.0) or 12.0), 1.0, ) near_limit_ms = int(max(hard_limit_sec * 0.85, 4.0) * 1000) if total_ms >= near_limit_ms: _, prev_end_ms = valid_segments[-2] if prev_end_ms > 0: return min(audio.size, prev_end_ms * 16) return None async def _transcribe_audio_text( self, engine: Qwen3ASREngine, ctx: ConnectionContext, audio: np.ndarray, *, final_pass: bool = False, ) -> str: temp_path: Optional[str] = None try: temp_path = get_speaker_registry_service().save_audio_array_to_temp( audio, sample_rate=16000, ) text = await run_sync( engine.transcribe_file, temp_path, ctx.params.get("context", ""), False, False, False, 16000, ) normalized = self._normalize_output_text( text or "", ctx, enable_itn=False, ) if final_pass and normalized.strip(): return self._normalize_output_text( normalized, ctx, enable_itn=bool(ctx.params.get("enable_inverse_text_normalization", True)), ) return normalized except Exception as exc: logger.warning("Realtime segment audio transcription failed: %s", exc) return "" finally: get_speaker_registry_service().cleanup_file(temp_path) async def _transcribe_audio_longform_text( self, engine: Qwen3ASREngine, ctx: ConnectionContext, audio: np.ndarray, *, final_pass: bool = False, ) -> str: duration_sec = float(len(audio)) / 16000.0 if duration_sec < settings.REALTIME_LONGFORM_MIN_SEC: return await self._transcribe_audio_text( engine, ctx, audio, final_pass=final_pass, ) chunk_samples = int(max(settings.REALTIME_LONGFORM_CHUNK_SEC, 1.0) * 16000) overlap_samples = int(max(settings.REALTIME_LONGFORM_OVERLAP_SEC, 0.0) * 16000) step_samples = max(int(16000), chunk_samples - overlap_samples) assembler = RealtimeTranscriptAssembler() start = 0 while start < audio.size: end = min(audio.size, start + chunk_samples) chunk_audio = np.asarray(audio[start:end], dtype=np.float32) if chunk_audio.size == 0: break chunk_text = await self._transcribe_audio_text( engine, ctx, chunk_audio, final_pass=final_pass, ) assembler.push(chunk_text) if end >= audio.size: break start += step_samples return assembler.text() def _split_max_duration_audio( self, audio: np.ndarray, ) -> tuple[np.ndarray, np.ndarray]: tail_samples = int(max(settings.REALTIME_MAX_SEGMENT_TAIL_SEC, 0.0) * 16000) min_flush_samples = int(max(settings.REALTIME_SPEAKER_MIN_SEC, 1.0) * 16000) if tail_samples <= 0 or audio.size <= tail_samples + min_flush_samples: return audio, np.array([], dtype=np.float32) flush_audio = np.asarray(audio[:-tail_samples], dtype=np.float32) carry_audio = np.asarray(audio[-tail_samples:], dtype=np.float32) return flush_audio, carry_audio def _should_force_stable_segment(self, ctx: ConnectionContext, text: str) -> bool: force_sec = float( max( float( ctx.params.get( "force_stable_segment_sec", getattr(settings, "REALTIME_FORCE_STABLE_SEGMENT_SEC", 0.0), ) or 0.0 ), 0.0, ) ) if force_sec <= 0: return False duration_sec = float(ctx.segment_audio_buffer.size) / 16000.0 if duration_sec < force_sec: return False partial_text = text.strip() min_chars = max( int( ctx.params.get( "force_stable_min_chars", getattr(settings, "REALTIME_FORCE_STABLE_MIN_CHARS", 24), ) or 0 ), 8, ) if len(partial_text) < min_chars: return False # 只因为前文某处已经出现过句号,就立刻截断整段,容易把后半句切坏。 # 这里收紧条件:优先要求“当前尾部已经自然收句”。 if self._ends_with_sentence_punctuation(partial_text): return True # 如果句尾还没闭合,只在明显超时并且 partial 连续稳定时才软提交, # 避免把“我们需要…”这类正在继续的从句过早切成两段。 force_overtime_sec = max(force_sec + 2.0, force_sec * 1.35) soft_limit_samples = self._soft_segment_limit_samples(ctx) if soft_limit_samples > 0 and ctx.segment_audio_buffer.size >= soft_limit_samples: force_overtime_sec = min( force_overtime_sec, float(ctx.segment_audio_buffer.size) / 16000.0, ) if duration_sec >= force_overtime_sec and ctx.partial_stable_rounds >= 1: return True return ctx.partial_stable_rounds >= 2 def _choose_committed_segment_text( self, *, final_text: str, stable_text: str, ctx: ConnectionContext, duration_sec: float, language: str, ) -> str: """Prefer the more plausible candidate between final retranscription and stable partial text.""" _ = (ctx, language) final_candidate = self._sanitize_candidate_text(final_text) stable_candidate = self._sanitize_candidate_text(stable_text) if not final_candidate: return stable_candidate if not stable_candidate: return final_candidate if final_candidate == stable_candidate: return final_candidate suspicious_final = self._is_suspicious_segment_text( final_candidate, language=self._infer_text_language(final_candidate), dominant_language=self._get_session_dominant_language(ctx), duration_sec=max(duration_sec, 0.1), explicit_language=ctx.params.get("language"), ) suspicious_stable = self._is_suspicious_segment_text( stable_candidate, language=self._infer_text_language(stable_candidate), dominant_language=self._get_session_dominant_language(ctx), duration_sec=max(duration_sec, 0.1), explicit_language=ctx.params.get("language"), ) if suspicious_final and not suspicious_stable: return stable_candidate if ( len(stable_candidate) >= max(len(final_candidate) + 10, 24) and final_candidate in stable_candidate ): return stable_candidate if ( len(final_candidate) >= max(len(stable_candidate) + 12, 28) and stable_candidate in final_candidate and not suspicious_final ): return final_candidate if len(stable_candidate) > len(final_candidate) and not suspicious_stable: return stable_candidate return final_candidate async def _resolve_segment_speaker( self, ctx: ConnectionContext, audio: np.ndarray, *, reason: str, segment_start_ms: int, segment_end_ms: int, ) -> Optional[Dict[str, Any]]: if not bool(ctx.params.get("enable_speaker", True)): return None speaker_info = await get_realtime_speaker_clusterer().resolve_segment_speaker( self._speaker_match_records(ctx), np.asarray(audio, dtype=np.float32), segment_start_ms=segment_start_ms, segment_end_ms=segment_end_ms, enable_registry_match=bool(ctx.params.get("enable_speaker_identification", True)), speaker_threshold=ctx.params.get("speaker_threshold"), ) if not speaker_info: return None speaker_info["_segment_duration_ms"] = int(max(segment_end_ms - segment_start_ms, 0)) inherited_named_speaker = self._inherit_recent_named_speaker(ctx, speaker_info) if inherited_named_speaker is not None: speaker_info = inherited_named_speaker inherited_speaker = self._inherit_recent_anonymous_speaker(ctx, speaker_info) if inherited_speaker is not None: speaker_info = inherited_speaker stable = self._is_stable_speaker_info(speaker_info) speaker_info["speaker_pending"] = not stable cluster_index = speaker_info.get("_cluster_index") if ( stable and not speaker_info.get("registry_speaker_id") and not speaker_info.get("user_id") ): safe_cluster_index = self._safe_int(cluster_index) display_id = self._match_display_speaker_id( ctx, cluster_index=safe_cluster_index, embedding=speaker_info.get("_embedding"), ) if display_id: speaker_info["speaker_id"] = display_id speaker_info["speaker_name"] = display_id safe_cluster_index = self._safe_int(cluster_index) if ( stable and safe_cluster_index is not None and not speaker_info.get("registry_speaker_id") and not speaker_info.get("user_id") and speaker_info.get("speaker_id") ): ctx.speaker_display_map[safe_cluster_index] = str(speaker_info["speaker_id"]) if not stable: speaker_info["speaker_id"] = -1 speaker_info["speaker_name"] = "" return speaker_info async def _resolve_and_emit_segment_speaker( self, websocket: WebSocket, ctx: ConnectionContext, task_id: str, *, segment_index: int, audio: np.ndarray, reason: str, segment_start_ms: int, segment_end_ms: int, ) -> None: try: speaker_info = await self._resolve_segment_speaker( ctx, audio, reason=reason, segment_start_ms=segment_start_ms, segment_end_ms=segment_end_ms, ) if not speaker_info: return segment = next( (item for item in ctx.confirmed_segments if int(item.get("index", -1)) == int(segment_index)), None, ) if segment is None: return before_view = self._public_speaker_view(segment) segment.update(speaker_info) after_view = self._public_speaker_view(segment) self._record_segment_speaker( ctx, segment_index=segment_index, segment_start_ms=segment_start_ms, segment_end_ms=segment_end_ms, speaker_info=speaker_info, ) if before_view == after_view: return queue_backlog = ctx.speaker_job_queue.qsize() updates: List[Dict[str, Any]] = [] if queue_backlog <= 1: updates = self._maybe_recluster_recent_segments(ctx) if updates: for updated in updates: await self._emit_speaker_update( websocket, ctx, task_id, int(updated["segment_index"]), ) return await self._emit_speaker_update(websocket, ctx, task_id, segment_index) except Exception as exc: logger.warning("[%s] Async speaker resolution failed for segment %s: %s", task_id, segment_index, exc) async def _refine_segment_text( self, engine: Qwen3ASREngine, ctx: ConnectionContext, fallback_text: str, *, reason: str, force: bool = False, audio: Optional[np.ndarray] = None, ) -> str: target_audio = np.asarray( ctx.segment_audio_buffer if audio is None else audio, dtype=np.float32, ) duration_sec = float(len(target_audio)) / 16000.0 if ( not force and not bool( ctx.params.get("enable_segment_refine", settings.REALTIME_ENABLE_SEGMENT_REFINE) ) ): return fallback_text if not force and duration_sec < 6.0 and reason != "max_duration": return fallback_text temp_path: Optional[str] = None try: temp_path = get_speaker_registry_service().save_audio_array_to_temp( target_audio, sample_rate=16000, ) refined_text = await run_sync( engine.transcribe_file, temp_path, ctx.params.get("context", ""), False, False, False, 16000, ) refined_text = self._sanitize_candidate_text( self._normalize_output_text( refined_text or "", ctx, enable_itn=bool(ctx.params.get("enable_inverse_text_normalization", True)), ) ) if refined_text.strip(): return refined_text except Exception as exc: logger.warning("Realtime segment refinement failed: %s", exc) finally: get_speaker_registry_service().cleanup_file(temp_path) return fallback_text async def _commit_retranscribe_turn( self, websocket: WebSocket, ctx: ConnectionContext, task_id: str, reason: str, *, finalized_audio_override: Optional[np.ndarray] = None, carry_audio_override: Optional[np.ndarray] = None, emit_segment_start: bool = True, ) -> None: engine = await self._ensure_engine(ctx) current_audio = np.asarray(ctx.segment_audio_buffer, dtype=np.float32) finalized_audio = current_audio carry_audio = np.array([], dtype=np.float32) if finalized_audio_override is not None: finalized_audio = np.asarray(finalized_audio_override, dtype=np.float32) carry_audio = np.asarray( np.array([], dtype=np.float32) if carry_audio_override is None else carry_audio_override, dtype=np.float32, ) elif reason == "silence" and self._enable_realtime_vad_split(ctx): finalized_audio, carry_audio = await self._split_silence_audio_by_vad(current_audio) elif reason in {"max_duration", "long_speech"} and self._enable_realtime_vad_split(ctx): finalized_audio, carry_audio = self._split_max_duration_audio(current_audio) if carry_audio.size == 0: stream_text, stream_language = await self._finish_realtime_stream_text(ctx, engine) if stream_text.strip(): observed_text = self._update_segment_observed_text( ctx, stream_text, stream_language or self._infer_text_language(stream_text), ) self._update_best_partial( ctx, observed_text, stream_language or self._infer_text_language(observed_text), ) else: ctx.realtime_stream_state = None if self._enable_realtime_longform(ctx): segment_text = await self._transcribe_audio_longform_text( engine, ctx, finalized_audio, final_pass=True, ) else: segment_text = await self._transcribe_audio_text( engine, ctx, finalized_audio, final_pass=True, ) segment_text = self._sanitize_candidate_text(segment_text) if not segment_text.strip(): segment_text = self._prefer_segment_text( ctx.best_partial_text, ctx.segment_observed_text, ) segment_language = self._infer_text_language(segment_text) if not segment_language and ctx.segment_observed_language: segment_language = ctx.segment_observed_language if not segment_language and ctx.best_partial_language: segment_language = ctx.best_partial_language duration_sec = float(finalized_audio.size) / 16000.0 dominant_language = self._get_session_dominant_language(ctx) suspicious = self._is_suspicious_segment_text( segment_text, language=segment_language, dominant_language=dominant_language, duration_sec=max(duration_sec, 0.1), explicit_language=ctx.params.get("language"), ) if suspicious and finalized_audio.size > 0 and self._enable_realtime_refine(ctx): refined_text = await self._refine_segment_text( engine, ctx, segment_text, reason=reason, force=True, audio=finalized_audio, ) if refined_text.strip(): segment_text = refined_text inferred = self._infer_text_language(segment_text) if inferred: segment_language = inferred stable_text = self._prefer_segment_text( ctx.best_partial_text, ctx.segment_observed_text, ) segment_text = self._choose_committed_segment_text( final_text=segment_text, stable_text=stable_text, ctx=ctx, duration_sec=duration_sec, language=segment_language, ) segment_text = self._trim_previous_segment_overlap(ctx, segment_text) is_valid = self._is_valid_committed_segment_text(segment_text, duration_sec) segment_duration_ms = int(finalized_audio.size / 16) segment_start_ms = int(ctx.timeline_cursor_ms) segment_end_ms = int(segment_start_ms + segment_duration_ms) if is_valid: segment_payload = { "index": ctx.segment_index, "text": segment_text, "language": segment_language, "reason": reason, "duration_ms": segment_duration_ms, "start_ms": segment_start_ms, "end_ms": segment_end_ms, "sentence_type": 1, "speaker_id": -1, "speaker_name": "", "user_id": None, } segment_payload["speaker_pending"] = True ctx.confirmed_segments.append(segment_payload) full_text = "\n".join(segment["text"] for segment in ctx.confirmed_segments) full_text_tail = ( full_text[-self._SEGMENT_EVENT_TEXT_TAIL_CHARS:] if len(full_text) > self._SEGMENT_EVENT_TEXT_TAIL_CHARS else full_text ) sentence_payload = self._build_tencent_sentence( segment_payload, sentence_type=1, ) await self._send_json_safe( websocket, ctx, { "type": "sentences", "code": 0, "voice_id": task_id, "final": 0, "result": { "slice_type": 2, "index": int(ctx.segment_index), "voice_text_str": full_text_tail, }, "sentences": [sentence_payload], }, ) self._ensure_speaker_worker(websocket, ctx, task_id) ctx.speaker_job_queue.put_nowait( { "segment_index": int(ctx.segment_index), "audio": np.asarray(finalized_audio, dtype=np.float32), "reason": reason, "segment_start_ms": segment_start_ms, "segment_end_ms": segment_end_ms, } ) ctx.segment_index += 1 ctx.timeline_cursor_ms = segment_end_ms ctx.segment_audio_buffer = np.asarray(carry_audio, dtype=np.float32) ctx.stream_window_buffer = np.asarray(carry_audio, dtype=np.float32) ctx.sentence_active = bool(carry_audio.size > 0) ctx.silence_samples = 0 ctx.total_samples = int(carry_audio.size) self._reset_partial_state(ctx) if ctx.sentence_active and carry_audio.size > 0: await self._push_realtime_stream_audio( ctx, engine, np.asarray(carry_audio, dtype=np.float32), ) async def _send_error( self, websocket: WebSocket, message: str, task_id: str, code: str = "DEFAULT_SERVER_ERROR", ) -> None: try: error = create_error_response(error_code=code, message=message, task_id=task_id) await websocket.send_json( { "type": "error", "code": error.get("code", -1), "message": error.get("message", message), "voice_id": task_id, } ) except Exception: pass async def handle_connection(self, websocket: WebSocket, task_id: str) -> None: await websocket.accept() logger.info("[%s] Qwen3 websocket connected", task_id) ctx = ConnectionContext() session_id = task_id keep_for_resume = False try: while True: message = await websocket.receive() if "text" in message: data = json.loads(message["text"]) msg_type = data.get("type", "") if msg_type == "start": if ctx.state != ConnectionState.READY: await self._send_error(websocket, "识别已在进行中", task_id, "INVALID_STATE") continue payload = data.get("payload", {}) requested_session_id = str(payload.get("session_id") or task_id).strip() or task_id resumed_ctx, resumed = await self._acquire_session(requested_session_id) if resumed_ctx is not ctx: ctx = resumed_ctx session_id = requested_session_id task_id = session_id keep_for_resume = True if resumed and ctx.state != ConnectionState.READY: await self._send_json_safe( websocket, ctx, { "type": "voice_id", "voice_id": task_id, "session_id": session_id, } ) await self._send_json_safe( websocket, ctx, { "type": "start", "session_id": session_id, } ) logger.info("[%s] Recognition resumed", task_id) continue match_speaker_registry = bool( payload.get( "match_speaker_registry", payload.get("enable_speaker_identification", False), ) ) ctx.params = { "format": payload.get("format", "pcm"), "sample_rate": payload.get("sample_rate", 16000), "language": payload.get("language"), "context": payload.get("context", ""), "enable_inverse_text_normalization": payload.get( "enable_inverse_text_normalization", True, ), "chunk_size_sec": payload.get( "chunk_size_sec", settings.REALTIME_STREAM_CHUNK_SEC, ), "unfixed_chunk_num": payload.get( "unfixed_chunk_num", settings.REALTIME_STREAM_MAX_PENDING_CHUNKS, ), "unfixed_token_num": payload.get("unfixed_token_num", 5), "silence_duration_ms": payload.get("silence_duration_ms", 800), "min_partial_sec": payload.get( "min_partial_sec", 0.9, ), "partial_window_sec": payload.get( "partial_window_sec", settings.REALTIME_PARTIAL_WINDOW_SEC, ), "pre_roll_ms": payload.get("pre_roll_ms", 240), "max_sentence_count": payload.get("max_sentence_count", 8), "max_partial_text_chars": payload.get( "max_partial_text_chars", ctx.DEFAULT_MAX_PARTIAL_TEXT_CHARS, ), "partial_holdback_chars": payload.get( "partial_holdback_chars", settings.REALTIME_PARTIAL_HOLDBACK_CHARS, ), "enable_native_partial_stream": payload.get( "enable_native_partial_stream", False, ), # 与离线会议接口保持一致:前端优先传 enable_speaker / match_speaker_registry。 "enable_speaker": payload.get("enable_speaker", True), "match_speaker_registry": match_speaker_registry, "enable_speaker_identification": match_speaker_registry, "speaker_threshold": payload.get("speaker_threshold"), "enable_segment_refine": payload.get( "enable_segment_refine", settings.REALTIME_ENABLE_SEGMENT_REFINE, ), "enable_realtime_longform": payload.get("enable_realtime_longform", False), "enable_realtime_vad_split": payload.get("enable_realtime_vad_split", False), "force_stable_segment_sec": payload.get( "force_stable_segment_sec", settings.REALTIME_FORCE_STABLE_SEGMENT_SEC, ), "force_stable_min_chars": payload.get( "force_stable_min_chars", settings.REALTIME_FORCE_STABLE_MIN_CHARS, ), "soft_limit_sec": payload.get( "soft_limit_sec", settings.REALTIME_FINAL_SEGMENT_SOFT_LIMIT_SEC, ), "hard_limit_sec": payload.get( "hard_limit_sec", payload.get( "max_segment_sec", settings.REALTIME_FINAL_SEGMENT_HARD_LIMIT_SEC, ), ), } await self._ensure_engine(ctx) ctx.stream_window_buffer = np.array([], dtype=np.float32) self._reset_partial_state(ctx) await self._send_json_safe( websocket, ctx, { "type": "voice_id", "voice_id": task_id, "session_id": session_id, } ) await self._send_json_safe( websocket, ctx, { "type": "start", "session_id": session_id, }, ) ctx.state = ConnectionState.STARTED logger.info("[%s] Recognition started: %s", task_id, ctx.params) elif msg_type == "stop": if ctx.state in (ConnectionState.STARTED, ConnectionState.STREAMING): await self._stop(websocket, ctx, task_id) keep_for_resume = False break else: await self._send_error( websocket, f"未知消息类型: {msg_type}", task_id, "INVALID_MESSAGE", ) elif "bytes" in message: if ctx.state not in (ConnectionState.STARTED, ConnectionState.STREAMING): await self._send_error(websocket, "请先发送 start", task_id, "INVALID_STATE") continue audio = _convert_audio( message["bytes"], ctx.params["format"], ctx.params["sample_rate"], ) if audio is None: continue has_voice = self._has_voice(audio) if not ctx.sentence_active: if not has_voice: self._append_pre_roll(ctx, audio) continue self._start_turn(ctx, audio) engine = await self._ensure_engine(ctx) await self._push_realtime_stream_audio( ctx, engine, np.asarray(ctx.segment_audio_buffer, dtype=np.float32), ) else: self._append_turn_audio(ctx, audio, has_voice=has_voice) engine = await self._ensure_engine(ctx) await self._push_realtime_stream_audio(ctx, engine, audio) if not ctx.sentence_active: continue if self._should_decode_turn_partial(ctx): current, current_language = await self._decode_turn_partial_text( engine, ctx, ) ctx.last_partial_decode_samples = int(ctx.segment_audio_buffer.size) current = self._sanitize_candidate_text(current) if current and not self._is_degenerate_repetition(current): current_language = current_language or self._infer_text_language(current) observed = self._update_segment_observed_text( ctx, current, current_language, ) visible_observed = self._trim_previous_segment_overlap(ctx, observed) partial_display = self._clip_partial_text(visible_observed, ctx) if ( partial_display and ( partial_display != ctx.last_partial_display_text or current_language != ctx.last_partial_language ) ): ctx.last_partial_chunk_id += 1 ctx.last_partial_text = visible_observed.strip() ctx.last_partial_display_text = partial_display.strip() ctx.last_partial_language = current_language self._update_best_partial(ctx, visible_observed, current_language) sentence_payload = self._build_tencent_sentence( { "index": ctx.segment_index, "start_ms": int(ctx.timeline_cursor_ms), "end_ms": int(ctx.timeline_cursor_ms + (ctx.segment_audio_buffer.size / 16)), "text": partial_display, "speaker_id": -1, "sentence_type": 0, }, sentence_type=0, sentence_text=partial_display, ) await self._send_json_safe( websocket, ctx, { "type": "sentences", "code": 0, "voice_id": task_id, "final": 0, "result": { "slice_type": 1, "index": int(ctx.segment_index), "voice_text_str": partial_display, }, "sentences": [sentence_payload], }, ) ctx.state = ConnectionState.STREAMING if self._sentence_count(visible_observed) >= self._max_sentence_count(ctx): await self._commit_retranscribe_turn( websocket, ctx, task_id, "sentence_limit", ) continue if self._should_commit_complete_sentence(ctx, visible_observed): await self._commit_retranscribe_turn( websocket, ctx, task_id, "complete_sentence", ) continue silence_threshold = self._get_dynamic_silence_threshold_samples(ctx) if ctx.silence_samples >= silence_threshold: await self._commit_retranscribe_turn( websocket, ctx, task_id, "silence", ) ctx.state = ConnectionState.STREAMING continue max_segment_samples = self._max_segment_samples(ctx) if max_segment_samples > 0 and ctx.segment_audio_buffer.size >= max_segment_samples: if not self._enable_realtime_vad_split(ctx): await self._commit_retranscribe_turn( websocket, ctx, task_id, "max_duration", ) ctx.state = ConnectionState.STREAMING continue split_sample = await self._find_completed_segment_split_sample( ctx.segment_audio_buffer, ) if split_sample is not None and 0 < split_sample < ctx.segment_audio_buffer.size: finalized_audio = np.asarray( ctx.segment_audio_buffer[:split_sample], dtype=np.float32, ) carry_audio = np.asarray( ctx.segment_audio_buffer[split_sample:], dtype=np.float32, ) await self._commit_retranscribe_turn( websocket, ctx, task_id, "long_speech", finalized_audio_override=finalized_audio, carry_audio_override=carry_audio, ) else: await self._commit_retranscribe_turn( websocket, ctx, task_id, "max_duration", ) ctx.state = ConnectionState.STREAMING continue except WebSocketDisconnect: logger.info("[%s] WebSocket disconnected", task_id) except Exception as exc: logger.error("[%s] Connection error: %s", task_id, exc) await self._send_error(websocket, str(exc), task_id) finally: await self._detach_session(session_id, keep_for_resume=keep_for_resume) logger.info("[%s] Connection closed", task_id) async def _stop( self, websocket: WebSocket, ctx: ConnectionContext, task_id: str, ) -> None: try: if ctx.sentence_active and ctx.segment_audio_buffer.size > 0: await self._commit_retranscribe_turn( websocket, ctx, task_id, "final", emit_segment_start=False, ) final_updates = self._recluster_confirmed_segments(ctx) for updated in final_updates: sentence_payload = self._build_tencent_sentence( updated["segment"], sentence_type=1, ) await websocket.send_json( { "type": "sentences", "code": 0, "voice_id": task_id, "final": 0, "result": { "slice_type": 2, "index": int(updated["segment_index"]), "voice_text_str": str(updated["segment"].get("text") or ""), }, "sentences": [sentence_payload], } ) all_texts = [ segment["text"] for segment in ctx.confirmed_segments if segment["text"].strip() ] full_text = "\n".join(all_texts) public_segments = [ self._public_payload(segment) or {} for segment in ctx.confirmed_segments ] await websocket.send_json( { "type": "end", "code": 0, "message": "", "voice_id": task_id, "session_id": task_id, "final": 1, "result": { "slice_type": 2, "index": max(len(ctx.confirmed_segments) - 1, 0), "voice_text_str": full_text, }, "sentences": [ self._build_tencent_sentence(segment, sentence_type=1) for segment in public_segments ], } ) logger.info("[%s] Recognition completed, segments=%s", task_id, len(ctx.confirmed_segments)) except Exception as exc: logger.error("[%s] Stop failed: %s", task_id, exc) await self._send_error(websocket, f"结束识别失败: {exc}", task_id)