ForcePilot/backend/package/yuxi/knowledge/parser/asr.py
Kris b3bdb4fa2f feat: 新增TTS语音合成、音视频解析能力
1. 新增内置TTS工具,支持OpenAI兼容语音合成服务
2. 新增ASR音频解析和视频解析处理器
3. 扩展文件解析支持范围,添加音视频格式支持
4. 新增ffmpeg音视频处理基础工具库
5. 补充相关单元测试用例
2026-06-13 20:33:40 +08:00

180 lines
7.3 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters

This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.

"""ASR 解析器 — 将音频文件转写为 Markdown 文本。"""
import os
import tempfile
from pathlib import Path
from typing import Any
import httpx
from yuxi.knowledge.parser.base import BaseDocumentProcessor, DocumentProcessorException
from yuxi.knowledge.parser.ffmpeg import (
FFMPEG_MAX_AUDIO_DURATION_SECS,
MAX_AUDIO_BYTES,
convert_to_wav,
probe_audio_duration,
)
class ASRProcessor(BaseDocumentProcessor):
"""音频文件 ASR 解析器,将音频转写为 Markdown 文本。"""
def get_service_name(self) -> str:
return "asr"
def get_supported_extensions(self) -> list[str]:
return [".mp3", ".wav", ".m4a", ".flac", ".ogg", ".wma"]
def process_file(self, file_path: str, params: dict[str, Any] | None = None) -> str:
"""
解析音频文件,返回 Markdown 文本。
流程:
1. 检查文件大小和音频时长
2. 预处理:转换为 WAV16kHz, mono
3. ASR 转写
4. 格式化为 Markdown带时间戳段落
"""
params = params or {}
# 文件大小检查
file_size = Path(file_path).stat().st_size
if file_size > MAX_AUDIO_BYTES:
raise DocumentProcessorException(
f"音频文件大小 {file_size / 1024 / 1024:.1f}MB 超过限制 {MAX_AUDIO_BYTES / 1024 / 1024:.0f}MB",
service_name=self.get_service_name(),
)
# 时长检查
duration = probe_audio_duration(file_path)
if duration and duration > FFMPEG_MAX_AUDIO_DURATION_SECS:
raise DocumentProcessorException(
f"音频时长 {duration:.0f}s 超过限制 {FFMPEG_MAX_AUDIO_DURATION_SECS}s",
service_name=self.get_service_name(),
)
# 预处理:转换为 WAV
with tempfile.NamedTemporaryFile(suffix=".wav", delete=False) as tmp:
wav_path = tmp.name
try:
convert_to_wav(file_path, wav_path)
segments = self._transcribe(wav_path, params)
finally:
Path(wav_path).unlink(missing_ok=True)
return self._segments_to_markdown(segments, file_path)
def _transcribe(self, wav_path: str, params: dict) -> list[dict]:
"""ASR 转写,通过环境变量配置的 ASR 服务调用。"""
asr_provider = params.get("asr_provider", os.getenv("ASR_PROVIDER", "funasr"))
if asr_provider == "funasr":
return self._transcribe_funasr(wav_path)
elif asr_provider == "whisper":
return self._transcribe_whisper(wav_path)
else:
raise DocumentProcessorException(
f"不支持的 ASR 提供者: {asr_provider}",
service_name=self.get_service_name(),
)
def _transcribe_funasr(self, wav_path: str) -> list[dict]:
"""调用 FunASR 服务进行转写。"""
api_uri = os.getenv("FUNASR_API_URI", "http://funasr:10095")
try:
with open(wav_path, "rb") as f:
with httpx.Client(timeout=120) as client:
resp = client.post(
f"{api_uri}/asr",
files={"file": (Path(wav_path).name, f, "audio/wav")},
)
resp.raise_for_status()
data = resp.json()
# FunASR 返回格式适配
# 常见格式: {"text": "...", "segments": [...]} 或 {"result": [...]}
if isinstance(data, dict):
if "segments" in data:
return data["segments"]
# 单条结果包装为 segments
text = data.get("text", data.get("result", ""))
if isinstance(text, list):
return [{"start": 0, "end": 0, "text": t} for t in text]
if text:
return [{"start": 0, "end": 0, "text": str(text)}]
elif isinstance(data, list):
return data
return []
except httpx.HTTPError as e:
raise DocumentProcessorException(
f"FunASR 服务调用失败: {e}",
service_name=self.get_service_name(),
)
def _transcribe_whisper(self, wav_path: str) -> list[dict]:
"""调用 Whisper APIOpenAI 兼容)进行转写。"""
api_key = os.getenv("WHISPER_API_KEY", "")
base_url = os.getenv("WHISPER_BASE_URL", "https://api.openai.com/v1")
if not api_key:
raise DocumentProcessorException(
"Whisper API Key 未配置 (WHISPER_API_KEY)",
service_name=self.get_service_name(),
)
try:
with open(wav_path, "rb") as f:
with httpx.Client(timeout=120) as client:
resp = client.post(
f"{base_url}/audio/transcriptions",
headers={"Authorization": f"Bearer {api_key}"},
files={"file": (Path(wav_path).name, f, "audio/wav")},
data={"response_format": "verbose_json", "timestamp_granularities[]": "segment"},
)
resp.raise_for_status()
data = resp.json()
segments = data.get("segments", [])
if not segments and data.get("text"):
return [{"start": 0, "end": 0, "text": data["text"]}]
return [
{"start": s.get("start", 0), "end": s.get("end", 0), "text": s.get("text", "")}
for s in segments
]
except httpx.HTTPError as e:
raise DocumentProcessorException(
f"Whisper API 调用失败: {e}",
service_name=self.get_service_name(),
)
@staticmethod
def _segments_to_markdown(segments: list[dict], source: str) -> str:
"""将 ASR 片段格式化为 Markdown。"""
lines = ["# 音频转写", "", f"> 来源: {source}", ""]
for seg in segments:
start = seg.get("start", 0)
text = seg.get("text", "").strip()
if text:
mm_ss = f"{int(start) // 60:02d}:{int(start) % 60:02d}"
lines.append(f"**[{mm_ss}]** {text}")
lines.append("")
return "\n".join(lines)
def check_health(self) -> dict[str, Any]:
"""检查 ASR 服务健康状态。"""
asr_provider = os.getenv("ASR_PROVIDER", "funasr")
try:
if asr_provider == "funasr":
api_uri = os.getenv("FUNASR_API_URI", "http://funasr:10095")
with httpx.Client(timeout=5) as client:
resp = client.get(f"{api_uri}/health")
if resp.status_code == 200:
return {"status": "healthy", "message": "FunASR 服务正常"}
return {"status": "unhealthy", "message": f"FunASR 返回 {resp.status_code}"}
elif asr_provider == "whisper":
if os.getenv("WHISPER_API_KEY"):
return {"status": "healthy", "message": "Whisper API 已配置"}
return {"status": "unavailable", "message": "Whisper API Key 未配置"}
except Exception as e:
return {"status": "unhealthy", "message": str(e)}
return {"status": "unavailable", "message": f"ASR 提供者 {asr_provider} 不可用"}