feat: 新增TTS语音合成、音视频解析能力

1. 新增内置TTS工具,支持OpenAI兼容语音合成服务
2. 新增ASR音频解析和视频解析处理器
3. 扩展文件解析支持范围,添加音视频格式支持
4. 新增ffmpeg音视频处理基础工具库
5. 补充相关单元测试用例
This commit is contained in:
Kris 2026-06-13 20:33:40 +08:00
parent 436c09fb90
commit b3bdb4fa2f
15 changed files with 1237 additions and 0 deletions

View File

@ -1,9 +1,11 @@
# buildin 工具包
from .install_skill import install_skill
from .tools import ask_user_question, present_artifacts
from .tts import text_to_speech
__all__ = [
"ask_user_question",
"install_skill",
"present_artifacts",
"text_to_speech",
]

View File

@ -0,0 +1,44 @@
"""TTS 语音合成工具。"""
from uuid import uuid4
from yuxi.agents.toolkits.registry import tool
from yuxi.storage.minio import aupload_file_to_minio
from yuxi.agents.toolkits.buildin.tts import get_tts_provider
from yuxi.utils import logger
@tool(
category="buildin",
tags=["语音"],
display_name="语音合成",
config_guide="使用前需配置 TTS_API_KEY 环境变量",
)
async def text_to_speech(text: str, voice: str = "default") -> str:
"""
将文本转换为语音并返回音频文件 URL
Args:
text: 要转换的文本不超过 4096 字符
voice: 语音类型default / alloy / echo / fable / onyx / nova / shimmer
Returns:
音频文件的访问 URL
"""
if len(text) > 4096:
return "错误:文本长度超过 4096 字符限制"
provider = get_tts_provider()
if not provider:
return "错误TTS 服务未配置,请设置 TTS_API_KEY 环境变量"
try:
result = await provider.synthesize(text, voice=voice)
except Exception as e:
logger.error("TTS 合成失败: %s", e)
return f"错误:语音合成失败 — {e}"
# 存储到 MinIO
file_name = f"tts/{uuid4().hex}.mp3"
url = await aupload_file_to_minio("public", file_name, result.audio_data)
return url

View File

@ -0,0 +1,17 @@
"""TTS 语音合成模块"""
import yuxi.agents.toolkits.buildin.tts.openai_compatible # noqa: F401
from yuxi.agents.toolkits.buildin.tts.provider import (
SpeechSynthesisResult,
TTSProvider,
get_tts_provider,
register_tts_provider,
)
__all__ = [
"SpeechSynthesisResult",
"TTSProvider",
"get_tts_provider",
"register_tts_provider",
]

View File

@ -0,0 +1,48 @@
import os
import httpx
from yuxi.agents.toolkits.buildin.tts.provider import TTSProvider, SpeechSynthesisResult, register_tts_provider
@register_tts_provider("openai_compatible")
class OpenAICompatibleTTSProvider(TTSProvider):
"""OpenAI 兼容 TTS 提供者(支持 OpenAI / DashScope / SiliconFlow 等)"""
def __init__(self):
self._api_key = os.getenv("TTS_API_KEY", "")
self._base_url = os.getenv("TTS_BASE_URL", "https://api.openai.com/v1")
self._model = os.getenv("TTS_MODEL", "tts-1")
self._default_voice = os.getenv("TTS_DEFAULT_VOICE", "alloy")
@property
def provider_id(self) -> str:
return "openai_compatible"
@property
def display_name(self) -> str:
return "OpenAI 兼容 TTS"
def is_configured(self) -> bool:
return bool(self._api_key)
async def synthesize(self, text: str, voice: str = "default") -> SpeechSynthesisResult:
voice = voice if voice != "default" else self._default_voice
async with httpx.AsyncClient(timeout=60) as client:
resp = await client.post(
f"{self._base_url}/audio/speech",
headers={"Authorization": f"Bearer {self._api_key}"},
json={"model": self._model, "input": text, "voice": voice},
)
resp.raise_for_status()
return SpeechSynthesisResult(audio_data=resp.content)
async def list_voices(self) -> list[dict]:
return [
{"id": "alloy", "name": "Alloy"},
{"id": "echo", "name": "Echo"},
{"id": "fable", "name": "Fable"},
{"id": "onyx", "name": "Onyx"},
{"id": "nova", "name": "Nova"},
{"id": "shimmer", "name": "Shimmer"},
]

View File

