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