feat: 新增TTS语音合成、音视频解析能力
1. 新增内置TTS工具,支持OpenAI兼容语音合成服务 2. 新增ASR音频解析和视频解析处理器 3. 扩展文件解析支持范围,添加音视频格式支持 4. 新增ffmpeg音视频处理基础工具库 5. 补充相关单元测试用例
This commit is contained in:
parent
436c09fb90
commit
b3bdb4fa2f
@ -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",
|
||||
]
|
||||
|
||||
44
backend/package/yuxi/agents/toolkits/buildin/tts.py
Normal file
44
backend/package/yuxi/agents/toolkits/buildin/tts.py
Normal 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
|
||||
17
backend/package/yuxi/agents/toolkits/buildin/tts/__init__.py
Normal file
17
backend/package/yuxi/agents/toolkits/buildin/tts/__init__.py
Normal 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",
|
||||
]
|
||||
@ -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"},
|
||||
]
|
||||
67
backend/package/yuxi/agents/toolkits/buildin/tts/provider.py
Normal file
67
backend/package/yuxi/agents/toolkits/buildin/tts/provider.py
Normal 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
|
||||
179
backend/package/yuxi/knowledge/parser/asr.py
Normal file
179
backend/package/yuxi/knowledge/parser/asr.py
Normal 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. 预处理:转换为 WAV(16kHz, 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 API(OpenAI 兼容)进行转写。"""
|
||||
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} 不可用"}
|
||||
@ -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
|
||||
|
||||
206
backend/package/yuxi/knowledge/parser/ffmpeg.py
Normal file
206
backend/package/yuxi/knowledge/parser/ffmpeg.py
Normal 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
|
||||
@ -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()
|
||||
|
||||
135
backend/package/yuxi/knowledge/parser/video.py
Normal file
135
backend/package/yuxi/knowledge/parser/video.py
Normal 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 服务不可用"}
|
||||
@ -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."""
|
||||
|
||||
123
backend/test/unit/agents/toolkits/buildin/test_asr.py
Normal file
123
backend/test/unit/agents/toolkits/buildin/test_asr.py
Normal 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"
|
||||
130
backend/test/unit/agents/toolkits/buildin/test_ffmpeg.py
Normal file
130
backend/test/unit/agents/toolkits/buildin/test_ffmpeg.py
Normal 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
|
||||
83
backend/test/unit/agents/toolkits/buildin/test_tts.py
Normal file
83
backend/test/unit/agents/toolkits/buildin/test_tts.py
Normal 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"
|
||||
126
backend/test/unit/agents/toolkits/buildin/test_video.py
Normal file
126
backend/test/unit/agents/toolkits/buildin/test_video.py
Normal 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"
|
||||
Loading…
Reference in New Issue
Block a user