@ -0,0 +1,67 @@
from __future__ import annotations
from abc import ABC, abstractmethod
from dataclasses import dataclass
_TTS_PROVIDERS: dict[str, type[TTSProvider]] = {}
@dataclass(frozen=True)
class SpeechSynthesisResult:
audio_data: bytes
content_type: str = "audio/mpeg"
class TTSProvider(ABC):
"""TTS 提供者抽象基类"""
@property
@abstractmethod
def provider_id(self) -> str: ...
@property
@abstractmethod
def display_name(self) -> str: ...
@abstractmethod
async def synthesize(self, text: str, voice: str = "default") -> SpeechSynthesisResult: ...
@abstractmethod
async def list_voices(self) -> list[dict]: ...
@abstractmethod
def is_configured(self) -> bool: ...
def register_tts_provider(provider_id: str | None = None):
"""装饰器工厂:将 TTSProvider 子类注册到全局注册表
可显式指定 provider_id否则通过临时实例获取
"""
def decorator(cls: type[TTSProvider]) -> type[TTSProvider]:
key = provider_id or cls().provider_id
_TTS_PROVIDERS[key] = cls
return cls
return decorator
def get_tts_provider(provider_id: str | None = None) -> TTSProvider | None:
"""获取已配置的 TTS 提供者实例
- 指定 provider_id 查找对应提供者并检查是否已配置
- 未指定时返回第一个已配置的提供者
- 无可用提供者时返回 None
"""
if provider_id is not None:
cls = _TTS_PROVIDERS.get(provider_id)
if cls is None:
return None
provider = cls()
return provider if provider.is_configured() else None
for cls in _TTS_PROVIDERS.values():
provider = cls()
if provider.is_configured():
return provider
return None

View File

