test/app/services/qwen3_websocket_asr.py

3455 lines
132 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters!

This file contains ambiguous Unicode characters that may be confused with others in your current locale. If your use case is intentional and legitimate, you can safely ignore this warning. Use the Escape button to highlight these characters.

# -*- 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))
pre_roll_start_sample: int = 0
audio_samples_received: int = 0
segment_buffer_start_sample: int = 0
# FSMN-VAD cache and pending samples are private to this WebSocket stream.
vad_cache: Dict[str, Any] = field(default_factory=dict)
vad_pending_audio: np.ndarray = field(default_factory=lambda: np.array([], dtype=np.float32))
vad_speech_active: bool = False
vad_speech_start_sample: Optional[int] = None
vad_failed: bool = False
# 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 = ""
partial_preview_text: str = ""
stable_partial_prefix: str = ""
pending_partial_revision_text: str = ""
pending_partial_revision_rounds: int = 0
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 "<asr_text>" in raw:
_prefix, raw = raw.split("<asr_text>", 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",
getattr(settings, "REALTIME_ENABLE_VAD", True),
)
)
def _enable_realtime_vad(self, ctx: ConnectionContext) -> bool:
if str(ctx.params.get("format", "pcm")).lower() == "wav":
# Uploaded WAV is delivered as one complete message, not live frames.
return False
return bool(
ctx.params.get(
"enable_realtime_vad",
getattr(settings, "REALTIME_ENABLE_VAD", True),
)
) and not ctx.vad_failed
async def _warm_realtime_vad(self, ctx: ConnectionContext) -> None:
"""Load the shared FSMN-VAD model before the first speech boundary is needed."""
if not self._enable_realtime_vad(ctx):
return
try:
await run_sync(get_global_vad_model, settings.DEVICE)
except Exception as exc:
# Realtime energy gating remains available if the optional VAD model fails.
ctx.vad_failed = True
logger.warning("Realtime FSMN-VAD unavailable; using energy fallback: %s", exc)
@staticmethod
def _generate_realtime_vad_chunks(
chunks: List[np.ndarray],
cache: Dict[str, Any],
chunk_size_ms: int,
*,
is_final: bool = False,
) -> List[List[List[int]]]:
"""Run cached FSMN-VAD inference over fixed 200 ms mono audio frames."""
vad_model = get_global_vad_model(settings.DEVICE)
results: List[List[List[int]]] = []
with get_vad_inference_lock():
for index, chunk in enumerate(chunks):
output = vad_model.generate(
input=np.asarray(chunk, dtype=np.float32),
cache=cache,
is_final=is_final and index == len(chunks) - 1,
chunk_size=chunk_size_ms,
disable_pbar=True,
)
values = output[0].get("value", []) if output else []
results.append(values or [])
return results
async def _push_realtime_vad_audio(
self,
ctx: ConnectionContext,
audio: np.ndarray,
) -> List[tuple[int, int]]:
"""Feed each stream's FSMN-VAD cache and return completed absolute boundaries."""
if audio.size == 0 or not self._enable_realtime_vad(ctx):
return []
chunk_size_ms = max(
int(getattr(settings, "REALTIME_VAD_CHUNK_MS", 200) or 200),
100,
)
chunk_samples = max(int(chunk_size_ms * 16), 1600)
ctx.vad_pending_audio = np.concatenate([ctx.vad_pending_audio, audio])
complete_count = int(ctx.vad_pending_audio.size // chunk_samples)
if complete_count <= 0:
return []
split_at = complete_count * chunk_samples
chunks = [
np.asarray(ctx.vad_pending_audio[offset:offset + chunk_samples], dtype=np.float32)
for offset in range(0, split_at, chunk_samples)
]
ctx.vad_pending_audio = np.asarray(ctx.vad_pending_audio[split_at:], dtype=np.float32)
try:
outputs = await run_sync(
self._generate_realtime_vad_chunks,
chunks,
ctx.vad_cache,
chunk_size_ms,
)
except Exception as exc:
ctx.vad_failed = True
ctx.vad_speech_active = False
logger.warning("Realtime FSMN-VAD inference failed; using energy fallback: %s", exc)
return []
completed: List[tuple[int, int]] = []
for values in outputs:
for value in values:
if not isinstance(value, (list, tuple)) or len(value) < 2:
continue
start_ms, end_ms = int(value[0]), int(value[1])
if start_ms >= 0:
ctx.vad_speech_start_sample = start_ms * 16
ctx.vad_speech_active = True
if end_ms >= 0:
start_sample = (
start_ms * 16
if start_ms >= 0
else ctx.vad_speech_start_sample
)
if start_sample is not None and end_ms * 16 > start_sample:
completed.append((int(start_sample), end_ms * 16))
ctx.vad_speech_active = False
ctx.vad_speech_start_sample = None
return completed
async def _flush_realtime_vad(self, ctx: ConnectionContext) -> List[tuple[int, int]]:
"""Flush the final short VAD frame so stop can use its last boundary."""
if not self._enable_realtime_vad(ctx) or ctx.vad_pending_audio.size == 0:
return []
chunk_size_ms = max(
int(getattr(settings, "REALTIME_VAD_CHUNK_MS", 200) or 200),
100,
)
final_chunk = np.asarray(ctx.vad_pending_audio, dtype=np.float32)
ctx.vad_pending_audio = np.array([], dtype=np.float32)
try:
outputs = await run_sync(
self._generate_realtime_vad_chunks,
[final_chunk],
ctx.vad_cache,
chunk_size_ms,
is_final=True,
)
except Exception as exc:
ctx.vad_failed = True
logger.warning("Realtime FSMN-VAD flush failed; using energy fallback: %s", exc)
return []
completed: List[tuple[int, int]] = []
for values in outputs:
for value in values:
if not isinstance(value, (list, tuple)) or len(value) < 2:
continue
start_ms, end_ms = int(value[0]), int(value[1])
if start_ms >= 0:
ctx.vad_speech_start_sample = start_ms * 16
ctx.vad_speech_active = True
if end_ms >= 0:
start_sample = (
start_ms * 16
if start_ms >= 0
else ctx.vad_speech_start_sample
)
if start_sample is not None and end_ms * 16 > start_sample:
completed.append((int(start_sample), end_ms * 16))
ctx.vad_speech_active = False
ctx.vad_speech_start_sample = None
return completed
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)
# 默认使用 Qwen3 原生流式推理;显式传 false 时仍保留旧式 partial 兼容路径。
return True
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,
)
# 流状态有效但当前还没有解码文本时,等待下一个流式块,避免反复整段重转写。
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,
*,
start_sample: Optional[int] = None,
) -> None:
if audio.size == 0:
return
if ctx.pre_roll_audio.size == 0 and start_sample is not None:
ctx.pre_roll_start_sample = int(start_sample)
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)
if start_sample is not None:
ctx.pre_roll_start_sample = int(start_sample) + int(audio.size)
elif ctx.pre_roll_audio.size > max_samples:
removed = int(ctx.pre_roll_audio.size - max_samples)
ctx.pre_roll_audio = ctx.pre_roll_audio[removed:]
ctx.pre_roll_start_sample += removed
def _start_turn(
self,
ctx: ConnectionContext,
audio: np.ndarray,
*,
start_sample: Optional[int] = None,
) -> None:
parts: List[np.ndarray] = []
turn_start_sample = int(start_sample or 0)
if ctx.pre_roll_audio.size > 0:
parts.append(np.asarray(ctx.pre_roll_audio, dtype=np.float32))
turn_start_sample = int(ctx.pre_roll_start_sample)
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.segment_buffer_start_sample = turn_start_sample
ctx.pre_roll_audio = np.array([], dtype=np.float32)
ctx.pre_roll_start_sample = int(turn_start_sample + ctx.segment_audio_buffer.size)
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.partial_preview_text = ""
ctx.stable_partial_prefix = ""
ctx.pending_partial_revision_text = ""
ctx.pending_partial_revision_rounds = 0
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 _replace_segment_observed_text(
self,
ctx: ConnectionContext,
text: str,
language: str,
) -> str:
"""用 Qwen 流式接口当前返回的整句快照更新 partial,允许修订句尾。"""
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
ctx.segment_observed_text = candidate
if language:
ctx.segment_observed_language = language
return candidate
def _update_best_partial(
self,
ctx: ConnectionContext,
text: str,
language: str,
*,
replace_snapshot: bool = False,
) -> None:
candidate = self._sanitize_candidate_text(text)
if not candidate:
return
if self._is_unstable_expansion(ctx.best_partial_text, candidate):
return
if replace_snapshot:
# Qwen partial 是整句快照;修订句尾时应淘汰旧快照,避免旧文本被 final 选回。
ctx.best_partial_text = candidate
ctx.best_partial_language = language or ctx.best_partial_language
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 = ""
ctx.pending_partial_revision_text = ""
ctx.pending_partial_revision_rounds = 0
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:
# A single shorter or divergent snapshot may be a transient model
# revision. Hold the visible text until the same correction repeats.
if text == ctx.pending_partial_revision_text:
ctx.pending_partial_revision_rounds += 1
else:
ctx.pending_partial_revision_text = text
ctx.pending_partial_revision_rounds = 1
if ctx.pending_partial_revision_rounds < 2:
ctx.last_partial_raw_text = text
return previous_emitted
# The repeated revision is stable enough to replace the old prefix.
stable_prefix = text[: self._stable_prefix_cutoff(text, stable_common)]
ctx.pending_partial_revision_text = ""
ctx.pending_partial_revision_rounds = 0
else:
ctx.pending_partial_revision_text = ""
ctx.pending_partial_revision_rounds = 0
stable_prefix = text[:stable_common]
stable_prefix = text[: self._stable_prefix_cutoff(text, len(stable_prefix))]
else:
ctx.pending_partial_revision_text = ""
ctx.pending_partial_revision_rounds = 0
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 _append_only_partial_preview(
self,
ctx: ConnectionContext,
text: str,
) -> str:
"""Expose only newly stabilized prefixes; the final event carries revisions."""
candidate = self._partial_display_text(ctx, text)
previous = ctx.partial_preview_text
if not previous or candidate.startswith(previous):
ctx.partial_preview_text = candidate
# Qwen may rewrite its rolling snapshot. Keep the already shown prefix
# fixed until finalization so the browser never retracts earlier words.
return ctx.partial_preview_text
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 []
try:
def _run_vad() -> list[list[int]]:
vad_model = get_global_vad_model(settings.DEVICE)
with get_vad_inference_lock():
# FunASR's FSMN-VAD accepts mono float32 NumPy audio directly.
result = vad_model.generate(
input=np.asarray(audio, dtype=np.float32),
cache={},
disable_pbar=True,
)
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
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,
segment_start_ms_override: Optional[int] = None,
segment_end_ms_override: Optional[int] = 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)
has_native_stream_snapshot = (
carry_audio.size == 0
and self._should_use_native_partial_stream(ctx, engine)
and ctx.realtime_stream_state is not None
)
if carry_audio.size == 0:
stream_text, stream_language = await self._finish_realtime_stream_text(ctx, engine)
if stream_text.strip():
update_observed_text = (
self._replace_segment_observed_text
if has_native_stream_snapshot
else self._update_segment_observed_text
)
observed_text = update_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),
replace_snapshot=has_native_stream_snapshot,
)
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
if segment_start_ms_override is None
else segment_start_ms_override
)
segment_end_ms = int(
segment_start_ms + segment_duration_ms
if segment_end_ms_override is None
else segment_end_ms_override
)
if segment_end_ms < segment_start_ms:
segment_end_ms = 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
if carry_audio.size > 0:
ctx.segment_buffer_start_sample += max(
int(current_audio.size - carry_audio.size),
0,
)
else:
ctx.segment_buffer_start_sample = int(ctx.audio_samples_received)
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 _commit_vad_boundary(
self,
websocket: WebSocket,
ctx: ConnectionContext,
task_id: str,
boundary: tuple[int, int],
*,
reason: str = "vad_end",
emit_segment_start: bool = True,
) -> bool:
"""Trim one turn to FSMN-VAD's absolute start/end samples before commit."""
current_audio = np.asarray(ctx.segment_audio_buffer, dtype=np.float32)
turn_start_sample = int(ctx.segment_buffer_start_sample)
start_sample = max(int(boundary[0]), turn_start_sample)
stream_end_sample = turn_start_sample + int(current_audio.size)
end_sample = min(int(boundary[1]), stream_end_sample)
if end_sample <= start_sample:
return False
start_offset = start_sample - turn_start_sample
end_offset = end_sample - turn_start_sample
finalized_audio = np.asarray(current_audio[start_offset:end_offset], dtype=np.float32)
trailing_audio = np.asarray(current_audio[end_offset:], dtype=np.float32)
if finalized_audio.size == 0:
return False
await self._commit_retranscribe_turn(
websocket,
ctx,
task_id,
reason,
finalized_audio_override=finalized_audio,
segment_start_ms_override=int(start_sample / 16),
segment_end_ms_override=int(end_sample / 16),
emit_segment_start=emit_segment_start,
)
# Preserve post-end audio as pre-roll for a possible next utterance in
# the same network frame, while the bubble duration follows VAD bounds.
self._append_pre_roll(
ctx,
trailing_audio,
start_sample=end_sample,
)
ctx.vad_speech_start_sample = None
return True
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",
settings.REALTIME_VAD_PRE_ROLL_MS,
),
"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",
True,
),
# 与离线会议接口保持一致:前端优先传 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": payload.get(
"enable_realtime_vad",
settings.REALTIME_ENABLE_VAD,
),
"enable_realtime_vad_split": payload.get(
"enable_realtime_vad_split",
settings.REALTIME_ENABLE_VAD,
),
"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,
),
),
}
ctx.vad_cache = {}
ctx.vad_pending_audio = np.array([], dtype=np.float32)
ctx.vad_speech_active = False
ctx.vad_speech_start_sample = None
ctx.vad_failed = False
ctx.audio_samples_received = 0
ctx.segment_buffer_start_sample = 0
await self._ensure_engine(ctx)
await self._warm_realtime_vad(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
audio_start_sample = int(ctx.audio_samples_received)
ctx.audio_samples_received += int(audio.size)
# The bundled demo sends a whole WAV in one frame; keep it
# intact for final ASR instead of treating it as microphone pre-roll.
bulk_audio_frame = (
str(ctx.params.get("format", "pcm")).lower() == "wav"
or audio.size > int(1.5 * 16000)
)
vad_boundaries = (
[]
if str(ctx.params.get("format", "pcm")).lower() == "wav"
else await self._push_realtime_vad_audio(ctx, audio)
)
energy_voice = self._has_voice(audio)
has_voice = (
bulk_audio_frame
or energy_voice
or ctx.vad_speech_active
or bool(vad_boundaries)
)
if not ctx.sentence_active:
if not has_voice:
self._append_pre_roll(
ctx,
audio,
start_sample=audio_start_sample,
)
continue
if bulk_audio_frame:
self._start_turn(
ctx,
audio,
start_sample=audio_start_sample,
)
else:
self._append_pre_roll(
ctx,
audio,
start_sample=audio_start_sample,
)
self._start_turn(
ctx,
(
np.array([], dtype=np.float32)
if ctx.pre_roll_audio.size > 0
else audio
),
start_sample=audio_start_sample,
)
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=energy_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):
has_native_stream_snapshot = (
self._should_use_native_partial_stream(ctx, engine)
and ctx.realtime_stream_state is not None
)
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)
update_observed_text = (
self._replace_segment_observed_text
if has_native_stream_snapshot
else self._update_segment_observed_text
)
observed = update_observed_text(ctx, current, current_language)
visible_observed = self._trim_previous_segment_overlap(ctx, observed)
if has_native_stream_snapshot:
# Qwen returns full-utterance snapshots; stabilize the
# displayed prefix while allowing repeated corrections.
visible_observed = self._stabilize_partial_text(
ctx,
visible_observed,
)
partial_display = self._append_only_partial_preview(
ctx,
visible_observed,
)
partial_display = self._clip_partial_text(partial_display, 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,
replace_snapshot=has_native_stream_snapshot,
)
if bulk_audio_frame:
partial_start_sample = ctx.segment_buffer_start_sample
partial_end_sample = ctx.audio_samples_received
else:
partial_start_sample = ctx.vad_speech_start_sample
if partial_start_sample is None and vad_boundaries:
partial_start_sample = vad_boundaries[0][0]
if partial_start_sample is None:
partial_start_sample = ctx.segment_buffer_start_sample
partial_end_sample = (
vad_boundaries[0][1]
if vad_boundaries
else ctx.audio_samples_received
)
sentence_payload = self._build_tencent_sentence(
{
"index": ctx.segment_index,
"start_ms": int(partial_start_sample / 16),
"end_ms": int(partial_end_sample / 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
if vad_boundaries and ctx.sentence_active and not bulk_audio_frame:
if await self._commit_vad_boundary(
websocket,
ctx,
task_id,
vad_boundaries[0],
):
ctx.state = ConnectionState.STREAMING
continue
silence_threshold = self._get_dynamic_silence_threshold_samples(ctx)
if self._enable_realtime_vad(ctx) and ctx.vad_speech_active:
# Give FSMN-VAD time to confirm an endpoint before the
# legacy energy-based safety fallback closes the turn.
vad_fallback_sec = max(
float(settings.REALTIME_VAD_FINALIZE_SILENCE_SEC) + 0.8,
1.2,
)
silence_threshold = max(
silence_threshold,
int(vad_fallback_sec * 16000),
)
if ctx.silence_samples >= silence_threshold:
if self._enable_realtime_vad(ctx):
fallback_segments = await self._run_vad_segments(
np.asarray(ctx.segment_audio_buffer, dtype=np.float32)
)
valid_segments = [
(int(item[0]), int(item[1]))
for item in (fallback_segments or [])
if len(item) >= 2 and int(item[1]) > int(item[0])
]
if valid_segments:
fallback_boundary = (
int(ctx.segment_buffer_start_sample + valid_segments[0][0] * 16),
int(ctx.segment_buffer_start_sample + valid_segments[-1][1] * 16),
)
if await self._commit_vad_boundary(
websocket,
ctx,
task_id,
fallback_boundary,
reason="vad_fallback",
):
ctx.state = ConnectionState.STREAMING
continue
await self._commit_retranscribe_turn(
websocket,
ctx,
task_id,
"silence",
finalized_audio_override=np.asarray(
ctx.segment_audio_buffer,
dtype=np.float32,
),
carry_audio_override=np.array([], dtype=np.float32),
)
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:
final_vad_boundaries = await self._flush_realtime_vad(ctx)
if ctx.sentence_active and ctx.segment_audio_buffer.size > 0:
final_boundary = (
final_vad_boundaries[-1]
if final_vad_boundaries
else None
)
if (
final_boundary is None
and self._enable_realtime_vad(ctx)
and ctx.vad_speech_active
and ctx.vad_speech_start_sample is not None
):
final_boundary = (
int(ctx.vad_speech_start_sample),
int(ctx.audio_samples_received),
)
if final_boundary is not None and await self._commit_vad_boundary(
websocket,
ctx,
task_id,
final_boundary,
reason="final",
emit_segment_start=False,
):
pass
else:
await self._commit_retranscribe_turn(
websocket,
ctx,
task_id,
"final",
emit_segment_start=False,
)
# Drain pending diarization before the final event; the browser closes
# the socket on `end`, so later speaker updates would otherwise be lost.
await ctx.speaker_job_queue.join()
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)