zcbot/core/asr_xfyun.py

299 lines
12 KiB
Python

"""讯飞语音听写(IAT,流式版)WebSocket 客户端 — 整段音频转写。
接口文档: https://www.xfyun.cn/doc/asr/voicedictation/API.html
- 鉴权: HMAC-SHA256 签名拼进 wss URL query(authorization/date/host)
- 音频: 仅收 16kHz/16bit/单声道 PCM(raw),单会话上限 60 秒
- 凭据: env XFYUN_APPID / XFYUN_API_KEY / XFYUN_API_SECRET(换正式 key 只改 .env)
两条使用路径:
- `XfyunStream`:流式会话(web /v1/asr/stream 代理用)—— 边录边喂 PCM 分片,开
dwa=wpgs 动态修正,**服务端**合并 sn/pgs/rg 增量,partial 全文经回调推给前端,
松手后只等最终帧,消掉"整段回放"的等待。
- `transcribe()`:整段 PCM → 最终文本(POST /v1/asr/transcribe,流式失败的兜底)。
"""
from __future__ import annotations
import asyncio
import base64
import hashlib
import hmac
import json
import os
import time
from contextlib import suppress
from typing import Any
from urllib.parse import urlencode
from wsgiref.handlers import format_date_time
XFYUN_HOST = "iat-api.xfyun.cn"
XFYUN_PATH = "/v2/iat"
SAMPLE_RATE = 16000
BYTES_PER_SECOND = SAMPLE_RATE * 2 # 16bit mono
MAX_SECONDS = 60 # 讯飞单会话硬上限
MAX_PCM_BYTES = MAX_SECONDS * BYTES_PER_SECOND
# 每帧 9600B(=300ms 音频),base64 后 12800B < 讯飞单帧上限 13000B。
# 帧间隔:文档推荐 40ms 是给实时麦克风流的节奏;整段转写(transcribe 兜底路径)实测
# 2026-07-07(16.8s 音频):40ms→4.6s / 20ms→2.6s / 8ms→1.75s / 0ms→1.77s(识别文本
# 全一致,0ms 也不报错)—— 取 8ms:已到讯飞处理耗时地板,再压没有收益。
_FRAME_BYTES = 9600
_FRAME_INTERVAL = 0.008
_RECV_TIMEOUT = 30 # 单次等响应上限;讯飞 10s 无数据也会主动断
# 常见错误码 → 给用户/运维看的中文提示(全量见文档错误码页)
_ERR_HINTS = {
10105: "appid 不合法(检查 XFYUN_APPID)",
10107: "参数值非法",
10110: "无授权许可(服务未开通或已过期)",
10114: "会话超时(单次音频超 60 秒)",
10163: "业务参数校验失败",
10165: "会话句柄无效",
10313: "appid 不能为空或与 apikey 不匹配",
11200: "无权限或用量已用完(检查讯飞控制台服务量)",
11201: "当日流控超限",
}
class XfyunASRError(RuntimeError):
"""转写失败(讯飞返回错误码 / 连接异常),code 为讯飞错误码(连接类异常为 None)。"""
def __init__(self, message: str, code: int | None = None):
super().__init__(message)
self.code = code
class XfyunASRNotConfigured(XfyunASRError):
"""XFYUN_* env 未配置 — 上层应返回 501 而非 502。"""
def build_auth_url(api_key: str, api_secret: str) -> str:
"""按讯飞规范生成带签名的 wss URL(date 用 RFC1123 GMT)。"""
date = format_date_time(time.time())
signature_origin = f"host: {XFYUN_HOST}\ndate: {date}\nGET {XFYUN_PATH} HTTP/1.1"
signature = base64.b64encode(
hmac.new(api_secret.encode(), signature_origin.encode(), hashlib.sha256).digest()
).decode()
authorization_origin = (
f'api_key="{api_key}", algorithm="hmac-sha256", '
f'headers="host date request-line", signature="{signature}"'
)
authorization = base64.b64encode(authorization_origin.encode()).decode()
query = urlencode({"authorization": authorization, "date": date, "host": XFYUN_HOST})
return f"wss://{XFYUN_HOST}{XFYUN_PATH}?{query}"
def _load_credentials() -> tuple[str, str, str]:
appid = (os.getenv("XFYUN_APPID") or "").strip()
api_key = (os.getenv("XFYUN_API_KEY") or "").strip()
api_secret = (os.getenv("XFYUN_API_SECRET") or "").strip()
if not (appid and api_key and api_secret):
raise XfyunASRNotConfigured(
"语音识别未配置:需在 .env 设 XFYUN_APPID / XFYUN_API_KEY / XFYUN_API_SECRET"
)
return appid, api_key, api_secret
class XfyunStream:
"""流式会话:实时喂 PCM 分片,partial 全文经 on_text 异步回调推出(wpgs 已合并)。
用法(web /v1/asr/stream 代理):
st = XfyunStream(on_text=cb) # await cb(text, final) 从内部 recv task 调
await st.start()
await st.feed(chunk); ... # 16k/16bit/mono PCM,任意大小(内部按帧上限分片)
text = await st.finish() # 发 status=2,等最终文本
await st.close() # 幂等清理(finally 里调)
引擎可能先于 finish 收尾(静音超 vad_eos / 60s 上限)→ on_text(text, True) 已推,
此后 feed 静默丢弃,finish 直接返缓存结果。
"""
def __init__(self, on_text, *, language: str = "zh_cn"):
self._on_text = on_text
self._language = language
self._ws: Any = None
self._recv_task: asyncio.Task | None = None
self._segs: dict[int, str] = {} # sn → 该段文本;pgs=rpl 按 rg 区间删旧段
self._done = asyncio.Event()
self._error: XfyunASRError | None = None
self._final_text: str | None = None
self._first = True
self._appid = ""
async def start(self) -> None:
self._appid, api_key, api_secret = _load_credentials()
import websockets
try:
self._ws = await websockets.connect(
build_auth_url(api_key, api_secret), open_timeout=10
)
except Exception as e:
raise XfyunASRError(f"讯飞连接失败:{type(e).__name__}: {e}")
self._recv_task = asyncio.create_task(self._recv_loop())
def _assemble(self, result: dict) -> str:
"""合并一帧 wpgs 结果,返回当前完整文本。"""
sn = result.get("sn") or 0
if result.get("pgs") == "rpl":
a, b = (result.get("rg") or [sn, sn])[:2]
for k in [k for k in self._segs if a <= k <= b]:
del self._segs[k]
self._segs[sn] = "".join(
cw.get("w", "")
for w in (result.get("ws") or [])
for cw in (w.get("cw") or [])
)
return "".join(self._segs[k] for k in sorted(self._segs)).strip()
async def _recv_loop(self) -> None:
try:
while True:
msg = json.loads(await self._ws.recv())
code = msg.get("code")
if code:
hint = _ERR_HINTS.get(code, msg.get("message") or "未知错误")
raise XfyunASRError(f"讯飞识别失败({code}):{hint}", code=code)
data = msg.get("data") or {}
text = self._assemble(data.get("result") or {})
final = data.get("status") == 2
if final:
self._final_text = text
await self._on_text(text, final)
if final:
return
except XfyunASRError as e:
self._error = e
except asyncio.CancelledError:
raise
except Exception as e: # 连接中断 / 回调方(客户端 ws)挂了
self._error = XfyunASRError(f"讯飞连接中断:{type(e).__name__}: {e}")
finally:
self._done.set()
def _frame(self, chunk: bytes) -> str:
frame: dict = {"data": {
"status": 0 if self._first else 1,
"format": f"audio/L16;rate={SAMPLE_RATE}",
"encoding": "raw",
"audio": base64.b64encode(chunk).decode(),
}}
if self._first:
frame["common"] = {"app_id": self._appid}
frame["business"] = {
"language": self._language,
"domain": "iat",
"accent": "mandarin",
"vad_eos": 10000,
"ptt": 1,
"dwa": "wpgs", # 动态修正:边说边出字,回改在服务端 _assemble 消化
}
self._first = False
return json.dumps(frame)
async def feed(self, chunk: bytes) -> None:
if self._done.is_set() or not chunk:
return # 引擎已收尾(提前结束/出错)→ 静默丢弃
try:
for off in range(0, len(chunk), _FRAME_BYTES):
await self._ws.send(self._frame(chunk[off : off + _FRAME_BYTES]))
except Exception:
pass # 发送失败由 recv loop 报连接错误,这里不重复抛
async def finish(self, timeout: float = 15) -> str:
if not self._done.is_set():
try:
if self._first:
await self._ws.send(self._frame(b"")) # 一帧音频都没喂过:先补参数帧
await self._ws.send(json.dumps({"data": {
"status": 2,
"format": f"audio/L16;rate={SAMPLE_RATE}",
"encoding": "raw",
"audio": "",
}}))
except Exception:
pass
try:
await asyncio.wait_for(self._done.wait(), timeout=timeout)
except asyncio.TimeoutError:
raise XfyunASRError("讯飞识别超时(结束后无最终结果)")
if self._error:
raise self._error
return self._final_text or ""
async def close(self) -> None:
if self._recv_task and not self._recv_task.done():
self._recv_task.cancel()
with suppress(asyncio.CancelledError):
await self._recv_task
if self._ws:
with suppress(Exception):
await self._ws.close()
async def transcribe(pcm: bytes, *, language: str = "zh_cn") -> str:
"""整段 PCM(16k/16bit/mono, 小端)→ 识别文本。空音频/纯静音返回 """""
appid, api_key, api_secret = _load_credentials()
if not pcm:
return ""
if len(pcm) > MAX_PCM_BYTES:
raise XfyunASRError(f"音频超长:最多 {MAX_SECONDS}")
import websockets # uvicorn[standard] 附带;requirements 亦显式声明
audio_meta = {"format": f"audio/L16;rate={SAMPLE_RATE}", "encoding": "raw"}
async def _send_frames(ws) -> None:
first = True
for off in range(0, len(pcm), _FRAME_BYTES):
chunk = base64.b64encode(pcm[off : off + _FRAME_BYTES]).decode()
frame: dict = {"data": {"status": 1 if not first else 0, **audio_meta, "audio": chunk}}
if first:
frame["common"] = {"app_id": appid}
# vad_eos 拉到上限 10s:整段转写不希望句间停顿被判"说完了"截尾
frame["business"] = {
"language": language,
"domain": "iat",
"accent": "mandarin",
"vad_eos": 10000,
"ptt": 1,
}
first = False
await ws.send(json.dumps(frame))
await asyncio.sleep(_FRAME_INTERVAL)
await ws.send(json.dumps({"data": {"status": 2, **audio_meta, "audio": ""}}))
pieces: list[str] = []
try:
async with websockets.connect(build_auth_url(api_key, api_secret), open_timeout=10) as ws:
send_task = asyncio.create_task(_send_frames(ws))
try:
while True:
raw = await asyncio.wait_for(ws.recv(), timeout=_RECV_TIMEOUT)
msg = json.loads(raw)
code = msg.get("code")
if code:
hint = _ERR_HINTS.get(code, msg.get("message") or "未知错误")
raise XfyunASRError(f"讯飞识别失败({code}):{hint}", code=code)
data = msg.get("data") or {}
result = data.get("result") or {}
pieces.append(
"".join(
cw.get("w", "")
for w in (result.get("ws") or [])
for cw in (w.get("cw") or [])
)
)
if data.get("status") == 2:
break
finally:
send_task.cancel()
with suppress(asyncio.CancelledError):
await send_task
except XfyunASRError:
raise
except asyncio.TimeoutError:
raise XfyunASRError("讯飞识别超时(30s 无响应)")
except Exception as e: # 握手 401/连接断开等
raise XfyunASRError(f"讯飞连接失败:{type(e).__name__}: {e}")
return "".join(pieces).strip()