643 lines
25 KiB
Python
643 lines
25 KiB
Python
"""Browser-facing WebSocket adapter for the migrated FunASR engine.
|
||
|
||
The frontend contract remains the existing start/binary/stop protocol. All
|
||
speech activity detection and streaming ASR state now come from FunASR; this
|
||
module only translates engine events into the existing sentence snapshots.
|
||
"""
|
||
|
||
from __future__ import annotations
|
||
|
||
import argparse
|
||
import asyncio
|
||
import json
|
||
import logging
|
||
import os
|
||
import time
|
||
import webbrowser
|
||
from dataclasses import dataclass, replace
|
||
from pathlib import Path
|
||
from typing import Any
|
||
from uuid import uuid4
|
||
|
||
from aiohttp import WSMsgType, web
|
||
from dotenv import load_dotenv
|
||
|
||
try:
|
||
from .auxiliary_service import AuxiliaryModelService, AuxiliaryServiceConfig
|
||
from .funasr_engine import FunASRModelService, FunASRSegment, FunASRServiceConfig
|
||
from .speaker_assembler import SegmentAssembler
|
||
except ImportError:
|
||
from auxiliary_service import AuxiliaryModelService, AuxiliaryServiceConfig
|
||
from funasr_engine import FunASRModelService, FunASRSegment, FunASRServiceConfig
|
||
from speaker_assembler import SegmentAssembler
|
||
|
||
|
||
PROJECT_ROOT = Path(__file__).resolve().parents[1]
|
||
load_dotenv(PROJECT_ROOT / ".env")
|
||
|
||
WEB_HOST = os.getenv("WEB_HOST", "0.0.0.0")
|
||
WEB_PORT = int(os.getenv("WEB_PORT", "8082"))
|
||
WEB_DISPLAY_HOST = os.getenv("WEB_DISPLAY_HOST", "127.0.0.1")
|
||
LOCAL_ENGINE_URL = "local://funasr"
|
||
PARTIAL_BYTES_PER_SECOND = 16000 * 2
|
||
MIN_SPEAKER_VOICE_MS = 800
|
||
LOGGER = logging.getLogger(__name__)
|
||
|
||
|
||
class EndOfStream:
|
||
"""Queue marker that cannot be confused with an audio frame."""
|
||
|
||
|
||
EOF = EndOfStream()
|
||
MODEL_SERVICE_KEY = web.AppKey("model_service", FunASRModelService)
|
||
AUXILIARY_SERVICE_KEY = web.AppKey("auxiliary_service", AuxiliaryModelService)
|
||
|
||
|
||
@dataclass(frozen=True)
|
||
class SpeakerJob:
|
||
"""Completed FunASR turn waiting for optional CAM++ speaker matching."""
|
||
|
||
sentence_id: int
|
||
audio: bytes
|
||
start_time_ms: float
|
||
end_time_ms: float
|
||
voiced_ms: float
|
||
|
||
|
||
@dataclass
|
||
class SessionMetrics:
|
||
"""Small runtime snapshot shown in the existing frontend."""
|
||
|
||
started_at: float
|
||
audio_bytes: int = 0
|
||
input_chunks: int = 0
|
||
partial_count: int = 0
|
||
partial_revisions: int = 0
|
||
first_partial_ms: float | None = None
|
||
final_ms: float | None = None
|
||
|
||
def snapshot(self) -> dict[str, Any]:
|
||
"""Return JSON-safe metrics relative to session start."""
|
||
return {
|
||
"audio_bytes": self.audio_bytes,
|
||
"input_chunks": self.input_chunks,
|
||
"partial_count": self.partial_count,
|
||
"partial_revisions": self.partial_revisions,
|
||
"first_partial_ms": self.first_partial_ms,
|
||
"final_ms": self.final_ms,
|
||
"elapsed_ms": round((time.perf_counter() - self.started_at) * 1000, 1),
|
||
}
|
||
|
||
|
||
class IncrementalWavDecoder:
|
||
"""Strip a streamed RIFF header before forwarding PCM16 to FunASR."""
|
||
|
||
def __init__(self) -> None:
|
||
self.buffer = bytearray()
|
||
self.payload_started = False
|
||
self.riff_read = False
|
||
self.format_valid = False
|
||
self.data_remaining: int | None = None
|
||
|
||
def feed(self, chunk: bytes) -> bytes:
|
||
"""Parse complete RIFF chunks without buffering the whole recording."""
|
||
if self.payload_started:
|
||
if self.data_remaining is None:
|
||
return chunk
|
||
payload = chunk[: self.data_remaining]
|
||
self.data_remaining -= len(payload)
|
||
return payload
|
||
|
||
self.buffer.extend(chunk)
|
||
if not self.riff_read:
|
||
if len(self.buffer) < 12:
|
||
return b""
|
||
if self.buffer[:4] != b"RIFF" or self.buffer[8:12] != b"WAVE":
|
||
raise ValueError("文件不是有效的 RIFF/WAV 音频")
|
||
del self.buffer[:12]
|
||
self.riff_read = True
|
||
|
||
while len(self.buffer) >= 8:
|
||
kind = bytes(self.buffer[:4])
|
||
size = int.from_bytes(self.buffer[4:8], "little")
|
||
if size > 1024 * 1024:
|
||
raise ValueError("WAV 元数据头过大,请转换为标准 PCM WAV")
|
||
chunk_size = 8 + size + (size % 2)
|
||
if len(self.buffer) < chunk_size:
|
||
return b""
|
||
body = self.buffer[8 : 8 + size]
|
||
if kind == b"fmt ":
|
||
if size < 16:
|
||
raise ValueError("WAV fmt 区块不完整")
|
||
fields = (
|
||
int.from_bytes(body[0:2], "little"),
|
||
int.from_bytes(body[2:4], "little"),
|
||
int.from_bytes(body[4:8], "little"),
|
||
int.from_bytes(body[14:16], "little"),
|
||
)
|
||
if fields != (1, 1, 16000, 16):
|
||
raise ValueError("WAV 必须为 16kHz、单声道、PCM16")
|
||
self.format_valid = True
|
||
if kind == b"data":
|
||
if not self.format_valid or size % 2:
|
||
raise ValueError("WAV 必须为 16kHz、单声道、PCM16")
|
||
self.payload_started = True
|
||
self.data_remaining = size
|
||
del self.buffer[:8]
|
||
payload = bytes(self.buffer[:size])
|
||
del self.buffer[: min(size, len(self.buffer))]
|
||
self.data_remaining -= len(payload)
|
||
return payload
|
||
del self.buffer[:chunk_size]
|
||
return b""
|
||
|
||
def finish(self) -> None:
|
||
"""Reject a truncated or header-only WAV before final inference."""
|
||
if not self.payload_started:
|
||
raise ValueError("WAV 文件不完整,未收到 data 音频区块")
|
||
if self.data_remaining not in (None, 0):
|
||
raise ValueError("WAV 文件不完整,未收到全部音频数据")
|
||
|
||
|
||
class RealtimeSession:
|
||
"""Translate FunASR events into the existing sentence/display protocol."""
|
||
|
||
def __init__(
|
||
self,
|
||
ws: web.WebSocketResponse,
|
||
model_service: FunASRModelService,
|
||
auxiliary_service: AuxiliaryModelService | None,
|
||
start: dict[str, Any],
|
||
) -> None:
|
||
self.ws = ws
|
||
self.model_service = model_service
|
||
self.auxiliary_service = auxiliary_service
|
||
self.start = start
|
||
self.engine = model_service.create_session()
|
||
self.session_id = uuid4().hex
|
||
self.send_lock = asyncio.Lock()
|
||
self.audio_queue: asyncio.Queue[bytes | EndOfStream] = asyncio.Queue(maxsize=128)
|
||
self.speaker_queue: asyncio.Queue[SpeakerJob | EndOfStream] = asyncio.Queue(maxsize=64)
|
||
self.assembler = SegmentAssembler()
|
||
self.metrics = SessionMetrics(time.perf_counter())
|
||
self.source = str(start.get("source") or "mic")
|
||
self.file_name = str(start.get("file_name") or "audio.pcm")
|
||
self.wav_decoder = (
|
||
IncrementalWavDecoder()
|
||
if self.source == "file" and Path(self.file_name).suffix.lower() == ".wav"
|
||
else None
|
||
)
|
||
self.merge_adjacent = self._parse_flag(start.get("display_merge"), True)
|
||
# Speaker labels are required for every browser session.
|
||
self.speaker_enabled = True
|
||
self.speaker_warning_sent = False
|
||
self.input_stopped = False
|
||
|
||
@staticmethod
|
||
def _parse_flag(value: Any, default: bool) -> bool:
|
||
"""Accept booleans, 0/1 and string flags from old frontend clients."""
|
||
if value is None:
|
||
return default
|
||
if isinstance(value, str):
|
||
return value.strip().lower() not in {"", "0", "false", "no", "off"}
|
||
return bool(value)
|
||
|
||
async def emit(self, payload: dict[str, Any]) -> None:
|
||
"""Send ordered JSON while the browser connection is still alive."""
|
||
async with self.send_lock:
|
||
if not self.ws.closed:
|
||
await self.ws.send_json(payload)
|
||
|
||
async def emit_state(self, sentence: dict[str, Any] | None = None) -> None:
|
||
"""Send both the legacy sentence event and the full display snapshot."""
|
||
state = {
|
||
"type": "display_state",
|
||
"raw_segments": self.assembler.raw_snapshot(),
|
||
"display_blocks": self.assembler.display_blocks(self.merge_adjacent),
|
||
"metrics": self.metrics.snapshot(),
|
||
}
|
||
if sentence is not None:
|
||
await self.emit(
|
||
{
|
||
"type": "sentences",
|
||
"sentences": [sentence],
|
||
"metrics": self.metrics.snapshot(),
|
||
}
|
||
)
|
||
await self.emit(state)
|
||
|
||
async def warn_speaker(self, message: str) -> None:
|
||
"""Expose speaker-service failures without interrupting ASR."""
|
||
if self.speaker_warning_sent:
|
||
return
|
||
await self.emit(
|
||
{
|
||
"type": "speaker_warning",
|
||
"session_id": self.session_id,
|
||
"speaker_service_url": getattr(
|
||
getattr(self.auxiliary_service, "config", None), "base_url", None
|
||
),
|
||
"message": message,
|
||
}
|
||
)
|
||
self.speaker_warning_sent = True
|
||
|
||
async def _emit_engine_segment(self, segment: FunASRSegment) -> None:
|
||
"""Map one FunASR partial/final event to a stable sentence ID."""
|
||
if not segment.text:
|
||
if segment.is_final and self.assembler.segments.pop(segment.sentence_id, None) is not None:
|
||
await self.emit_state()
|
||
return
|
||
|
||
sentence_type = 1 if segment.is_final else 0
|
||
sentence = self.assembler.apply_sentence(
|
||
{
|
||
"sentence_id": segment.sentence_id,
|
||
"sentence": segment.text,
|
||
"sentence_type": sentence_type,
|
||
"start_time": segment.start_time_ms,
|
||
"end_time": segment.end_time_ms,
|
||
"speaker_id": -1,
|
||
"speaker_name": "",
|
||
"speaker_evidence": "pending",
|
||
"speaker_confidence": 0.0,
|
||
"speaker_strategy": "funasr_pending",
|
||
"commit_reason": segment.reason,
|
||
"speaker_status": (
|
||
"queued" if segment.is_final else "waiting_final"
|
||
)
|
||
if self.speaker_enabled
|
||
else "disabled",
|
||
"speaker_reason": (
|
||
"等待 CAM++ 声纹处理"
|
||
if segment.is_final
|
||
else "FunASR 流式结果,等待片段结束"
|
||
)
|
||
if self.speaker_enabled
|
||
else "说话人分离已关闭",
|
||
}
|
||
)
|
||
if sentence_type == 0:
|
||
self.metrics.partial_count += 1
|
||
if sentence["revision_count"] > 0:
|
||
self.metrics.partial_revisions += 1
|
||
if self.metrics.first_partial_ms is None:
|
||
self.metrics.first_partial_ms = round(
|
||
(time.perf_counter() - self.metrics.started_at) * 1000, 1
|
||
)
|
||
else:
|
||
self.metrics.final_ms = round(
|
||
(time.perf_counter() - self.metrics.started_at) * 1000, 1
|
||
)
|
||
await self.emit_state(sentence)
|
||
|
||
if segment.is_final and self.speaker_enabled:
|
||
await self.speaker_queue.put(
|
||
SpeakerJob(
|
||
sentence_id=segment.sentence_id,
|
||
audio=segment.audio,
|
||
start_time_ms=segment.start_time_ms,
|
||
end_time_ms=segment.end_time_ms,
|
||
voiced_ms=segment.voiced_ms,
|
||
)
|
||
)
|
||
|
||
async def _resolve_speaker(self, job: SpeakerJob) -> None:
|
||
"""Keep the existing optional CAM++ display integration."""
|
||
if not self.speaker_enabled:
|
||
return
|
||
|
||
async def update_status(status: str, reason: str) -> None:
|
||
updated = self.assembler.apply_speaker_update(
|
||
{
|
||
"sentence_id": job.sentence_id,
|
||
"speaker_id": -1,
|
||
"speaker_evidence": "pending",
|
||
"speaker_confidence": 0.0,
|
||
"speaker_status": status,
|
||
"speaker_reason": reason,
|
||
}
|
||
)
|
||
await self.emit_state(updated)
|
||
|
||
if job.voiced_ms < MIN_SPEAKER_VOICE_MS:
|
||
await update_status(
|
||
"insufficient_audio",
|
||
f"有效语音不足 {MIN_SPEAKER_VOICE_MS}ms,不继承上一位说话人",
|
||
)
|
||
return
|
||
if self.auxiliary_service is None:
|
||
await update_status("service_unavailable", "未配置说话人辅助模型服务")
|
||
await self.warn_speaker("未配置辅助模型服务,无法执行说话人分离")
|
||
return
|
||
|
||
await update_status("processing", "正在提取 CAM++ 声纹并匹配说话人")
|
||
try:
|
||
speaker = await self.auxiliary_service.resolve_speaker(
|
||
job.audio,
|
||
self.session_id,
|
||
job.start_time_ms,
|
||
job.end_time_ms,
|
||
)
|
||
except Exception as exc:
|
||
LOGGER.exception("speaker resolve failed: session=%s sentence=%s", self.session_id, job.sentence_id)
|
||
await update_status("service_error", str(exc))
|
||
await self.warn_speaker(str(exc))
|
||
return
|
||
|
||
if not speaker:
|
||
await update_status("no_embedding", "辅助服务未返回可用声纹结果")
|
||
return
|
||
update = dict(speaker)
|
||
update["sentence_id"] = job.sentence_id
|
||
update["speaker_name"] = str(update.get("speaker_name") or "")
|
||
updated = self.assembler.apply_speaker_update(update)
|
||
if updated is not None:
|
||
await self.emit_state(updated)
|
||
|
||
async def process_speakers(self) -> None:
|
||
"""Process completed turns in order so speaker clusters stay stable."""
|
||
while True:
|
||
item = await self.speaker_queue.get()
|
||
if isinstance(item, EndOfStream):
|
||
return
|
||
await self._resolve_speaker(item)
|
||
|
||
async def _feed(self, chunk: bytes) -> None:
|
||
"""Forward raw PCM to FunASR and publish all returned events."""
|
||
if not chunk:
|
||
return
|
||
self.metrics.audio_bytes += len(chunk)
|
||
for segment in await self.engine.feed(chunk):
|
||
await self._emit_engine_segment(segment)
|
||
|
||
async def process_audio(self) -> None:
|
||
"""Consume audio until EOF, with no local RMS/VLLM segmentation path."""
|
||
while True:
|
||
item = await self.audio_queue.get()
|
||
if isinstance(item, EndOfStream):
|
||
break
|
||
self.metrics.input_chunks += 1
|
||
chunk = self.wav_decoder.feed(item) if self.wav_decoder is not None else item
|
||
await self._feed(chunk)
|
||
|
||
if self.wav_decoder is not None:
|
||
self.wav_decoder.finish()
|
||
for segment in await self.engine.finish():
|
||
await self._emit_engine_segment(segment)
|
||
await self.emit({"type": "metrics", "metrics": self.metrics.snapshot()})
|
||
|
||
async def close(self) -> None:
|
||
"""Drop per-session FunASR caches and temporary state."""
|
||
self.engine = None # type: ignore[assignment]
|
||
|
||
|
||
async def index_handler(_: web.Request) -> web.FileResponse:
|
||
"""Serve the browser page without stale-cache surprises."""
|
||
return web.FileResponse(
|
||
Path(__file__).parent / "static" / "index.html",
|
||
headers={"Cache-Control": "no-store"},
|
||
)
|
||
|
||
|
||
async def config_handler(request: web.Request) -> web.Response:
|
||
"""Expose local FunASR settings using the old response field names."""
|
||
model_service: FunASRModelService = request.app[MODEL_SERVICE_KEY]
|
||
auxiliary = request.app.get(AUXILIARY_SERVICE_KEY)
|
||
response = web.json_response(
|
||
{
|
||
"model_service_url": LOCAL_ENGINE_URL,
|
||
"model": model_service.config.model,
|
||
"engine": "funasr",
|
||
"speaker_service_url": getattr(
|
||
getattr(auxiliary, "config", None), "base_url", None
|
||
),
|
||
}
|
||
)
|
||
# The standalone frontend reads this API from its own origin.
|
||
response.headers["Access-Control-Allow-Origin"] = os.getenv(
|
||
"FRONTEND_ORIGIN", "http://127.0.0.1:8080"
|
||
)
|
||
return response
|
||
|
||
|
||
async def websocket_handler(request: web.Request) -> web.WebSocketResponse:
|
||
"""Keep the old browser protocol while using FunASR internally."""
|
||
ws = web.WebSocketResponse(max_msg_size=64 * 1024 * 1024)
|
||
await ws.prepare(request)
|
||
model_service: FunASRModelService = request.app[MODEL_SERVICE_KEY]
|
||
auxiliary_service = request.app.get(AUXILIARY_SERVICE_KEY)
|
||
processing: asyncio.Task[None] | None = None
|
||
speaker_processing: asyncio.Task[None] | None = None
|
||
session: RealtimeSession | None = None
|
||
input_finished = False
|
||
|
||
try:
|
||
first = await ws.receive()
|
||
if first.type != WSMsgType.TEXT:
|
||
await ws.send_json({"type": "error", "message": "first message must be JSON start"})
|
||
return ws
|
||
try:
|
||
start = json.loads(first.data)
|
||
except json.JSONDecodeError:
|
||
await ws.send_json({"type": "error", "message": "invalid start JSON"})
|
||
return ws
|
||
if not isinstance(start, dict) or start.get("type") != "start":
|
||
await ws.send_json({"type": "error", "message": "first message must have type=start"})
|
||
return ws
|
||
|
||
source = str(start.get("source") or "mic")
|
||
suffix = Path(str(start.get("file_name") or "")).suffix.lower()
|
||
if source == "file" and suffix not in {".pcm", ".wav"}:
|
||
await ws.send_json(
|
||
{
|
||
"type": "error",
|
||
"message": "实时流式测试的文件模式只支持 PCM 或 WAV,请改用麦克风、PCM 或 WAV",
|
||
}
|
||
)
|
||
return ws
|
||
|
||
session = RealtimeSession(ws, model_service, auxiliary_service, start)
|
||
speaker_health: dict[str, Any] | None = None
|
||
speaker_health_error: str | None = None
|
||
if session.speaker_enabled and auxiliary_service is not None:
|
||
try:
|
||
speaker_health = await asyncio.wait_for(auxiliary_service.health(), timeout=5)
|
||
if speaker_health.get("speaker_embedding_ready", speaker_health.get("ready")) is False:
|
||
speaker_health_error = "辅助模型服务未就绪,请检查 /health 返回的 models 状态"
|
||
except Exception as exc:
|
||
speaker_health_error = f"说话人辅助服务不可用:{exc}"
|
||
|
||
await session.emit(
|
||
{
|
||
"type": "start",
|
||
"model_service_url": LOCAL_ENGINE_URL,
|
||
"model": model_service.config.model,
|
||
"engine": "funasr",
|
||
"session_id": session.session_id,
|
||
"enable_native_partial_stream": True,
|
||
"native_partial_supported": True,
|
||
"partial_mode": "funasr_streaming_cache",
|
||
"speaker_diarization_enabled": session.speaker_enabled,
|
||
"speaker_service_url": getattr(
|
||
getattr(auxiliary_service, "config", None), "base_url", None
|
||
),
|
||
"speaker_service_health": speaker_health,
|
||
"speaker_gap_enabled": False,
|
||
"sentence_strategy": start.get("sentence_strategy", 0),
|
||
"display_state_supported": True,
|
||
}
|
||
)
|
||
if speaker_health_error:
|
||
await session.warn_speaker(speaker_health_error)
|
||
|
||
processing = asyncio.create_task(session.process_audio())
|
||
if session.speaker_enabled:
|
||
speaker_processing = asyncio.create_task(session.process_speakers())
|
||
|
||
async def receive_or_raise() -> Any:
|
||
"""Wake up immediately when the engine worker fails."""
|
||
receive_task = asyncio.create_task(ws.receive())
|
||
workers = [task for task in (processing, speaker_processing) if task is not None]
|
||
done, _ = await asyncio.wait(
|
||
[receive_task, *workers],
|
||
return_when=asyncio.FIRST_COMPLETED,
|
||
)
|
||
if receive_task in done:
|
||
return await receive_task
|
||
receive_task.cancel()
|
||
await asyncio.gather(receive_task, return_exceptions=True)
|
||
for worker in workers:
|
||
if worker in done:
|
||
await worker
|
||
raise RuntimeError("FunASR 实时处理任务意外结束")
|
||
|
||
while not ws.closed:
|
||
message = await receive_or_raise()
|
||
if message.type == WSMsgType.BINARY:
|
||
await session.audio_queue.put(bytes(message.data))
|
||
continue
|
||
if message.type == WSMsgType.TEXT:
|
||
try:
|
||
control = json.loads(message.data)
|
||
except json.JSONDecodeError:
|
||
continue
|
||
if not isinstance(control, dict):
|
||
continue
|
||
if control.get("type") in {"eof", "stop"}:
|
||
input_finished = True
|
||
session.input_stopped = control.get("type") == "stop"
|
||
await session.emit({"type": "draining", "message": "FunASR 正在完成最终识别"})
|
||
await session.audio_queue.put(EOF)
|
||
break
|
||
if control.get("type") == "abort":
|
||
return ws
|
||
if message.type in {WSMsgType.ERROR, WSMsgType.CLOSE, WSMsgType.CLOSED}:
|
||
return ws
|
||
|
||
if input_finished:
|
||
await processing
|
||
if speaker_processing is not None:
|
||
await session.speaker_queue.put(EOF)
|
||
await speaker_processing
|
||
await session.emit_state()
|
||
await session.emit(
|
||
{
|
||
"type": "end",
|
||
"metrics": session.metrics.snapshot(),
|
||
"sentences": session.assembler.raw_snapshot(),
|
||
"display_blocks": session.assembler.display_blocks(session.merge_adjacent),
|
||
}
|
||
)
|
||
except asyncio.CancelledError:
|
||
raise
|
||
except Exception as exc:
|
||
LOGGER.exception("FunASR WebSocket session failed")
|
||
if not ws.closed:
|
||
await ws.send_json({"type": "error", "message": str(exc)})
|
||
finally:
|
||
for task in (processing, speaker_processing):
|
||
if task is not None and not task.done():
|
||
task.cancel()
|
||
await asyncio.gather(
|
||
*(task for task in (processing, speaker_processing) if task is not None),
|
||
return_exceptions=True,
|
||
)
|
||
if session is not None:
|
||
reset = getattr(session.auxiliary_service, "reset_speaker_session", None)
|
||
await session.close()
|
||
if reset is not None:
|
||
try:
|
||
await asyncio.wait_for(reset(session.session_id), timeout=5)
|
||
except Exception:
|
||
LOGGER.warning("speaker session cleanup failed: %s", session.session_id, exc_info=True)
|
||
if not ws.closed:
|
||
await ws.close()
|
||
return ws
|
||
|
||
|
||
async def start_app(
|
||
model: str | None = None,
|
||
device: str | None = None,
|
||
) -> web.Application:
|
||
"""Create the FunASR-backed frontend application."""
|
||
config = FunASRServiceConfig.from_env()
|
||
if model:
|
||
config = replace(config, model=model)
|
||
if device:
|
||
config = replace(config, device=device)
|
||
|
||
app = web.Application()
|
||
app[MODEL_SERVICE_KEY] = FunASRModelService(config)
|
||
app[AUXILIARY_SERVICE_KEY] = AuxiliaryModelService(
|
||
AuxiliaryServiceConfig(
|
||
base_url=os.getenv("AUXILIARY_SERVICE_URL", "http://127.0.0.1:8010")
|
||
)
|
||
)
|
||
|
||
async def lifecycle(application: web.Application):
|
||
await application[MODEL_SERVICE_KEY].start()
|
||
await application[AUXILIARY_SERVICE_KEY].start()
|
||
try:
|
||
health = await application[AUXILIARY_SERVICE_KEY].health()
|
||
if not health.get("speaker_embedding_ready"):
|
||
raise RuntimeError("CAM++ speaker service is not ready")
|
||
except Exception:
|
||
await application[AUXILIARY_SERVICE_KEY].close()
|
||
await application[MODEL_SERVICE_KEY].close()
|
||
raise
|
||
yield
|
||
await application[AUXILIARY_SERVICE_KEY].close()
|
||
await application[MODEL_SERVICE_KEY].close()
|
||
|
||
app.cleanup_ctx.append(lifecycle)
|
||
app.router.add_get("/", index_handler)
|
||
app.router.add_get("/api/config", config_handler)
|
||
app.router.add_static("/static/", Path(__file__).parent / "static")
|
||
app.router.add_get("/ws", websocket_handler)
|
||
app.router.add_static("/", Path(__file__).parent / "static", show_index=False)
|
||
return app
|
||
|
||
|
||
def main() -> None:
|
||
"""Start the FunASR-backed browser demo."""
|
||
parser = argparse.ArgumentParser(description=__doc__)
|
||
parser.add_argument("--model", default=os.getenv("FUNASR_ASR_MODEL"))
|
||
parser.add_argument("--device", default=os.getenv("FUNASR_DEVICE"))
|
||
parser.add_argument("--no-browser", action="store_true")
|
||
args = parser.parse_args()
|
||
logging.basicConfig(level=logging.INFO)
|
||
if not args.no_browser:
|
||
webbrowser.open(f"http://{WEB_DISPLAY_HOST}:{WEB_PORT}/")
|
||
print(f"FunASR demo: http://{WEB_DISPLAY_HOST}:{WEB_PORT}/", flush=True)
|
||
print(f"FunASR model: {args.model or FunASRServiceConfig.from_env().model}", flush=True)
|
||
web.run_app(
|
||
start_app(model=args.model, device=args.device),
|
||
host=WEB_HOST,
|
||
port=WEB_PORT,
|
||
)
|
||
|
||
|
||
if __name__ == "__main__":
|
||
main()
|