@ -0,0 +1,179 @@
"""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} 不可用"}

View File

@ -25,6 +25,8 @@ class DocumentProcessorFactory:
"mineru_official": ("yuxi.knowledge.parser.mineru_official", "MinerUOfficialParser"),
"pp_structure_v3_ocr": ("yuxi.knowledge.parser.pp_structure_v3", "PPStructureV3Parser"),
"deepseek_ocr": ("yuxi.knowledge.parser.deepseek_ocr", "DeepSeekOCRParser"),
"asr": ("yuxi.knowledge.parser.asr", "ASRProcessor"),
"video": ("yuxi.knowledge.parser.video", "VideoProcessor"),
}
@classmethod

View File

@ -0,0 +1,206 @@
"""ffmpeg/ffprobe 执行基础设施,提供音视频处理的底层能力。"""
import json
import subprocess
import threading
from pathlib import Path
from yuxi.utils import logger
# ── 限制常量 ──────────────────────────────────────────────
FFMPEG_MAX_BUFFER_BYTES = 10 * 1024 * 1024 # 10MB
FFPROBE_TIMEOUT_MS = 10_000 # 10s
FFMPEG_TIMEOUT_MS = 45_000 # 45s
FFMPEG_MAX_AUDIO_DURATION_SECS = 1200 # 20 minutes
MAX_AUDIO_BYTES = 16 * 1024 * 1024 # 16MB
MAX_VIDEO_BYTES = 16 * 1024 * 1024 # 16MB
FFMPEG_MAX_CONCURRENT = 4
# ── 并发信号量 ────────────────────────────────────────────
_semaphore = threading.Semaphore(FFMPEG_MAX_CONCURRENT)
# ── 异常 ──────────────────────────────────────────────────
class FFmpegError(Exception):
"""ffmpeg/ffprobe 执行错误。"""
def __init__(self, message: str, returncode: int = -1):
super().__init__(message)
self.message = message
self.returncode = returncode
# ── 核心执行函数 ──────────────────────────────────────────
def run_ffmpeg(args: list[str], timeout_ms: int = FFMPEG_TIMEOUT_MS, *, return_stderr: bool = False) -> str:
"""同步执行 ffmpeg 命令,默认返回 stdout 内容return_stderr=True 时返回 stderr。"""
cmd = ["ffmpeg", *args]
timeout_sec = timeout_ms / 1000
with _semaphore:
try:
result = subprocess.run(
cmd,
capture_output=True,
timeout=timeout_sec,
)
except subprocess.TimeoutExpired:
raise FFmpegError(
f"ffmpeg 执行超时 ({timeout_ms}ms): {' '.join(cmd)}",
returncode=-1,
)
except OSError as e:
raise FFmpegError(
f"ffmpeg 执行失败(请确认 ffmpeg 已安装): {e}",
returncode=-1,
)
output = result.stderr if return_stderr else result.stdout
if len(output) > FFMPEG_MAX_BUFFER_BYTES:
raise FFmpegError(
f"ffmpeg 输出超过缓冲区限制 ({FFMPEG_MAX_BUFFER_BYTES} bytes)",
returncode=result.returncode,
)
if result.returncode != 0:
stderr = result.stderr.decode(errors="replace")[:512]
raise FFmpegError(
f"ffmpeg 返回非零退出码 {result.returncode}: {stderr}",
returncode=result.returncode,
)
return output.decode(errors="replace")
def run_ffprobe(args: list[str], timeout_ms: int = FFPROBE_TIMEOUT_MS) -> str:
"""同步执行 ffprobe 命令,返回 stdout 内容。"""
cmd = ["ffprobe", *args]
timeout_sec = timeout_ms / 1000
with _semaphore:
try:
result = subprocess.run(
cmd,
capture_output=True,
timeout=timeout_sec,
)
except subprocess.TimeoutExpired:
raise FFmpegError(
f"ffprobe 执行超时 ({timeout_ms}ms): {' '.join(cmd)}",
returncode=-1,
)
except OSError as e:
raise FFmpegError(
f"ffprobe 执行失败(请确认 ffprobe 已安装): {e}",
returncode=-1,
)
if len(result.stdout) > FFMPEG_MAX_BUFFER_BYTES:
raise FFmpegError(
f"ffprobe 输出超过缓冲区限制 ({FFMPEG_MAX_BUFFER_BYTES} bytes)",
returncode=result.returncode,
)
if result.returncode != 0:
stderr = result.stderr.decode(errors="replace")[:512]
raise FFmpegError(
f"ffprobe 返回非零退出码 {result.returncode}: {stderr}",
returncode=result.returncode,
)
return result.stdout.decode(errors="replace")
# ── 业务函数 ──────────────────────────────────────────────
def probe_audio_duration(file_path: str) -> float | None:
"""使用 ffprobe 获取音频时长(秒),失败返回 None。"""
try:
output = run_ffprobe(
["-v", "quiet", "-print_format", "json", "-show_format", file_path],
)
info = json.loads(output)
return float(info["format"]["duration"])
except (FFmpegError, KeyError, ValueError, json.JSONDecodeError) as e:
logger.warning("获取音频时长失败: %s%s", file_path, e)
return None
def convert_to_wav(input_path: str, output_path: str) -> None:
"""将音频转换为 WAV 格式16kHz, 单声道, PCM s16le"""
run_ffmpeg([
"-i", input_path,
"-vn",
"-acodec", "pcm_s16le",
"-ar", "16000",
"-ac", "1",
output_path,
])
def extract_audio_from_video(video_path: str, output_path: str) -> None:
"""从视频中提取音轨,输出为 WAV 格式16kHz, 单声道, PCM s16le"""
run_ffmpeg([
"-i", video_path,
"-vn",
"-acodec", "pcm_s16le",
"-ar", "16000",
"-ac", "1",
output_path,
])
def extract_keyframes(
video_path: str,
output_dir: str,
scene_threshold: float = 0.3,
) -> list[tuple[str, float]]:
"""通过场景变化检测提取关键帧,返回 [(路径, 时间戳秒), ...] 按时间排序。
分两步执行
1. ffmpeg select+showinfo 检测场景切换帧的 PTS
2. -ss 精确跳转到每个时间点提取帧
"""
Path(output_dir).mkdir(parents=True, exist_ok=True)
# 第一步:检测场景切换帧的 PTS 时间戳
timestamps = _detect_scene_change_pts(video_path, scene_threshold)
if not timestamps:
return []
# 第二步:在每个时间点提取帧
frames: list[tuple[str, float]] = []
for i, ts in enumerate(timestamps, start=1):
output_path = str(Path(output_dir) / f"frame_{i:04d}.png")
try:
run_ffmpeg([
"-ss", f"{ts:.6f}",
"-i", video_path,
"-frames:v", "1",
output_path,
])
frames.append((output_path, ts))
except FFmpegError:
logger.debug("关键帧提取跳过: ts=%.2f", ts)
return frames
def _detect_scene_change_pts(video_path: str, scene_threshold: float) -> list[float]:
"""使用 ffmpeg select+showinfo 检测场景切换帧的 PTS 时间戳。"""
import re
stderr_text = run_ffmpeg([
"-i", video_path,
"-vf", f"select=gt(scene\\,{scene_threshold}),showinfo",
"-vsync", "vfr",
"-f", "null", "-",
], return_stderr=True)
timestamps: list[float] = []
for match in re.finditer(r"pts_time:(\d+\.?\d*)", stderr_text):
try:
timestamps.append(float(match.group(1)))
except ValueError:
pass
return timestamps

View File

@ -41,8 +41,25 @@ SUPPORTED_FILE_EXTENSIONS: tuple[str, ...] = (
".tiff",
".tif",
".zip",
# 音频
".mp3",
".wav",
".m4a",
".flac",
".ogg",
".wma",
# 视频
".mp4",
".avi",
".mov",
".mkv",
".webm",
".flv",
)
AUDIO_EXTENSIONS = {".mp3", ".wav", ".m4a", ".flac", ".ogg", ".wma"}
VIDEO_EXTENSIONS = {".mp4", ".avi", ".mov", ".mkv", ".webm", ".flv"}
def is_supported_file_extension(file_name: str | os.PathLike[str]) -> bool:
"""Check whether the given file path has a supported extension."""
@ -269,6 +286,44 @@ async def parse_image_async(file, params=None):
return await asyncio.to_thread(parse_image, file, params=params)
def parse_audio(file, params=None):
"""解析音频文件,使用 ASR 转写为文本。"""
from yuxi.knowledge.parser.base import DocumentProcessorException
from yuxi.knowledge.parser.factory import DocumentProcessorFactory
try:
return DocumentProcessorFactory.process_file("asr", file, params)
except DocumentProcessorException as e:
logger.error(f"音频处理失败: {e.service_name} - {str(e)}")
raise
except Exception as e: # noqa: BLE001
logger.error(f"音频解析失败: {str(e)}")
raise DocumentProcessorException(f"音频解析失败: {str(e)}", "asr", "parsing_failed")
def parse_video(file, params=None):
"""解析视频文件,提取音频 ASR + 关键帧 OCR。"""
from yuxi.knowledge.parser.base import DocumentProcessorException
from yuxi.knowledge.parser.factory import DocumentProcessorFactory
try:
return DocumentProcessorFactory.process_file("video", file, params)
except DocumentProcessorException as e:
logger.error(f"视频处理失败: {e.service_name} - {str(e)}")
raise
except Exception as e: # noqa: BLE001
logger.error(f"视频解析失败: {str(e)}")
raise DocumentProcessorException(f"视频解析失败: {str(e)}", "video", "parsing_failed")
async def parse_audio_async(file, params=None):
return await asyncio.to_thread(parse_audio, file, params=params)
async def parse_video_async(file, params=None):
return await asyncio.to_thread(parse_video, file, params=params)
async def _process_file_to_markdown_core(
file_path: str, params: dict | None = None
) -> tuple[str, str | None, dict[str, Any]]:
@ -345,6 +400,14 @@ async def _process_file_to_markdown_core(
text = await parse_image_async(str(file_path_obj), params=params)
result = f"{text}"
elif file_ext in AUDIO_EXTENSIONS:
text = await parse_audio_async(str(file_path_obj), params=params)
result = f"{text}"
elif file_ext in VIDEO_EXTENSIONS:
text = await parse_video_async(str(file_path_obj), params=params)
result = f"{text}"
elif file_ext in [".html", ".htm"]:
with open(file_path_obj, encoding="utf-8") as f:
content = f.read()

View File

@ -0,0 +1,135 @@
"""视频解析器 — 提取音频 ASR + 关键帧 OCR输出 Markdown 文本。"""
import tempfile
from pathlib import Path
from typing import Any
from yuxi.knowledge.parser.base import BaseDocumentProcessor, DocumentProcessorException
from yuxi.knowledge.parser.ffmpeg import (
MAX_VIDEO_BYTES,
FFmpegError,
extract_audio_from_video,
extract_keyframes,
)
from yuxi.utils import logger
class VideoProcessor(BaseDocumentProcessor):
"""视频文件解析器:提取音频 → ASR + 关键帧 OCR。"""
def get_service_name(self) -> str:
return "video"
def get_supported_extensions(self) -> list[str]:
return [".mp4", ".avi", ".mov", ".mkv", ".webm", ".flv"]
def process_file(self, file_path: str, params: dict[str, Any] | None = None) -> str:
"""
解析视频文件返回 Markdown 文本
流程:
1. 提取音频轨道 ASR 转写
2. 提取关键帧 OCR复用现有 DocumentProcessorFactory
3. 合并为 Markdown
"""
params = params or {}
# 文件大小检查
file_size = Path(file_path).stat().st_size
if file_size > MAX_VIDEO_BYTES:
raise DocumentProcessorException(
f"视频文件大小 {file_size / 1024 / 1024:.1f}MB 超过限制 {MAX_VIDEO_BYTES / 1024 / 1024:.0f}MB",
service_name=self.get_service_name(),
)
sections: list[str] = []
# 1. 提取音频 → ASR
asr_markdown = self._process_audio_track(file_path, params)
if asr_markdown:
sections.append(asr_markdown)
# 2. 提取关键帧 → OCR
ocr_markdown = self._process_keyframes(file_path, params)
if ocr_markdown:
sections.append(ocr_markdown)
if not sections:
raise DocumentProcessorException(
"视频解析未产生任何内容",
service_name=self.get_service_name(),
)
return "\n\n---\n\n".join(sections)
def _process_audio_track(self, video_path: str, params: dict) -> str | None:
"""提取音频轨道并 ASR 转写。"""
with tempfile.NamedTemporaryFile(suffix=".wav", delete=False) as tmp:
wav_path = tmp.name
try:
extract_audio_from_video(video_path, wav_path)
from yuxi.knowledge.parser.factory import DocumentProcessorFactory
return DocumentProcessorFactory.process_file("asr", wav_path, params)
except FFmpegError:
logger.info("视频无音轨或音频提取失败: %s", video_path)
return None
except DocumentProcessorException:
logger.warning("视频音频 ASR 转写失败: %s", video_path)
return None
finally:
Path(wav_path).unlink(missing_ok=True)
def _process_keyframes(self, video_path: str, params: dict) -> str | None:
"""提取关键帧并 OCR。"""
from yuxi.knowledge.parser.factory import DocumentProcessorFactory
ocr_engine = params.get("ocr_engine", "rapid_ocr")
if ocr_engine == "disable":
return None
with tempfile.TemporaryDirectory() as tmpdir:
try:
frames = extract_keyframes(video_path, tmpdir)
except FFmpegError as e:
logger.warning("关键帧提取失败: %s%s", video_path, e)
return None
if not frames:
return None
ocr_texts: list[str] = []
processor = DocumentProcessorFactory.get_processor(ocr_engine)
for frame_path, timestamp in frames:
try:
text = processor.process_file(frame_path, params)
if text.strip():
mm_ss = f"{int(timestamp) // 60:02d}:{int(timestamp) % 60:02d}"
ocr_texts.append(f"**[{mm_ss}]** {text.strip()}")
except Exception as e:
logger.debug("关键帧 OCR 失败: %s%s", frame_path, e)
continue
if not ocr_texts:
return None
return "## 视频画面文字\n\n" + "\n\n".join(ocr_texts)
def check_health(self) -> dict[str, Any]:
"""检查视频处理依赖ffmpeg + ASR + OCR"""
# 检查 ffmpeg
try:
from yuxi.knowledge.parser.ffmpeg import run_ffmpeg
run_ffmpeg(["-version"], timeout_ms=3000)
except Exception:
return {"status": "unavailable", "message": "ffmpeg 不可用"}
# 检查 ASR
from yuxi.knowledge.parser.factory import DocumentProcessorFactory
asr_health = DocumentProcessorFactory.check_health("asr")
# 检查 OCR
ocr_health = DocumentProcessorFactory.check_health("rapid_ocr")
if asr_health["status"] == "healthy" and ocr_health["status"] == "healthy":
return {"status": "healthy", "message": "视频解析服务正常"}
return {"status": "unhealthy", "message": "ASR 或 OCR 服务不可用"}

View File

@ -13,6 +13,18 @@ if str(PROJECT_ROOT) not in sys.path:
# Avoid package-level knowledge graph initialization during pytest collection.
os.environ.setdefault("YUXI_SKIP_APP_INIT", "1")
# Pre-populate sys.modules for packages with heavy __init__.py imports
# (sklearn/scipy/numpy chain) so that submodules can be imported directly.
import types
_PACKAGE_DIR = PROJECT_ROOT / "package"
for _mod_name in ("yuxi.knowledge", "yuxi.knowledge.parser"):
if _mod_name not in sys.modules:
_m = types.ModuleType(_mod_name)
_m.__path__ = [str(_PACKAGE_DIR / _mod_name.replace(".", os.sep))]
_m.__package__ = _mod_name
sys.modules[_mod_name] = _m
def pytest_configure(config: pytest.Config) -> None:
"""Register shared markers without binding every test to a live environment."""

View File

@ -0,0 +1,123 @@
"""ASR 解析器单元测试。"""
import os
from unittest.mock import MagicMock, patch
import pytest
from yuxi.knowledge.parser.asr import ASRProcessor
from yuxi.knowledge.parser.base import DocumentProcessorException
@pytest.fixture
def asr():
return ASRProcessor()
class TestASRProcessorBasics:
def test_service_name(self, asr):
assert asr.get_service_name() == "asr"
def test_supported_extensions(self, asr):
exts = asr.get_supported_extensions()
assert ".mp3" in exts
assert ".wav" in exts
assert ".m4a" in exts
assert ".flac" in exts
def test_supports_file_type(self, asr):
assert asr.supports_file_type(".mp3")
assert asr.supports_file_type(".wav")
assert not asr.supports_file_type(".pdf")
class TestSegmentsToMarkdown:
def test_basic_formatting(self):
segments = [
{"start": 0, "end": 5, "text": "Hello world"},
{"start": 65, "end": 70, "text": "Second segment"},
]
result = ASRProcessor._segments_to_markdown(segments, "test.mp3")
assert "# 音频转写" in result
assert "> 来源: test.mp3" in result
assert "**[00:00]** Hello world" in result
assert "**[01:05]** Second segment" in result
def test_empty_segments(self):
result = ASRProcessor._segments_to_markdown([], "test.mp3")
assert "# 音频转写" in result
assert "> 来源: test.mp3" in result
def test_skip_empty_text(self):
segments = [
{"start": 0, "end": 5, "text": ""},
{"start": 5, "end": 10, "text": " "},
{"start": 10, "end": 15, "text": "valid"},
]
result = ASRProcessor._segments_to_markdown(segments, "test.mp3")
assert "valid" in result
assert result.count("**[") == 1
class TestProcessFile:
@patch("yuxi.knowledge.parser.asr.probe_audio_duration", return_value=60.0)
@patch("yuxi.knowledge.parser.asr.convert_to_wav")
@patch("yuxi.knowledge.parser.asr.ASRProcessor._transcribe")
def test_success(self, mock_transcribe, mock_convert, mock_duration, asr, tmp_path):
test_file = tmp_path / "test.mp3"
test_file.write_bytes(b"fake audio data")
mock_transcribe.return_value = [{"start": 0, "end": 5, "text": "Hello"}]
result = asr.process_file(str(test_file))
assert "# 音频转写" in result
assert "Hello" in result
@patch("yuxi.knowledge.parser.asr.probe_audio_duration", return_value=60.0)
def test_file_too_large(self, mock_duration, asr, tmp_path):
test_file = tmp_path / "big.mp3"
test_file.write_bytes(b"x" * (16 * 1024 * 1024 + 1))
with pytest.raises(DocumentProcessorException, match="大小"):
asr.process_file(str(test_file))
@patch("yuxi.knowledge.parser.asr.probe_audio_duration", return_value=1500.0)
def test_duration_too_long(self, mock_duration, asr, tmp_path):
test_file = tmp_path / "long.mp3"
test_file.write_bytes(b"fake audio")
with pytest.raises(DocumentProcessorException, match="时长"):
asr.process_file(str(test_file))
@patch("yuxi.knowledge.parser.asr.probe_audio_duration", return_value=60.0)
@patch("yuxi.knowledge.parser.asr.convert_to_wav")
def test_unsupported_asr_provider(self, mock_convert, mock_duration, asr, tmp_path):
test_file = tmp_path / "test.mp3"
test_file.write_bytes(b"fake audio")
with pytest.raises(DocumentProcessorException, match="不支持的 ASR 提供者"):
asr.process_file(str(test_file), {"asr_provider": "unknown"})
class TestCheckHealth:
@patch.dict(os.environ, {"ASR_PROVIDER": "whisper", "WHISPER_API_KEY": "test-key"})
def test_whisper_configured(self, asr):
result = asr.check_health()
assert result["status"] == "healthy"
@patch.dict(os.environ, {"ASR_PROVIDER": "whisper"}, clear=False)
def test_whisper_not_configured(self, asr):
with patch.dict(os.environ, {"WHISPER_API_KEY": ""}):
result = asr.check_health()
assert result["status"] == "unavailable"
@patch.dict(os.environ, {"ASR_PROVIDER": "funasr"})
@patch("yuxi.knowledge.parser.asr.httpx.Client")
def test_funasr_unreachable(self, mock_client_cls, asr):
mock_client = MagicMock()
mock_client.__enter__ = MagicMock(return_value=mock_client)
mock_client.__exit__ = MagicMock(return_value=False)
mock_client.get.side_effect = Exception("connection refused")
mock_client_cls.return_value = mock_client
result = asr.check_health()
assert result["status"] == "unhealthy"

View File

@ -0,0 +1,130 @@
"""ffmpeg 基础设施模块单元测试。"""
import json
import subprocess
from unittest.mock import MagicMock, patch
import pytest
from yuxi.knowledge.parser.ffmpeg import (
FFMPEG_MAX_AUDIO_DURATION_SECS,
FFMPEG_MAX_BUFFER_BYTES,
FFMPEG_MAX_CONCURRENT,
FFMPEG_TIMEOUT_MS,
FFPROBE_TIMEOUT_MS,
FFmpegError,
MAX_AUDIO_BYTES,
MAX_VIDEO_BYTES,
convert_to_wav,
extract_audio_from_video,
extract_keyframes,
probe_audio_duration,
run_ffmpeg,
run_ffprobe,
)
class TestFFmpegError:
def test_error_attributes(self):
err = FFmpegError("test error", returncode=1)
assert str(err) == "test error"
assert err.returncode == 1
assert err.message == "test error"
def test_default_returncode(self):
err = FFmpegError("test")
assert err.returncode == -1
class TestRunFFmpeg:
@patch("yuxi.knowledge.parser.ffmpeg.subprocess.run")
def test_success(self, mock_run):
mock_run.return_value = MagicMock(returncode=0, stdout=b"output", stderr=b"")
result = run_ffmpeg(["-version"])
assert result == "output"
mock_run.assert_called_once()
@patch("yuxi.knowledge.parser.ffmpeg.subprocess.run")
def test_timeout(self, mock_run):
mock_run.side_effect = subprocess.TimeoutExpired(cmd="ffmpeg", timeout=45)
with pytest.raises(FFmpegError, match="超时"):
run_ffmpeg(["-i", "input.mp4"])
@patch("yuxi.knowledge.parser.ffmpeg.subprocess.run")
def test_nonzero_returncode(self, mock_run):
mock_run.return_value = MagicMock(returncode=1, stdout=b"", stderr=b"error msg")
with pytest.raises(FFmpegError, match="非零退出码"):
run_ffmpeg(["-i", "bad.mp4"])
class TestRunFFprobe:
@patch("yuxi.knowledge.parser.ffmpeg.subprocess.run")
def test_success(self, mock_run):
mock_run.return_value = MagicMock(returncode=0, stdout=b'{"format":{}}', stderr=b"")
result = run_ffprobe(["-show_format", "test.mp3"])
assert result == '{"format":{}}'
class TestProbeAudioDuration:
@patch("yuxi.knowledge.parser.ffmpeg.run_ffprobe")
def test_success(self, mock_ffprobe):
mock_ffprobe.return_value = json.dumps({"format": {"duration": "120.5"}})
duration = probe_audio_duration("test.mp3")
assert duration == 120.5
@patch("yuxi.knowledge.parser.ffmpeg.run_ffprobe")
def test_failure_returns_none(self, mock_ffprobe):
mock_ffprobe.side_effect = FFmpegError("failed")
duration = probe_audio_duration("bad.mp3")
assert duration is None
class TestConvertToWav:
@patch("yuxi.knowledge.parser.ffmpeg.run_ffmpeg")
def test_calls_ffmpeg_with_correct_args(self, mock_ffmpeg):
convert_to_wav("input.mp3", "output.wav")
mock_ffmpeg.assert_called_once()
args = mock_ffmpeg.call_args[0][0]
assert "-i" in args
assert "input.mp3" in args
assert "-ar" in args
assert "16000" in args
assert "-ac" in args
assert "1" in args
class TestExtractAudioFromVideo:
@patch("yuxi.knowledge.parser.ffmpeg.run_ffmpeg")
def test_calls_ffmpeg_with_correct_args(self, mock_ffmpeg):
extract_audio_from_video("video.mp4", "audio.wav")
mock_ffmpeg.assert_called_once()
args = mock_ffmpeg.call_args[0][0]
assert "-i" in args
assert "video.mp4" in args
assert "-vn" in args
class TestExtractKeyframes:
@patch("yuxi.knowledge.parser.ffmpeg._detect_scene_change_pts", return_value=[])
def test_no_frames(self, mock_detect, tmp_path):
frames = extract_keyframes("video.mp4", str(tmp_path))
assert frames == []
@patch("yuxi.knowledge.parser.ffmpeg.run_ffmpeg")
@patch("yuxi.knowledge.parser.ffmpeg._detect_scene_change_pts", return_value=[5.0, 12.5])
def test_with_frames(self, mock_detect, mock_ffmpeg, tmp_path):
frames = extract_keyframes("video.mp4", str(tmp_path))
assert len(frames) == 2
assert frames[0][1] == 5.0
assert frames[1][1] == 12.5
class TestConstants:
def test_constants_values(self):
assert FFMPEG_MAX_BUFFER_BYTES == 10 * 1024 * 1024
assert FFPROBE_TIMEOUT_MS == 10_000
assert FFMPEG_TIMEOUT_MS == 45_000
assert FFMPEG_MAX_AUDIO_DURATION_SECS == 1200
assert MAX_AUDIO_BYTES == 16 * 1024 * 1024
assert MAX_VIDEO_BYTES == 16 * 1024 * 1024
assert FFMPEG_MAX_CONCURRENT == 4

View File

@ -0,0 +1,83 @@
"""TTS 模块单元测试。"""
import os
from unittest.mock import patch
import pytest
# 导入 tts 模块触发 __init__.py 中的 provider 注册
import yuxi.agents.toolkits.buildin.tts # noqa: F401
from yuxi.agents.toolkits.buildin.tts.openai_compatible import OpenAICompatibleTTSProvider
from yuxi.agents.toolkits.buildin.tts.provider import (
SpeechSynthesisResult,
TTSProvider,
get_tts_provider,
register_tts_provider,
)
class TestSpeechSynthesisResult:
def test_creation(self):
result = SpeechSynthesisResult(audio_data=b"audio", content_type="audio/mpeg")
assert result.audio_data == b"audio"
assert result.content_type == "audio/mpeg"
def test_default_content_type(self):
result = SpeechSynthesisResult(audio_data=b"audio")
assert result.content_type == "audio/mpeg"
def test_frozen(self):
result = SpeechSynthesisResult(audio_data=b"audio")
with pytest.raises(AttributeError):
result.audio_data = b"new"
class TestGetTTSProvider:
@patch.dict(os.environ, {"TTS_API_KEY": ""})
def test_no_configured_provider(self):
provider = get_tts_provider()
assert provider is None
@patch.dict(os.environ, {"TTS_API_KEY": "test-key"})
def test_configured_provider_by_id(self):
provider = get_tts_provider("openai_compatible")
assert provider is not None
assert provider.provider_id == "openai_compatible"
@patch.dict(os.environ, {"TTS_API_KEY": "test-key"})
def test_configured_provider_auto(self):
provider = get_tts_provider()
assert provider is not None
assert provider.provider_id == "openai_compatible"
def test_unknown_provider(self):
provider = get_tts_provider("nonexistent")
assert provider is None
@patch.dict(os.environ, {"TTS_API_KEY": ""})
def test_unconfigured_provider(self):
provider = get_tts_provider()
assert provider is None
class TestOpenAICompatibleTTSProvider:
@patch.dict(os.environ, {"TTS_API_KEY": "test-key", "TTS_BASE_URL": "https://api.test.com/v1", "TTS_MODEL": "tts-1"})
def test_is_configured(self):
provider = OpenAICompatibleTTSProvider()
assert provider.is_configured() is True
@patch.dict(os.environ, {"TTS_API_KEY": ""})
def test_not_configured(self):
provider = OpenAICompatibleTTSProvider()
assert provider.is_configured() is False
def test_provider_id(self):
provider = OpenAICompatibleTTSProvider()
assert provider.provider_id == "openai_compatible"
@pytest.mark.asyncio
async def test_list_voices(self):
provider = OpenAICompatibleTTSProvider()
voices = await provider.list_voices()
assert len(voices) == 6
assert voices[0]["id"] == "alloy"

View File

@ -0,0 +1,126 @@
"""视频解析器单元测试。"""
from unittest.mock import MagicMock, patch
import pytest
from yuxi.knowledge.parser.base import DocumentProcessorException
from yuxi.knowledge.parser.ffmpeg import FFmpegError
from yuxi.knowledge.parser.video import VideoProcessor
@pytest.fixture
def video():
return VideoProcessor()
class TestVideoProcessorBasics:
def test_service_name(self, video):
assert video.get_service_name() == "video"
def test_supported_extensions(self, video):
exts = video.get_supported_extensions()
assert ".mp4" in exts
assert ".avi" in exts
assert ".mov" in exts
assert ".mkv" in exts
def test_supports_file_type(self, video):
assert video.supports_file_type(".mp4")
assert not video.supports_file_type(".mp3")
class TestProcessFile:
@patch(
"yuxi.knowledge.parser.video.VideoProcessor._process_keyframes",
return_value="## 视频画面文字\n\n**[00:00]** text",
)
@patch(
"yuxi.knowledge.parser.video.VideoProcessor._process_audio_track",
return_value="# 音频转写\n\n**[00:00]** hello",
)
def test_success_both_tracks(self, mock_audio, mock_keyframes, video, tmp_path):
test_file = tmp_path / "test.mp4"
test_file.write_bytes(b"fake video data")
result = video.process_file(str(test_file))
assert "音频转写" in result
assert "视频画面文字" in result
assert "---" in result
@patch("yuxi.knowledge.parser.video.VideoProcessor._process_keyframes", return_value=None)
@patch("yuxi.knowledge.parser.video.VideoProcessor._process_audio_track", return_value="# 音频转写")
def test_audio_only(self, mock_audio, mock_keyframes, video, tmp_path):
test_file = tmp_path / "test.mp4"
test_file.write_bytes(b"fake video data")
result = video.process_file(str(test_file))
assert "音频转写" in result
@patch("yuxi.knowledge.parser.video.VideoProcessor._process_keyframes", return_value="## 视频画面文字")
@patch("yuxi.knowledge.parser.video.VideoProcessor._process_audio_track", return_value=None)
def test_keyframes_only(self, mock_audio, mock_keyframes, video, tmp_path):
test_file = tmp_path / "test.mp4"
test_file.write_bytes(b"fake video data")
result = video.process_file(str(test_file))
assert "视频画面文字" in result
@patch("yuxi.knowledge.parser.video.VideoProcessor._process_keyframes", return_value=None)
@patch("yuxi.knowledge.parser.video.VideoProcessor._process_audio_track", return_value=None)
def test_no_content_raises(self, mock_audio, mock_keyframes, video, tmp_path):
test_file = tmp_path / "test.mp4"
test_file.write_bytes(b"fake video data")
with pytest.raises(DocumentProcessorException, match="未产生任何内容"):
video.process_file(str(test_file))
def test_file_too_large(self, video, tmp_path):
test_file = tmp_path / "big.mp4"
test_file.write_bytes(b"x" * (16 * 1024 * 1024 + 1))
with pytest.raises(DocumentProcessorException, match="大小"):
video.process_file(str(test_file))
class TestProcessAudioTrack:
@patch("yuxi.knowledge.parser.factory.DocumentProcessorFactory")
@patch("yuxi.knowledge.parser.video.extract_audio_from_video")
def test_success(self, mock_extract, mock_factory, video):
mock_factory.process_file.return_value = "# 音频转写"
result = video._process_audio_track("video.mp4", {})
assert result == "# 音频转写"
@patch("yuxi.knowledge.parser.video.extract_audio_from_video", side_effect=FFmpegError("no audio"))
def test_no_audio_track_returns_none(self, mock_extract, video):
result = video._process_audio_track("video.mp4", {})
assert result is None
class TestProcessKeyframes:
@patch("yuxi.knowledge.parser.video.extract_keyframes", return_value=[])
def test_no_frames(self, mock_extract, video):
result = video._process_keyframes("video.mp4", {})
assert result is None
def test_ocr_disabled(self, video):
result = video._process_keyframes("video.mp4", {"ocr_engine": "disable"})
assert result is None
@patch("yuxi.knowledge.parser.factory.DocumentProcessorFactory")
@patch("yuxi.knowledge.parser.video.extract_keyframes")
def test_with_frames(self, mock_extract, mock_factory, video, tmp_path):
frame_path = str(tmp_path / "frame_0001.png")
mock_extract.return_value = [(frame_path, 1.0)]
mock_processor = MagicMock()
mock_processor.process_file.return_value = "OCR text"
mock_factory.get_processor.return_value = mock_processor
result = video._process_keyframes("video.mp4", {})
assert result is not None
assert "## 视频画面文字" in result
assert "**[00:01]** OCR text" in result
class TestCheckHealth:
@patch("yuxi.knowledge.parser.ffmpeg.run_ffmpeg", side_effect=FFmpegError("not found"))
def test_ffmpeg_unavailable(self, mock_ffmpeg, video):
result = video.check_health()
assert result["status"] == "unavailable"