Compare commits
15 Commits
| Author | SHA1 | Date | |
|---|---|---|---|
| 1da3ce95e2 | |||
| 683c0aba41 | |||
| a6636033f4 | |||
| a1c6074b3d | |||
| 783fe0bab3 | |||
| e6cf7548b1 | |||
| bc1c454fad | |||
| 9e2daf0022 | |||
| f61648bb0f | |||
| f4791e726e | |||
| 6dbbb5b71a | |||
| a6ebb762e6 | |||
| 38fda9aa9e | |||
| 77f9a908ad | |||
| 190d9e46f7 |
@@ -0,0 +1,16 @@
|
||||
import React from "react"
|
||||
import { Navigate } from "react-router-dom"
|
||||
import { useAuthStore } from "@/store/authStore"
|
||||
|
||||
/** 受保护的路由组件 */
|
||||
// eslint-disable-next-line react-refresh/only-export-components
|
||||
export const ProtectedRoute = ({ children }: { children: React.ReactNode }) => {
|
||||
const isAuthenticated = useAuthStore((state) => state.isAuthenticated)
|
||||
const hasAccessToken = Boolean(localStorage.getItem("access_token"))
|
||||
|
||||
if (!isAuthenticated || !hasAccessToken) {
|
||||
return <Navigate to="/login" replace />
|
||||
}
|
||||
|
||||
return <>{children}</>
|
||||
}
|
||||
@@ -0,0 +1,225 @@
|
||||
import { Navigate, type RouteObject } from "react-router-dom"
|
||||
import MainLayout from "@/components/layout/MainLayout"
|
||||
import { ProtectedRoute } from "./ProtectedRoute"
|
||||
|
||||
/**
|
||||
* 受保护的 /app 子路由
|
||||
* 所有页面使用 lazy 懒加载
|
||||
*/
|
||||
const appChildren: RouteObject[] = [
|
||||
{
|
||||
index: true,
|
||||
element: <Navigate to="/app/dashboard" replace />,
|
||||
},
|
||||
{
|
||||
path: "dashboard",
|
||||
lazy: () =>
|
||||
import("@/pages/dashboard/Dashboard").then((m) => ({
|
||||
Component: m.default,
|
||||
})),
|
||||
},
|
||||
{
|
||||
path: "assets",
|
||||
lazy: () =>
|
||||
import("@/pages/assets/AssetLibrary").then((m) => ({
|
||||
Component: m.default,
|
||||
})),
|
||||
},
|
||||
{
|
||||
path: "titles",
|
||||
lazy: () =>
|
||||
import("@/pages/titles/TitleLibrary").then((m) => ({
|
||||
Component: m.default,
|
||||
})),
|
||||
},
|
||||
{
|
||||
path: "voices",
|
||||
lazy: () =>
|
||||
import("@/pages/voices/VoiceLibrary").then((m) => ({
|
||||
Component: m.default,
|
||||
})),
|
||||
},
|
||||
{
|
||||
path: "templates",
|
||||
lazy: () =>
|
||||
import("@/pages/templates/TemplateLibrary").then((m) => ({
|
||||
Component: m.default,
|
||||
})),
|
||||
},
|
||||
{
|
||||
path: "generate",
|
||||
lazy: () =>
|
||||
import("@/pages/generate/GeneratePage").then((m) => ({
|
||||
Component: m.default,
|
||||
})),
|
||||
},
|
||||
{
|
||||
path: "history",
|
||||
lazy: () =>
|
||||
import("@/pages/history/TaskHistory").then((m) => ({
|
||||
Component: m.default,
|
||||
})),
|
||||
},
|
||||
{
|
||||
path: "products",
|
||||
lazy: () =>
|
||||
import("@/pages/products/ProductLibrary").then((m) => ({
|
||||
Component: m.default,
|
||||
})),
|
||||
},
|
||||
{
|
||||
path: "products/:id",
|
||||
lazy: () =>
|
||||
import("@/pages/products/ProductDetail").then((m) => ({
|
||||
Component: m.default,
|
||||
})),
|
||||
},
|
||||
{
|
||||
path: "tasks",
|
||||
lazy: () =>
|
||||
import("@/pages/tasks/TaskCenter").then((m) => ({
|
||||
Component: m.default,
|
||||
})),
|
||||
},
|
||||
{
|
||||
path: "editing-planner",
|
||||
lazy: () =>
|
||||
import("@/pages/editing-planner/EditingPlanner").then((m) => ({
|
||||
Component: m.default,
|
||||
})),
|
||||
},
|
||||
{
|
||||
path: "my-templates",
|
||||
lazy: () =>
|
||||
import("@/pages/my-templates/MyTemplates").then((m) => ({
|
||||
Component: m.default,
|
||||
})),
|
||||
},
|
||||
{
|
||||
path: "voice-clone",
|
||||
lazy: () =>
|
||||
import("@/pages/voice-clone/VoiceClone").then((m) => ({
|
||||
Component: m.default,
|
||||
})),
|
||||
},
|
||||
{
|
||||
path: "voice-materials",
|
||||
lazy: () =>
|
||||
import("@/pages/voice-materials/VoiceMaterialLibrary").then((m) => ({
|
||||
Component: m.default,
|
||||
})),
|
||||
},
|
||||
{
|
||||
path: "my-voices",
|
||||
lazy: () =>
|
||||
import("@/pages/my-voices/MyVoices").then((m) => ({
|
||||
Component: m.default,
|
||||
})),
|
||||
},
|
||||
{
|
||||
path: "accounts",
|
||||
lazy: () =>
|
||||
import("@/pages/accounts/Accounts").then((m) => ({
|
||||
Component: m.default,
|
||||
})),
|
||||
},
|
||||
{
|
||||
path: "duplication",
|
||||
lazy: () =>
|
||||
import("@/pages/duplication/DuplicationUpload").then((m) => ({
|
||||
Component: m.default,
|
||||
})),
|
||||
},
|
||||
{
|
||||
path: "duplication/results",
|
||||
lazy: () =>
|
||||
import("@/pages/duplication/DuplicationResults").then((m) => ({
|
||||
Component: m.default,
|
||||
})),
|
||||
},
|
||||
{
|
||||
path: "duplication/:id",
|
||||
lazy: () =>
|
||||
import("@/pages/duplication/DuplicationDetail").then((m) => ({
|
||||
Component: m.default,
|
||||
})),
|
||||
},
|
||||
{
|
||||
path: "subscription",
|
||||
lazy: () =>
|
||||
import("@/pages/subscription/Plans").then((m) => ({
|
||||
Component: m.default,
|
||||
})),
|
||||
},
|
||||
{
|
||||
path: "subscription/upgrade",
|
||||
lazy: () =>
|
||||
import("@/pages/subscription/UpgradeSubscription").then((m) => ({
|
||||
Component: m.default,
|
||||
})),
|
||||
},
|
||||
{
|
||||
path: "subscription/billing",
|
||||
lazy: () =>
|
||||
import("@/pages/subscription/Billing").then((m) => ({
|
||||
Component: m.default,
|
||||
})),
|
||||
},
|
||||
{
|
||||
path: "profile",
|
||||
lazy: () =>
|
||||
import("@/pages/profile/Settings").then((m) => ({
|
||||
Component: m.default,
|
||||
})),
|
||||
},
|
||||
{
|
||||
path: "admin",
|
||||
children: [
|
||||
{
|
||||
index: true,
|
||||
lazy: () =>
|
||||
import("@/pages/admin/AdminComingSoon").then((m) => ({
|
||||
Component: m.default,
|
||||
})),
|
||||
},
|
||||
{
|
||||
path: "users",
|
||||
lazy: () =>
|
||||
import("@/pages/admin/AdminComingSoon").then((m) => ({
|
||||
Component: m.default,
|
||||
})),
|
||||
},
|
||||
{
|
||||
path: "analytics",
|
||||
lazy: () =>
|
||||
import("@/pages/admin/AdminComingSoon").then((m) => ({
|
||||
Component: m.default,
|
||||
})),
|
||||
},
|
||||
{
|
||||
path: "monitor",
|
||||
lazy: () =>
|
||||
import("@/pages/admin/AdminComingSoon").then((m) => ({
|
||||
Component: m.default,
|
||||
})),
|
||||
},
|
||||
{
|
||||
path: "logs",
|
||||
lazy: () =>
|
||||
import("@/pages/admin/AdminComingSoon").then((m) => ({
|
||||
Component: m.default,
|
||||
})),
|
||||
},
|
||||
],
|
||||
},
|
||||
]
|
||||
|
||||
export const appRoutes: RouteObject = {
|
||||
path: "/app",
|
||||
element: (
|
||||
<ProtectedRoute>
|
||||
<MainLayout />
|
||||
</ProtectedRoute>
|
||||
),
|
||||
children: appChildren,
|
||||
}
|
||||
@@ -3,283 +3,13 @@
|
||||
* 扁平化路由:去掉 Project 层级,所有资源直接归属用户
|
||||
*/
|
||||
import { createBrowserRouter, Navigate } from "react-router-dom"
|
||||
import React from "react"
|
||||
import MainLayout from "@/components/layout/MainLayout"
|
||||
import Login from "@/pages/auth/Login"
|
||||
import Register from "@/pages/auth/Register"
|
||||
import ForgotPassword from "@/pages/auth/ForgotPassword"
|
||||
import ResetPassword from "@/pages/auth/ResetPassword"
|
||||
import WechatCallback from "@/pages/auth/WechatCallback"
|
||||
import HomePage from "@/pages/home/HomePage"
|
||||
import { useAuthStore } from "@/store/authStore"
|
||||
|
||||
/** 受保护的路由组件 */
|
||||
// eslint-disable-next-line react-refresh/only-export-components
|
||||
const ProtectedRoute = ({ children }: { children: React.ReactNode }) => {
|
||||
const isAuthenticated = useAuthStore((state) => state.isAuthenticated)
|
||||
const hasAccessToken = Boolean(localStorage.getItem("access_token"))
|
||||
|
||||
if (!isAuthenticated || !hasAccessToken) {
|
||||
return <Navigate to="/login" replace />
|
||||
}
|
||||
|
||||
return <>{children}</>
|
||||
}
|
||||
|
||||
/** 首页路由组件:已登录跳 dashboard,未登录显示落地页 */
|
||||
// eslint-disable-next-line react-refresh/only-export-components
|
||||
const HomeRoute: React.FC = () => {
|
||||
const isAuthenticated = useAuthStore((state) => state.isAuthenticated)
|
||||
const hasAccessToken = Boolean(localStorage.getItem("access_token"))
|
||||
|
||||
if (isAuthenticated && hasAccessToken) {
|
||||
return <Navigate to="/app/dashboard" replace />
|
||||
}
|
||||
|
||||
return <HomePage />
|
||||
}
|
||||
import { publicRoutes } from "./publicRoutes"
|
||||
import { appRoutes } from "./appRoutes"
|
||||
|
||||
/** 路由配置 */
|
||||
export const router = createBrowserRouter([
|
||||
{
|
||||
path: "/",
|
||||
element: <HomeRoute />,
|
||||
},
|
||||
{
|
||||
path: "/login",
|
||||
element: <Login />,
|
||||
},
|
||||
{
|
||||
path: "/register",
|
||||
element: <Register />,
|
||||
},
|
||||
{
|
||||
path: "/forgot-password",
|
||||
element: <ForgotPassword />,
|
||||
},
|
||||
{
|
||||
path: "/reset-password",
|
||||
element: <ResetPassword />,
|
||||
},
|
||||
{
|
||||
path: "/auth/wechat/callback",
|
||||
element: <WechatCallback />,
|
||||
},
|
||||
{
|
||||
path: "/app",
|
||||
element: (
|
||||
<ProtectedRoute>
|
||||
<MainLayout />
|
||||
</ProtectedRoute>
|
||||
),
|
||||
children: [
|
||||
{
|
||||
index: true,
|
||||
element: <Navigate to="/app/dashboard" replace />,
|
||||
},
|
||||
{
|
||||
path: "dashboard",
|
||||
lazy: () =>
|
||||
import("@/pages/dashboard/Dashboard").then((m) => ({
|
||||
Component: m.default,
|
||||
})),
|
||||
},
|
||||
{
|
||||
path: "assets",
|
||||
lazy: () =>
|
||||
import("@/pages/assets/AssetLibrary").then((m) => ({
|
||||
Component: m.default,
|
||||
})),
|
||||
},
|
||||
{
|
||||
path: "titles",
|
||||
lazy: () =>
|
||||
import("@/pages/titles/TitleLibrary").then((m) => ({
|
||||
Component: m.default,
|
||||
})),
|
||||
},
|
||||
{
|
||||
path: "voices",
|
||||
lazy: () =>
|
||||
import("@/pages/voices/VoiceLibrary").then((m) => ({
|
||||
Component: m.default,
|
||||
})),
|
||||
},
|
||||
{
|
||||
path: "templates",
|
||||
lazy: () =>
|
||||
import("@/pages/templates/TemplateLibrary").then((m) => ({
|
||||
Component: m.default,
|
||||
})),
|
||||
},
|
||||
{
|
||||
path: "generate",
|
||||
lazy: () =>
|
||||
import("@/pages/generate/GeneratePage").then((m) => ({
|
||||
Component: m.default,
|
||||
})),
|
||||
},
|
||||
{
|
||||
path: "history",
|
||||
lazy: () =>
|
||||
import("@/pages/history/TaskHistory").then((m) => ({
|
||||
Component: m.default,
|
||||
})),
|
||||
},
|
||||
{
|
||||
path: "products",
|
||||
lazy: () =>
|
||||
import("@/pages/products/ProductLibrary").then((m) => ({
|
||||
Component: m.default,
|
||||
})),
|
||||
},
|
||||
{
|
||||
path: "products/:id",
|
||||
lazy: () =>
|
||||
import("@/pages/products/ProductDetail").then((m) => ({
|
||||
Component: m.default,
|
||||
})),
|
||||
},
|
||||
{
|
||||
path: "tasks",
|
||||
lazy: () =>
|
||||
import("@/pages/tasks/TaskCenter").then((m) => ({
|
||||
Component: m.default,
|
||||
})),
|
||||
},
|
||||
{
|
||||
path: "editing-planner",
|
||||
lazy: () =>
|
||||
import("@/pages/editing-planner/EditingPlanner").then((m) => ({
|
||||
Component: m.default,
|
||||
})),
|
||||
},
|
||||
{
|
||||
path: "my-templates",
|
||||
lazy: () =>
|
||||
import("@/pages/my-templates/MyTemplates").then((m) => ({
|
||||
Component: m.default,
|
||||
})),
|
||||
},
|
||||
{
|
||||
path: "voice-clone",
|
||||
lazy: () =>
|
||||
import("@/pages/voice-clone/VoiceClone").then((m) => ({
|
||||
Component: m.default,
|
||||
})),
|
||||
},
|
||||
{
|
||||
path: "voice-materials",
|
||||
lazy: () =>
|
||||
import("@/pages/voice-materials/VoiceMaterialLibrary").then((m) => ({
|
||||
Component: m.default,
|
||||
})),
|
||||
},
|
||||
{
|
||||
path: "my-voices",
|
||||
lazy: () =>
|
||||
import("@/pages/my-voices/MyVoices").then((m) => ({
|
||||
Component: m.default,
|
||||
})),
|
||||
},
|
||||
{
|
||||
path: "accounts",
|
||||
lazy: () =>
|
||||
import("@/pages/accounts/Accounts").then((m) => ({
|
||||
Component: m.default,
|
||||
})),
|
||||
},
|
||||
{
|
||||
path: "duplication",
|
||||
lazy: () =>
|
||||
import("@/pages/duplication/DuplicationUpload").then((m) => ({
|
||||
Component: m.default,
|
||||
})),
|
||||
},
|
||||
{
|
||||
path: "duplication/results",
|
||||
lazy: () =>
|
||||
import("@/pages/duplication/DuplicationResults").then((m) => ({
|
||||
Component: m.default,
|
||||
})),
|
||||
},
|
||||
{
|
||||
path: "duplication/:id",
|
||||
lazy: () =>
|
||||
import("@/pages/duplication/DuplicationDetail").then((m) => ({
|
||||
Component: m.default,
|
||||
})),
|
||||
},
|
||||
{
|
||||
path: "subscription",
|
||||
lazy: () =>
|
||||
import("@/pages/subscription/Plans").then((m) => ({
|
||||
Component: m.default,
|
||||
})),
|
||||
},
|
||||
{
|
||||
path: "subscription/upgrade",
|
||||
lazy: () =>
|
||||
import("@/pages/subscription/UpgradeSubscription").then((m) => ({
|
||||
Component: m.default,
|
||||
})),
|
||||
},
|
||||
{
|
||||
path: "subscription/billing",
|
||||
lazy: () =>
|
||||
import("@/pages/subscription/Billing").then((m) => ({
|
||||
Component: m.default,
|
||||
})),
|
||||
},
|
||||
{
|
||||
path: "profile",
|
||||
lazy: () =>
|
||||
import("@/pages/profile/Settings").then((m) => ({
|
||||
Component: m.default,
|
||||
})),
|
||||
},
|
||||
{
|
||||
path: "admin",
|
||||
children: [
|
||||
{
|
||||
index: true,
|
||||
lazy: () =>
|
||||
import("@/pages/admin/AdminComingSoon").then((m) => ({
|
||||
Component: m.default,
|
||||
})),
|
||||
},
|
||||
{
|
||||
path: "users",
|
||||
lazy: () =>
|
||||
import("@/pages/admin/AdminComingSoon").then((m) => ({
|
||||
Component: m.default,
|
||||
})),
|
||||
},
|
||||
{
|
||||
path: "analytics",
|
||||
lazy: () =>
|
||||
import("@/pages/admin/AdminComingSoon").then((m) => ({
|
||||
Component: m.default,
|
||||
})),
|
||||
},
|
||||
{
|
||||
path: "monitor",
|
||||
lazy: () =>
|
||||
import("@/pages/admin/AdminComingSoon").then((m) => ({
|
||||
Component: m.default,
|
||||
})),
|
||||
},
|
||||
{
|
||||
path: "logs",
|
||||
lazy: () =>
|
||||
import("@/pages/admin/AdminComingSoon").then((m) => ({
|
||||
Component: m.default,
|
||||
})),
|
||||
},
|
||||
],
|
||||
},
|
||||
],
|
||||
},
|
||||
...publicRoutes,
|
||||
appRoutes,
|
||||
{
|
||||
path: "*",
|
||||
element: <Navigate to="/" replace />,
|
||||
|
||||
@@ -0,0 +1,48 @@
|
||||
import { Navigate, type RouteObject } from "react-router-dom"
|
||||
import HomePage from "@/pages/home/HomePage"
|
||||
import Login from "@/pages/auth/Login"
|
||||
import Register from "@/pages/auth/Register"
|
||||
import ForgotPassword from "@/pages/auth/ForgotPassword"
|
||||
import ResetPassword from "@/pages/auth/ResetPassword"
|
||||
import WechatCallback from "@/pages/auth/WechatCallback"
|
||||
import { useAuthStore } from "@/store/authStore"
|
||||
|
||||
/** 首页路由组件:已登录跳 dashboard,未登录显示落地页 */
|
||||
// eslint-disable-next-line react-refresh/only-export-components
|
||||
const HomeRoute: React.FC = () => {
|
||||
const isAuthenticated = useAuthStore((state) => state.isAuthenticated)
|
||||
const hasAccessToken = Boolean(localStorage.getItem("access_token"))
|
||||
|
||||
if (isAuthenticated && hasAccessToken) {
|
||||
return <Navigate to="/app/dashboard" replace />
|
||||
}
|
||||
|
||||
return <HomePage />
|
||||
}
|
||||
|
||||
export const publicRoutes: RouteObject[] = [
|
||||
{
|
||||
path: "/",
|
||||
element: <HomeRoute />,
|
||||
},
|
||||
{
|
||||
path: "/login",
|
||||
element: <Login />,
|
||||
},
|
||||
{
|
||||
path: "/register",
|
||||
element: <Register />,
|
||||
},
|
||||
{
|
||||
path: "/forgot-password",
|
||||
element: <ForgotPassword />,
|
||||
},
|
||||
{
|
||||
path: "/reset-password",
|
||||
element: <ResetPassword />,
|
||||
},
|
||||
{
|
||||
path: "/auth/wechat/callback",
|
||||
element: <WechatCallback />,
|
||||
},
|
||||
]
|
||||
@@ -1,4 +1,4 @@
|
||||
"""视频调速引擎 — 基于 FFmpeg setpts + atempo 的速度调整能力。
|
||||
"""视频调速引擎 — 基于 FFmpeg setpts + atempo 的速度调整能力.
|
||||
|
||||
支持:
|
||||
- 0.25x ~ 4x 变速范围
|
||||
@@ -6,147 +6,57 @@
|
||||
- 音频调速(atempo,多级串联处理超范围值)
|
||||
- 音调修正(pitch_correct,默认开启)
|
||||
- 边界自动钳制,不阻断渲染
|
||||
|
||||
注:核心领域模型已抽离到 packages/domain/speed_config.py,
|
||||
本模块保留薄包装层,确保向后兼容。
|
||||
"""
|
||||
|
||||
from dataclasses import dataclass
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import Optional
|
||||
|
||||
# ─── 常量 ───────────────────────────────────────────────
|
||||
MIN_SPEED = 0.25
|
||||
MAX_SPEED = 4.0
|
||||
DEFAULT_SPEED = 1.0
|
||||
|
||||
# atempo 单级有效范围
|
||||
_ATEMPO_MIN = 0.5
|
||||
_ATEMPO_MAX = 2.0
|
||||
|
||||
|
||||
@dataclass
|
||||
class SpeedConfig:
|
||||
"""调速配置。
|
||||
|
||||
Attributes:
|
||||
speed: 播放速度,0.25~4.0,1.0 为原速
|
||||
pitch_correct: 是否保持音调(默认 True,用 atempo 时间拉伸算法)
|
||||
"""
|
||||
|
||||
speed: float = DEFAULT_SPEED
|
||||
pitch_correct: bool = True
|
||||
|
||||
@classmethod
|
||||
def parse(cls, data: Optional[dict]) -> "SpeedConfig":
|
||||
"""从 dict 解析配置,无效值回退到默认。"""
|
||||
if not data or not isinstance(data, dict):
|
||||
return cls()
|
||||
|
||||
speed = data.get("speed", DEFAULT_SPEED)
|
||||
if not isinstance(speed, (int, float)):
|
||||
speed = DEFAULT_SPEED
|
||||
|
||||
pitch_correct = data.get("pitch_correct", True)
|
||||
if not isinstance(pitch_correct, bool):
|
||||
pitch_correct = True
|
||||
|
||||
config = cls(speed=float(speed), pitch_correct=pitch_correct)
|
||||
config.clamp()
|
||||
return config
|
||||
|
||||
def clamp(self) -> None:
|
||||
"""将速度钳制到合法范围。"""
|
||||
if self.speed <= 0:
|
||||
self.speed = DEFAULT_SPEED
|
||||
elif self.speed < MIN_SPEED:
|
||||
self.speed = MIN_SPEED
|
||||
elif self.speed > MAX_SPEED:
|
||||
self.speed = MAX_SPEED
|
||||
|
||||
@property
|
||||
def is_original(self) -> bool:
|
||||
"""是否原速(无需调速)。"""
|
||||
return abs(self.speed - 1.0) < 1e-6
|
||||
from packages.domain.speed_config import ( # noqa: F401 — 向后兼容
|
||||
DEFAULT_SPEED,
|
||||
MAX_SPEED,
|
||||
MIN_SPEED,
|
||||
SpeedConfig,
|
||||
_split_atempo_stages,
|
||||
adjust_duration as _adjust_duration_base,
|
||||
build_audio_filter as _build_audio_filter_base,
|
||||
build_video_filter as _build_video_filter_base,
|
||||
resolve_clip_speed as _resolve_clip_speed_base,
|
||||
)
|
||||
|
||||
|
||||
class SpeedEngine:
|
||||
"""调速引擎 — 生成 FFmpeg 调速滤镜链。
|
||||
"""调速引擎 — 生成 FFmpeg 调速滤镜链.
|
||||
|
||||
用法:
|
||||
engine = SpeedEngine()
|
||||
video_filter = engine.build_video_filter(config)
|
||||
audio_filter = engine.build_audio_filter(config)
|
||||
new_duration = engine.adjust_duration(duration, config)
|
||||
薄包装层,实际逻辑委托给 packages.domain.speed_config。
|
||||
"""
|
||||
|
||||
def build_video_filter(self, config: SpeedConfig) -> str:
|
||||
"""生成视频调速滤镜字符串。
|
||||
|
||||
返回 setpts 滤镜表达式,原速时返回空字符串。
|
||||
"""
|
||||
if config.is_original:
|
||||
return ""
|
||||
# setpts=PTS/speed — speed>1 加速,speed<1 减速
|
||||
return f"setpts=PTS/{config.speed:.4f}"
|
||||
"""生成视频调速滤镜字符串."""
|
||||
return _build_video_filter_base(config)
|
||||
|
||||
def build_audio_filter(self, config: SpeedConfig) -> str:
|
||||
"""生成音频调速滤镜字符串。
|
||||
|
||||
atempo 单级范围 0.5~2.0,超出范围时自动多级串联:
|
||||
- 0.25x → atempo=0.5,atempo=0.5
|
||||
- 4x → atempo=2.0,atempo=2.0
|
||||
- 0.3x → atempo=0.5,atempo=0.6
|
||||
- 3x → atempo=2.0,atempo=1.5
|
||||
|
||||
原速时返回空字符串。
|
||||
"""
|
||||
if config.is_original:
|
||||
return ""
|
||||
|
||||
speed = config.speed
|
||||
stages: list[float] = self._split_atempo_stages(speed)
|
||||
return ",".join(f"atempo={s:.4f}" for s in stages)
|
||||
"""生成音频调速滤镜字符串."""
|
||||
return _build_audio_filter_base(config)
|
||||
|
||||
@staticmethod
|
||||
def _split_atempo_stages(speed: float) -> list[float]:
|
||||
"""将速度拆分为多级 atempo 串联,每级都在 [0.5, 2.0] 范围内。"""
|
||||
if _ATEMPO_MIN <= speed <= _ATEMPO_MAX:
|
||||
return [speed]
|
||||
|
||||
stages: list[float] = []
|
||||
remaining = speed
|
||||
|
||||
# 加速场景(speed > 2.0)
|
||||
if speed > _ATEMPO_MAX:
|
||||
while remaining > _ATEMPO_MAX:
|
||||
stages.append(_ATEMPO_MAX)
|
||||
remaining /= _ATEMPO_MAX
|
||||
stages.append(remaining)
|
||||
|
||||
# 减速场景(speed < 0.5)
|
||||
else:
|
||||
while remaining < _ATEMPO_MIN:
|
||||
stages.append(_ATEMPO_MIN)
|
||||
remaining /= _ATEMPO_MIN
|
||||
stages.append(remaining)
|
||||
|
||||
return stages
|
||||
"""将速度拆分为多级 atempo 串联(内部方法,向后兼容)."""
|
||||
return _split_atempo_stages(speed)
|
||||
|
||||
def adjust_duration(self, original_duration: float, config: SpeedConfig) -> float:
|
||||
"""计算调速后的时长。
|
||||
|
||||
加速 → 时长变短;减速 → 时长变长。
|
||||
"""
|
||||
if config.is_original or original_duration <= 0:
|
||||
return original_duration
|
||||
return original_duration / config.speed
|
||||
"""计算调速后的时长."""
|
||||
return _adjust_duration_base(original_duration, config)
|
||||
|
||||
def build_clip_speed_filter(
|
||||
self,
|
||||
speed: float,
|
||||
pitch_correct: bool = True,
|
||||
) -> tuple[str, str, SpeedConfig]:
|
||||
"""便捷方法:从单一 speed 值生成视频+音频滤镜。
|
||||
|
||||
返回 (video_filter, audio_filter, config)。
|
||||
"""
|
||||
"""便捷方法:从单一 speed 值生成视频+音频滤镜."""
|
||||
config = SpeedConfig(speed=speed, pitch_correct=pitch_correct)
|
||||
config.clamp()
|
||||
return (
|
||||
@@ -160,8 +70,5 @@ class SpeedEngine:
|
||||
clip_config: dict,
|
||||
global_speed: float = DEFAULT_SPEED,
|
||||
) -> float:
|
||||
"""从 clip config 中解析 playback_speed,0 或缺失则使用全局速度。"""
|
||||
speed = clip_config.get("playback_speed", 0) if clip_config else 0
|
||||
if not isinstance(speed, (int, float)) or speed <= 0:
|
||||
return global_speed
|
||||
return float(speed)
|
||||
"""从 clip config 中解析 playback_speed,0 或缺失则使用全局速度."""
|
||||
return _resolve_clip_speed_base(clip_config, global_speed)
|
||||
|
||||
@@ -0,0 +1,179 @@
|
||||
"""调速配置领域模型 — 纯逻辑,无FFmpeg依赖.
|
||||
|
||||
抽离自 speed_engine.py,包含:
|
||||
- SpeedConfig 数据类(解析/钳制/原速判断)
|
||||
- 视频/音频调速滤镜构建
|
||||
- atempo 多级拆分算法
|
||||
- 时长计算
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from dataclasses import dataclass
|
||||
from typing import Any
|
||||
|
||||
# ─── 常量 ───────────────────────────────────────────────
|
||||
MIN_SPEED = 0.25
|
||||
MAX_SPEED = 4.0
|
||||
DEFAULT_SPEED = 1.0
|
||||
|
||||
# atempo 单级有效范围
|
||||
_ATEMPO_MIN = 0.5
|
||||
_ATEMPO_MAX = 2.0
|
||||
|
||||
|
||||
@dataclass
|
||||
class SpeedConfig:
|
||||
"""调速配置.
|
||||
|
||||
Attributes:
|
||||
speed: 播放速度,0.25~4.0,1.0 为原速
|
||||
pitch_correct: 是否保持音调(默认 True,用 atempo 时间拉伸算法)
|
||||
"""
|
||||
|
||||
speed: float = DEFAULT_SPEED
|
||||
pitch_correct: bool = True
|
||||
|
||||
@classmethod
|
||||
def parse(cls, data: dict[str, Any] | None) -> SpeedConfig:
|
||||
"""从 dict 解析配置,无效值回退到默认."""
|
||||
if not data or not isinstance(data, dict):
|
||||
return cls()
|
||||
|
||||
speed = data.get("speed", DEFAULT_SPEED)
|
||||
if not isinstance(speed, (int, float)):
|
||||
speed = DEFAULT_SPEED
|
||||
|
||||
pitch_correct = data.get("pitch_correct", True)
|
||||
if not isinstance(pitch_correct, bool):
|
||||
pitch_correct = True
|
||||
|
||||
config = cls(speed=float(speed), pitch_correct=pitch_correct)
|
||||
config.clamp()
|
||||
return config
|
||||
|
||||
def clamp(self) -> None:
|
||||
"""将速度钳制到合法范围."""
|
||||
if self.speed <= 0:
|
||||
self.speed = DEFAULT_SPEED
|
||||
elif self.speed < MIN_SPEED:
|
||||
self.speed = MIN_SPEED
|
||||
elif self.speed > MAX_SPEED:
|
||||
self.speed = MAX_SPEED
|
||||
|
||||
@property
|
||||
def is_original(self) -> bool:
|
||||
"""是否原速(无需调速)."""
|
||||
return abs(self.speed - 1.0) < 1e-6
|
||||
|
||||
@property
|
||||
def is_fast(self) -> bool:
|
||||
"""是否加速播放."""
|
||||
return self.speed > 1.0
|
||||
|
||||
@property
|
||||
def is_slow(self) -> bool:
|
||||
"""是否减速播放."""
|
||||
return self.speed < 1.0
|
||||
|
||||
|
||||
# ── 滤镜构建 ────────────────────────────────────────────────────────────────
|
||||
|
||||
|
||||
def build_video_filter(config: SpeedConfig) -> str:
|
||||
"""生成视频调速滤镜字符串.
|
||||
|
||||
返回 setpts 滤镜表达式,原速时返回空字符串。
|
||||
"""
|
||||
if config.is_original:
|
||||
return ""
|
||||
# setpts=PTS/speed — speed>1 加速,speed<1 减速
|
||||
return f"setpts=PTS/{config.speed:.4f}"
|
||||
|
||||
|
||||
def build_audio_filter(config: SpeedConfig) -> str:
|
||||
"""生成音频调速滤镜字符串.
|
||||
|
||||
atempo 单级范围 0.5~2.0,超出范围时自动多级串联:
|
||||
- 0.25x → atempo=0.5,atempo=0.5
|
||||
- 4x → atempo=2.0,atempo=2.0
|
||||
- 0.3x → atempo=0.5,atempo=0.6
|
||||
- 3x → atempo=2.0,atempo=1.5
|
||||
|
||||
原速时返回空字符串。
|
||||
"""
|
||||
if config.is_original:
|
||||
return ""
|
||||
|
||||
speed = config.speed
|
||||
stages: list[float] = _split_atempo_stages(speed)
|
||||
return ",".join(f"atempo={s:.4f}" for s in stages)
|
||||
|
||||
|
||||
def _split_atempo_stages(speed: float) -> list[float]:
|
||||
"""将速度拆分为多级 atempo 串联,每级都在 [0.5, 2.0] 范围内."""
|
||||
if _ATEMPO_MIN <= speed <= _ATEMPO_MAX:
|
||||
return [speed]
|
||||
|
||||
stages: list[float] = []
|
||||
remaining = speed
|
||||
|
||||
# 加速场景(speed > 2.0)
|
||||
if speed > _ATEMPO_MAX:
|
||||
while remaining > _ATEMPO_MAX:
|
||||
stages.append(_ATEMPO_MAX)
|
||||
remaining /= _ATEMPO_MAX
|
||||
stages.append(remaining)
|
||||
|
||||
# 减速场景(speed < 0.5)
|
||||
else:
|
||||
while remaining < _ATEMPO_MIN:
|
||||
stages.append(_ATEMPO_MIN)
|
||||
remaining /= _ATEMPO_MIN
|
||||
stages.append(remaining)
|
||||
|
||||
return stages
|
||||
|
||||
|
||||
# ── 时长计算 ────────────────────────────────────────────────────────────────
|
||||
|
||||
|
||||
def adjust_duration(original_duration: float, config: SpeedConfig) -> float:
|
||||
"""计算调速后的时长.
|
||||
|
||||
加速 → 时长变短;减速 → 时长变长。
|
||||
"""
|
||||
if config.is_original or original_duration <= 0:
|
||||
return original_duration
|
||||
return original_duration / config.speed
|
||||
|
||||
|
||||
# ── 便捷方法 ────────────────────────────────────────────────────────────────
|
||||
|
||||
|
||||
def build_clip_speed_filter(
|
||||
speed: float,
|
||||
pitch_correct: bool = True,
|
||||
) -> tuple[str, str, SpeedConfig]:
|
||||
"""便捷方法:从单一 speed 值生成视频+音频滤镜.
|
||||
|
||||
返回 (video_filter, audio_filter, config)。
|
||||
"""
|
||||
config = SpeedConfig(speed=speed, pitch_correct=pitch_correct)
|
||||
config.clamp()
|
||||
return (
|
||||
build_video_filter(config),
|
||||
build_audio_filter(config),
|
||||
config,
|
||||
)
|
||||
|
||||
|
||||
def resolve_clip_speed(
|
||||
clip_config: dict[str, Any] | None,
|
||||
global_speed: float = DEFAULT_SPEED,
|
||||
) -> float:
|
||||
"""从 clip config 中解析 playback_speed,0 或缺失则使用全局速度."""
|
||||
speed = clip_config.get("playback_speed", 0) if clip_config else 0
|
||||
if not isinstance(speed, (int, float)) or speed <= 0:
|
||||
return global_speed
|
||||
return float(speed)
|
||||
@@ -0,0 +1,313 @@
|
||||
"""speed_config 领域模型单测."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import pytest
|
||||
|
||||
from packages.domain.speed_config import (
|
||||
DEFAULT_SPEED,
|
||||
MAX_SPEED,
|
||||
MIN_SPEED,
|
||||
SpeedConfig,
|
||||
adjust_duration,
|
||||
build_audio_filter,
|
||||
build_video_filter,
|
||||
build_clip_speed_filter,
|
||||
resolve_clip_speed,
|
||||
)
|
||||
|
||||
# ── 常量测试 ────────────────────────────────────────────────────────────────
|
||||
|
||||
|
||||
class TestConstants:
|
||||
def test_min_speed(self):
|
||||
assert MIN_SPEED == 0.25
|
||||
|
||||
def test_max_speed(self):
|
||||
assert MAX_SPEED == 4.0
|
||||
|
||||
def test_default_speed(self):
|
||||
assert DEFAULT_SPEED == 1.0
|
||||
|
||||
|
||||
# ── SpeedConfig.parse 测试 ────────────────────────────────────────────────
|
||||
|
||||
|
||||
class TestSpeedConfigParse:
|
||||
def test_none_returns_default(self):
|
||||
cfg = SpeedConfig.parse(None)
|
||||
assert cfg.speed == DEFAULT_SPEED
|
||||
assert cfg.pitch_correct is True
|
||||
|
||||
def test_empty_dict_returns_default(self):
|
||||
cfg = SpeedConfig.parse({})
|
||||
assert cfg.speed == DEFAULT_SPEED
|
||||
|
||||
def test_invalid_type_returns_default(self):
|
||||
cfg = SpeedConfig.parse("not_a_dict")
|
||||
assert cfg.speed == DEFAULT_SPEED
|
||||
|
||||
def test_valid_speed(self):
|
||||
cfg = SpeedConfig.parse({"speed": 2.0})
|
||||
assert cfg.speed == 2.0
|
||||
|
||||
def test_speed_clamped_low(self):
|
||||
cfg = SpeedConfig.parse({"speed": 0.1})
|
||||
assert cfg.speed == MIN_SPEED
|
||||
|
||||
def test_speed_clamped_high(self):
|
||||
cfg = SpeedConfig.parse({"speed": 5.0})
|
||||
assert cfg.speed == MAX_SPEED
|
||||
|
||||
def test_zero_speed_returns_default(self):
|
||||
cfg = SpeedConfig.parse({"speed": 0})
|
||||
assert cfg.speed == DEFAULT_SPEED
|
||||
|
||||
def test_negative_speed_returns_default(self):
|
||||
cfg = SpeedConfig.parse({"speed": -1.0})
|
||||
assert cfg.speed == DEFAULT_SPEED
|
||||
|
||||
def test_pitch_correct_false(self):
|
||||
cfg = SpeedConfig.parse({"pitch_correct": False})
|
||||
assert cfg.pitch_correct is False
|
||||
|
||||
def test_pitch_correct_invalid_type_defaults_true(self):
|
||||
cfg = SpeedConfig.parse({"pitch_correct": "yes"})
|
||||
assert cfg.pitch_correct is True
|
||||
|
||||
def test_string_speed_invalid_uses_default(self):
|
||||
cfg = SpeedConfig.parse({"speed": "fast"})
|
||||
assert cfg.speed == DEFAULT_SPEED
|
||||
|
||||
|
||||
# ── SpeedConfig.clamp 测试 ────────────────────────────────────────────────
|
||||
|
||||
|
||||
class TestClamp:
|
||||
def test_already_valid_unchanged(self):
|
||||
cfg = SpeedConfig(speed=1.5)
|
||||
cfg.clamp()
|
||||
assert cfg.speed == 1.5
|
||||
|
||||
def test_below_min_clamped(self):
|
||||
cfg = SpeedConfig(speed=0.1)
|
||||
cfg.clamp()
|
||||
assert cfg.speed == MIN_SPEED
|
||||
|
||||
def test_above_max_clamped(self):
|
||||
cfg = SpeedConfig(speed=10.0)
|
||||
cfg.clamp()
|
||||
assert cfg.speed == MAX_SPEED
|
||||
|
||||
def test_zero_defaults(self):
|
||||
cfg = SpeedConfig(speed=0.0)
|
||||
cfg.clamp()
|
||||
assert cfg.speed == DEFAULT_SPEED
|
||||
|
||||
def test_negative_defaults(self):
|
||||
cfg = SpeedConfig(speed=-2.0)
|
||||
cfg.clamp()
|
||||
assert cfg.speed == DEFAULT_SPEED
|
||||
|
||||
def test_exact_min_stays(self):
|
||||
cfg = SpeedConfig(speed=MIN_SPEED)
|
||||
cfg.clamp()
|
||||
assert cfg.speed == MIN_SPEED
|
||||
|
||||
def test_exact_max_stays(self):
|
||||
cfg = SpeedConfig(speed=MAX_SPEED)
|
||||
cfg.clamp()
|
||||
assert cfg.speed == MAX_SPEED
|
||||
|
||||
|
||||
# ── is_original / is_fast / is_slow 测试 ─────────────────────────────────
|
||||
|
||||
|
||||
class TestSpeedProperties:
|
||||
def test_is_original_true(self):
|
||||
cfg = SpeedConfig(speed=1.0)
|
||||
assert cfg.is_original is True
|
||||
|
||||
def test_is_original_false_fast(self):
|
||||
cfg = SpeedConfig(speed=2.0)
|
||||
assert cfg.is_original is False
|
||||
|
||||
def test_is_original_false_slow(self):
|
||||
cfg = SpeedConfig(speed=0.5)
|
||||
assert cfg.is_original is False
|
||||
|
||||
def test_is_original_near_one(self):
|
||||
cfg = SpeedConfig(speed=1.0000001)
|
||||
assert cfg.is_original is True
|
||||
|
||||
def test_is_fast_true(self):
|
||||
cfg = SpeedConfig(speed=2.0)
|
||||
assert cfg.is_fast is True
|
||||
|
||||
def test_is_fast_false(self):
|
||||
cfg = SpeedConfig(speed=0.5)
|
||||
assert cfg.is_fast is False
|
||||
|
||||
def test_is_false_for_original(self):
|
||||
cfg = SpeedConfig(speed=1.0)
|
||||
assert cfg.is_fast is False
|
||||
assert cfg.is_slow is False
|
||||
|
||||
def test_is_slow_true(self):
|
||||
cfg = SpeedConfig(speed=0.5)
|
||||
assert cfg.is_slow is True
|
||||
|
||||
def test_is_slow_false(self):
|
||||
cfg = SpeedConfig(speed=2.0)
|
||||
assert cfg.is_slow is False
|
||||
|
||||
|
||||
# ── build_video_filter 测试 ───────────────────────────────────────────────
|
||||
|
||||
|
||||
class TestBuildVideoFilter:
|
||||
def test_original_speed_empty(self):
|
||||
cfg = SpeedConfig(speed=1.0)
|
||||
assert build_video_filter(cfg) == ""
|
||||
|
||||
def test_fast_speed_setpts(self):
|
||||
cfg = SpeedConfig(speed=2.0)
|
||||
result = build_video_filter(cfg)
|
||||
assert "setpts=PTS/2.0000" in result
|
||||
|
||||
def test_slow_speed_setpts(self):
|
||||
cfg = SpeedConfig(speed=0.5)
|
||||
result = build_video_filter(cfg)
|
||||
assert "setpts=PTS/0.5000" in result
|
||||
|
||||
def test_format_four_decimals(self):
|
||||
cfg = SpeedConfig(speed=1.5)
|
||||
result = build_video_filter(cfg)
|
||||
assert "1.5000" in result
|
||||
|
||||
|
||||
# ── build_audio_filter 测试 ───────────────────────────────────────────────
|
||||
|
||||
|
||||
class TestBuildAudioFilter:
|
||||
def test_original_speed_empty(self):
|
||||
cfg = SpeedConfig(speed=1.0)
|
||||
assert build_audio_filter(cfg) == ""
|
||||
|
||||
def test_single_stage_within_range(self):
|
||||
cfg = SpeedConfig(speed=1.5)
|
||||
result = build_audio_filter(cfg)
|
||||
assert result == "atempo=1.5000"
|
||||
assert result.count("atempo") == 1
|
||||
|
||||
def test_fast_two_stages(self):
|
||||
cfg = SpeedConfig(speed=3.0)
|
||||
result = build_audio_filter(cfg)
|
||||
assert result.count("atempo") == 2
|
||||
# 2.0 * 1.5 = 3.0
|
||||
assert "atempo=2.0000" in result
|
||||
assert "atempo=1.5000" in result
|
||||
|
||||
def test_max_speed_two_stages(self):
|
||||
cfg = SpeedConfig(speed=4.0)
|
||||
result = build_audio_filter(cfg)
|
||||
assert result.count("atempo") == 2
|
||||
# 2.0 * 2.0 = 4.0
|
||||
assert result == "atempo=2.0000,atempo=2.0000"
|
||||
|
||||
def test_slow_two_stages(self):
|
||||
cfg = SpeedConfig(speed=0.25)
|
||||
result = build_audio_filter(cfg)
|
||||
assert result.count("atempo") == 2
|
||||
# 0.5 * 0.5 = 0.25
|
||||
assert result == "atempo=0.5000,atempo=0.5000"
|
||||
|
||||
def test_slow_single_stage(self):
|
||||
cfg = SpeedConfig(speed=0.8)
|
||||
result = build_audio_filter(cfg)
|
||||
assert result == "atempo=0.8000"
|
||||
assert result.count("atempo") == 1
|
||||
|
||||
def test_exactly_two_point_zero_single(self):
|
||||
cfg = SpeedConfig(speed=2.0)
|
||||
result = build_audio_filter(cfg)
|
||||
assert result.count("atempo") == 1
|
||||
assert "atempo=2.0000" in result
|
||||
|
||||
def test_exactly_half_single(self):
|
||||
cfg = SpeedConfig(speed=0.5)
|
||||
result = build_audio_filter(cfg)
|
||||
assert result.count("atempo") == 1
|
||||
assert "atempo=0.5000" in result
|
||||
|
||||
|
||||
# ── adjust_duration 测试 ──────────────────────────────────────────────────
|
||||
|
||||
|
||||
class TestAdjustDuration:
|
||||
def test_original_speed_unchanged(self):
|
||||
cfg = SpeedConfig(speed=1.0)
|
||||
assert adjust_duration(10.0, cfg) == 10.0
|
||||
|
||||
def test_double_speed_halved(self):
|
||||
cfg = SpeedConfig(speed=2.0)
|
||||
assert adjust_duration(10.0, cfg) == 5.0
|
||||
|
||||
def test_half_speed_doubled(self):
|
||||
cfg = SpeedConfig(speed=0.5)
|
||||
assert adjust_duration(10.0, cfg) == 20.0
|
||||
|
||||
def test_zero_duration_unchanged(self):
|
||||
cfg = SpeedConfig(speed=2.0)
|
||||
assert adjust_duration(0.0, cfg) == 0.0
|
||||
|
||||
def test_negative_duration_unchanged(self):
|
||||
cfg = SpeedConfig(speed=2.0)
|
||||
assert adjust_duration(-5.0, cfg) == -5.0
|
||||
|
||||
|
||||
# ── build_clip_speed_filter 测试 ─────────────────────────────────────────
|
||||
|
||||
|
||||
class TestBuildClipSpeedFilter:
|
||||
def test_normal_speed(self):
|
||||
vf, af, cfg = build_clip_speed_filter(2.0)
|
||||
assert vf == "setpts=PTS/2.0000"
|
||||
assert "atempo=2.0000" in af
|
||||
assert cfg.speed == 2.0
|
||||
|
||||
def test_clamped_speed(self):
|
||||
vf, af, cfg = build_clip_speed_filter(10.0)
|
||||
assert cfg.speed == MAX_SPEED
|
||||
|
||||
def test_pitch_correct_param(self):
|
||||
vf, af, cfg = build_clip_speed_filter(1.5, pitch_correct=False)
|
||||
assert cfg.pitch_correct is False
|
||||
|
||||
def test_original_speed_empty_filters(self):
|
||||
vf, af, cfg = build_clip_speed_filter(1.0)
|
||||
assert vf == ""
|
||||
assert af == ""
|
||||
|
||||
|
||||
# ── resolve_clip_speed 测试 ──────────────────────────────────────────────
|
||||
|
||||
|
||||
class TestResolveClipSpeed:
|
||||
def test_none_config_uses_global(self):
|
||||
assert resolve_clip_speed(None, 1.5) == 1.5
|
||||
|
||||
def test_no_playback_speed_uses_global(self):
|
||||
assert resolve_clip_speed({}, 1.5) == 1.5
|
||||
|
||||
def test_zero_speed_uses_global(self):
|
||||
assert resolve_clip_speed({"playback_speed": 0}, 1.5) == 1.5
|
||||
|
||||
def test_valid_speed_returns_speed(self):
|
||||
assert resolve_clip_speed({"playback_speed": 2.0}, 1.0) == 2.0
|
||||
|
||||
def test_invalid_type_uses_global(self):
|
||||
assert resolve_clip_speed({"playback_speed": "fast"}, 1.0) == 1.0
|
||||
|
||||
def test_default_global_speed(self):
|
||||
assert resolve_clip_speed({}) == DEFAULT_SPEED
|
||||
@@ -1,689 +0,0 @@
|
||||
"""TTS 相关领域模块单测 — text_splitter + tts_config + tts_job."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import pytest
|
||||
|
||||
from packages.application.tts_job.text_splitter import split_text
|
||||
from packages.domain.tts_config import TtsConfig
|
||||
from packages.domain.tts_job import TTSJob, TTSJobStatus, TERMINAL_STATUSES
|
||||
|
||||
# ═══════════════════════════════════════════════════════════════════════════════
|
||||
# text_splitter 文本分段
|
||||
# ═══════════════════════════════════════════════════════════════════════════════
|
||||
|
||||
|
||||
class TestSplitTextBasic:
|
||||
"""基础分段功能."""
|
||||
|
||||
def test_empty_text_returns_empty_list(self):
|
||||
assert split_text("", max_chars=500) == []
|
||||
|
||||
def test_whitespace_only_returns_empty_list(self):
|
||||
assert split_text(" \n\n ", max_chars=500) == []
|
||||
|
||||
def test_short_text_returns_single_segment(self):
|
||||
text = "你好世界"
|
||||
result = split_text(text, max_chars=500)
|
||||
assert len(result) == 1
|
||||
assert result[0] == text
|
||||
|
||||
def test_text_equal_to_max_chars_single_segment(self):
|
||||
text = "a" * 500
|
||||
result = split_text(text, max_chars=500)
|
||||
assert len(result) == 1
|
||||
assert len(result[0]) == 500
|
||||
|
||||
def test_none_max_chars_uses_default(self):
|
||||
"""默认 max_chars=500."""
|
||||
text = "你好"
|
||||
result = split_text(text) # 使用默认值
|
||||
assert len(result) == 1
|
||||
|
||||
|
||||
class TestSplitTextSentenceBoundary:
|
||||
"""按句子边界分段."""
|
||||
|
||||
def test_splits_at_period(self):
|
||||
text = "第一句。第二句。第三句。"
|
||||
result = split_text(text, max_chars=10)
|
||||
# 每句都比较短,会在句子边界处合并
|
||||
assert len(result) >= 2
|
||||
assert "".join(result) == text.strip()
|
||||
|
||||
def test_splits_at_question_mark(self):
|
||||
text = "你是谁?我是AI。你好吗?很好。"
|
||||
result = split_text(text, max_chars=15)
|
||||
assert len(result) >= 2
|
||||
assert "".join(result) == text.strip()
|
||||
|
||||
def test_splits_at_exclamation_mark(self):
|
||||
text = "太棒了!真厉害!好厉害!"
|
||||
result = split_text(text, max_chars=10)
|
||||
assert len(result) >= 2
|
||||
assert "".join(result) == text.strip()
|
||||
|
||||
def test_splits_at_newline(self):
|
||||
text = "第一段\n第二段\n第三段"
|
||||
result = split_text(text, max_chars=10)
|
||||
assert len(result) >= 2
|
||||
|
||||
def test_splits_at_semicolon(self):
|
||||
text = "第一部分;第二部分;第三部分。"
|
||||
result = split_text(text, max_chars=15)
|
||||
assert len(result) >= 1
|
||||
assert "".join(result) == text.strip()
|
||||
|
||||
|
||||
class TestSplitTextLongSentence:
|
||||
"""长句子(超过 max_chars)强制切段."""
|
||||
|
||||
def test_very_long_sentence_hard_cut(self):
|
||||
"""单个超长句子会被强制切段."""
|
||||
text = "我" * 600 # 没有标点
|
||||
result = split_text(text, max_chars=500)
|
||||
assert len(result) >= 2
|
||||
total = sum(len(seg) for seg in result)
|
||||
assert total == len(text)
|
||||
|
||||
def test_each_segment_leq_max_chars(self):
|
||||
"""每个分段都不超过 max_chars."""
|
||||
text = "测试句子。" * 100
|
||||
result = split_text(text, max_chars=100)
|
||||
for seg in result:
|
||||
assert len(seg) <= 100
|
||||
|
||||
def test_no_empty_segments(self):
|
||||
"""不产生空分段."""
|
||||
text = "测试。" * 50
|
||||
result = split_text(text, max_chars=50)
|
||||
for seg in result:
|
||||
assert len(seg) > 0
|
||||
|
||||
|
||||
class TestSplitTextMergeShortSegments:
|
||||
"""合并过短的分段."""
|
||||
|
||||
def test_short_segments_get_merged(self):
|
||||
"""< 50 字符的段会被合并(如果不超限)."""
|
||||
# 多个短句子,应该会被合并
|
||||
text = "你好。我是。他是。她是。它是。"
|
||||
result = split_text(text, max_chars=50)
|
||||
# 每段6字符左右,应该被合并成一段
|
||||
assert len(result) < 5
|
||||
|
||||
def test_last_short_segment_merged_to_previous(self):
|
||||
"""最后一段如果很短,会合并到前一段."""
|
||||
text = "a" * 48 + "。" + "b" * 48 + "。" + "cc"
|
||||
result = split_text(text, max_chars=100)
|
||||
# 最后的 "cc" 很短,应该被合并
|
||||
assert result[-1] != "cc"
|
||||
|
||||
|
||||
class TestSplitTextEdgeCases:
|
||||
"""边界情况."""
|
||||
|
||||
def test_single_character(self):
|
||||
result = split_text("一", max_chars=500)
|
||||
assert result == ["一"]
|
||||
|
||||
def test_only_punctuation(self):
|
||||
text = "。。。"
|
||||
result = split_text(text, max_chars=500)
|
||||
assert len(result) == 1
|
||||
|
||||
def test_mixed_chinese_english(self):
|
||||
text = "Hello世界。Hello世界。" * 20
|
||||
result = split_text(text, max_chars=50)
|
||||
assert len(result) >= 2
|
||||
assert "".join(result) == text.strip()
|
||||
|
||||
def test_max_chars_one(self):
|
||||
"""极端情况:max_chars=1."""
|
||||
text = "abc"
|
||||
result = split_text(text, max_chars=1)
|
||||
assert len(result) == 3
|
||||
assert result == ["a", "b", "c"]
|
||||
|
||||
|
||||
# ═══════════════════════════════════════════════════════════════════════════════
|
||||
# TtsConfig 配置解析 + 边界钳制
|
||||
# ═══════════════════════════════════════════════════════════════════════════════
|
||||
|
||||
|
||||
class TestTtsConfigDefaults:
|
||||
"""默认值."""
|
||||
|
||||
def test_default_config_disabled(self):
|
||||
config = TtsConfig()
|
||||
assert config.enabled is False
|
||||
assert config.voice_id == ""
|
||||
assert config.speed == 1.0
|
||||
assert config.pitch == 0.0
|
||||
assert config.volume == 0.8
|
||||
assert config.text == ""
|
||||
assert config.align_mode == "full"
|
||||
assert config.overlap_mode == "replace"
|
||||
|
||||
def test_parse_none_returns_default(self):
|
||||
config = TtsConfig.parse(None)
|
||||
assert config.enabled is False
|
||||
|
||||
def test_parse_empty_dict_returns_default(self):
|
||||
config = TtsConfig.parse({})
|
||||
assert config.enabled is False
|
||||
|
||||
def test_parse_non_dict_returns_default(self):
|
||||
config = TtsConfig.parse("not a dict")
|
||||
assert config.enabled is False
|
||||
|
||||
|
||||
class TestTtsConfigParseEnabled:
|
||||
"""enabled 字段解析."""
|
||||
|
||||
def test_parse_enabled_true(self):
|
||||
config = TtsConfig.parse({"enabled": True})
|
||||
assert config.enabled is True
|
||||
|
||||
def test_parse_enabled_false(self):
|
||||
config = TtsConfig.parse({"enabled": False})
|
||||
assert config.enabled is False
|
||||
|
||||
def test_parse_enabled_invalid_type(self):
|
||||
"""enabled 不是 bool 时回退到 False."""
|
||||
config = TtsConfig.parse({"enabled": "true"})
|
||||
assert config.enabled is False
|
||||
|
||||
def test_disabled_ignores_other_fields(self):
|
||||
"""enabled=False 时其他字段都用默认值."""
|
||||
config = TtsConfig.parse(
|
||||
{
|
||||
"enabled": False,
|
||||
"voice_id": "test",
|
||||
"speed": 2.0,
|
||||
}
|
||||
)
|
||||
assert config.enabled is False
|
||||
assert config.voice_id == ""
|
||||
assert config.speed == 1.0
|
||||
|
||||
|
||||
class TestTtsConfigParseFields:
|
||||
"""各字段解析."""
|
||||
|
||||
def test_parse_voice_id(self):
|
||||
config = TtsConfig.parse({"enabled": True, "voice_id": "voice_001"})
|
||||
assert config.voice_id == "voice_001"
|
||||
|
||||
def test_parse_voice_id_invalid_type(self):
|
||||
config = TtsConfig.parse({"enabled": True, "voice_id": 123})
|
||||
assert config.voice_id == ""
|
||||
|
||||
def test_parse_speed(self):
|
||||
config = TtsConfig.parse({"enabled": True, "speed": 1.5})
|
||||
assert config.speed == 1.5
|
||||
|
||||
def test_parse_speed_int(self):
|
||||
config = TtsConfig.parse({"enabled": True, "speed": 2})
|
||||
assert config.speed == 2.0
|
||||
|
||||
def test_parse_speed_invalid_type(self):
|
||||
config = TtsConfig.parse({"enabled": True, "speed": "fast"})
|
||||
assert config.speed == 1.0
|
||||
|
||||
def test_parse_pitch(self):
|
||||
config = TtsConfig.parse({"enabled": True, "pitch": 5})
|
||||
assert config.pitch == 5.0
|
||||
|
||||
def test_parse_pitch_invalid_type(self):
|
||||
config = TtsConfig.parse({"enabled": True, "pitch": "high"})
|
||||
assert config.pitch == 0.0
|
||||
|
||||
def test_parse_volume(self):
|
||||
config = TtsConfig.parse({"enabled": True, "volume": 0.5})
|
||||
assert config.volume == 0.5
|
||||
|
||||
def test_parse_volume_invalid_type(self):
|
||||
config = TtsConfig.parse({"enabled": True, "volume": "loud"})
|
||||
assert config.volume == 0.8
|
||||
|
||||
def test_parse_text(self):
|
||||
config = TtsConfig.parse({"enabled": True, "text": "你好世界"})
|
||||
assert config.text == "你好世界"
|
||||
|
||||
def test_parse_text_invalid_type(self):
|
||||
config = TtsConfig.parse({"enabled": True, "text": 12345})
|
||||
assert config.text == ""
|
||||
|
||||
def test_parse_align_mode_valid(self):
|
||||
config = TtsConfig.parse({"enabled": True, "align_mode": "subtitle"})
|
||||
assert config.align_mode == "subtitle"
|
||||
|
||||
def test_parse_align_mode_invalid(self):
|
||||
config = TtsConfig.parse({"enabled": True, "align_mode": "invalid"})
|
||||
assert config.align_mode == "full"
|
||||
|
||||
def test_parse_overlap_mode_mix(self):
|
||||
config = TtsConfig.parse({"enabled": True, "overlap_mode": "mix"})
|
||||
assert config.overlap_mode == "mix"
|
||||
|
||||
def test_parse_overlap_mode_invalid(self):
|
||||
config = TtsConfig.parse({"enabled": True, "overlap_mode": "invalid"})
|
||||
assert config.overlap_mode == "replace"
|
||||
|
||||
|
||||
class TestTtsConfigClamp:
|
||||
"""边界钳制."""
|
||||
|
||||
def test_speed_below_minimum_clamped(self):
|
||||
config = TtsConfig.parse({"enabled": True, "speed": 0.1})
|
||||
assert config.speed == 0.5
|
||||
|
||||
def test_speed_above_maximum_clamped(self):
|
||||
config = TtsConfig.parse({"enabled": True, "speed": 3.0})
|
||||
assert config.speed == 2.0
|
||||
|
||||
def test_speed_at_minimum_ok(self):
|
||||
config = TtsConfig.parse({"enabled": True, "speed": 0.5})
|
||||
assert config.speed == 0.5
|
||||
|
||||
def test_speed_at_maximum_ok(self):
|
||||
config = TtsConfig.parse({"enabled": True, "speed": 2.0})
|
||||
assert config.speed == 2.0
|
||||
|
||||
def test_pitch_below_minimum_clamped(self):
|
||||
config = TtsConfig.parse({"enabled": True, "pitch": -20})
|
||||
assert config.pitch == -12
|
||||
|
||||
def test_pitch_above_maximum_clamped(self):
|
||||
config = TtsConfig.parse({"enabled": True, "pitch": 20})
|
||||
assert config.pitch == 12
|
||||
|
||||
def test_pitch_at_minimum_ok(self):
|
||||
config = TtsConfig.parse({"enabled": True, "pitch": -12})
|
||||
assert config.pitch == -12
|
||||
|
||||
def test_pitch_at_maximum_ok(self):
|
||||
config = TtsConfig.parse({"enabled": True, "pitch": 12})
|
||||
assert config.pitch == 12
|
||||
|
||||
def test_volume_below_minimum_clamped(self):
|
||||
config = TtsConfig.parse({"enabled": True, "volume": -0.5})
|
||||
assert config.volume == 0.0
|
||||
|
||||
def test_volume_above_maximum_clamped(self):
|
||||
config = TtsConfig.parse({"enabled": True, "volume": 2.0})
|
||||
assert config.volume == 1.0
|
||||
|
||||
def test_volume_at_minimum_ok(self):
|
||||
config = TtsConfig.parse({"enabled": True, "volume": 0.0})
|
||||
assert config.volume == 0.0
|
||||
|
||||
def test_volume_at_maximum_ok(self):
|
||||
config = TtsConfig.parse({"enabled": True, "volume": 1.0})
|
||||
assert config.volume == 1.0
|
||||
|
||||
|
||||
# ═══════════════════════════════════════════════════════════════════════════════
|
||||
# TTSJob 领域模型 — 状态机
|
||||
# ═══════════════════════════════════════════════════════════════════════════════
|
||||
|
||||
|
||||
class TestTTSJobCreate:
|
||||
"""创建任务."""
|
||||
|
||||
def test_create_basic(self):
|
||||
job = TTSJob.create(user_id="user1", input_text="你好")
|
||||
assert job.id is not None
|
||||
assert len(job.id) > 0
|
||||
assert job.user_id == "user1"
|
||||
assert job.input_text == "你好"
|
||||
assert job.status == TTSJobStatus.PENDING
|
||||
assert job.retry_count == 0
|
||||
assert job.max_retries == 3
|
||||
assert job.started_at is None
|
||||
assert job.completed_at is None
|
||||
|
||||
def test_create_with_voice_id(self):
|
||||
job = TTSJob.create(user_id="user1", input_text="你好", voice_id="voice_001")
|
||||
assert job.voice_id == "voice_001"
|
||||
|
||||
def test_create_with_project_id(self):
|
||||
job = TTSJob.create(user_id="user1", input_text="你好", project_id="proj_001")
|
||||
assert job.project_id == "proj_001"
|
||||
|
||||
def test_create_with_voice_clone_profile_id(self):
|
||||
job = TTSJob.create(
|
||||
user_id="user1",
|
||||
input_text="你好",
|
||||
voice_clone_profile_id="clone_001",
|
||||
)
|
||||
assert job.voice_clone_profile_id == "clone_001"
|
||||
|
||||
def test_create_with_custom_max_retries(self):
|
||||
job = TTSJob.create(user_id="user1", input_text="你好", max_retries=5)
|
||||
assert job.max_retries == 5
|
||||
|
||||
def test_create_with_format(self):
|
||||
job = TTSJob.create(user_id="user1", input_text="你好", format="wav")
|
||||
assert job.format == "wav"
|
||||
|
||||
def test_create_with_sample_rate(self):
|
||||
job = TTSJob.create(user_id="user1", input_text="你好", sample_rate=44100)
|
||||
assert job.sample_rate == 44100
|
||||
|
||||
|
||||
class TestTTSJobStatusProperties:
|
||||
"""状态查询属性."""
|
||||
|
||||
def test_pending_not_terminal(self):
|
||||
job = TTSJob.create(user_id="u1", input_text="hi")
|
||||
assert job.is_terminal is False
|
||||
|
||||
def test_processing_not_terminal(self):
|
||||
job = TTSJob.create(user_id="u1", input_text="hi")
|
||||
job.mark_processing()
|
||||
assert job.is_terminal is False
|
||||
|
||||
def test_completed_is_terminal(self):
|
||||
job = TTSJob.create(user_id="u1", input_text="hi")
|
||||
job.mark_processing()
|
||||
job.mark_completed(output_audio_url="url")
|
||||
assert job.is_terminal is True
|
||||
assert job.is_completed is True
|
||||
|
||||
def test_failed_is_terminal(self):
|
||||
job = TTSJob.create(user_id="u1", input_text="hi")
|
||||
job.mark_processing()
|
||||
job.mark_failed("error")
|
||||
assert job.is_terminal is True
|
||||
assert job.status == TTSJobStatus.FAILED
|
||||
|
||||
def test_cancelled_is_terminal(self):
|
||||
job = TTSJob.create(user_id="u1", input_text="hi")
|
||||
job.mark_cancelled()
|
||||
assert job.is_terminal is True
|
||||
|
||||
def test_terminal_statuses_contains_all_three(self):
|
||||
assert TTSJobStatus.COMPLETED in TERMINAL_STATUSES
|
||||
assert TTSJobStatus.FAILED in TERMINAL_STATUSES
|
||||
assert TTSJobStatus.CANCELLED in TERMINAL_STATUSES
|
||||
|
||||
|
||||
class TestTTSJobTransitions:
|
||||
"""状态转换."""
|
||||
|
||||
def test_pending_to_processing(self):
|
||||
job = TTSJob.create(user_id="u1", input_text="hi")
|
||||
job.mark_processing()
|
||||
assert job.status == TTSJobStatus.PROCESSING
|
||||
assert job.started_at is not None
|
||||
assert job.error_message == ""
|
||||
|
||||
def test_pending_to_failed(self):
|
||||
job = TTSJob.create(user_id="u1", input_text="hi")
|
||||
job.mark_failed("网络错误")
|
||||
assert job.status == TTSJobStatus.FAILED
|
||||
assert job.error_message == "网络错误"
|
||||
|
||||
def test_pending_to_cancelled(self):
|
||||
job = TTSJob.create(user_id="u1", input_text="hi")
|
||||
job.mark_cancelled()
|
||||
assert job.status == TTSJobStatus.CANCELLED
|
||||
|
||||
def test_processing_to_completed(self):
|
||||
job = TTSJob.create(user_id="u1", input_text="hi")
|
||||
job.mark_processing()
|
||||
job.mark_completed(output_audio_url="https://example.com/audio.mp3")
|
||||
assert job.status == TTSJobStatus.COMPLETED
|
||||
assert job.output_audio_url == "https://example.com/audio.mp3"
|
||||
assert job.completed_at is not None
|
||||
|
||||
def test_processing_to_failed(self):
|
||||
job = TTSJob.create(user_id="u1", input_text="hi")
|
||||
job.mark_processing()
|
||||
job.mark_failed("API超时")
|
||||
assert job.status == TTSJobStatus.FAILED
|
||||
assert job.error_message == "API超时"
|
||||
|
||||
def test_processing_to_cancelled(self):
|
||||
job = TTSJob.create(user_id="u1", input_text="hi")
|
||||
job.mark_processing()
|
||||
job.mark_cancelled()
|
||||
assert job.status == TTSJobStatus.CANCELLED
|
||||
|
||||
def test_failed_to_pending_retry(self):
|
||||
job = TTSJob.create(user_id="u1", input_text="hi")
|
||||
job.mark_failed("error")
|
||||
job.prepare_retry()
|
||||
assert job.status == TTSJobStatus.PENDING
|
||||
assert job.retry_count == 1
|
||||
assert job.error_message == ""
|
||||
assert job.started_at is None
|
||||
|
||||
def test_completed_cannot_transition_back(self):
|
||||
job = TTSJob.create(user_id="u1", input_text="hi")
|
||||
job.mark_processing()
|
||||
job.mark_completed(output_audio_url="url")
|
||||
with pytest.raises(ValueError, match="非法状态转换"):
|
||||
job.mark_failed("test")
|
||||
|
||||
def test_cancelled_cannot_retry(self):
|
||||
job = TTSJob.create(user_id="u1", input_text="hi")
|
||||
job.mark_cancelled()
|
||||
with pytest.raises(ValueError, match="不可重试"):
|
||||
job.prepare_retry()
|
||||
|
||||
def test_mark_failed_sets_error_message(self):
|
||||
job = TTSJob.create(user_id="u1", input_text="hi")
|
||||
job.mark_failed("服务端错误")
|
||||
assert job.error_message == "服务端错误"
|
||||
|
||||
def test_failed_status_after_mark_failed(self):
|
||||
job = TTSJob.create(user_id="u1", input_text="hi")
|
||||
job.mark_processing()
|
||||
job.mark_failed("API超时")
|
||||
assert job.status == TTSJobStatus.FAILED
|
||||
assert job.error_message == "API超时"
|
||||
|
||||
|
||||
class TestTTSJobRetry:
|
||||
"""重试逻辑."""
|
||||
|
||||
def test_retry_increments_retry_count(self):
|
||||
job = TTSJob.create(user_id="u1", input_text="hi")
|
||||
job.mark_failed("error1")
|
||||
job.prepare_retry()
|
||||
assert job.retry_count == 1
|
||||
|
||||
job.mark_processing()
|
||||
job.mark_failed("error2")
|
||||
job.prepare_retry()
|
||||
assert job.retry_count == 2
|
||||
|
||||
def test_retry_clears_error_message(self):
|
||||
job = TTSJob.create(user_id="u1", input_text="hi")
|
||||
job.mark_failed("error")
|
||||
job.prepare_retry()
|
||||
assert job.error_message == ""
|
||||
|
||||
def test_retry_resets_timestamps(self):
|
||||
job = TTSJob.create(user_id="u1", input_text="hi")
|
||||
job.mark_processing()
|
||||
job.mark_failed("error")
|
||||
job.prepare_retry()
|
||||
assert job.started_at is None
|
||||
|
||||
def test_can_retry_while_below_max_retries(self):
|
||||
job = TTSJob.create(user_id="u1", input_text="hi", max_retries=2)
|
||||
# 第一次失败+重试
|
||||
job.mark_failed("e1")
|
||||
job.prepare_retry()
|
||||
assert job.retry_count == 1
|
||||
# 第二次失败+重试
|
||||
job.mark_processing()
|
||||
job.mark_failed("e2")
|
||||
job.prepare_retry()
|
||||
assert job.retry_count == 2
|
||||
|
||||
def test_cannot_retry_when_exceeded_max_retries(self):
|
||||
job = TTSJob.create(user_id="u1", input_text="hi", max_retries=1)
|
||||
job.mark_failed("e1")
|
||||
job.prepare_retry()
|
||||
assert job.retry_count == 1
|
||||
# 再次失败就不能重试了(已经用完1次重试)
|
||||
job.mark_processing()
|
||||
job.mark_failed("e2")
|
||||
with pytest.raises(ValueError, match="不可重试"):
|
||||
job.prepare_retry()
|
||||
|
||||
def test_is_retryable_true_when_failed_and_under_limit(self):
|
||||
job = TTSJob.create(user_id="u1", input_text="hi", max_retries=3)
|
||||
job.mark_failed("error")
|
||||
assert job.is_retryable is True
|
||||
|
||||
def test_is_retryable_false_when_not_failed(self):
|
||||
job = TTSJob.create(user_id="u1", input_text="hi")
|
||||
assert job.is_retryable is False
|
||||
|
||||
def test_prepare_retry_fails_when_not_failed(self):
|
||||
job = TTSJob.create(user_id="u1", input_text="hi")
|
||||
with pytest.raises(ValueError, match="不可重试"):
|
||||
job.prepare_retry()
|
||||
|
||||
|
||||
class TestTTSJobCompleted:
|
||||
"""完成时的字段."""
|
||||
|
||||
def test_mark_completed_sets_output_url(self):
|
||||
job = TTSJob.create(user_id="u1", input_text="hi")
|
||||
job.mark_processing()
|
||||
job.mark_completed(
|
||||
output_audio_url="https://cdn.example.com/audio.mp3",
|
||||
output_audio_key="audio/xxx.mp3",
|
||||
duration=10.5,
|
||||
file_size=102400,
|
||||
)
|
||||
assert job.output_audio_url == "https://cdn.example.com/audio.mp3"
|
||||
assert job.output_audio_key == "audio/xxx.mp3"
|
||||
assert job.duration == 10.5
|
||||
assert job.file_size == 102400
|
||||
assert job.completed_at is not None
|
||||
|
||||
def test_mark_completed_default_values(self):
|
||||
job = TTSJob.create(user_id="u1", input_text="hi")
|
||||
job.mark_processing()
|
||||
job.mark_completed(output_audio_url="url")
|
||||
assert job.duration == 0.0
|
||||
assert job.file_size == 0
|
||||
|
||||
|
||||
class TestTTSJobMetadata:
|
||||
"""元数据."""
|
||||
|
||||
def test_create_with_metadata(self):
|
||||
meta = {"source": "api", "priority": "high"}
|
||||
job = TTSJob.create(user_id="u1", input_text="hi", metadata=meta)
|
||||
assert job.metadata["source"] == "api"
|
||||
assert job.metadata["priority"] == "high"
|
||||
|
||||
def test_default_metadata_empty_dict(self):
|
||||
job = TTSJob.create(user_id="u1", input_text="hi")
|
||||
assert job.metadata == {}
|
||||
|
||||
|
||||
class TestTTSJobCreateValidation:
|
||||
"""创建时的参数校验."""
|
||||
|
||||
def test_empty_user_id_raises(self):
|
||||
with pytest.raises(ValueError, match="user_id"):
|
||||
TTSJob.create(user_id="", input_text="hi")
|
||||
|
||||
def test_whitespace_user_id_raises(self):
|
||||
with pytest.raises(ValueError, match="user_id"):
|
||||
TTSJob.create(user_id=" ", input_text="hi")
|
||||
|
||||
def test_empty_input_text_raises(self):
|
||||
with pytest.raises(ValueError, match="input_text"):
|
||||
TTSJob.create(user_id="u1", input_text="")
|
||||
|
||||
def test_input_text_too_long_raises(self):
|
||||
long_text = "a" * 10001
|
||||
with pytest.raises(ValueError, match="10000"):
|
||||
TTSJob.create(user_id="u1", input_text=long_text)
|
||||
|
||||
def test_invalid_format_raises(self):
|
||||
with pytest.raises(ValueError, match="不支持的输出格式"):
|
||||
TTSJob.create(user_id="u1", input_text="hi", format="flac")
|
||||
|
||||
def test_valid_format_wav(self):
|
||||
job = TTSJob.create(user_id="u1", input_text="hi", format="wav")
|
||||
assert job.format == "wav"
|
||||
|
||||
def test_valid_format_pcm(self):
|
||||
job = TTSJob.create(user_id="u1", input_text="hi", format="pcm")
|
||||
assert job.format == "pcm"
|
||||
|
||||
def test_input_text_stripped(self):
|
||||
job = TTSJob.create(user_id="u1", input_text=" 你好 ")
|
||||
assert job.input_text == "你好"
|
||||
|
||||
|
||||
class TestTTSJobToDict:
|
||||
"""序列化."""
|
||||
|
||||
def test_to_dict_contains_key_fields(self):
|
||||
job = TTSJob.create(user_id="u1", input_text="hi", voice_id="v1")
|
||||
d = job.to_dict()
|
||||
assert d["id"] == job.id
|
||||
assert d["user_id"] == "u1"
|
||||
assert d["input_text"] == "hi"
|
||||
assert d["voice_id"] == "v1"
|
||||
assert d["status"] == "pending"
|
||||
assert d["retry_count"] == 0
|
||||
assert d["is_retryable"] is False
|
||||
|
||||
def test_to_dict_completed_status(self):
|
||||
job = TTSJob.create(user_id="u1", input_text="hi")
|
||||
job.mark_processing()
|
||||
job.mark_completed(output_audio_url="https://example.com/a.mp3", duration=10.5, file_size=1024)
|
||||
d = job.to_dict()
|
||||
assert d["status"] == "completed"
|
||||
assert d["output_audio_url"] == "https://example.com/a.mp3"
|
||||
assert d["duration"] == 10.5
|
||||
assert d["file_size"] == 1024
|
||||
assert d["is_completed"] is True
|
||||
assert d["started_at"] is not None
|
||||
assert d["completed_at"] is not None
|
||||
|
||||
def test_to_dict_failed_status(self):
|
||||
job = TTSJob.create(user_id="u1", input_text="hi")
|
||||
job.mark_failed("error msg")
|
||||
d = job.to_dict()
|
||||
assert d["status"] == "failed"
|
||||
assert d["error_message"] == "error msg"
|
||||
assert d["is_retryable"] is True
|
||||
|
||||
|
||||
class TestTTSJobCompletedValidation:
|
||||
"""完成时的校验."""
|
||||
|
||||
def test_mark_completed_empty_url_raises(self):
|
||||
job = TTSJob.create(user_id="u1", input_text="hi")
|
||||
job.mark_processing()
|
||||
with pytest.raises(ValueError, match="output_audio_url"):
|
||||
job.mark_completed(output_audio_url="")
|
||||
|
||||
def test_is_completed_requires_url(self):
|
||||
"""is_completed 属性需要 output_audio_url."""
|
||||
job = TTSJob.create(user_id="u1", input_text="hi")
|
||||
job.mark_processing()
|
||||
# 直接设置状态为 completed 但不给 URL(模拟异常情况)
|
||||
# 正常流程 mark_completed 会校验 URL,所以这里不会出现
|
||||
# 但确认属性逻辑:没有 URL 时 is_completed 为 False
|
||||
job.output_audio_url = ""
|
||||
# 直接绕过状态机
|
||||
from packages.domain.tts_job import _VALID_TRANSITIONS # noqa
|
||||
|
||||
job.status = TTSJobStatus.COMPLETED
|
||||
assert job.is_completed is False
|
||||
Reference in New Issue
Block a user