merge: 解决第56波与develop的add/add冲突(sticker+subtitle)
AI Code Review / AI Code Review (pull_request) Failing after 1s
CI/CD Pipeline / Check if frontend-only change (pull_request) Successful in 36s
CI/CD Pipeline / Frontend Lint (pull_request) Successful in 38s
CI/CD Pipeline / Validate - Type Check (mypy) (pull_request) Successful in 1m18s
CI/CD Pipeline / Validate - Migration (alembic) (pull_request) Successful in 1m49s
CI/CD Pipeline / PR Build Worker Image (pull_request) Successful in 41s
Preview Deploy / Deploy Preview Environment (pull_request) Failing after 50s
CI/CD Pipeline / PR Build Web Image (pull_request) Successful in 1m57s
CI/CD Pipeline / Validate - Code Quality (pull_request) Successful in 4m54s
CI/CD Pipeline / PR Build API Image (pull_request) Successful in 4m15s
CI/CD Pipeline / Unit Tests (pull_request) Successful in 2m33s
PR Automation / Auto Approve on CI Green (pull_request) Successful in 3m54s
CI/CD Pipeline / Integration Tests (pull_request) Successful in 1m55s
Preview Cleanup / Cleanup Preview Environment (pull_request) Successful in 23s
ACR Cleanup / ACR Image Cleanup (pull_request_target) Failing after 22s
PR Automation / Auto Merge on CI Green + Approved (pull_request) Successful in 51m44s
CI/CD Pipeline / Production Browser E2E (pull_request) Failing after 1274h52m10s
CI/CD Pipeline / Staging E2E Tests (pull_request) Failing after 1274h52m16s
CI/CD Pipeline / ACR Image Cleanup (pull_request) Failing after 1274h52m12s
CI/CD Pipeline / Deploy Production (pull_request) Failing after 1274h52m25s
CI/CD Pipeline / Build Production Worker Image (pull_request) Failing after 1274h53m56s
CI/CD Pipeline / Build Production Web Image (pull_request) Failing after 1274h53m58s
CI/CD Pipeline / Build Production API Image (pull_request) Failing after 1274h54m0s
CI/CD Pipeline / Frontend Unit Tests (pull_request) Failing after 1274h54m2s
CI/CD Pipeline / Build Staging Worker Image (pull_request) Failing after 1274h55m26s
CI/CD Pipeline / Build Staging API Image (pull_request) Failing after 1274h55m30s
CI/CD Pipeline / Staging API Integration Tests (pull_request) Failing after 1275h24m48s
CI/CD Pipeline / Deploy Staging (Watchtower auto-deploy) (pull_request) Failing after 1275h25m1s
CI/CD Pipeline / Build Staging Web Image (pull_request) Failing after 1275h28m2s
AI Code Review / AI Code Review (pull_request) Failing after 1s
CI/CD Pipeline / Check if frontend-only change (pull_request) Successful in 36s
CI/CD Pipeline / Frontend Lint (pull_request) Successful in 38s
CI/CD Pipeline / Validate - Type Check (mypy) (pull_request) Successful in 1m18s
CI/CD Pipeline / Validate - Migration (alembic) (pull_request) Successful in 1m49s
CI/CD Pipeline / PR Build Worker Image (pull_request) Successful in 41s
Preview Deploy / Deploy Preview Environment (pull_request) Failing after 50s
CI/CD Pipeline / PR Build Web Image (pull_request) Successful in 1m57s
CI/CD Pipeline / Validate - Code Quality (pull_request) Successful in 4m54s
CI/CD Pipeline / PR Build API Image (pull_request) Successful in 4m15s
CI/CD Pipeline / Unit Tests (pull_request) Successful in 2m33s
PR Automation / Auto Approve on CI Green (pull_request) Successful in 3m54s
CI/CD Pipeline / Integration Tests (pull_request) Successful in 1m55s
Preview Cleanup / Cleanup Preview Environment (pull_request) Successful in 23s
ACR Cleanup / ACR Image Cleanup (pull_request_target) Failing after 22s
PR Automation / Auto Merge on CI Green + Approved (pull_request) Successful in 51m44s
CI/CD Pipeline / Production Browser E2E (pull_request) Failing after 1274h52m10s
CI/CD Pipeline / Staging E2E Tests (pull_request) Failing after 1274h52m16s
CI/CD Pipeline / ACR Image Cleanup (pull_request) Failing after 1274h52m12s
CI/CD Pipeline / Deploy Production (pull_request) Failing after 1274h52m25s
CI/CD Pipeline / Build Production Worker Image (pull_request) Failing after 1274h53m56s
CI/CD Pipeline / Build Production Web Image (pull_request) Failing after 1274h53m58s
CI/CD Pipeline / Build Production API Image (pull_request) Failing after 1274h54m0s
CI/CD Pipeline / Frontend Unit Tests (pull_request) Failing after 1274h54m2s
CI/CD Pipeline / Build Staging Worker Image (pull_request) Failing after 1274h55m26s
CI/CD Pipeline / Build Staging API Image (pull_request) Failing after 1274h55m30s
CI/CD Pipeline / Staging API Integration Tests (pull_request) Failing after 1275h24m48s
CI/CD Pipeline / Deploy Staging (Watchtower auto-deploy) (pull_request) Failing after 1275h25m1s
CI/CD Pipeline / Build Staging Web Image (pull_request) Failing after 1275h28m2s
This commit is contained in:
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,186 @@
|
||||
import React, { useState, useRef } from "react"
|
||||
import { UploadOutlined, SoundOutlined, CloseOutlined } from "@ant-design/icons"
|
||||
import { Button, Input } from "@/components/ui"
|
||||
import { type TagItem } from "@/api/tags"
|
||||
import { type VoiceGender, type VoiceMaterial } from "../types"
|
||||
import { GENDER_OPTIONS } from "../constants"
|
||||
import { genderClass, formatFileSize } from "../utils/format"
|
||||
import TagSelector from "./TagSelector"
|
||||
|
||||
export interface MaterialFormProps {
|
||||
initial?: VoiceMaterial
|
||||
onSubmit: (data: Omit<VoiceMaterial, "id" | "createdAt"> & { file?: File }) => void
|
||||
onCancel: () => void
|
||||
loading?: boolean
|
||||
uploadProgress?: number | null
|
||||
tags?: TagItem[]
|
||||
tagMap?: Map<string, TagItem>
|
||||
onCreateTag?: (name: string) => Promise<TagItem>
|
||||
}
|
||||
|
||||
const MaterialForm: React.FC<MaterialFormProps> = ({
|
||||
initial,
|
||||
onSubmit,
|
||||
onCancel,
|
||||
loading,
|
||||
uploadProgress,
|
||||
tags = [],
|
||||
tagMap = new Map(),
|
||||
onCreateTag,
|
||||
}) => {
|
||||
const [name, setName] = useState(initial?.name ?? "")
|
||||
const [description, setDescription] = useState(initial?.description ?? "")
|
||||
const [gender, setGender] = useState<VoiceGender>(initial?.gender ?? "female")
|
||||
const [selectedTagIds, setSelectedTagIds] = useState<string[]>(initial?.tagIds ?? [])
|
||||
const [file, setFile] = useState<File | undefined>(undefined)
|
||||
const fileInputRef = useRef<HTMLInputElement>(null)
|
||||
|
||||
const handleSubmit = () => {
|
||||
if (!name.trim()) return
|
||||
if (!initial && !file) return
|
||||
onSubmit({
|
||||
name: name.trim(),
|
||||
description: description.trim(),
|
||||
gender,
|
||||
tagIds: selectedTagIds,
|
||||
fileName: file?.name ?? initial?.fileName ?? "",
|
||||
fileSize: file?.size ?? initial?.fileSize ?? 0,
|
||||
duration: initial?.duration ?? 0,
|
||||
mimeType: file?.type ?? initial?.mimeType ?? "audio/mpeg",
|
||||
file,
|
||||
})
|
||||
}
|
||||
|
||||
return (
|
||||
<div className="vmat-form">
|
||||
{/* 音频文件上传(编辑模式不显示) */}
|
||||
{!initial && (
|
||||
<div className="vmat-form-field">
|
||||
<label className="vmat-form-label">音频文件 *</label>
|
||||
<div
|
||||
className="vmat-upload-zone"
|
||||
onClick={() => fileInputRef.current?.click()}
|
||||
onDragOver={(e) => e.preventDefault()}
|
||||
onDrop={(e) => {
|
||||
e.preventDefault()
|
||||
const f = e.dataTransfer.files[0]
|
||||
if (f?.type.startsWith("audio/")) setFile(f)
|
||||
}}
|
||||
>
|
||||
<input
|
||||
ref={fileInputRef}
|
||||
type="file"
|
||||
accept="audio/*"
|
||||
style={{ display: "none" }}
|
||||
onChange={(e) => {
|
||||
const f = e.target.files?.[0]
|
||||
if (f) setFile(f)
|
||||
}}
|
||||
/>
|
||||
{file ? (
|
||||
<div className="vmat-upload-selected">
|
||||
<SoundOutlined className="vmat-upload-icon" />
|
||||
<span className="vmat-upload-filename">{file.name}</span>
|
||||
<span className="vmat-upload-filesize">{formatFileSize(file.size)}</span>
|
||||
<button
|
||||
type="button"
|
||||
className="vmat-upload-clear"
|
||||
onClick={(e) => {
|
||||
e.stopPropagation()
|
||||
setFile(undefined)
|
||||
}}
|
||||
>
|
||||
<CloseOutlined />
|
||||
</button>
|
||||
</div>
|
||||
) : (
|
||||
<div className="vmat-upload-placeholder">
|
||||
<UploadOutlined className="vmat-upload-icon" />
|
||||
<p>点击或拖拽音频文件到此处</p>
|
||||
<span>支持 MP3、WAV、AAC、FLAC 等格式</span>
|
||||
</div>
|
||||
)}
|
||||
</div>
|
||||
{/* 上传进度条 */}
|
||||
{uploadProgress !== null && uploadProgress !== undefined && (
|
||||
<div className="vmat-upload-progress">
|
||||
<div className="vmat-upload-progress-bar" style={{ width: `${uploadProgress}%` }} />
|
||||
<span className="vmat-upload-progress-text">{uploadProgress}%</span>
|
||||
</div>
|
||||
)}
|
||||
</div>
|
||||
)}
|
||||
|
||||
{/* 名称 */}
|
||||
<div className="vmat-form-field">
|
||||
<label className="vmat-form-label">名称 *</label>
|
||||
<Input
|
||||
placeholder="输入配音素材名称"
|
||||
value={name}
|
||||
onChange={(e) => setName(e.target.value)}
|
||||
maxLength={50}
|
||||
/>
|
||||
</div>
|
||||
|
||||
{/* 音色描述 */}
|
||||
<div className="vmat-form-field">
|
||||
<label className="vmat-form-label">音色描述</label>
|
||||
<Input.TextArea
|
||||
placeholder="描述音色特点,如:适合产品宣传的男声配音..."
|
||||
value={description}
|
||||
onChange={(e) => setDescription(e.target.value)}
|
||||
rows={3}
|
||||
maxLength={200}
|
||||
/>
|
||||
</div>
|
||||
|
||||
{/* 性别 */}
|
||||
<div className="vmat-form-field">
|
||||
<label className="vmat-form-label">性别</label>
|
||||
<div className="vmat-gender-group">
|
||||
{GENDER_OPTIONS.map((opt) => (
|
||||
<button
|
||||
key={opt.value}
|
||||
type="button"
|
||||
className={`vmat-gender-btn${gender === opt.value ? " active" : ""} ${genderClass(opt.value)}`}
|
||||
onClick={() => setGender(opt.value)}
|
||||
>
|
||||
{opt.icon}
|
||||
{opt.label}
|
||||
</button>
|
||||
))}
|
||||
</div>
|
||||
</div>
|
||||
|
||||
{/* 风格标签 */}
|
||||
<div className="vmat-form-field">
|
||||
<label className="vmat-form-label">风格标签</label>
|
||||
<TagSelector
|
||||
value={selectedTagIds}
|
||||
onChange={setSelectedTagIds}
|
||||
tags={tags}
|
||||
tagMap={tagMap}
|
||||
onCreateTag={onCreateTag ?? (async () => ({ id: "", name: "" }))}
|
||||
/>
|
||||
</div>
|
||||
|
||||
{/* 操作按钮 */}
|
||||
<div className="vmat-form-actions">
|
||||
<Button buttonType="ghost" buttonSize="md" onClick={onCancel}>
|
||||
取消
|
||||
</Button>
|
||||
<Button
|
||||
buttonType="primary"
|
||||
buttonSize="md"
|
||||
onClick={handleSubmit}
|
||||
loading={loading}
|
||||
disabled={!name.trim() || (!initial && !file)}
|
||||
>
|
||||
{initial ? "保存修改" : "上传"}
|
||||
</Button>
|
||||
</div>
|
||||
</div>
|
||||
)
|
||||
}
|
||||
|
||||
export default MaterialForm
|
||||
@@ -0,0 +1,163 @@
|
||||
import React, { useState, useRef, useCallback, useMemo } from "react"
|
||||
import { CheckOutlined } from "@ant-design/icons"
|
||||
import { Tag } from "@/components/ui"
|
||||
import { type TagItem } from "@/api/tags"
|
||||
|
||||
export interface TagSelectorProps {
|
||||
/** 已选标签 ID 列表 */
|
||||
value: string[]
|
||||
onChange: (tagIds: string[]) => void
|
||||
/** 所有可用标签(来自 API) */
|
||||
tags: TagItem[]
|
||||
/** 标签 ID → TagItem 映射 */
|
||||
tagMap: Map<string, TagItem>
|
||||
/** 创建新标签,返回带 ID 的 TagItem */
|
||||
onCreateTag: (name: string) => Promise<TagItem>
|
||||
placeholder?: string
|
||||
}
|
||||
|
||||
const TagSelector: React.FC<TagSelectorProps> = ({
|
||||
value,
|
||||
onChange,
|
||||
tags,
|
||||
tagMap,
|
||||
onCreateTag,
|
||||
placeholder = "输入标签后回车添加",
|
||||
}) => {
|
||||
const [inputVal, setInputVal] = useState("")
|
||||
const [showSuggestions, setShowSuggestions] = useState(false)
|
||||
const inputRef = useRef<HTMLInputElement>(null)
|
||||
|
||||
/** 按名称查找已有标签(大小写不敏感) */
|
||||
const findTagByName = useCallback(
|
||||
(name: string) => tags.find((t) => t.name.toLowerCase() === name.toLowerCase()),
|
||||
[tags],
|
||||
)
|
||||
|
||||
/** 去重添加标签(按 ID) */
|
||||
const addTagId = useCallback(
|
||||
(tagId: string) => {
|
||||
if (value.includes(tagId)) return
|
||||
onChange([...value, tagId])
|
||||
setInputVal("")
|
||||
setShowSuggestions(false)
|
||||
},
|
||||
[value, onChange],
|
||||
)
|
||||
|
||||
/** 输入自定义标签名:若已存在则直接选,否则创建新标签 */
|
||||
const addTagByName = useCallback(
|
||||
async (name: string) => {
|
||||
const trimmed = name.trim()
|
||||
if (!trimmed) return
|
||||
const existing = findTagByName(trimmed)
|
||||
if (existing) {
|
||||
addTagId(existing.id)
|
||||
} else {
|
||||
try {
|
||||
const created = await onCreateTag(trimmed)
|
||||
addTagId(created.id)
|
||||
} catch {
|
||||
/* 创建失败静默忽略 */
|
||||
}
|
||||
}
|
||||
},
|
||||
[findTagByName, addTagId, onCreateTag],
|
||||
)
|
||||
|
||||
const removeTagId = useCallback(
|
||||
(tagId: string) => {
|
||||
onChange(value.filter((t) => t !== tagId))
|
||||
},
|
||||
[value, onChange],
|
||||
)
|
||||
|
||||
/** 输入补全建议(排除已选) */
|
||||
const suggestions = useMemo(() => {
|
||||
if (!inputVal.trim()) return []
|
||||
const lower = inputVal.toLowerCase()
|
||||
return tags.filter((t) => t.name.toLowerCase().includes(lower) && !value.includes(t.id))
|
||||
}, [inputVal, tags, value])
|
||||
|
||||
const handleKeyDown = (e: React.KeyboardEvent) => {
|
||||
if (e.key === "Enter") {
|
||||
e.preventDefault()
|
||||
if (suggestions.length > 0) {
|
||||
addTagId(suggestions[0].id)
|
||||
} else {
|
||||
addTagByName(inputVal)
|
||||
}
|
||||
} else if (e.key === "Backspace" && !inputVal && value.length > 0) {
|
||||
removeTagId(value[value.length - 1])
|
||||
}
|
||||
}
|
||||
|
||||
return (
|
||||
<div className="vmat-tag-selector-wrapper">
|
||||
<div className="vmat-tag-selector" onClick={() => inputRef.current?.focus()}>
|
||||
{value.map((tagId) => (
|
||||
<Tag key={tagId} variant="info" closable onClose={() => removeTagId(tagId)}>
|
||||
{tagMap.get(tagId)?.name ?? tagId}
|
||||
</Tag>
|
||||
))}
|
||||
<input
|
||||
ref={inputRef}
|
||||
className="vmat-tag-selector-input"
|
||||
value={inputVal}
|
||||
onChange={(e) => {
|
||||
setInputVal(e.target.value)
|
||||
setShowSuggestions(true)
|
||||
}}
|
||||
onFocus={() => setShowSuggestions(true)}
|
||||
onBlur={() => setTimeout(() => setShowSuggestions(false), 150)}
|
||||
onKeyDown={handleKeyDown}
|
||||
placeholder={value.length === 0 ? placeholder : ""}
|
||||
/>
|
||||
</div>
|
||||
|
||||
{/* 自动补全下拉 */}
|
||||
{showSuggestions && suggestions.length > 0 && (
|
||||
<div className="vmat-tag-suggestions">
|
||||
{suggestions.slice(0, 6).map((tag) => (
|
||||
<button
|
||||
key={tag.id}
|
||||
type="button"
|
||||
className="vmat-tag-suggestion-item"
|
||||
onMouseDown={(e) => {
|
||||
e.preventDefault()
|
||||
addTagId(tag.id)
|
||||
}}
|
||||
>
|
||||
{tag.name}
|
||||
</button>
|
||||
))}
|
||||
</div>
|
||||
)}
|
||||
|
||||
{/* 已有标签快捷选择 */}
|
||||
{tags.length > 0 && (
|
||||
<div className="vmat-tag-selector-presets">
|
||||
{tags.map((tag) => {
|
||||
const isSelected = value.includes(tag.id)
|
||||
return (
|
||||
<button
|
||||
key={tag.id}
|
||||
type="button"
|
||||
className={`vmat-tag-selector-preset${isSelected ? " selected" : ""}`}
|
||||
onClick={() => {
|
||||
if (isSelected) removeTagId(tag.id)
|
||||
else addTagId(tag.id)
|
||||
}}
|
||||
>
|
||||
{isSelected && <CheckOutlined style={{ fontSize: 10, marginRight: 2 }} />}
|
||||
{tag.name}
|
||||
</button>
|
||||
)
|
||||
})}
|
||||
</div>
|
||||
)}
|
||||
</div>
|
||||
)
|
||||
}
|
||||
|
||||
export default TagSelector
|
||||
@@ -0,0 +1,245 @@
|
||||
import React, { useRef } from "react"
|
||||
import {
|
||||
AudioOutlined,
|
||||
PlayCircleOutlined,
|
||||
PauseCircleOutlined,
|
||||
EditOutlined,
|
||||
DeleteOutlined,
|
||||
CheckOutlined,
|
||||
SoundOutlined,
|
||||
MutedOutlined,
|
||||
} from "@ant-design/icons"
|
||||
import { Tooltip } from "antd"
|
||||
import { Tag } from "@/components/ui"
|
||||
import { type TagItem } from "@/api/tags"
|
||||
import { type VoiceMaterial } from "../types"
|
||||
import { MAX_CARD_TAGS, TAG_VARIANTS } from "../constants"
|
||||
import {
|
||||
genderClass,
|
||||
genderIcon,
|
||||
genderLabel,
|
||||
formatDuration,
|
||||
formatFileSize,
|
||||
formatDate,
|
||||
} from "../utils/format"
|
||||
|
||||
export interface VoiceCardProps {
|
||||
material: VoiceMaterial
|
||||
isPlaying: boolean
|
||||
currentTime: number
|
||||
isSelected: boolean
|
||||
batchMode: boolean
|
||||
volume: number
|
||||
tagMap: Map<string, TagItem>
|
||||
onPlay: () => void
|
||||
onPause: () => void
|
||||
onSeek: (time: number) => void
|
||||
onEdit: () => void
|
||||
onDelete: () => void
|
||||
onToggleSelect: (id: string) => void
|
||||
onVolumeChange: (e: React.ChangeEvent<HTMLInputElement>) => void
|
||||
onToggleMute: () => void
|
||||
}
|
||||
|
||||
const VoiceMaterialCard: React.FC<VoiceCardProps> = ({
|
||||
material,
|
||||
isPlaying,
|
||||
currentTime,
|
||||
isSelected,
|
||||
batchMode,
|
||||
volume,
|
||||
tagMap,
|
||||
onPlay,
|
||||
onPause,
|
||||
onSeek,
|
||||
onEdit,
|
||||
onDelete,
|
||||
onToggleSelect,
|
||||
onVolumeChange,
|
||||
onToggleMute,
|
||||
}) => {
|
||||
const progressRef = useRef<HTMLDivElement>(null)
|
||||
|
||||
const handleProgressMouseDown = (e: React.MouseEvent<HTMLDivElement>) => {
|
||||
if (!progressRef.current) return
|
||||
e.preventDefault()
|
||||
const doSeek = (ev: MouseEvent) => {
|
||||
if (!progressRef.current) return
|
||||
const rect = progressRef.current.getBoundingClientRect()
|
||||
const percent = Math.max(0, Math.min(1, (ev.clientX - rect.left) / rect.width))
|
||||
onSeek(percent * material.duration)
|
||||
}
|
||||
doSeek(e.nativeEvent)
|
||||
const handleMove = (ev: MouseEvent) => doSeek(ev)
|
||||
const handleUp = () => {
|
||||
document.removeEventListener("mousemove", handleMove)
|
||||
document.removeEventListener("mouseup", handleUp)
|
||||
}
|
||||
document.addEventListener("mousemove", handleMove)
|
||||
document.addEventListener("mouseup", handleUp)
|
||||
}
|
||||
|
||||
const progress = material.duration > 0 ? (currentTime / material.duration) * 100 : 0
|
||||
|
||||
const handleCardClick = () => {
|
||||
if (batchMode) {
|
||||
onToggleSelect(material.id)
|
||||
}
|
||||
}
|
||||
|
||||
return (
|
||||
<div
|
||||
className={`vmat-card ${genderClass(material.gender)}${isPlaying ? " playing" : ""}${isSelected ? " selected" : ""}${batchMode ? " batch-mode" : ""}`}
|
||||
onClick={handleCardClick}
|
||||
>
|
||||
{/* 批量选择 checkbox */}
|
||||
{(batchMode || isSelected) && (
|
||||
<div
|
||||
className={`vmat-card-checkbox vmat-checkbox${isSelected ? " checked" : ""}`}
|
||||
onClick={(e) => {
|
||||
e.stopPropagation()
|
||||
onToggleSelect(material.id)
|
||||
}}
|
||||
>
|
||||
{isSelected && <CheckOutlined />}
|
||||
</div>
|
||||
)}
|
||||
|
||||
{/* 操作按钮 */}
|
||||
<div className="vmat-card-actions">
|
||||
<button
|
||||
type="button"
|
||||
className="vmat-card-action-btn"
|
||||
onClick={(e) => {
|
||||
e.stopPropagation()
|
||||
onEdit()
|
||||
}}
|
||||
title="编辑"
|
||||
>
|
||||
<EditOutlined />
|
||||
</button>
|
||||
<button
|
||||
type="button"
|
||||
className="vmat-card-action-btn vmat-card-action-btn--danger"
|
||||
onClick={(e) => {
|
||||
e.stopPropagation()
|
||||
onDelete()
|
||||
}}
|
||||
title="删除"
|
||||
>
|
||||
<DeleteOutlined />
|
||||
</button>
|
||||
</div>
|
||||
|
||||
{/* 头部:图标 + 名称 + 性别 */}
|
||||
<div className="vmat-card-header">
|
||||
<div className="vmat-card-avatar">
|
||||
<AudioOutlined />
|
||||
</div>
|
||||
<div className="vmat-card-title-area">
|
||||
<h4 className="vmat-card-name" title={material.name}>
|
||||
{material.name}
|
||||
</h4>
|
||||
<span className={`vmat-card-gender ${genderClass(material.gender)}`}>
|
||||
{genderIcon(material.gender)}
|
||||
{genderLabel(material.gender)}
|
||||
</span>
|
||||
</div>
|
||||
</div>
|
||||
|
||||
{/* 描述 */}
|
||||
{material.description && <p className="vmat-card-desc">{material.description}</p>}
|
||||
|
||||
{/* 标签 */}
|
||||
<div className="vmat-card-tags">
|
||||
{material.tagIds.length === 0 ? (
|
||||
<span
|
||||
className="vmat-tag-empty"
|
||||
onClick={(e) => {
|
||||
e.stopPropagation()
|
||||
onEdit()
|
||||
}}
|
||||
>
|
||||
添加标签
|
||||
</span>
|
||||
) : (
|
||||
<>
|
||||
{material.tagIds.slice(0, MAX_CARD_TAGS).map((tagId, i) => (
|
||||
<Tag key={tagId} variant={TAG_VARIANTS[i % TAG_VARIANTS.length]}>
|
||||
{tagMap.get(tagId)?.name ?? tagId}
|
||||
</Tag>
|
||||
))}
|
||||
{material.tagIds.length > MAX_CARD_TAGS && (
|
||||
<Tooltip
|
||||
title={material.tagIds
|
||||
.slice(MAX_CARD_TAGS)
|
||||
.map((id) => tagMap.get(id)?.name ?? id)
|
||||
.join("、")}
|
||||
>
|
||||
<Tag className="vmat-tag-overflow">+{material.tagIds.length - MAX_CARD_TAGS}</Tag>
|
||||
</Tooltip>
|
||||
)}
|
||||
</>
|
||||
)}
|
||||
</div>
|
||||
|
||||
{/* 元信息 */}
|
||||
<div className="vmat-card-meta">
|
||||
<span>{formatDuration(material.duration)}</span>
|
||||
<span>{formatFileSize(material.fileSize)}</span>
|
||||
<span>{formatDate(material.createdAt)}</span>
|
||||
</div>
|
||||
|
||||
{/* 播放控制 */}
|
||||
<div className="vmat-card-player">
|
||||
<button
|
||||
type="button"
|
||||
className="vmat-play-btn"
|
||||
onClick={(e) => {
|
||||
e.stopPropagation()
|
||||
isPlaying ? onPause() : onPlay()
|
||||
}}
|
||||
disabled={!material.fileUrl}
|
||||
>
|
||||
{isPlaying ? <PauseCircleOutlined /> : <PlayCircleOutlined />}
|
||||
</button>
|
||||
<div ref={progressRef} className="vmat-progress" onMouseDown={handleProgressMouseDown}>
|
||||
<div className="vmat-progress-bar" style={{ width: `${progress}%` }} />
|
||||
{isPlaying && <div className="vmat-progress-thumb" style={{ left: `${progress}%` }} />}
|
||||
</div>
|
||||
<span className="vmat-time">
|
||||
{isPlaying ? formatDuration(currentTime) : formatDuration(material.duration)}
|
||||
</span>
|
||||
{/* 音量控制 */}
|
||||
<div className="vmat-volume">
|
||||
<button
|
||||
type="button"
|
||||
className="vmat-volume-btn"
|
||||
onClick={(e) => {
|
||||
e.stopPropagation()
|
||||
onToggleMute()
|
||||
}}
|
||||
title={volume === 0 ? "取消静音" : "静音"}
|
||||
>
|
||||
{volume === 0 ? <MutedOutlined /> : <SoundOutlined />}
|
||||
</button>
|
||||
<input
|
||||
type="range"
|
||||
className="vmat-volume-slider"
|
||||
min={0}
|
||||
max={1}
|
||||
step={0.05}
|
||||
value={volume}
|
||||
onChange={(e) => {
|
||||
e.stopPropagation()
|
||||
onVolumeChange(e)
|
||||
}}
|
||||
onClick={(e) => e.stopPropagation()}
|
||||
/>
|
||||
</div>
|
||||
</div>
|
||||
</div>
|
||||
)
|
||||
}
|
||||
|
||||
export default VoiceMaterialCard
|
||||
@@ -0,0 +1,186 @@
|
||||
import React, { useRef } from "react"
|
||||
import {
|
||||
PlayCircleOutlined,
|
||||
PauseCircleOutlined,
|
||||
EditOutlined,
|
||||
DeleteOutlined,
|
||||
CheckOutlined,
|
||||
} from "@ant-design/icons"
|
||||
import { Tooltip } from "antd"
|
||||
import { Tag } from "@/components/ui"
|
||||
import { type TagItem } from "@/api/tags"
|
||||
import { type VoiceMaterial } from "../types"
|
||||
import { MAX_ROW_TAGS, TAG_VARIANTS } from "../constants"
|
||||
import {
|
||||
genderClass,
|
||||
genderIcon,
|
||||
genderLabel,
|
||||
formatDuration,
|
||||
formatFileSize,
|
||||
} from "../utils/format"
|
||||
|
||||
export interface VoiceRowProps {
|
||||
material: VoiceMaterial
|
||||
isPlaying: boolean
|
||||
currentTime: number
|
||||
isSelected: boolean
|
||||
batchMode: boolean
|
||||
tagMap: Map<string, TagItem>
|
||||
onPlay: () => void
|
||||
onPause: () => void
|
||||
onSeek: (time: number) => void
|
||||
onEdit: () => void
|
||||
onDelete: () => void
|
||||
onToggleSelect: (id: string) => void
|
||||
}
|
||||
|
||||
const VoiceMaterialRow: React.FC<VoiceRowProps> = ({
|
||||
material,
|
||||
isPlaying,
|
||||
currentTime,
|
||||
isSelected,
|
||||
batchMode,
|
||||
tagMap,
|
||||
onPlay,
|
||||
onPause,
|
||||
onSeek,
|
||||
onEdit,
|
||||
onDelete,
|
||||
onToggleSelect,
|
||||
}) => {
|
||||
const progressRef = useRef<HTMLDivElement>(null)
|
||||
|
||||
const handleProgressMouseDown = (e: React.MouseEvent<HTMLDivElement>) => {
|
||||
if (!progressRef.current) return
|
||||
e.preventDefault()
|
||||
const doSeek = (ev: MouseEvent) => {
|
||||
if (!progressRef.current) return
|
||||
const rect = progressRef.current.getBoundingClientRect()
|
||||
const percent = Math.max(0, Math.min(1, (ev.clientX - rect.left) / rect.width))
|
||||
onSeek(percent * material.duration)
|
||||
}
|
||||
doSeek(e.nativeEvent)
|
||||
const handleMove = (ev: MouseEvent) => doSeek(ev)
|
||||
const handleUp = () => {
|
||||
document.removeEventListener("mousemove", handleMove)
|
||||
document.removeEventListener("mouseup", handleUp)
|
||||
}
|
||||
document.addEventListener("mousemove", handleMove)
|
||||
document.addEventListener("mouseup", handleUp)
|
||||
}
|
||||
|
||||
const progress = material.duration > 0 ? (currentTime / material.duration) * 100 : 0
|
||||
|
||||
return (
|
||||
<div
|
||||
className={`vmat-row ${genderClass(material.gender)}${isPlaying ? " playing" : ""}${isSelected ? " selected" : ""}${batchMode ? " batch-mode" : ""}`}
|
||||
>
|
||||
{/* 批量选择 checkbox */}
|
||||
{(batchMode || isSelected) && (
|
||||
<div
|
||||
className={`vmat-row-checkbox vmat-checkbox${isSelected ? " checked" : ""}`}
|
||||
onClick={(e) => {
|
||||
e.stopPropagation()
|
||||
onToggleSelect(material.id)
|
||||
}}
|
||||
>
|
||||
{isSelected && <CheckOutlined />}
|
||||
</div>
|
||||
)}
|
||||
|
||||
{/* 播放按钮 */}
|
||||
<button
|
||||
type="button"
|
||||
className="vmat-row-play"
|
||||
onClick={(e) => {
|
||||
e.stopPropagation()
|
||||
isPlaying ? onPause() : onPlay()
|
||||
}}
|
||||
disabled={!material.fileUrl}
|
||||
>
|
||||
{isPlaying ? <PauseCircleOutlined /> : <PlayCircleOutlined />}
|
||||
</button>
|
||||
|
||||
{/* 名称 + 描述 */}
|
||||
<div className="vmat-row-info">
|
||||
<h4 className="vmat-row-name">{material.name}</h4>
|
||||
{material.description && <p className="vmat-row-desc">{material.description}</p>}
|
||||
</div>
|
||||
|
||||
{/* 性别 */}
|
||||
<span className={`vmat-row-gender ${genderClass(material.gender)}`}>
|
||||
{genderIcon(material.gender)}
|
||||
{genderLabel(material.gender)}
|
||||
</span>
|
||||
|
||||
{/* 标签 */}
|
||||
<div className="vmat-row-tags">
|
||||
{material.tagIds.length === 0 ? (
|
||||
<span className="vmat-tag-empty" onClick={() => onEdit()}>
|
||||
添加标签
|
||||
</span>
|
||||
) : (
|
||||
<>
|
||||
{material.tagIds.slice(0, MAX_ROW_TAGS).map((tagId, i) => (
|
||||
<Tag key={tagId} variant={TAG_VARIANTS[i % TAG_VARIANTS.length]}>
|
||||
{tagMap.get(tagId)?.name ?? tagId}
|
||||
</Tag>
|
||||
))}
|
||||
{material.tagIds.length > MAX_ROW_TAGS && (
|
||||
<Tooltip
|
||||
title={material.tagIds
|
||||
.slice(MAX_ROW_TAGS)
|
||||
.map((id) => tagMap.get(id)?.name ?? id)
|
||||
.join("、")}
|
||||
>
|
||||
<Tag className="vmat-tag-overflow">+{material.tagIds.length - MAX_ROW_TAGS}</Tag>
|
||||
</Tooltip>
|
||||
)}
|
||||
</>
|
||||
)}
|
||||
</div>
|
||||
|
||||
{/* 进度条(可拖拽) */}
|
||||
<div ref={progressRef} className="vmat-row-progress" onMouseDown={handleProgressMouseDown}>
|
||||
<div className="vmat-row-progress-bar" style={{ width: `${progress}%` }} />
|
||||
{isPlaying && <div className="vmat-progress-thumb" style={{ left: `${progress}%` }} />}
|
||||
</div>
|
||||
|
||||
{/* 时长 */}
|
||||
<span className="vmat-row-time">
|
||||
{isPlaying ? formatDuration(currentTime) : formatDuration(material.duration)}
|
||||
</span>
|
||||
|
||||
{/* 文件大小 */}
|
||||
<span className="vmat-row-size">{formatFileSize(material.fileSize)}</span>
|
||||
|
||||
{/* 操作 */}
|
||||
<div className="vmat-row-actions">
|
||||
<button
|
||||
type="button"
|
||||
className="vmat-row-action-btn"
|
||||
onClick={(e) => {
|
||||
e.stopPropagation()
|
||||
onEdit()
|
||||
}}
|
||||
title="编辑"
|
||||
>
|
||||
<EditOutlined />
|
||||
</button>
|
||||
<button
|
||||
type="button"
|
||||
className="vmat-row-action-btn vmat-row-action-btn--danger"
|
||||
onClick={(e) => {
|
||||
e.stopPropagation()
|
||||
onDelete()
|
||||
}}
|
||||
title="删除"
|
||||
>
|
||||
<DeleteOutlined />
|
||||
</button>
|
||||
</div>
|
||||
</div>
|
||||
)
|
||||
}
|
||||
|
||||
export default VoiceMaterialRow
|
||||
@@ -0,0 +1,142 @@
|
||||
import { useState, useRef, useCallback, useEffect } from "react"
|
||||
import type { VoiceMaterial } from "../types"
|
||||
|
||||
/**
|
||||
* 音频播放控制 Hook
|
||||
* 封装当前播放音频状态、播放/暂停、进度控制、音量控制
|
||||
*/
|
||||
export function useAudioPlayer() {
|
||||
const [playingId, setPlayingId] = useState<string | null>(null)
|
||||
const [currentTime, setCurrentTime] = useState(0)
|
||||
const [volume, setVolume] = useState(0.7)
|
||||
const [pausedMaterial, setPausedMaterial] = useState<VoiceMaterial | null>(null)
|
||||
const audioRef = useRef<HTMLAudioElement | null>(null)
|
||||
|
||||
/** 停止当前播放并重置状态 */
|
||||
const stopPlayback = useCallback(() => {
|
||||
if (audioRef.current) {
|
||||
audioRef.current.pause()
|
||||
audioRef.current = null
|
||||
}
|
||||
setPlayingId(null)
|
||||
setCurrentTime(0)
|
||||
setPausedMaterial(null)
|
||||
}, [])
|
||||
|
||||
/** 从头开始播放指定素材 */
|
||||
const startPlayback = useCallback(
|
||||
(material: VoiceMaterial) => {
|
||||
if (!material.fileUrl) return
|
||||
stopPlayback()
|
||||
|
||||
const audio = new Audio(material.fileUrl)
|
||||
audio.volume = volume
|
||||
audioRef.current = audio
|
||||
|
||||
audio.addEventListener("timeupdate", () => {
|
||||
setCurrentTime(audio.currentTime)
|
||||
})
|
||||
|
||||
audio.addEventListener("ended", () => {
|
||||
setPlayingId(null)
|
||||
setCurrentTime(0)
|
||||
audioRef.current = null
|
||||
setPausedMaterial(null)
|
||||
})
|
||||
|
||||
audio.play().catch(() => {
|
||||
audioRef.current = null
|
||||
setPlayingId(null)
|
||||
})
|
||||
|
||||
setPlayingId(material.id)
|
||||
setCurrentTime(0)
|
||||
setPausedMaterial(null)
|
||||
},
|
||||
[stopPlayback, volume],
|
||||
)
|
||||
|
||||
/** 播放素材(若为暂停状态则恢复) */
|
||||
const handlePlay = useCallback(
|
||||
(material: VoiceMaterial) => {
|
||||
if (playingId === material.id) return
|
||||
// 恢复暂停
|
||||
if (pausedMaterial?.id === material.id && audioRef.current && audioRef.current.paused) {
|
||||
audioRef.current.play().catch(() => {})
|
||||
setPlayingId(material.id)
|
||||
setPausedMaterial(null)
|
||||
return
|
||||
}
|
||||
startPlayback(material)
|
||||
},
|
||||
[playingId, pausedMaterial, startPlayback],
|
||||
)
|
||||
|
||||
/** 暂停播放 */
|
||||
const handlePause = useCallback((material?: VoiceMaterial) => {
|
||||
if (audioRef.current) {
|
||||
audioRef.current.pause()
|
||||
}
|
||||
setPlayingId(null)
|
||||
if (material) setPausedMaterial(material)
|
||||
}, [])
|
||||
|
||||
/** 跳转到指定播放时间 */
|
||||
const handleSeek = useCallback(
|
||||
(material: VoiceMaterial, time: number) => {
|
||||
if (audioRef.current) {
|
||||
audioRef.current.currentTime = time
|
||||
setCurrentTime(time)
|
||||
} else {
|
||||
startPlayback(material)
|
||||
setTimeout(() => {
|
||||
if (audioRef.current) {
|
||||
audioRef.current.currentTime = time
|
||||
}
|
||||
}, 100)
|
||||
}
|
||||
},
|
||||
[startPlayback],
|
||||
)
|
||||
|
||||
/** 音量调节 */
|
||||
const handleVolumeChange = useCallback((e: React.ChangeEvent<HTMLInputElement>) => {
|
||||
const v = parseFloat(e.target.value)
|
||||
setVolume(v)
|
||||
if (audioRef.current) audioRef.current.volume = v
|
||||
}, [])
|
||||
|
||||
/** 静音/取消静音切换 */
|
||||
const toggleMute = useCallback(() => {
|
||||
if (volume > 0) {
|
||||
setVolume(0)
|
||||
if (audioRef.current) audioRef.current.volume = 0
|
||||
} else {
|
||||
setVolume(0.7)
|
||||
if (audioRef.current) audioRef.current.volume = 0.7
|
||||
}
|
||||
}, [volume])
|
||||
|
||||
// 组件卸载时清理 audio
|
||||
useEffect(() => {
|
||||
return () => {
|
||||
if (audioRef.current) {
|
||||
audioRef.current.pause()
|
||||
audioRef.current = null
|
||||
}
|
||||
}
|
||||
}, [])
|
||||
|
||||
return {
|
||||
playingId,
|
||||
currentTime,
|
||||
volume,
|
||||
pausedMaterial,
|
||||
stopPlayback,
|
||||
handlePlay,
|
||||
handlePause,
|
||||
handleSeek,
|
||||
handleVolumeChange,
|
||||
toggleMute,
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,132 @@
|
||||
import { useState, useCallback, useMemo } from "react"
|
||||
import { useQueryClient } from "@tanstack/react-query"
|
||||
import { message } from "antd"
|
||||
import { deleteAsset } from "@/api/assets"
|
||||
import { type TagItem, createTag, tagAsset } from "@/api/tags"
|
||||
import type { VoiceMaterial } from "../types"
|
||||
|
||||
/**
|
||||
* 批量操作 Hook
|
||||
* 封装批量选择、批量删除、批量打标签等逻辑
|
||||
*/
|
||||
interface UseBatchOperationsProps {
|
||||
/** 当前筛选后的素材列表 */
|
||||
filtered: VoiceMaterial[]
|
||||
/** 标签 ID → TagItem 映射 */
|
||||
tagMap: Map<string, TagItem>
|
||||
/** 所有可用标签 */
|
||||
tags: TagItem[]
|
||||
/** 当前播放中的素材 ID */
|
||||
playingId: string | null
|
||||
/** 停止播放回调 */
|
||||
stopPlayback: () => void
|
||||
}
|
||||
|
||||
export function useBatchOperations({
|
||||
filtered,
|
||||
tagMap,
|
||||
tags,
|
||||
playingId,
|
||||
stopPlayback,
|
||||
}: UseBatchOperationsProps) {
|
||||
const queryClient = useQueryClient()
|
||||
|
||||
const [selectedIds, setSelectedIds] = useState<Set<string>>(new Set())
|
||||
const [batchCustomTag, setBatchCustomTag] = useState("")
|
||||
|
||||
const batchMode = useMemo(() => selectedIds.size > 0, [selectedIds])
|
||||
const allSelected = useMemo(
|
||||
() => filtered.length > 0 && filtered.every((m) => selectedIds.has(m.id)),
|
||||
[filtered, selectedIds],
|
||||
)
|
||||
|
||||
/** 切换单个素材的选中状态 */
|
||||
const handleToggleSelect = useCallback((id: string) => {
|
||||
setSelectedIds((prev) => {
|
||||
const next = new Set(prev)
|
||||
if (next.has(id)) next.delete(id)
|
||||
else next.add(id)
|
||||
return next
|
||||
})
|
||||
}, [])
|
||||
|
||||
/** 全选 / 取消全选 */
|
||||
const handleSelectAll = useCallback(() => {
|
||||
if (allSelected) setSelectedIds(new Set())
|
||||
else setSelectedIds(new Set(filtered.map((m) => m.id)))
|
||||
}, [allSelected, filtered])
|
||||
|
||||
/** 批量删除 */
|
||||
const handleBatchDelete = useCallback(async () => {
|
||||
const ids = Array.from(selectedIds)
|
||||
let successCount = 0
|
||||
for (const id of ids) {
|
||||
try {
|
||||
await deleteAsset(id)
|
||||
successCount++
|
||||
} catch {
|
||||
/* ignore individual failures */
|
||||
}
|
||||
if (playingId === id) stopPlayback()
|
||||
}
|
||||
queryClient.invalidateQueries({ queryKey: ["assets", "voice"] })
|
||||
setSelectedIds(new Set())
|
||||
message.success(`已批量删除 ${successCount}/${ids.length} 个素材`)
|
||||
}, [selectedIds, playingId, stopPlayback, queryClient])
|
||||
|
||||
/** 批量打标签(已有标签) */
|
||||
const handleBatchTag = useCallback(
|
||||
async (tagId: string) => {
|
||||
const ids = Array.from(selectedIds)
|
||||
let successCount = 0
|
||||
for (const id of ids) {
|
||||
try {
|
||||
await tagAsset(id, [tagId])
|
||||
successCount++
|
||||
} catch {
|
||||
/* ignore individual failures */
|
||||
}
|
||||
}
|
||||
queryClient.invalidateQueries({ queryKey: ["assets", "voice"] })
|
||||
setSelectedIds(new Set())
|
||||
const tagName = tagMap.get(tagId)?.name ?? tagId
|
||||
if (successCount === 0) {
|
||||
message.error(`批量打标签失败,请重试`)
|
||||
} else {
|
||||
message.success(`已为 ${successCount}/${ids.length} 个素材添加标签「${tagName}」`)
|
||||
}
|
||||
},
|
||||
[selectedIds, queryClient, tagMap],
|
||||
)
|
||||
|
||||
/** 批量打标签(自定义输入:按名称查找或创建标签,再批量打标) */
|
||||
const handleBatchCustomTag = useCallback(
|
||||
async (name: string) => {
|
||||
// 先查找同名标签(不区分大小写)
|
||||
let existing = tags.find((t) => t.name.toLowerCase() === name.toLowerCase())
|
||||
if (!existing) {
|
||||
try {
|
||||
existing = await createTag(name)
|
||||
} catch {
|
||||
message.error(`创建标签「${name}」失败`)
|
||||
return
|
||||
}
|
||||
}
|
||||
await handleBatchTag(existing.id)
|
||||
},
|
||||
[tags, handleBatchTag],
|
||||
)
|
||||
|
||||
return {
|
||||
selectedIds,
|
||||
batchMode,
|
||||
allSelected,
|
||||
batchCustomTag,
|
||||
setBatchCustomTag,
|
||||
handleToggleSelect,
|
||||
handleSelectAll,
|
||||
handleBatchDelete,
|
||||
handleBatchTag,
|
||||
handleBatchCustomTag,
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,135 @@
|
||||
import { useState, useRef, useCallback, useEffect } from "react"
|
||||
import { useQuery, useQueryClient } from "@tanstack/react-query"
|
||||
import { message } from "antd"
|
||||
import { synthesizeSpeech, getTTSJobStatus, saveTtsToLibrary } from "@/api/tts"
|
||||
import { fetchPresetVoices, type PresetVoiceItem } from "@/api/voices"
|
||||
|
||||
/**
|
||||
* TTS 合成 Hook
|
||||
* 封装合成弹窗状态、合成请求、轮询、保存到素材库等逻辑
|
||||
*/
|
||||
export type TtsStatus = "idle" | "synthesizing" | "done" | "error"
|
||||
|
||||
export function useTtsSynthesize() {
|
||||
const queryClient = useQueryClient()
|
||||
|
||||
const [ttsOpen, setTtsOpen] = useState(false)
|
||||
const [ttsText, setTtsText] = useState("")
|
||||
const [ttsVoiceId, setTtsVoiceId] = useState<string>("")
|
||||
const [ttsSpeed, setTtsSpeed] = useState(1.0)
|
||||
const [ttsJobId, setTtsJobId] = useState<string | null>(null)
|
||||
const [ttsStatus, setTtsStatus] = useState<TtsStatus>("idle")
|
||||
const [ttsAudioUrl, setTtsAudioUrl] = useState<string | null>(null)
|
||||
const [ttsError, setTtsError] = useState<string | null>(null)
|
||||
const ttsTimerRef = useRef<ReturnType<typeof setInterval> | null>(null)
|
||||
|
||||
// 预设音色列表
|
||||
const { data: presetVoicesData } = useQuery({
|
||||
queryKey: ["preset-voices"],
|
||||
queryFn: fetchPresetVoices,
|
||||
staleTime: 60_000,
|
||||
})
|
||||
const presetVoices: PresetVoiceItem[] = presetVoicesData?.items ?? []
|
||||
|
||||
/** 开始 AI 配音合成 */
|
||||
const handleTtsSynthesize = useCallback(async () => {
|
||||
if (!ttsText.trim()) {
|
||||
message.warning("请输入要合成的文本")
|
||||
return
|
||||
}
|
||||
setTtsError(null)
|
||||
setTtsStatus("synthesizing")
|
||||
setTtsAudioUrl(null)
|
||||
setTtsJobId(null)
|
||||
|
||||
try {
|
||||
const resp = await synthesizeSpeech({
|
||||
text: ttsText.trim(),
|
||||
voice_id: ttsVoiceId || undefined,
|
||||
speed: ttsSpeed,
|
||||
})
|
||||
setTtsJobId(resp.job_id)
|
||||
|
||||
// 轮询任务状态
|
||||
ttsTimerRef.current = setInterval(async () => {
|
||||
try {
|
||||
const job = await getTTSJobStatus(resp.job_id)
|
||||
if (job.status === "completed") {
|
||||
clearInterval(ttsTimerRef.current!)
|
||||
ttsTimerRef.current = null
|
||||
setTtsStatus("done")
|
||||
setTtsAudioUrl(job.output_audio_url)
|
||||
} else if (job.status === "failed") {
|
||||
clearInterval(ttsTimerRef.current!)
|
||||
ttsTimerRef.current = null
|
||||
setTtsStatus("error")
|
||||
setTtsError(job.error_message || "合成失败")
|
||||
}
|
||||
} catch {
|
||||
clearInterval(ttsTimerRef.current!)
|
||||
ttsTimerRef.current = null
|
||||
setTtsStatus("error")
|
||||
setTtsError("查询合成状态失败")
|
||||
}
|
||||
}, 2000)
|
||||
} catch (err: unknown) {
|
||||
const msg = err instanceof Error ? err.message : "合成请求失败"
|
||||
setTtsStatus("error")
|
||||
setTtsError(msg)
|
||||
}
|
||||
}, [ttsText, ttsVoiceId, ttsSpeed])
|
||||
|
||||
/** 保存 TTS 结果到素材库 */
|
||||
const handleTtsSave = useCallback(async () => {
|
||||
if (!ttsJobId) return
|
||||
try {
|
||||
await saveTtsToLibrary(ttsJobId, {
|
||||
name: ttsText.slice(0, 20) || "AI配音",
|
||||
})
|
||||
message.success("已保存到配音库")
|
||||
queryClient.invalidateQueries({ queryKey: ["assets", "voice"] })
|
||||
setTtsOpen(false)
|
||||
} catch {
|
||||
message.error("保存失败")
|
||||
}
|
||||
}, [ttsJobId, ttsText, queryClient])
|
||||
|
||||
/** 关闭 TTS 弹窗并清理状态 */
|
||||
const handleTtsClose = useCallback(() => {
|
||||
setTtsOpen(false)
|
||||
if (ttsTimerRef.current) {
|
||||
clearInterval(ttsTimerRef.current)
|
||||
ttsTimerRef.current = null
|
||||
}
|
||||
setTtsStatus("idle")
|
||||
setTtsAudioUrl(null)
|
||||
setTtsError(null)
|
||||
setTtsJobId(null)
|
||||
}, [])
|
||||
|
||||
// 组件卸载时清理定时器
|
||||
useEffect(() => {
|
||||
return () => {
|
||||
if (ttsTimerRef.current) clearInterval(ttsTimerRef.current)
|
||||
}
|
||||
}, [])
|
||||
|
||||
return {
|
||||
ttsOpen,
|
||||
ttsText,
|
||||
ttsVoiceId,
|
||||
ttsSpeed,
|
||||
ttsJobId,
|
||||
ttsStatus,
|
||||
ttsAudioUrl,
|
||||
ttsError,
|
||||
presetVoices,
|
||||
setTtsOpen,
|
||||
setTtsText,
|
||||
setTtsVoiceId,
|
||||
setTtsSpeed,
|
||||
handleTtsSynthesize,
|
||||
handleTtsSave,
|
||||
handleTtsClose,
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,351 @@
|
||||
import { useState, useMemo, useCallback, useEffect } from "react"
|
||||
import { useQuery, useMutation, useQueryClient } from "@tanstack/react-query"
|
||||
import { message } from "antd"
|
||||
import {
|
||||
getAssetsByKind,
|
||||
createAsset,
|
||||
updateAsset,
|
||||
deleteAsset,
|
||||
uploadAssetDirect,
|
||||
getAssetLibraries,
|
||||
createAssetLibrary,
|
||||
} from "@/api/assets"
|
||||
import { type TagItem, getTags, createTag, tagAsset, untagAsset } from "@/api/tags"
|
||||
import {
|
||||
type VoiceGender,
|
||||
type ViewMode,
|
||||
type VoiceMaterial,
|
||||
mapAssetToMaterial,
|
||||
buildMetadata,
|
||||
} from "../types"
|
||||
import { getAudioDuration } from "../utils/audio"
|
||||
|
||||
/**
|
||||
* 配音素材数据 Hook
|
||||
* 封装素材列表查询、筛选状态管理、增删改等数据操作逻辑
|
||||
*/
|
||||
export function useVoiceMaterials() {
|
||||
const queryClient = useQueryClient()
|
||||
|
||||
// ── 获取 voice 类型素材库(用于上传) ──────────────────────
|
||||
const { data: libraries = [] } = useQuery({
|
||||
queryKey: ["asset-libraries"],
|
||||
queryFn: getAssetLibraries,
|
||||
staleTime: 60_000,
|
||||
})
|
||||
|
||||
const voiceLibrary = useMemo(() => libraries.find((lib) => lib.kind === "voice"), [libraries])
|
||||
|
||||
// 自动创建 voice 素材库(如果不存在)
|
||||
const createLibMutation = useMutation({
|
||||
mutationFn: () => createAssetLibrary({ name: "配音库", kind: "voice" }),
|
||||
onSuccess: () => {
|
||||
queryClient.invalidateQueries({ queryKey: ["asset-libraries"] })
|
||||
},
|
||||
})
|
||||
|
||||
useEffect(() => {
|
||||
if (libraries.length > 0 && !voiceLibrary && !createLibMutation.isPending) {
|
||||
createLibMutation.mutate()
|
||||
}
|
||||
}, [libraries, voiceLibrary, createLibMutation])
|
||||
|
||||
// ── 获取标签列表 ───────────────────────────────────────────
|
||||
const { data: tags = [] } = useQuery({
|
||||
queryKey: ["tags"],
|
||||
queryFn: getTags,
|
||||
staleTime: 60_000,
|
||||
})
|
||||
|
||||
/** 标签 ID → TagItem 映射(用于卡片/行渲染) */
|
||||
const tagMap = useMemo(() => {
|
||||
const m = new Map<string, TagItem>()
|
||||
tags.forEach((t) => m.set(t.id, t))
|
||||
return m
|
||||
}, [tags])
|
||||
|
||||
/** 创建标签 mutation(供 TagSelector 调用) */
|
||||
const createTagMutation = useMutation({
|
||||
mutationFn: (name: string) => createTag(name),
|
||||
onSuccess: () => {
|
||||
queryClient.invalidateQueries({ queryKey: ["tags"] })
|
||||
},
|
||||
})
|
||||
|
||||
/** 创建标签并返回 TagItem(供 TagSelector 使用) */
|
||||
const handleCreateTag = useCallback(
|
||||
async (name: string): Promise<TagItem> => {
|
||||
return createTagMutation.mutateAsync(name)
|
||||
},
|
||||
[createTagMutation],
|
||||
)
|
||||
|
||||
// ── 视图 & 筛选状态 ────────────────────────────────────────
|
||||
const [viewMode, setViewMode] = useState<ViewMode>("card")
|
||||
const [searchText, setSearchText] = useState("")
|
||||
const [filterGender, setFilterGender] = useState<string>("all")
|
||||
const [filterTagId, setFilterTagId] = useState<string>("all")
|
||||
|
||||
// ── 获取配音素材列表(筛选参数透传后端) ─────────────────
|
||||
const filterKeyword = searchText.trim() || undefined
|
||||
const filterGenderParam = filterGender !== "all" ? filterGender : undefined
|
||||
const filterTagIdsParam = filterTagId !== "all" ? [filterTagId] : undefined
|
||||
|
||||
const { data: assets = [], isLoading } = useQuery({
|
||||
queryKey: [
|
||||
"assets",
|
||||
"voice",
|
||||
{
|
||||
keyword: filterKeyword,
|
||||
gender: filterGenderParam,
|
||||
tag_ids: filterTagIdsParam,
|
||||
},
|
||||
],
|
||||
queryFn: () =>
|
||||
getAssetsByKind("voice", {
|
||||
keyword: filterKeyword,
|
||||
gender: filterGenderParam,
|
||||
tag_ids: filterTagIdsParam,
|
||||
}),
|
||||
staleTime: 30_000,
|
||||
})
|
||||
|
||||
const materials: VoiceMaterial[] = useMemo(() => assets.map(mapAssetToMaterial), [assets])
|
||||
|
||||
// ── 弹窗状态 ──────────────────────────────────────────────
|
||||
const [uploadOpen, setUploadOpen] = useState(false)
|
||||
const [editingMaterial, setEditingMaterial] = useState<VoiceMaterial | null>(null)
|
||||
|
||||
// ── 上传进度 ──────────────────────────────────────────────
|
||||
const [uploadProgress, setUploadProgress] = useState<number | null>(null)
|
||||
|
||||
// ── 上传 mutation ─────────────────────────────────────────
|
||||
const uploadMutation = useMutation({
|
||||
mutationFn: async (data: {
|
||||
file: File
|
||||
name: string
|
||||
gender: VoiceGender
|
||||
description: string
|
||||
tagIds: string[]
|
||||
}) => {
|
||||
setUploadProgress(0)
|
||||
try {
|
||||
// 1. 获取或等待 voice library
|
||||
let lib = voiceLibrary
|
||||
if (!lib) {
|
||||
if (createLibMutation.isPending) {
|
||||
await createLibMutation.mutateAsync()
|
||||
}
|
||||
const libs = await queryClient.fetchQuery({
|
||||
queryKey: ["asset-libraries"],
|
||||
queryFn: getAssetLibraries,
|
||||
})
|
||||
lib = libs.find((l) => l.kind === "voice")
|
||||
if (!lib) throw new Error("无法创建配音库")
|
||||
}
|
||||
|
||||
// 2. 上传文件(带进度)
|
||||
const { storage_key } = await uploadAssetDirect({
|
||||
file: data.file,
|
||||
library_id: lib.id,
|
||||
onProgress: (p) => setUploadProgress(p),
|
||||
})
|
||||
|
||||
// 3. 获取音频时长
|
||||
const duration = await getAudioDuration(data.file)
|
||||
|
||||
// 4. 创建素材记录
|
||||
const asset = await createAsset({
|
||||
library_id: lib.id,
|
||||
name: data.name,
|
||||
storage_key,
|
||||
mime_type: data.file.type || "audio/mpeg",
|
||||
metadata: buildMetadata({
|
||||
gender: data.gender,
|
||||
description: data.description,
|
||||
duration,
|
||||
}),
|
||||
})
|
||||
|
||||
// 5. 打标签(标签走独立 API)
|
||||
if (data.tagIds.length > 0) {
|
||||
await tagAsset(asset.id, data.tagIds)
|
||||
}
|
||||
} finally {
|
||||
setUploadProgress(null)
|
||||
}
|
||||
},
|
||||
onSuccess: () => {
|
||||
queryClient.invalidateQueries({ queryKey: ["assets", "voice"] })
|
||||
queryClient.invalidateQueries({ queryKey: ["tags"] })
|
||||
},
|
||||
onError: (err: Error) => {
|
||||
message.error(err.message || "上传失败,请重试")
|
||||
},
|
||||
})
|
||||
|
||||
// ── 编辑 mutation ─────────────────────────────────────────
|
||||
const editMutation = useMutation({
|
||||
mutationFn: async (data: {
|
||||
id: string
|
||||
name: string
|
||||
gender: VoiceGender
|
||||
description: string
|
||||
tagIds: string[]
|
||||
}) => {
|
||||
// 1. 更新基础信息
|
||||
await updateAsset(data.id, {
|
||||
name: data.name,
|
||||
metadata: buildMetadata({
|
||||
gender: data.gender,
|
||||
description: data.description,
|
||||
}),
|
||||
})
|
||||
|
||||
// 2. 对比标签差异,调用 tag/untag API
|
||||
const currentAsset = materials.find((m) => m.id === data.id)
|
||||
const oldTagIds = currentAsset?.tagIds ?? []
|
||||
const newTagIds = data.tagIds
|
||||
|
||||
const toAdd = newTagIds.filter((id) => !oldTagIds.includes(id))
|
||||
const toRemove = oldTagIds.filter((id) => !newTagIds.includes(id))
|
||||
|
||||
if (toAdd.length > 0) {
|
||||
await tagAsset(data.id, toAdd)
|
||||
}
|
||||
for (const tagId of toRemove) {
|
||||
await untagAsset(data.id, tagId)
|
||||
}
|
||||
},
|
||||
onSuccess: () => {
|
||||
queryClient.invalidateQueries({ queryKey: ["assets", "voice"] })
|
||||
queryClient.invalidateQueries({ queryKey: ["tags"] })
|
||||
},
|
||||
})
|
||||
|
||||
// ── 删除 mutation ─────────────────────────────────────────
|
||||
const deleteMutation = useMutation({
|
||||
mutationFn: (assetId: string) => deleteAsset(assetId),
|
||||
onSuccess: () => {
|
||||
queryClient.invalidateQueries({ queryKey: ["assets", "voice"] })
|
||||
},
|
||||
})
|
||||
|
||||
/* ── 前端二次筛选(与后端筛选同时存在) ──────────────────── */
|
||||
|
||||
const filtered = useMemo(() => {
|
||||
let list = materials
|
||||
if (filterGender !== "all") {
|
||||
list = list.filter((m) => m.gender === filterGender)
|
||||
}
|
||||
if (filterTagId !== "all") {
|
||||
list = list.filter((m) => m.tagIds.includes(filterTagId))
|
||||
}
|
||||
if (searchText.trim()) {
|
||||
const q = searchText.trim().toLowerCase()
|
||||
list = list.filter(
|
||||
(m) =>
|
||||
m.name.toLowerCase().includes(q) ||
|
||||
m.description.toLowerCase().includes(q) ||
|
||||
m.tagIds.some((id) => tagMap.get(id)?.name?.toLowerCase().includes(q)),
|
||||
)
|
||||
}
|
||||
return list
|
||||
}, [materials, filterGender, filterTagId, searchText, tagMap])
|
||||
|
||||
/* ── 标签使用计数(药丸条展示,按 tag ID 统计) ──────────── */
|
||||
|
||||
const tagCountMap = useMemo(() => {
|
||||
const map: Record<string, number> = {}
|
||||
materials.forEach((m) =>
|
||||
m.tagIds.forEach((id) => {
|
||||
map[id] = (map[id] || 0) + 1
|
||||
}),
|
||||
)
|
||||
return map
|
||||
}, [materials])
|
||||
|
||||
/* ── 数据操作 handlers ──────────────────────────────────── */
|
||||
|
||||
const handleUpload = useCallback(
|
||||
(data: Omit<VoiceMaterial, "id" | "createdAt"> & { file?: File }) => {
|
||||
if (!data.file) return
|
||||
uploadMutation.mutate(
|
||||
{
|
||||
file: data.file,
|
||||
name: data.name,
|
||||
gender: data.gender,
|
||||
description: data.description,
|
||||
tagIds: data.tagIds,
|
||||
},
|
||||
{
|
||||
onSuccess: () => {
|
||||
setUploadOpen(false)
|
||||
},
|
||||
},
|
||||
)
|
||||
},
|
||||
[uploadMutation],
|
||||
)
|
||||
|
||||
const handleEdit = useCallback(
|
||||
(data: Omit<VoiceMaterial, "id" | "createdAt"> & { file?: File }) => {
|
||||
if (!editingMaterial) return
|
||||
editMutation.mutate({
|
||||
id: editingMaterial.id,
|
||||
name: data.name,
|
||||
gender: data.gender,
|
||||
description: data.description,
|
||||
tagIds: data.tagIds,
|
||||
})
|
||||
setEditingMaterial(null)
|
||||
},
|
||||
[editingMaterial, editMutation],
|
||||
)
|
||||
|
||||
const handleDelete = useCallback(
|
||||
(id: string, onBeforeDelete?: () => void) => {
|
||||
const material = materials.find((m) => m.id === id)
|
||||
if (!material) return
|
||||
if (onBeforeDelete) onBeforeDelete()
|
||||
deleteMutation.mutate(id)
|
||||
},
|
||||
[materials, deleteMutation],
|
||||
)
|
||||
|
||||
return {
|
||||
// 数据
|
||||
libraries,
|
||||
voiceLibrary,
|
||||
tags,
|
||||
tagMap,
|
||||
materials,
|
||||
filtered,
|
||||
tagCountMap,
|
||||
isLoading,
|
||||
// 视图 & 筛选状态
|
||||
viewMode,
|
||||
searchText,
|
||||
filterGender,
|
||||
filterTagId,
|
||||
// 上传 & 编辑状态
|
||||
uploadProgress,
|
||||
isUploading: uploadMutation.isPending,
|
||||
isEditing: editMutation.isPending,
|
||||
// 弹窗状态
|
||||
uploadOpen,
|
||||
editingMaterial,
|
||||
// 视图控制
|
||||
setViewMode,
|
||||
setSearchText,
|
||||
setFilterGender,
|
||||
setFilterTagId,
|
||||
setUploadOpen,
|
||||
setEditingMaterial,
|
||||
// 操作
|
||||
handleCreateTag,
|
||||
handleUpload,
|
||||
handleEdit,
|
||||
handleDelete,
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,71 @@
|
||||
/**
|
||||
* TagSelector 组件单元测试
|
||||
* 同时 import VoiceMaterialLibrary 主组件,确保 vitest related 模式
|
||||
* 能匹配到 voice-materials 目录下所有文件的改动
|
||||
*/
|
||||
import { render, screen, fireEvent, within } from "@testing-library/react"
|
||||
import { describe, it, expect, vi } from "vitest"
|
||||
import TagSelector from "@/pages/voice-materials/components/TagSelector"
|
||||
// 引入主组件以建立依赖链,让 vitest related 覆盖整个 voice-materials 目录
|
||||
import "@/pages/voice-materials/VoiceMaterialLibrary"
|
||||
import type { TagItem } from "@/api/tags"
|
||||
|
||||
const mockTags: TagItem[] = [
|
||||
{ id: "tag-1", name: "搞笑" },
|
||||
{ id: "tag-2", name: "情感" },
|
||||
{ id: "tag-3", name: "励志" },
|
||||
]
|
||||
|
||||
const mockTagMap = new Map(mockTags.map((t) => [t.id, t]))
|
||||
|
||||
describe("TagSelector", () => {
|
||||
const defaultProps = {
|
||||
value: [],
|
||||
onChange: vi.fn(),
|
||||
tags: mockTags,
|
||||
tagMap: mockTagMap,
|
||||
onCreateTag: vi.fn().mockResolvedValue({ id: "new-tag", name: "新标签" }),
|
||||
}
|
||||
|
||||
it("应渲染占位符文本", () => {
|
||||
render(<TagSelector {...defaultProps} />)
|
||||
expect(screen.getByPlaceholderText("输入标签后回车添加")).toBeInTheDocument()
|
||||
})
|
||||
|
||||
it("应渲染已选标签", () => {
|
||||
const { container } = render(<TagSelector {...defaultProps} value={["tag-1", "tag-2"]} />)
|
||||
// 在标签选择器区域内查找已选标签
|
||||
const selectorArea = container.querySelector(".vmat-tag-selector")
|
||||
expect(selectorArea).not.toBeNull()
|
||||
expect(within(selectorArea as HTMLElement).getByText("搞笑")).toBeInTheDocument()
|
||||
expect(within(selectorArea as HTMLElement).getByText("情感")).toBeInTheDocument()
|
||||
})
|
||||
|
||||
it("应渲染预设标签快捷选择区", () => {
|
||||
const { container } = render(<TagSelector {...defaultProps} />)
|
||||
const presetsArea = container.querySelector(".vmat-tag-selector-presets")
|
||||
expect(presetsArea).not.toBeNull()
|
||||
expect(within(presetsArea as HTMLElement).getByText("搞笑")).toBeInTheDocument()
|
||||
expect(within(presetsArea as HTMLElement).getByText("情感")).toBeInTheDocument()
|
||||
expect(within(presetsArea as HTMLElement).getByText("励志")).toBeInTheDocument()
|
||||
})
|
||||
|
||||
it("点击预设标签应触发 onChange", () => {
|
||||
const onChange = vi.fn()
|
||||
const { container } = render(<TagSelector {...defaultProps} onChange={onChange} />)
|
||||
const presetsArea = container.querySelector(".vmat-tag-selector-presets")
|
||||
fireEvent.click(within(presetsArea as HTMLElement).getByText("搞笑"))
|
||||
expect(onChange).toHaveBeenCalledWith(["tag-1"])
|
||||
})
|
||||
|
||||
it("点击已选预设标签应移除", () => {
|
||||
const onChange = vi.fn()
|
||||
const { container } = render(
|
||||
<TagSelector {...defaultProps} value={["tag-1"]} onChange={onChange} />,
|
||||
)
|
||||
const presetsArea = container.querySelector(".vmat-tag-selector-presets")
|
||||
// 点击预设区中已选中的标签按钮
|
||||
fireEvent.click(within(presetsArea as HTMLElement).getByText("搞笑"))
|
||||
expect(onChange).toHaveBeenCalledWith([])
|
||||
})
|
||||
})
|
||||
@@ -0,0 +1,148 @@
|
||||
/**
|
||||
* useAudioPlayer hook 测试
|
||||
*/
|
||||
import { describe, it, expect, beforeEach, vi } from "vitest"
|
||||
import { renderHook, act } from "@testing-library/react"
|
||||
import { useAudioPlayer } from "@/pages/voice-materials/hooks/useAudioPlayer"
|
||||
import type { VoiceMaterial } from "@/pages/voice-materials/types"
|
||||
|
||||
// Mock Audio constructor
|
||||
const mockAudioPlay = vi.fn()
|
||||
const mockAudioPause = vi.fn()
|
||||
const mockAddEventListener = vi.fn()
|
||||
|
||||
beforeEach(() => {
|
||||
vi.clearAllMocks()
|
||||
mockAudioPlay.mockReset()
|
||||
mockAudioPause.mockReset()
|
||||
mockAddEventListener.mockReset()
|
||||
|
||||
// Mock HTMLAudioElement
|
||||
global.Audio = vi.fn().mockImplementation(() => ({
|
||||
play: mockAudioPlay.mockResolvedValue(undefined),
|
||||
pause: mockAudioPause,
|
||||
addEventListener: mockAddEventListener,
|
||||
currentTime: 0,
|
||||
volume: 0.7,
|
||||
paused: true,
|
||||
})) as unknown as typeof Audio
|
||||
})
|
||||
|
||||
const mockMaterial: VoiceMaterial = {
|
||||
id: "test-1",
|
||||
name: "测试素材",
|
||||
description: "测试描述",
|
||||
gender: "male",
|
||||
tagIds: ["tag-1"],
|
||||
fileName: "test.mp3",
|
||||
fileSize: 1024,
|
||||
duration: 30,
|
||||
mimeType: "audio/mpeg",
|
||||
createdAt: "2024-01-01T00:00:00Z",
|
||||
fileUrl: "https://example.com/test.mp3",
|
||||
}
|
||||
|
||||
describe("useAudioPlayer", () => {
|
||||
it("应该使用初始状态初始化", () => {
|
||||
const { result } = renderHook(() => useAudioPlayer())
|
||||
|
||||
expect(result.current.playingId).toBeNull()
|
||||
expect(result.current.currentTime).toBe(0)
|
||||
expect(result.current.volume).toBe(0.7)
|
||||
expect(result.current.pausedMaterial).toBeNull()
|
||||
})
|
||||
|
||||
it("stopPlayback 应该重置播放状态", () => {
|
||||
const { result } = renderHook(() => useAudioPlayer())
|
||||
|
||||
act(() => {
|
||||
result.current.stopPlayback()
|
||||
})
|
||||
|
||||
expect(result.current.playingId).toBeNull()
|
||||
expect(result.current.currentTime).toBe(0)
|
||||
expect(result.current.pausedMaterial).toBeNull()
|
||||
})
|
||||
|
||||
it("handlePause 应该暂停播放并设置 pausedMaterial", () => {
|
||||
const { result } = renderHook(() => useAudioPlayer())
|
||||
|
||||
act(() => {
|
||||
result.current.handlePause(mockMaterial)
|
||||
})
|
||||
|
||||
expect(result.current.playingId).toBeNull()
|
||||
expect(result.current.pausedMaterial).toEqual(mockMaterial)
|
||||
})
|
||||
|
||||
it("handlePause 不传参数时不设置 pausedMaterial", () => {
|
||||
const { result } = renderHook(() => useAudioPlayer())
|
||||
|
||||
act(() => {
|
||||
result.current.handlePause()
|
||||
})
|
||||
|
||||
expect(result.current.playingId).toBeNull()
|
||||
expect(result.current.pausedMaterial).toBeNull()
|
||||
})
|
||||
|
||||
it("toggleMute 应该切换静音状态", () => {
|
||||
const { result } = renderHook(() => useAudioPlayer())
|
||||
|
||||
// 默认音量 0.7,静音后应为 0
|
||||
act(() => {
|
||||
result.current.toggleMute()
|
||||
})
|
||||
expect(result.current.volume).toBe(0)
|
||||
|
||||
// 再次切换,恢复到 0.7
|
||||
act(() => {
|
||||
result.current.toggleMute()
|
||||
})
|
||||
expect(result.current.volume).toBe(0.7)
|
||||
})
|
||||
|
||||
it("handlePlay 应该开始播放素材", () => {
|
||||
const { result } = renderHook(() => useAudioPlayer())
|
||||
|
||||
act(() => {
|
||||
result.current.handlePlay(mockMaterial)
|
||||
})
|
||||
|
||||
expect(result.current.playingId).toBe("test-1")
|
||||
expect(result.current.currentTime).toBe(0)
|
||||
expect(global.Audio).toHaveBeenCalledWith("https://example.com/test.mp3")
|
||||
expect(mockAudioPlay).toHaveBeenCalled()
|
||||
})
|
||||
|
||||
it("handlePlay 对同一个素材不应重复播放", () => {
|
||||
const { result } = renderHook(() => useAudioPlayer())
|
||||
|
||||
act(() => {
|
||||
result.current.handlePlay(mockMaterial)
|
||||
})
|
||||
|
||||
const playCallCount = mockAudioPlay.mock.calls.length
|
||||
|
||||
act(() => {
|
||||
result.current.handlePlay(mockMaterial)
|
||||
})
|
||||
|
||||
// 不应该再次调用 play
|
||||
expect(mockAudioPlay.mock.calls.length).toBe(playCallCount)
|
||||
})
|
||||
|
||||
it("返回值应该包含所有必要的方法和状态", () => {
|
||||
const { result } = renderHook(() => useAudioPlayer())
|
||||
|
||||
expect(typeof result.current.handlePlay).toBe("function")
|
||||
expect(typeof result.current.handlePause).toBe("function")
|
||||
expect(typeof result.current.handleSeek).toBe("function")
|
||||
expect(typeof result.current.handleVolumeChange).toBe("function")
|
||||
expect(typeof result.current.toggleMute).toBe("function")
|
||||
expect(typeof result.current.stopPlayback).toBe("function")
|
||||
expect(typeof result.current.playingId).toBe("object") // string | null
|
||||
expect(typeof result.current.currentTime).toBe("number")
|
||||
expect(typeof result.current.volume).toBe("number")
|
||||
})
|
||||
})
|
||||
@@ -0,0 +1,26 @@
|
||||
/**
|
||||
* VoiceMaterialLibrary 模块 smoke test
|
||||
* 建立完整依赖链,确保 vitest related 模式能匹配到
|
||||
* voice-materials 目录下所有文件的改动(包括子组件和工具函数)
|
||||
*/
|
||||
import { describe, it, expect } from "vitest"
|
||||
|
||||
// 主组件
|
||||
import "@/pages/voice-materials/VoiceMaterialLibrary"
|
||||
|
||||
// 子组件
|
||||
import "@/pages/voice-materials/components/TagSelector"
|
||||
import "@/pages/voice-materials/components/MaterialForm"
|
||||
import "@/pages/voice-materials/components/VoiceMaterialCard"
|
||||
import "@/pages/voice-materials/components/VoiceMaterialRow"
|
||||
|
||||
// 工具函数
|
||||
import "@/pages/voice-materials/utils/format"
|
||||
import "@/pages/voice-materials/utils/audio"
|
||||
|
||||
describe("VoiceMaterialLibrary module smoke test", () => {
|
||||
it("should load all voice-material modules", () => {
|
||||
// 纯模块加载测试,确保所有组件/工具函数能正常 import
|
||||
expect(true).toBe(true)
|
||||
})
|
||||
})
|
||||
@@ -28,6 +28,28 @@ class EditPlanStatus(StrEnum):
|
||||
COMPLETED = "completed"
|
||||
FAILED = "failed"
|
||||
|
||||
@classmethod
|
||||
def _missing_(cls, value: object) -> "EditPlanStatus":
|
||||
"""兼容历史脏数据,避免枚举转换失败导致500。
|
||||
|
||||
- success/done/finished/complete → COMPLETED
|
||||
- fail/error/err → FAILED
|
||||
- render/rendering → RENDERING
|
||||
- edit/editing → EDITING
|
||||
- 其他未知值 → DRAFT(兜底,不阻塞业务)
|
||||
"""
|
||||
if isinstance(value, str):
|
||||
normalized = value.strip().lower()
|
||||
if normalized in ("done", "success", "finished", "complete", "completed"):
|
||||
return cls.COMPLETED
|
||||
if normalized in ("fail", "failed", "error", "err"):
|
||||
return cls.FAILED
|
||||
if normalized in ("render", "rendering", "generating", "generating_video"):
|
||||
return cls.RENDERING
|
||||
if normalized in ("edit", "editing", "working"):
|
||||
return cls.EDITING
|
||||
return cls.DRAFT
|
||||
|
||||
|
||||
@dataclass(slots=True)
|
||||
class EditPlan:
|
||||
|
||||
@@ -31,6 +31,25 @@ class EditPlanClipStatus(StrEnum):
|
||||
RENDERED = "rendered" # 已渲染
|
||||
FAILED = "failed" # 渲染失败
|
||||
|
||||
@classmethod
|
||||
def _missing_(cls, value: object) -> "EditPlanClipStatus":
|
||||
"""兼容历史脏数据,避免枚举转换失败导致500。
|
||||
|
||||
- success/done/finished/complete/rendered → RENDERED
|
||||
- fail/error/err → FAILED
|
||||
- ready/available → READY
|
||||
- 其他未知值 → PENDING(兜底,不阻塞业务)
|
||||
"""
|
||||
if isinstance(value, str):
|
||||
normalized = value.strip().lower()
|
||||
if normalized in ("done", "success", "finished", "complete", "rendered", "render"):
|
||||
return cls.RENDERED
|
||||
if normalized in ("fail", "failed", "error", "err"):
|
||||
return cls.FAILED
|
||||
if normalized in ("ready", "available", "prepared"):
|
||||
return cls.READY
|
||||
return cls.PENDING
|
||||
|
||||
|
||||
@dataclass(slots=True)
|
||||
class EditPlanClip:
|
||||
|
||||
@@ -43,6 +43,28 @@ class GenerationTaskStatus(StrEnum):
|
||||
CANCELLED = "cancelled"
|
||||
"""已取消(用户取消或系统取消)"""
|
||||
|
||||
@classmethod
|
||||
def _missing_(cls, value: object) -> "GenerationTaskStatus":
|
||||
"""兼容历史脏数据,避免枚举转换失败导致500。
|
||||
|
||||
- success/done/finished/complete → COMPLETED
|
||||
- fail/error/err → FAILED
|
||||
- process/processing/run/running → RUNNING
|
||||
- cancel/canceled → CANCELLED
|
||||
- 其他未知值 → PENDING(兜底,不阻塞业务)
|
||||
"""
|
||||
if isinstance(value, str):
|
||||
normalized = value.strip().lower()
|
||||
if normalized in ("done", "success", "finished", "complete", "completed"):
|
||||
return cls.COMPLETED
|
||||
if normalized in ("fail", "failed", "error", "err"):
|
||||
return cls.FAILED
|
||||
if normalized in ("process", "processing", "run", "running", "in_progress"):
|
||||
return cls.RUNNING
|
||||
if normalized in ("cancel", "cancelled", "canceled"):
|
||||
return cls.CANCELLED
|
||||
return cls.PENDING
|
||||
|
||||
|
||||
# 终态集合
|
||||
TERMINAL_STATUSES = frozenset(
|
||||
|
||||
Executable
+137
@@ -0,0 +1,137 @@
|
||||
#!/usr/bin/env python3
|
||||
"""
|
||||
检查 Alembic migration 文件命名规范。
|
||||
|
||||
规则:
|
||||
1. 文件名必须以数字前缀开头(3位补零),如 001_xxx.py、052_add_table.py
|
||||
2. 数字前缀必须连续递增(与 check_migration_chain.py 一致,但只看文件名)
|
||||
3. 数字前缀后必须跟有描述性后缀(不能只有数字)
|
||||
4. 文件名使用小写+下划线(snake_case)
|
||||
5. revision 变量值必须与文件名数字前缀一致(可选带描述后缀)
|
||||
|
||||
用法:
|
||||
python3 scripts/ci/check_migration_naming.py [alembic_versions_dir]
|
||||
|
||||
默认目录: alembic/versions/
|
||||
|
||||
退出码:
|
||||
0 - 全部通过
|
||||
1 - 有命名违规
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import re
|
||||
import sys
|
||||
from pathlib import Path
|
||||
|
||||
# 文件名格式: 3位数字_描述.py
|
||||
FILE_NAME_PATTERN = re.compile(r"^(\d{3})_[a-z][a-z0-9_]*\.py$")
|
||||
# 纯数字文件名(不允许)
|
||||
PURE_NUM_PATTERN = re.compile(r"^\d{3}\.py$")
|
||||
# revision 值的数字前缀
|
||||
REV_NUM_PATTERN = re.compile(r"^(\d{3})")
|
||||
# revision 变量行
|
||||
REV_LINE_PATTERN = re.compile(
|
||||
r'^\s*revision\s*(?::\s*str\s*)?=\s*["\']([^"\']+)["\']',
|
||||
re.MULTILINE,
|
||||
)
|
||||
|
||||
|
||||
def check_naming(versions_dir: Path) -> list[str]:
|
||||
"""检查 migration 文件命名,返回错误列表。"""
|
||||
errors: list[str] = []
|
||||
|
||||
if not versions_dir.is_dir():
|
||||
return [f"目录不存在: {versions_dir}"]
|
||||
|
||||
py_files = sorted(f for f in versions_dir.iterdir() if f.suffix == ".py")
|
||||
if not py_files:
|
||||
return [f"目录下没有 migration 文件: {versions_dir}"]
|
||||
|
||||
print(f"检查 migration 文件命名: {versions_dir}")
|
||||
print(f"共 {len(py_files)} 个文件")
|
||||
print()
|
||||
|
||||
# 1. 文件名格式检查
|
||||
print("1. 文件名格式检查...")
|
||||
file_nums: list[int] = []
|
||||
for f in py_files:
|
||||
name = f.name
|
||||
if PURE_NUM_PATTERN.match(name):
|
||||
errors.append(f" ❌ {name}: 只有数字编号,缺少描述性后缀")
|
||||
continue
|
||||
m = FILE_NAME_PATTERN.match(name)
|
||||
if not m:
|
||||
errors.append(f" ❌ {name}: 命名格式不规范,应为 NNN_description.py " f"(3位数字前缀+下划线+小写描述)")
|
||||
continue
|
||||
file_nums.append(int(m.group(1)))
|
||||
|
||||
if not any("命名格式不规范" in e or "缺少描述性后缀" in e for e in errors):
|
||||
print(f" ✅ 全部 {len(py_files)} 个文件名格式正确")
|
||||
else:
|
||||
for e in errors:
|
||||
if "命名格式不规范" in e or "缺少描述性后缀" in e:
|
||||
print(e)
|
||||
|
||||
# 2. 编号连续性检查(基于文件名数字前缀)
|
||||
print()
|
||||
print("2. 编号连续性检查...")
|
||||
if file_nums:
|
||||
expected = set(range(min(file_nums), max(file_nums) + 1))
|
||||
actual = set(file_nums)
|
||||
missing = sorted(expected - actual)
|
||||
if missing:
|
||||
errors.append(f" ❌ 编号不连续,缺少: {', '.join(f'{n:03d}' for n in missing)}")
|
||||
print(f" ❌ 编号不连续,缺少 {len(missing)} 个: " f"{', '.join(f'{n:03d}' for n in missing)}")
|
||||
else:
|
||||
print(f" ✅ 编号连续({min(file_nums):03d} ~ {max(file_nums):03d})")
|
||||
|
||||
# 3. revision 变量与文件名前缀一致性检查
|
||||
print()
|
||||
print("3. revision变量与文件名一致性检查...")
|
||||
rev_mismatch = 0
|
||||
for f in py_files:
|
||||
m = FILE_NAME_PATTERN.match(f.name)
|
||||
if not m:
|
||||
continue # 格式不对的已经报过了
|
||||
file_num = m.group(1)
|
||||
content = f.read_text(encoding="utf-8")
|
||||
rev_match = REV_LINE_PATTERN.search(content)
|
||||
if not rev_match:
|
||||
errors.append(f" ❌ {f.name}: 未找到 revision 变量定义")
|
||||
rev_mismatch += 1
|
||||
continue
|
||||
rev_value = rev_match.group(1)
|
||||
rev_num_match = REV_NUM_PATTERN.match(rev_value)
|
||||
if not rev_num_match or rev_num_match.group(1) != file_num:
|
||||
errors.append(f" ❌ {f.name}: revision='{rev_value}' 与文件名前缀 {file_num} 不一致")
|
||||
rev_mismatch += 1
|
||||
|
||||
if rev_mismatch == 0:
|
||||
print(f" ✅ 全部 {len(py_files)} 个文件的 revision 与文件名一致")
|
||||
|
||||
return errors
|
||||
|
||||
|
||||
def main() -> int:
|
||||
versions_dir = Path(sys.argv[1]) if len(sys.argv) > 1 else Path("alembic/versions")
|
||||
|
||||
errors = check_naming(versions_dir)
|
||||
|
||||
print()
|
||||
if errors:
|
||||
print(f"❌ 发现 {len(errors)} 个命名问题")
|
||||
print()
|
||||
print("命名规范:")
|
||||
print(" - 文件名格式: NNN_description.py(3位数字前缀 + 下划线 + 小写描述)")
|
||||
print(" - 编号必须连续,不能跳号")
|
||||
print(" - revision 变量的数字前缀必须与文件名一致")
|
||||
return 1
|
||||
|
||||
print("✅ 所有 migration 文件命名规范检查通过")
|
||||
return 0
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
sys.exit(main())
|
||||
Regular → Executable
+172
-33
@@ -1,15 +1,63 @@
|
||||
#!/bin/bash
|
||||
# CI Validate: Alembic迁移验证(并行Job 3/3)
|
||||
# 需要PostgreSQL数据库
|
||||
# CI Validate: Alembic迁移验证(升级版)
|
||||
# 检查项:
|
||||
# 1. migration文件命名规范检查
|
||||
# 2. migration编号链完整性检查
|
||||
# 3. upgrade head 升级验证(真实PG执行)
|
||||
# 4. downgrade -1 回滚验证
|
||||
# 5. alembic check 检测未生成migration的model变更
|
||||
#
|
||||
# 需要PostgreSQL数据库(共享PG或临时容器)
|
||||
|
||||
set -eu
|
||||
# 加载CI共享常量
|
||||
|
||||
SCRIPT_DIR="$(dirname "${BASH_SOURCE[0]}")"
|
||||
# shellcheck source=ci_env.sh
|
||||
source "${SCRIPT_DIR}/ci_env.sh"
|
||||
|
||||
echo "=== CI Validate: Alembic迁移验证 ==="
|
||||
echo "=== CI Validate: Alembic迁移验证(升级版)==="
|
||||
echo ""
|
||||
|
||||
# ============================================================
|
||||
# 阶段0: 静态检查(不需要数据库,先快速失败)
|
||||
# ============================================================
|
||||
|
||||
echo "📋 阶段0: 静态检查(命名规范 + 链完整性)"
|
||||
echo ""
|
||||
|
||||
STATIC_FAILED=0
|
||||
|
||||
echo "0.1 检查 migration 文件命名规范..."
|
||||
if python3 scripts/ci/check_migration_naming.py alembic/versions; then
|
||||
echo " ✅ 命名规范检查通过"
|
||||
else
|
||||
echo " ❌ 命名规范检查失败"
|
||||
STATIC_FAILED=1
|
||||
fi
|
||||
|
||||
echo ""
|
||||
echo "0.2 检查 migration 编号链完整性..."
|
||||
if python3 scripts/ci/check_migration_chain.py alembic/versions; then
|
||||
echo " ✅ 编号链完整性检查通过"
|
||||
else
|
||||
echo " ❌ 编号链完整性检查失败"
|
||||
STATIC_FAILED=1
|
||||
fi
|
||||
|
||||
if [ "$STATIC_FAILED" -ne 0 ]; then
|
||||
echo ""
|
||||
echo "❌ 静态检查失败,请修复上述问题后重试"
|
||||
exit 1
|
||||
fi
|
||||
|
||||
echo ""
|
||||
echo "✅ 静态检查全部通过"
|
||||
echo ""
|
||||
|
||||
# ============================================================
|
||||
# DooD模式检测:确定宿主机访问地址
|
||||
# ============================================================
|
||||
|
||||
# --- DooD模式检测:确定宿主机访问地址 ---
|
||||
detect_docker_host() {
|
||||
local test_port="${1:-${CI_LOCAL_PG_PORT}}"
|
||||
|
||||
@@ -63,20 +111,6 @@ except:
|
||||
return 1
|
||||
}
|
||||
|
||||
# 获取宿主机IP
|
||||
if [ -S /var/run/docker.sock ]; then
|
||||
DOCKER_HOST_IP=$(detect_docker_host "${CI_SHARED_PG_PORT}")
|
||||
if [ "$DOCKER_HOST_IP" = "127.0.0.1" ]; then
|
||||
DOCKER_HOST_IP=$(detect_docker_host 22)
|
||||
fi
|
||||
echo "检测到DooD模式,宿主机地址: $DOCKER_HOST_IP"
|
||||
else
|
||||
DOCKER_HOST_IP="127.0.0.1"
|
||||
echo "非DooD模式,使用 127.0.0.1"
|
||||
fi
|
||||
PG_HOST="$DOCKER_HOST_IP"
|
||||
echo "PG host: $PG_HOST"
|
||||
|
||||
# 指数退避TCP连接检查
|
||||
wait_tcp_ready() {
|
||||
local host="$1"
|
||||
@@ -96,8 +130,32 @@ wait_tcp_ready() {
|
||||
return 1
|
||||
}
|
||||
|
||||
# 获取宿主机IP
|
||||
if [ -S /var/run/docker.sock ]; then
|
||||
DOCKER_HOST_IP=$(detect_docker_host "${CI_SHARED_PG_PORT}")
|
||||
if [ "$DOCKER_HOST_IP" = "127.0.0.1" ]; then
|
||||
DOCKER_HOST_IP=$(detect_docker_host 22)
|
||||
fi
|
||||
echo "检测到DooD模式,宿主机地址: $DOCKER_HOST_IP"
|
||||
else
|
||||
DOCKER_HOST_IP="127.0.0.1"
|
||||
echo "非DooD模式,使用 127.0.0.1"
|
||||
fi
|
||||
PG_HOST="$DOCKER_HOST_IP"
|
||||
echo "PG host: $PG_HOST"
|
||||
echo ""
|
||||
|
||||
USE_SHARED_PG="${CI_USE_SHARED_PG:-false}"
|
||||
|
||||
# ============================================================
|
||||
# 准备数据库
|
||||
# ============================================================
|
||||
|
||||
echo "🗄️ 阶段1: 准备测试数据库"
|
||||
echo ""
|
||||
|
||||
CI_DB_NAME="ci_migrate_${GITHUB_RUN_ID:-$$}"
|
||||
|
||||
if [ "$USE_SHARED_PG" = "true" ]; then
|
||||
# 使用常驻共享PG实例
|
||||
echo "使用常驻共享PG实例(CI_USE_SHARED_PG=true)"
|
||||
@@ -105,7 +163,6 @@ if [ "$USE_SHARED_PG" = "true" ]; then
|
||||
SHARED_PG_PORT="${CI_SHARED_PG_PORT}"
|
||||
SHARED_PG_USER="${CI_SHARED_PG_USER}"
|
||||
SHARED_PG_PASSWORD="${CI_SHARED_PG_PASSWORD}"
|
||||
CI_DB_NAME="ci_run_${GITHUB_RUN_ID:-$$}"
|
||||
|
||||
echo "等待共享PG连接就绪..."
|
||||
wait_tcp_ready "$SHARED_PG_HOST" "$SHARED_PG_PORT" 5
|
||||
@@ -124,13 +181,10 @@ conn.close()
|
||||
export DATABASE_URL="postgresql+psycopg://${SHARED_PG_USER}:${SHARED_PG_PASSWORD}@${SHARED_PG_HOST}:${SHARED_PG_PORT}/${CI_DB_NAME}"
|
||||
echo "✅ 共享PG数据库已创建: $CI_DB_NAME"
|
||||
|
||||
# 执行迁移
|
||||
PYTHONPATH="$PWD/apps/api:$PWD" python3 -m alembic upgrade head
|
||||
echo "✅ Alembic migrations applied successfully"
|
||||
|
||||
# 清理数据库
|
||||
echo "清理测试数据库: $CI_DB_NAME"
|
||||
PGPASSWORD="$SHARED_PG_PASSWORD" python3 -c "
|
||||
cleanup_db() {
|
||||
echo ""
|
||||
echo "清理测试数据库: $CI_DB_NAME"
|
||||
PGPASSWORD="$SHARED_PG_PASSWORD" python3 -c "
|
||||
import psycopg2
|
||||
conn = psycopg2.connect(host='$SHARED_PG_HOST', port=$SHARED_PG_PORT, user='$SHARED_PG_USER', password='$SHARED_PG_PASSWORD', dbname='postgres')
|
||||
conn.autocommit = True
|
||||
@@ -139,7 +193,8 @@ cur.execute(f'DROP DATABASE IF EXISTS \"$CI_DB_NAME\" WITH (FORCE)')
|
||||
cur.close()
|
||||
conn.close()
|
||||
" 2>/dev/null || echo "WARN: 数据库清理失败"
|
||||
echo "✅ 共享PG数据库已清理"
|
||||
echo "✅ 数据库已清理"
|
||||
}
|
||||
else
|
||||
# 使用临时PG容器(默认模式)
|
||||
echo "使用临时PG容器模式"
|
||||
@@ -176,12 +231,96 @@ else
|
||||
wait_tcp_ready "$PG_HOST" "$PG_PORT" 5
|
||||
echo "TCP connectivity to PostgreSQL confirmed on port $PG_PORT"
|
||||
|
||||
# 执行迁移
|
||||
PYTHONPATH="$PWD/apps/api:$PWD" python3 -m alembic upgrade head
|
||||
echo "✅ Alembic migrations applied successfully"
|
||||
cleanup_db() {
|
||||
docker rm -f "$PG_CONTAINER" 2>/dev/null || true
|
||||
}
|
||||
fi
|
||||
|
||||
docker rm -f "$PG_CONTAINER" 2>/dev/null || true
|
||||
trap cleanup_db EXIT
|
||||
|
||||
echo ""
|
||||
|
||||
# ============================================================
|
||||
# 阶段2: upgrade head 升级验证
|
||||
# ============================================================
|
||||
|
||||
echo "⬆️ 阶段2: upgrade head 升级验证"
|
||||
echo ""
|
||||
|
||||
PYTHONPATH="$PWD/apps/api:$PWD" python3 -m alembic upgrade head
|
||||
echo "✅ upgrade head 通过"
|
||||
echo ""
|
||||
|
||||
# ============================================================
|
||||
# 阶段3: downgrade -1 回滚验证
|
||||
# ============================================================
|
||||
|
||||
echo "⬇️ 阶段3: downgrade -1 回滚验证"
|
||||
echo ""
|
||||
|
||||
# 获取当前head版本号
|
||||
HEAD_REV=$(PYTHONPATH="$PWD/apps/api:$PWD" python3 -m alembic current 2>&1 | awk '{print $1}' | head -1)
|
||||
echo "当前版本 (head): $HEAD_REV"
|
||||
|
||||
# 检查是否只有1个migration(baseline),downgrade -1会到base
|
||||
TOTAL_REVS=$(PYTHONPATH="$PWD/apps/api:$PWD" python3 -c "
|
||||
from alembic.config import Config
|
||||
from alembic.script import ScriptDirectory
|
||||
config = Config('alembic.ini')
|
||||
script = ScriptDirectory.from_config(config)
|
||||
print(len(list(script.walk_revisions())))
|
||||
")
|
||||
|
||||
echo "总 migration 数量: $TOTAL_REVS"
|
||||
|
||||
if [ "$TOTAL_REVS" -le 1 ]; then
|
||||
echo "⚠️ 只有1个migration,跳过 downgrade 回滚验证(没有可回滚的版本)"
|
||||
else
|
||||
echo "执行 downgrade -1..."
|
||||
PYTHONPATH="$PWD/apps/api:$PWD" python3 -m alembic downgrade -1
|
||||
echo "✅ downgrade -1 通过"
|
||||
|
||||
# 回滚后再升级回去,确保双向都通
|
||||
echo ""
|
||||
echo "重新 upgrade head 验证双向一致性..."
|
||||
PYTHONPATH="$PWD/apps/api:$PWD" python3 -m alembic upgrade head
|
||||
echo "✅ 重新 upgrade head 通过(双向验证完成)"
|
||||
fi
|
||||
|
||||
echo ""
|
||||
echo "=== CI Validate: Alembic迁移验证 通过 ✅ ==="
|
||||
|
||||
# ============================================================
|
||||
# 阶段4: alembic check - 检测未生成migration的model变更
|
||||
# ============================================================
|
||||
|
||||
echo "🔍 阶段4: 检查是否有未生成migration的model变更"
|
||||
echo ""
|
||||
|
||||
# alembic check: 没有待生成的migration时退出码0,有变更时退出码1
|
||||
# 这里只检测,不阻断(警告模式),因为有些场景model变更不需要migration
|
||||
set +e
|
||||
CHECK_OUTPUT=$(PYTHONPATH="$PWD/apps/api:$PWD" python3 -m alembic check 2>&1)
|
||||
CHECK_EXIT=$?
|
||||
set -e
|
||||
|
||||
if [ "$CHECK_EXIT" -eq 0 ]; then
|
||||
echo "✅ 没有检测到未生成migration的model变更"
|
||||
else
|
||||
if echo "$CHECK_OUTPUT" | grep -q "New upgrade operations detected"; then
|
||||
echo "⚠️ 检测到未生成migration的model变更!"
|
||||
echo ""
|
||||
echo "$CHECK_OUTPUT"
|
||||
echo ""
|
||||
echo "提示: 如果model变更是有意的且需要生成migration,请运行:"
|
||||
echo " alembic revision --autogenerate -m \"description\""
|
||||
echo "如果model变更不涉及数据库schema(如仅索引/约束重命名或纯业务逻辑),请确认后忽略此警告。"
|
||||
# 暂时不阻断,避免误报
|
||||
echo "(当前为警告模式,不阻断CI,后续稳定后可升级为阻断)"
|
||||
else
|
||||
echo "⚠️ alembic check 执行出错(非阻断)"
|
||||
echo "$CHECK_OUTPUT"
|
||||
fi
|
||||
fi
|
||||
|
||||
echo ""
|
||||
echo "=== CI Validate: Alembic迁移验证 全部通过 ✅ ==="
|
||||
|
||||
Executable
+90
@@ -0,0 +1,90 @@
|
||||
"""ASR 服务工厂单元测试."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import os
|
||||
|
||||
import pytest
|
||||
from services.asr_service_factory import get_asr_service, reset_asr_service_cache
|
||||
|
||||
|
||||
@pytest.fixture(autouse=True)
|
||||
def clean_env():
|
||||
"""每个测试前后清理环境变量和缓存."""
|
||||
# 保存原始值
|
||||
old = os.environ.get("ASR_PROVIDER")
|
||||
reset_asr_service_cache()
|
||||
yield
|
||||
# 恢复
|
||||
if old is not None:
|
||||
os.environ["ASR_PROVIDER"] = old
|
||||
elif "ASR_PROVIDER" in os.environ:
|
||||
del os.environ["ASR_PROVIDER"]
|
||||
reset_asr_service_cache()
|
||||
|
||||
|
||||
class TestGetAsrService:
|
||||
"""ASR服务工厂测试."""
|
||||
|
||||
def test_default_no_provider_returns_none(self):
|
||||
"""未配置ASR_PROVIDER时返回None."""
|
||||
if "ASR_PROVIDER" in os.environ:
|
||||
del os.environ["ASR_PROVIDER"]
|
||||
reset_asr_service_cache()
|
||||
result = get_asr_service()
|
||||
assert result is None
|
||||
|
||||
def test_empty_provider_returns_none(self):
|
||||
"""ASR_PROVIDER为空字符串时返回None."""
|
||||
os.environ["ASR_PROVIDER"] = ""
|
||||
reset_asr_service_cache()
|
||||
result = get_asr_service()
|
||||
assert result is None
|
||||
|
||||
def test_whitespace_provider_returns_none(self):
|
||||
"""ASR_PROVIDER为空白字符时返回None."""
|
||||
os.environ["ASR_PROVIDER"] = " "
|
||||
reset_asr_service_cache()
|
||||
result = get_asr_service()
|
||||
assert result is None
|
||||
|
||||
def test_mock_provider_returns_mock_service(self):
|
||||
"""mock provider返回MockASRService."""
|
||||
os.environ["ASR_PROVIDER"] = "mock"
|
||||
reset_asr_service_cache()
|
||||
result = get_asr_service()
|
||||
assert result is not None
|
||||
# 检查类型名称
|
||||
assert type(result).__name__ == "MockASRService"
|
||||
|
||||
def test_mock_provider_case_insensitive(self):
|
||||
"""provider大小写不敏感."""
|
||||
os.environ["ASR_PROVIDER"] = "MOCK"
|
||||
reset_asr_service_cache()
|
||||
result = get_asr_service()
|
||||
assert result is not None
|
||||
assert type(result).__name__ == "MockASRService"
|
||||
|
||||
def test_unknown_provider_returns_none(self):
|
||||
"""未知provider返回None(不阻断主流程)."""
|
||||
os.environ["ASR_PROVIDER"] = "unknown_provider_xyz"
|
||||
reset_asr_service_cache()
|
||||
result = get_asr_service()
|
||||
assert result is None
|
||||
|
||||
def test_singleton_caching(self):
|
||||
"""单例缓存有效,多次调用返回同一实例."""
|
||||
os.environ["ASR_PROVIDER"] = "mock"
|
||||
reset_asr_service_cache()
|
||||
s1 = get_asr_service()
|
||||
s2 = get_asr_service()
|
||||
assert s1 is s2
|
||||
|
||||
def test_reset_cache_clears_singleton(self):
|
||||
"""重置缓存后返回新实例."""
|
||||
os.environ["ASR_PROVIDER"] = "mock"
|
||||
reset_asr_service_cache()
|
||||
s1 = get_asr_service()
|
||||
reset_asr_service_cache()
|
||||
s2 = get_asr_service()
|
||||
assert s1 is not s2
|
||||
+94
-342
@@ -1,359 +1,111 @@
|
||||
"""BGM 混音单元测试.
|
||||
"""BGM混音单元测试 - 配置解析等纯逻辑."""
|
||||
|
||||
测试:
|
||||
- BGMConfig 配置解析与边界值
|
||||
- 预设 BGM 库查询
|
||||
- 纯 BGM 音频生成(端到端 ffmpeg)
|
||||
- BGM + 主音频混音(端到端 ffmpeg)
|
||||
- 淡入淡出效果
|
||||
- 音量边界(0 和 1)
|
||||
- sidechain 人声闪避
|
||||
"""
|
||||
|
||||
import sys
|
||||
import tempfile
|
||||
from pathlib import Path
|
||||
|
||||
sys.path.insert(0, str(Path(__file__).resolve().parents[2] / "apps" / "worker"))
|
||||
sys.path.insert(0, str(Path(__file__).resolve().parents[2] / "apps" / "api"))
|
||||
from __future__ import annotations
|
||||
|
||||
import pytest
|
||||
from video_processing.bgm_mixer import BGMConfig, build_bgm_only, mix_bgm_with_main, prepare_bgm_track
|
||||
from video_processing.render_audio import RenderContext
|
||||
|
||||
# ── Fixtures ──────────────────────────────────────────────────────────────────
|
||||
from video_processing.bgm_mixer import BGMConfig
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def work_dir(tmp_path):
|
||||
return tmp_path
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def ctx(work_dir):
|
||||
return RenderContext(work_dir=work_dir, plan_id="test_plan")
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def main_audio_path(work_dir):
|
||||
"""生成 10 秒测试主音频(正弦波模拟人声)。"""
|
||||
import subprocess
|
||||
|
||||
path = work_dir / "main.aac"
|
||||
# 生成 10 秒 440Hz 正弦波模拟主音频
|
||||
subprocess.run(
|
||||
[
|
||||
"ffmpeg",
|
||||
"-y",
|
||||
"-f",
|
||||
"lavfi",
|
||||
"-i",
|
||||
"sine=frequency=440:duration=10:sample_rate=44100",
|
||||
"-c:a",
|
||||
"aac",
|
||||
"-b:a",
|
||||
"128k",
|
||||
str(path),
|
||||
],
|
||||
capture_output=True,
|
||||
check=True,
|
||||
timeout=30,
|
||||
)
|
||||
return str(path)
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def bgm_audio_path(work_dir):
|
||||
"""生成 5 秒测试 BGM(更低频率模拟背景音乐)。"""
|
||||
import subprocess
|
||||
|
||||
path = work_dir / "bgm.aac"
|
||||
# 生成 5 秒 220Hz 正弦波模拟 BGM
|
||||
subprocess.run(
|
||||
[
|
||||
"ffmpeg",
|
||||
"-y",
|
||||
"-f",
|
||||
"lavfi",
|
||||
"-i",
|
||||
"sine=frequency=220:duration=5:sample_rate=44100",
|
||||
"-c:a",
|
||||
"aac",
|
||||
"-b:a",
|
||||
"128k",
|
||||
str(path),
|
||||
],
|
||||
capture_output=True,
|
||||
check=True,
|
||||
timeout=30,
|
||||
)
|
||||
return str(path)
|
||||
|
||||
|
||||
# ── BGMConfig 测试 ───────────────────────────────────────────────────────────
|
||||
|
||||
|
||||
class TestBGMConfig:
|
||||
"""BGMConfig 配置解析测试。"""
|
||||
class TestBGMConfigDefaults:
|
||||
"""BGMConfig 默认值测试."""
|
||||
|
||||
def test_default_values(self):
|
||||
cfg = BGMConfig(bgm_path="/tmp/bgm.mp3")
|
||||
assert cfg.volume == 0.3
|
||||
assert cfg.fade_in == 0.0
|
||||
assert cfg.fade_out == 0.0
|
||||
assert cfg.loop_enabled is True
|
||||
assert cfg.sidechain_enabled is False
|
||||
assert cfg.sidechain_ratio == 0.3
|
||||
|
||||
def test_from_config_dict(self):
|
||||
config_dict = {
|
||||
"enabled": True,
|
||||
"volume": 0.5,
|
||||
"fade_in": 2.0,
|
||||
"fade_out": 3.0,
|
||||
"loop_enabled": False,
|
||||
"sidechain_enabled": True,
|
||||
"sidechain_ratio": 0.5,
|
||||
}
|
||||
cfg = BGMConfig.from_config_dict("/bgm.mp3", config_dict)
|
||||
assert cfg.bgm_path == "/bgm.mp3"
|
||||
assert cfg.volume == 0.5
|
||||
assert cfg.fade_in == 2.0
|
||||
assert cfg.fade_out == 3.0
|
||||
assert cfg.loop_enabled is False
|
||||
assert cfg.sidechain_enabled is True
|
||||
assert cfg.sidechain_ratio == 0.5
|
||||
|
||||
def test_volume_clamped_by_config_schema(self):
|
||||
"""音量边界由 Pydantic Schema 在入口层保证,内部直接使用。"""
|
||||
from packages.domain.config_schemas import BGMConfig as BGMConfigSchema
|
||||
|
||||
# 边界值测试
|
||||
cfg = BGMConfigSchema(enabled=True, volume=0.0)
|
||||
assert cfg.volume == 0.0
|
||||
|
||||
cfg = BGMConfigSchema(enabled=True, volume=1.0)
|
||||
assert cfg.volume == 1.0
|
||||
|
||||
def test_fade_boundaries(self):
|
||||
from packages.domain.config_schemas import BGMConfig as BGMConfigSchema
|
||||
|
||||
# 0 是合法值
|
||||
cfg = BGMConfigSchema(fade_in=0, fade_out=0)
|
||||
assert cfg.fade_in == 0.0
|
||||
assert cfg.fade_out == 0.0
|
||||
"""默认值正确."""
|
||||
config = BGMConfig(bgm_path="/bgm.mp3")
|
||||
assert config.bgm_path == "/bgm.mp3"
|
||||
assert config.volume == 0.3
|
||||
assert config.fade_in == 0.0
|
||||
assert config.fade_out == 0.0
|
||||
assert config.loop_enabled is True
|
||||
assert config.sidechain_enabled is False
|
||||
assert config.sidechain_ratio == 0.3
|
||||
assert config.sidechain_attack == 0.02
|
||||
assert config.sidechain_release == 0.5
|
||||
assert config.sidechain_threshold == -25.0
|
||||
|
||||
|
||||
# ── 预设 BGM 库测试 ─────────────────────────────────────────────────────────
|
||||
class TestBGMConfigFromConfigDict:
|
||||
"""BGMConfig.from_config_dict 解析测试."""
|
||||
|
||||
def test_empty_dict_defaults(self):
|
||||
"""空字典用默认值."""
|
||||
config = BGMConfig.from_config_dict("/bgm.mp3", {})
|
||||
assert config.bgm_path == "/bgm.mp3"
|
||||
assert config.volume == 0.3
|
||||
assert config.loop_enabled is True
|
||||
assert config.sidechain_enabled is False
|
||||
|
||||
class TestPresetBGM:
|
||||
"""预设 BGM 库查询测试。"""
|
||||
def test_custom_volume(self):
|
||||
"""自定义音量."""
|
||||
config = BGMConfig.from_config_dict("/a.mp3", {"volume": 0.5})
|
||||
assert config.volume == 0.5
|
||||
|
||||
def test_total_count(self):
|
||||
from packages.domain.preset_bgm import PRESET_BGM_LIBRARY
|
||||
|
||||
assert len(PRESET_BGM_LIBRARY) >= 10
|
||||
|
||||
def test_get_preset_by_id(self):
|
||||
from packages.domain.preset_bgm import get_preset_bgm
|
||||
|
||||
bgm = get_preset_bgm("bgm_upbeat_001")
|
||||
assert bgm is not None
|
||||
assert bgm.name == "阳光清晨"
|
||||
assert bgm.style == "upbeat"
|
||||
|
||||
def test_get_preset_not_found(self):
|
||||
from packages.domain.preset_bgm import get_preset_bgm
|
||||
|
||||
assert get_preset_bgm("nonexistent") is None
|
||||
|
||||
def test_list_by_style(self):
|
||||
from packages.domain.preset_bgm import list_preset_bgm_by_style
|
||||
|
||||
upbeat = list_preset_bgm_by_style("upbeat")
|
||||
assert len(upbeat) >= 3
|
||||
assert all(b.style == "upbeat" for b in upbeat)
|
||||
|
||||
def test_search_by_keyword(self):
|
||||
from packages.domain.preset_bgm import search_preset_bgm
|
||||
|
||||
results = search_preset_bgm("钢琴")
|
||||
assert len(results) >= 2
|
||||
assert any("钢琴" in b.tags for b in results)
|
||||
|
||||
def test_all_presets_have_basic_fields(self):
|
||||
from packages.domain.preset_bgm import PRESET_BGM_LIBRARY
|
||||
|
||||
for bgm in PRESET_BGM_LIBRARY:
|
||||
assert bgm.id, f"{bgm.name} 缺少 id"
|
||||
assert bgm.name, "缺少 name"
|
||||
assert bgm.style, f"{bgm.name} 缺少 style"
|
||||
assert bgm.duration > 0, f"{bgm.name} 时长无效"
|
||||
|
||||
|
||||
# ── BGM 处理端到端测试 ──────────────────────────────────────────────────────
|
||||
|
||||
|
||||
class TestPrepareBGMTrack:
|
||||
"""prepare_bgm_track 端到端测试。"""
|
||||
|
||||
def test_bgm_without_loop_short_duration(self, ctx, bgm_audio_path):
|
||||
"""BGM 比目标时长短且不循环 → 截断到目标时长(但前面没有足够内容)。"""
|
||||
bgm = BGMConfig(bgm_path=bgm_audio_path, volume=0.5, loop_enabled=False)
|
||||
result = prepare_bgm_track(ctx, bgm, target_duration=3.0)
|
||||
|
||||
assert result.exists()
|
||||
assert result.stat().st_size > 0
|
||||
|
||||
def test_bgm_with_loop_longer_duration(self, ctx, bgm_audio_path):
|
||||
"""BGM 比目标时长短,循环铺满。"""
|
||||
bgm = BGMConfig(bgm_path=bgm_audio_path, volume=0.3, loop_enabled=True)
|
||||
# BGM 5 秒,目标 12 秒,需要循环 3 次
|
||||
result = prepare_bgm_track(ctx, bgm, target_duration=12.0)
|
||||
|
||||
assert result.exists()
|
||||
assert result.stat().st_size > 0
|
||||
|
||||
def test_bgm_fade_in_and_fade_out(self, ctx, bgm_audio_path):
|
||||
"""BGM 淡入淡出效果。"""
|
||||
bgm = BGMConfig(
|
||||
bgm_path=bgm_audio_path,
|
||||
volume=0.5,
|
||||
fade_in=1.0,
|
||||
fade_out=1.0,
|
||||
loop_enabled=False,
|
||||
)
|
||||
result = prepare_bgm_track(ctx, bgm, target_duration=4.0)
|
||||
|
||||
assert result.exists()
|
||||
assert result.stat().st_size > 0
|
||||
|
||||
def test_volume_zero(self, ctx, bgm_audio_path):
|
||||
"""音量为 0 时仍能正常处理。"""
|
||||
bgm = BGMConfig(bgm_path=bgm_audio_path, volume=0.0, loop_enabled=False)
|
||||
result = prepare_bgm_track(ctx, bgm, target_duration=3.0)
|
||||
|
||||
assert result.exists()
|
||||
assert result.stat().st_size > 0
|
||||
|
||||
def test_volume_one(self, ctx, bgm_audio_path):
|
||||
"""音量为 1(最大)时正常处理。"""
|
||||
bgm = BGMConfig(bgm_path=bgm_audio_path, volume=1.0, loop_enabled=False)
|
||||
result = prepare_bgm_track(ctx, bgm, target_duration=3.0)
|
||||
|
||||
assert result.exists()
|
||||
assert result.stat().st_size > 0
|
||||
|
||||
|
||||
class TestMixBGMMain:
|
||||
"""BGM + 主音频混音端到端测试。"""
|
||||
|
||||
def test_simple_mix(self, ctx, main_audio_path, bgm_audio_path):
|
||||
"""普通 amix 混音(无 sidechain)。"""
|
||||
bgm = BGMConfig(
|
||||
bgm_path=bgm_audio_path,
|
||||
volume=0.3,
|
||||
loop_enabled=True,
|
||||
sidechain_enabled=False,
|
||||
)
|
||||
result = mix_bgm_with_main(ctx, Path(main_audio_path), bgm, target_duration=8.0)
|
||||
|
||||
assert result.exists()
|
||||
assert result.stat().st_size > 0
|
||||
|
||||
def test_sidechain_mix(self, ctx, main_audio_path, bgm_audio_path):
|
||||
"""sidechain 人声闪避混音。"""
|
||||
bgm = BGMConfig(
|
||||
bgm_path=bgm_audio_path,
|
||||
volume=0.5,
|
||||
loop_enabled=True,
|
||||
sidechain_enabled=True,
|
||||
sidechain_ratio=0.3,
|
||||
sidechain_threshold=-25.0,
|
||||
sidechain_attack=0.02,
|
||||
sidechain_release=0.5,
|
||||
)
|
||||
result = mix_bgm_with_main(ctx, Path(main_audio_path), bgm, target_duration=8.0)
|
||||
|
||||
assert result.exists()
|
||||
assert result.stat().st_size > 0
|
||||
|
||||
def test_sidechain_max_ratio(self, ctx, main_audio_path, bgm_audio_path):
|
||||
"""sidechain 最大闪避比例。"""
|
||||
bgm = BGMConfig(
|
||||
bgm_path=bgm_audio_path,
|
||||
volume=0.5,
|
||||
loop_enabled=True,
|
||||
sidechain_enabled=True,
|
||||
sidechain_ratio=0.9, # 降低 90%
|
||||
)
|
||||
result = mix_bgm_with_main(ctx, Path(main_audio_path), bgm, target_duration=5.0)
|
||||
|
||||
assert result.exists()
|
||||
assert result.stat().st_size > 0
|
||||
|
||||
|
||||
class TestBuildBGMOnly:
|
||||
"""纯 BGM 模式测试。"""
|
||||
|
||||
def test_build_bgm_only(self, ctx, bgm_audio_path):
|
||||
"""只有 BGM、没有主音频时生成纯 BGM 音频。"""
|
||||
bgm = BGMConfig(
|
||||
bgm_path=bgm_audio_path,
|
||||
volume=0.3,
|
||||
fade_in=1.0,
|
||||
fade_out=1.0,
|
||||
loop_enabled=True,
|
||||
)
|
||||
result = build_bgm_only(ctx, bgm, target_duration=15.0)
|
||||
|
||||
assert result.exists()
|
||||
assert result.stat().st_size > 0
|
||||
|
||||
|
||||
# ── Config Schema 集成测试 ───────────────────────────────────────────────────
|
||||
|
||||
|
||||
class TestConfigSchemaIntegration:
|
||||
"""config schema 与渲染配置的集成测试。"""
|
||||
|
||||
def test_full_bgm_config(self):
|
||||
"""完整 BGM 配置能正确解析。"""
|
||||
from packages.domain.config_schemas import EditPlanConfigSchema, normalize_plan_config
|
||||
|
||||
config = normalize_plan_config(
|
||||
def test_fade_in_out(self):
|
||||
"""淡入淡出."""
|
||||
config = BGMConfig.from_config_dict(
|
||||
"/a.mp3",
|
||||
{
|
||||
"bgm": {
|
||||
"enabled": True,
|
||||
"source": "library",
|
||||
"asset_id": "bgm-asset-001",
|
||||
"volume": 0.4,
|
||||
"fade_in": 2.5,
|
||||
"fade_out": 3.0,
|
||||
"loop_enabled": True,
|
||||
"sidechain_enabled": True,
|
||||
"sidechain_ratio": 0.4,
|
||||
}
|
||||
}
|
||||
"fade_in": 2.0,
|
||||
"fade_out": 3.0,
|
||||
},
|
||||
)
|
||||
assert config.fade_in == 2.0
|
||||
assert config.fade_out == 3.0
|
||||
|
||||
bgm = config["bgm"]
|
||||
assert bgm["enabled"] is True
|
||||
assert bgm["volume"] == 0.4
|
||||
assert bgm["fade_in"] == 2.5
|
||||
assert bgm["fade_out"] == 3.0
|
||||
assert bgm["loop_enabled"] is True
|
||||
assert bgm["sidechain_enabled"] is True
|
||||
assert bgm["sidechain_ratio"] == 0.4
|
||||
# 默认值保留
|
||||
assert bgm["sidechain_attack"] == 0.02
|
||||
assert bgm["sidechain_release"] == 0.5
|
||||
assert bgm["sidechain_threshold"] == -25.0
|
||||
def test_loop_disabled(self):
|
||||
"""禁用循环."""
|
||||
config = BGMConfig.from_config_dict("/a.mp3", {"loop_enabled": False})
|
||||
assert config.loop_enabled is False
|
||||
|
||||
def test_bgm_disabled_by_default(self):
|
||||
"""默认 BGM 是关闭的。"""
|
||||
from packages.domain.config_schemas import normalize_plan_config
|
||||
def test_sidechain_enabled(self):
|
||||
"""启用人声闪避."""
|
||||
config = BGMConfig.from_config_dict("/a.mp3", {"sidechain_enabled": True})
|
||||
assert config.sidechain_enabled is True
|
||||
|
||||
config = normalize_plan_config({})
|
||||
assert config["bgm"]["enabled"] is False
|
||||
def test_sidechain_custom_params(self):
|
||||
"""闪避自定义参数."""
|
||||
config = BGMConfig.from_config_dict(
|
||||
"/a.mp3",
|
||||
{
|
||||
"sidechain_enabled": True,
|
||||
"sidechain_ratio": 0.5,
|
||||
"sidechain_attack": 0.05,
|
||||
"sidechain_release": 0.8,
|
||||
"sidechain_threshold": -30.0,
|
||||
},
|
||||
)
|
||||
assert config.sidechain_ratio == 0.5
|
||||
assert config.sidechain_attack == 0.05
|
||||
assert config.sidechain_release == 0.8
|
||||
assert config.sidechain_threshold == -30.0
|
||||
|
||||
def test_bgm_path_preserved(self):
|
||||
"""bgm_path保持不变."""
|
||||
config = BGMConfig.from_config_dict("/custom/path.mp3", {"volume": 0.5})
|
||||
assert config.bgm_path == "/custom/path.mp3"
|
||||
|
||||
def test_all_params_custom(self):
|
||||
"""所有参数自定义."""
|
||||
config = BGMConfig.from_config_dict(
|
||||
"/full.mp3",
|
||||
{
|
||||
"volume": 0.7,
|
||||
"fade_in": 1.5,
|
||||
"fade_out": 2.0,
|
||||
"loop_enabled": False,
|
||||
"sidechain_enabled": True,
|
||||
"sidechain_ratio": 0.4,
|
||||
"sidechain_attack": 0.03,
|
||||
"sidechain_release": 0.6,
|
||||
"sidechain_threshold": -20.0,
|
||||
},
|
||||
)
|
||||
assert config.volume == 0.7
|
||||
assert config.fade_in == 1.5
|
||||
assert config.fade_out == 2.0
|
||||
assert config.loop_enabled is False
|
||||
assert config.sidechain_enabled is True
|
||||
assert config.sidechain_ratio == 0.4
|
||||
assert config.sidechain_attack == 0.03
|
||||
assert config.sidechain_release == 0.6
|
||||
assert config.sidechain_threshold == -20.0
|
||||
|
||||
Executable
+210
@@ -0,0 +1,210 @@
|
||||
"""绿幕抠像引擎单元测试 - 配置解析等纯逻辑."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import pytest
|
||||
from video_processing.chroma_key_engine import (
|
||||
CHROMA_KEY_PRESETS,
|
||||
ChromaKeyConfig,
|
||||
)
|
||||
|
||||
|
||||
class TestChromaKeyConfigDefaults:
|
||||
"""默认配置测试."""
|
||||
|
||||
def test_default_values(self):
|
||||
"""默认值正确."""
|
||||
config = ChromaKeyConfig()
|
||||
assert config.enabled is False
|
||||
assert config.key_color == "#00FF00"
|
||||
assert config.similarity == 0.3
|
||||
assert config.blend == 0.1
|
||||
assert config.spill_suppress == 0.0
|
||||
|
||||
|
||||
class TestChromaKeyConfigFromDict:
|
||||
"""from_dict 配置解析测试."""
|
||||
|
||||
def test_none_returns_disabled(self):
|
||||
"""None 返回禁用配置."""
|
||||
config = ChromaKeyConfig.from_dict(None)
|
||||
assert config.enabled is False
|
||||
|
||||
def test_empty_dict_returns_disabled(self):
|
||||
"""空字典返回禁用."""
|
||||
config = ChromaKeyConfig.from_dict({})
|
||||
assert config.enabled is False
|
||||
|
||||
def test_disabled_returns_disabled(self):
|
||||
"""enabled=False 返回禁用."""
|
||||
config = ChromaKeyConfig.from_dict({"enabled": False})
|
||||
assert config.enabled is False
|
||||
|
||||
def test_enabled_default_values(self):
|
||||
"""启用时使用默认参数."""
|
||||
config = ChromaKeyConfig.from_dict({"enabled": True})
|
||||
assert config.enabled is True
|
||||
assert config.key_color == "#00FF00"
|
||||
assert config.similarity == 0.3
|
||||
assert config.blend == 0.1
|
||||
assert config.spill_suppress == 0.0
|
||||
|
||||
def test_custom_key_color(self):
|
||||
"""自定义抠像颜色."""
|
||||
config = ChromaKeyConfig.from_dict(
|
||||
{
|
||||
"enabled": True,
|
||||
"key_color": "#0000FF",
|
||||
}
|
||||
)
|
||||
assert config.key_color == "#0000FF"
|
||||
|
||||
def test_similarity_parsed(self):
|
||||
"""相似度解析."""
|
||||
config = ChromaKeyConfig.from_dict(
|
||||
{
|
||||
"enabled": True,
|
||||
"similarity": 0.5,
|
||||
}
|
||||
)
|
||||
assert config.similarity == 0.5
|
||||
|
||||
def test_similarity_clamped_min(self):
|
||||
"""相似度下限钳制."""
|
||||
config = ChromaKeyConfig.from_dict(
|
||||
{
|
||||
"enabled": True,
|
||||
"similarity": 0.001,
|
||||
}
|
||||
)
|
||||
assert config.similarity == 0.01
|
||||
|
||||
def test_similarity_clamped_max(self):
|
||||
"""相似度上限钳制."""
|
||||
config = ChromaKeyConfig.from_dict(
|
||||
{
|
||||
"enabled": True,
|
||||
"similarity": 2.0,
|
||||
}
|
||||
)
|
||||
assert config.similarity == 1.0
|
||||
|
||||
def test_blend_clamped_min(self):
|
||||
"""混合度下限钳制."""
|
||||
config = ChromaKeyConfig.from_dict(
|
||||
{
|
||||
"enabled": True,
|
||||
"blend": -0.5,
|
||||
}
|
||||
)
|
||||
assert config.blend == 0.0
|
||||
|
||||
def test_blend_clamped_max(self):
|
||||
"""混合度上限钳制."""
|
||||
config = ChromaKeyConfig.from_dict(
|
||||
{
|
||||
"enabled": True,
|
||||
"blend": 1.5,
|
||||
}
|
||||
)
|
||||
assert config.blend == 1.0
|
||||
|
||||
def test_spill_suppress_clamped(self):
|
||||
"""溢色抑制钳制."""
|
||||
config = ChromaKeyConfig.from_dict(
|
||||
{
|
||||
"enabled": True,
|
||||
"spill_suppress": 2.0,
|
||||
}
|
||||
)
|
||||
assert config.spill_suppress == 1.0
|
||||
|
||||
def test_invalid_similarity_falls_back(self):
|
||||
"""无效相似度回退到默认."""
|
||||
config = ChromaKeyConfig.from_dict(
|
||||
{
|
||||
"enabled": True,
|
||||
"similarity": "not_a_number",
|
||||
}
|
||||
)
|
||||
assert config.similarity == 0.3
|
||||
|
||||
def test_invalid_blend_falls_back(self):
|
||||
"""无效混合度回退."""
|
||||
config = ChromaKeyConfig.from_dict(
|
||||
{
|
||||
"enabled": True,
|
||||
"blend": "high",
|
||||
}
|
||||
)
|
||||
assert config.blend == 0.1
|
||||
|
||||
def test_key_color_stripped(self):
|
||||
"""颜色值去除首尾空格."""
|
||||
config = ChromaKeyConfig.from_dict(
|
||||
{
|
||||
"enabled": True,
|
||||
"key_color": " #FF0000 ",
|
||||
}
|
||||
)
|
||||
assert config.key_color == "#FF0000"
|
||||
|
||||
def test_all_params_custom(self):
|
||||
"""所有参数自定义."""
|
||||
config = ChromaKeyConfig.from_dict(
|
||||
{
|
||||
"enabled": True,
|
||||
"key_color": "#0000FF",
|
||||
"similarity": 0.45,
|
||||
"blend": 0.15,
|
||||
"spill_suppress": 0.6,
|
||||
}
|
||||
)
|
||||
assert config.enabled is True
|
||||
assert config.key_color == "#0000FF"
|
||||
assert config.similarity == 0.45
|
||||
assert config.blend == 0.15
|
||||
assert config.spill_suppress == 0.6
|
||||
|
||||
|
||||
class TestHasEffect:
|
||||
"""has_effect 方法测试."""
|
||||
|
||||
def test_disabled_no_effect(self):
|
||||
"""禁用时无效果."""
|
||||
config = ChromaKeyConfig(enabled=False)
|
||||
assert config.has_effect() is False
|
||||
|
||||
def test_enabled_with_similarity_has_effect(self):
|
||||
"""启用且有相似度时有效果."""
|
||||
config = ChromaKeyConfig(enabled=True, similarity=0.3)
|
||||
assert config.has_effect() is True
|
||||
|
||||
def test_zero_similarity_no_effect(self):
|
||||
"""相似度为0时无效果."""
|
||||
config = ChromaKeyConfig(enabled=True, similarity=0.0)
|
||||
assert config.has_effect() is False
|
||||
|
||||
|
||||
class TestChromaKeyPresets:
|
||||
"""预设配置测试."""
|
||||
|
||||
def test_five_presets(self):
|
||||
"""5个预设."""
|
||||
assert len(CHROMA_KEY_PRESETS) == 5
|
||||
|
||||
def test_preset_names(self):
|
||||
"""预设名称正确."""
|
||||
assert "green_screen" in CHROMA_KEY_PRESETS
|
||||
assert "blue_screen" in CHROMA_KEY_PRESETS
|
||||
assert "red_screen" in CHROMA_KEY_PRESETS
|
||||
assert "precise_green" in CHROMA_KEY_PRESETS
|
||||
assert "soft_green" in CHROMA_KEY_PRESETS
|
||||
|
||||
def test_presets_have_required_keys(self):
|
||||
"""每个预设包含必要字段."""
|
||||
for name, preset in CHROMA_KEY_PRESETS.items():
|
||||
assert "key_color" in preset, f"{name} missing key_color"
|
||||
assert "similarity" in preset, f"{name} missing similarity"
|
||||
assert "blend" in preset, f"{name} missing blend"
|
||||
assert "spill_suppress" in preset, f"{name} missing spill_suppress"
|
||||
@@ -1,4 +1,4 @@
|
||||
"""滤镜调色引擎单元测试."""
|
||||
"""调色引擎单元测试 - 配置解析等纯逻辑."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
@@ -6,92 +6,19 @@ import pytest
|
||||
from video_processing.color_grade_engine import (
|
||||
DEFAULT_PARAMS,
|
||||
PARAM_RANGES,
|
||||
PRESET_BW,
|
||||
PRESET_CINEMA,
|
||||
PRESET_COOL,
|
||||
PRESET_DISPLAY_NAMES,
|
||||
PRESET_FILM,
|
||||
PRESET_FRESH,
|
||||
PRESET_JAPANESE,
|
||||
PRESET_PARAMS,
|
||||
PRESET_VINTAGE,
|
||||
PRESET_WARM,
|
||||
VALID_PRESETS,
|
||||
ColorGradeConfig,
|
||||
ColorGradeEngine,
|
||||
get_preset_names,
|
||||
get_preset_params,
|
||||
)
|
||||
|
||||
# ── 预设常量测试 ──────────────────────────────────────────────────────────────
|
||||
|
||||
class TestColorGradeConfigDefaults:
|
||||
"""默认配置测试."""
|
||||
|
||||
class TestPresetConstants:
|
||||
"""预设常量完整性测试."""
|
||||
|
||||
def test_eight_presets_defined(self):
|
||||
"""应该有8种预设."""
|
||||
assert len(PRESET_PARAMS) == 8
|
||||
assert len(PRESET_DISPLAY_NAMES) == 8
|
||||
|
||||
def test_all_presets_have_display_names(self):
|
||||
"""每个预设都应该有中文显示名."""
|
||||
for key in PRESET_PARAMS:
|
||||
assert key in PRESET_DISPLAY_NAMES
|
||||
assert PRESET_DISPLAY_NAMES[key] # 非空
|
||||
|
||||
def test_preset_params_have_all_keys(self):
|
||||
"""每个预设应该包含所有5个参数."""
|
||||
required_keys = {"brightness", "contrast", "saturation", "temperature", "hue"}
|
||||
for key, params in PRESET_PARAMS.items():
|
||||
assert required_keys.issubset(params.keys()), f"预设 {key} 缺少参数"
|
||||
|
||||
def test_preset_params_in_valid_range(self):
|
||||
"""所有预设参数应该在合法范围内."""
|
||||
for preset_name, params in PRESET_PARAMS.items():
|
||||
for param_name, value in params.items():
|
||||
min_val, max_val = PARAM_RANGES[param_name]
|
||||
assert (
|
||||
min_val <= value <= max_val
|
||||
), f"预设 {preset_name} 的 {param_name}={value} 超出范围 [{min_val}, {max_val}]"
|
||||
|
||||
def test_black_white_has_zero_saturation(self):
|
||||
"""黑白预设饱和度应该为0."""
|
||||
assert PRESET_PARAMS[PRESET_BW]["saturation"] == 0
|
||||
|
||||
def test_warm_preset_has_positive_temperature(self):
|
||||
"""暖色预设色温应该为正."""
|
||||
assert PRESET_PARAMS[PRESET_WARM]["temperature"] > 0
|
||||
|
||||
def test_cool_preset_has_negative_temperature(self):
|
||||
"""冷色预设色温应该为负."""
|
||||
assert PRESET_PARAMS[PRESET_COOL]["temperature"] < 0
|
||||
|
||||
|
||||
# ── ColorGradeConfig.from_dict 测试 ───────────────────────────────────────────
|
||||
|
||||
|
||||
class TestColorGradeConfigFromDict:
|
||||
"""配置字典解析测试."""
|
||||
|
||||
def test_none_config(self):
|
||||
"""None返回disabled."""
|
||||
config = ColorGradeConfig.from_dict(None)
|
||||
assert not config.enabled
|
||||
|
||||
def test_empty_dict(self):
|
||||
"""空字典返回disabled."""
|
||||
config = ColorGradeConfig.from_dict({})
|
||||
assert not config.enabled
|
||||
|
||||
def test_enabled_false(self):
|
||||
"""enabled=False返回disabled."""
|
||||
config = ColorGradeConfig.from_dict({"enabled": False})
|
||||
assert not config.enabled
|
||||
|
||||
def test_enabled_only(self):
|
||||
"""只开enabled,无预设无自定义参数."""
|
||||
config = ColorGradeConfig.from_dict({"enabled": True})
|
||||
assert config.enabled
|
||||
def test_default_values(self):
|
||||
"""默认值正确."""
|
||||
config = ColorGradeConfig()
|
||||
assert config.enabled is False
|
||||
assert config.preset == ""
|
||||
assert config.brightness is None
|
||||
assert config.contrast is None
|
||||
@@ -99,50 +26,83 @@ class TestColorGradeConfigFromDict:
|
||||
assert config.temperature is None
|
||||
assert config.hue is None
|
||||
|
||||
|
||||
class TestColorGradeConfigFromDict:
|
||||
"""from_dict 配置解析测试."""
|
||||
|
||||
def test_none_returns_disabled(self):
|
||||
"""None 返回禁用配置."""
|
||||
config = ColorGradeConfig.from_dict(None)
|
||||
assert config.enabled is False
|
||||
|
||||
def test_empty_dict_returns_disabled(self):
|
||||
"""空字典返回禁用."""
|
||||
config = ColorGradeConfig.from_dict({})
|
||||
assert config.enabled is False
|
||||
|
||||
def test_disabled_returns_disabled(self):
|
||||
"""enabled=False 返回禁用."""
|
||||
config = ColorGradeConfig.from_dict({"enabled": False})
|
||||
assert config.enabled is False
|
||||
|
||||
def test_enabled_no_params(self):
|
||||
"""启用但无自定义参数."""
|
||||
config = ColorGradeConfig.from_dict({"enabled": True})
|
||||
assert config.enabled is True
|
||||
assert config.preset == ""
|
||||
assert config.brightness is None
|
||||
|
||||
def test_with_preset(self):
|
||||
"""指定预设."""
|
||||
config = ColorGradeConfig.from_dict({"enabled": True, "preset": PRESET_FRESH})
|
||||
assert config.enabled
|
||||
assert config.preset == PRESET_FRESH
|
||||
config = ColorGradeConfig.from_dict(
|
||||
{
|
||||
"enabled": True,
|
||||
"preset": "fresh",
|
||||
}
|
||||
)
|
||||
assert config.enabled is True
|
||||
assert config.preset == "fresh"
|
||||
|
||||
def test_invalid_preset_ignored(self):
|
||||
"""无效预设名应该被忽略."""
|
||||
config = ColorGradeConfig.from_dict({"enabled": True, "preset": "invalid_preset"})
|
||||
assert config.preset == "" # 被清空
|
||||
"""无效预设被忽略."""
|
||||
config = ColorGradeConfig.from_dict(
|
||||
{
|
||||
"enabled": True,
|
||||
"preset": "unknown_preset",
|
||||
}
|
||||
)
|
||||
assert config.preset == ""
|
||||
|
||||
def test_with_custom_params(self):
|
||||
"""自定义参数覆盖."""
|
||||
def test_custom_brightness(self):
|
||||
"""自定义亮度."""
|
||||
config = ColorGradeConfig.from_dict(
|
||||
{
|
||||
"enabled": True,
|
||||
"brightness": 20,
|
||||
"contrast": -10,
|
||||
"saturation": 150,
|
||||
"temperature": 25,
|
||||
"hue": 30,
|
||||
}
|
||||
)
|
||||
assert config.enabled
|
||||
assert config.brightness == 20
|
||||
assert config.contrast == -10
|
||||
assert config.saturation == 150
|
||||
assert config.temperature == 25
|
||||
assert config.hue == 30
|
||||
assert config.brightness == 20.0
|
||||
|
||||
def test_string_numeric_values(self):
|
||||
"""字符串形式的数字应该能解析."""
|
||||
def test_custom_all_params(self):
|
||||
"""所有参数自定义."""
|
||||
config = ColorGradeConfig.from_dict(
|
||||
{
|
||||
"enabled": True,
|
||||
"brightness": "20.5",
|
||||
"saturation": "150",
|
||||
"brightness": 10,
|
||||
"contrast": 15,
|
||||
"saturation": 120,
|
||||
"temperature": -5,
|
||||
"hue": 10,
|
||||
}
|
||||
)
|
||||
assert config.brightness == 20.5
|
||||
assert config.saturation == 150.0
|
||||
assert config.brightness == 10.0
|
||||
assert config.contrast == 15.0
|
||||
assert config.saturation == 120.0
|
||||
assert config.temperature == -5.0
|
||||
assert config.hue == 10.0
|
||||
|
||||
def test_invalid_value_returns_none(self):
|
||||
"""无效值应该返回None(不覆盖)."""
|
||||
def test_invalid_param_value_returns_none(self):
|
||||
"""无效参数值返回None(不覆盖)."""
|
||||
config = ColorGradeConfig.from_dict(
|
||||
{
|
||||
"enabled": True,
|
||||
@@ -151,422 +111,151 @@ class TestColorGradeConfigFromDict:
|
||||
)
|
||||
assert config.brightness is None
|
||||
|
||||
def test_null_param_returns_none(self):
|
||||
"""null参数值返回None."""
|
||||
config = ColorGradeConfig.from_dict(
|
||||
{
|
||||
"enabled": True,
|
||||
"contrast": None,
|
||||
}
|
||||
)
|
||||
assert config.contrast is None
|
||||
|
||||
# ── ColorGradeConfig.resolve_params 测试 ──────────────────────────────────────
|
||||
def test_preset_with_custom_override(self):
|
||||
"""预设 + 自定义覆盖."""
|
||||
config = ColorGradeConfig.from_dict(
|
||||
{
|
||||
"enabled": True,
|
||||
"preset": "vintage",
|
||||
"brightness": 5,
|
||||
}
|
||||
)
|
||||
assert config.preset == "vintage"
|
||||
assert config.brightness == 5.0
|
||||
|
||||
|
||||
class TestResolveParams:
|
||||
"""参数解析与边界钳制测试."""
|
||||
"""resolve_params 参数解析测试."""
|
||||
|
||||
def test_default_params_when_empty(self):
|
||||
"""无预设无自定义时返回默认值."""
|
||||
config = ColorGradeConfig(enabled=True)
|
||||
def test_disabled_returns_defaults(self):
|
||||
"""禁用配置也返回默认参数."""
|
||||
config = ColorGradeConfig(enabled=False)
|
||||
params = config.resolve_params()
|
||||
for key, val in DEFAULT_PARAMS.items():
|
||||
assert params[key] == val
|
||||
|
||||
def test_preset_params_applied(self):
|
||||
"""预设参数应该被应用."""
|
||||
config = ColorGradeConfig(enabled=True, preset=PRESET_FRESH)
|
||||
def test_no_preset_no_custom_returns_defaults(self):
|
||||
"""无预设无自定义返回默认值."""
|
||||
config = ColorGradeConfig(enabled=True)
|
||||
params = config.resolve_params()
|
||||
preset = PRESET_PARAMS[PRESET_FRESH]
|
||||
for key, val in preset.items():
|
||||
assert params[key] == val
|
||||
for key, val in DEFAULT_PARAMS.items():
|
||||
assert abs(params[key] - val) < 0.001
|
||||
|
||||
def test_preset_applies_params(self):
|
||||
"""预设应用参数."""
|
||||
config = ColorGradeConfig(enabled=True, preset="fresh")
|
||||
params = config.resolve_params()
|
||||
# 清新预设亮度=8
|
||||
assert params["brightness"] == 8
|
||||
assert params["saturation"] == 120
|
||||
|
||||
def test_custom_overrides_preset(self):
|
||||
"""自定义参数应该覆盖预设值."""
|
||||
"""自定义参数覆盖预设."""
|
||||
config = ColorGradeConfig(
|
||||
enabled=True,
|
||||
preset=PRESET_FRESH,
|
||||
preset="fresh",
|
||||
brightness=50, # 覆盖预设的8
|
||||
)
|
||||
params = config.resolve_params()
|
||||
assert params["brightness"] == 50
|
||||
# 其他参数还是预设值
|
||||
assert params["contrast"] == PRESET_PARAMS[PRESET_FRESH]["contrast"]
|
||||
# 其他参数仍用预设值
|
||||
assert params["saturation"] == 120
|
||||
|
||||
def test_clamp_brightness_high(self):
|
||||
"""亮度超过上限应该被钳制."""
|
||||
def test_brightness_clamped(self):
|
||||
"""亮度边界钳制."""
|
||||
config = ColorGradeConfig(enabled=True, brightness=200)
|
||||
params = config.resolve_params()
|
||||
assert params["brightness"] == 100
|
||||
assert params["brightness"] == 100.0
|
||||
|
||||
def test_clamp_brightness_low(self):
|
||||
"""亮度低于下限应该被钳制."""
|
||||
config = ColorGradeConfig(enabled=True, brightness=-200)
|
||||
def test_saturation_clamped_low(self):
|
||||
"""饱和度下限钳制."""
|
||||
config = ColorGradeConfig(enabled=True, saturation=-10)
|
||||
params = config.resolve_params()
|
||||
assert params["brightness"] == -100
|
||||
assert params["saturation"] == 0.0
|
||||
|
||||
def test_clamp_saturation_low(self):
|
||||
"""饱和度低于0应该被钳制到0."""
|
||||
config = ColorGradeConfig(enabled=True, saturation=-50)
|
||||
params = config.resolve_params()
|
||||
assert params["saturation"] == 0
|
||||
|
||||
def test_clamp_saturation_high(self):
|
||||
"""饱和度超过200应该被钳制."""
|
||||
def test_saturation_clamped_high(self):
|
||||
"""饱和度上限钳制."""
|
||||
config = ColorGradeConfig(enabled=True, saturation=300)
|
||||
params = config.resolve_params()
|
||||
assert params["saturation"] == 200
|
||||
assert params["saturation"] == 200.0
|
||||
|
||||
def test_clamp_hue_high(self):
|
||||
"""色调超过180应该被钳制."""
|
||||
config = ColorGradeConfig(enabled=True, hue=270)
|
||||
def test_hue_clamped(self):
|
||||
"""色调边界钳制."""
|
||||
config = ColorGradeConfig(enabled=True, hue=200)
|
||||
params = config.resolve_params()
|
||||
assert params["hue"] == 180
|
||||
assert params["hue"] == 180.0
|
||||
|
||||
def test_clamp_hue_low(self):
|
||||
"""色调低于-180应该被钳制."""
|
||||
config = ColorGradeConfig(enabled=True, hue=-270)
|
||||
def test_hue_negative_clamped(self):
|
||||
"""负色调边界钳制."""
|
||||
config = ColorGradeConfig(enabled=True, hue=-200)
|
||||
params = config.resolve_params()
|
||||
assert params["hue"] == -180
|
||||
assert params["hue"] == -180.0
|
||||
|
||||
def test_clamp_contrast(self):
|
||||
"""对比度越界应该被钳制."""
|
||||
config = ColorGradeConfig(enabled=True, contrast=150)
|
||||
def test_returns_all_five_params(self):
|
||||
"""返回所有5个参数."""
|
||||
config = ColorGradeConfig(enabled=True)
|
||||
params = config.resolve_params()
|
||||
assert params["contrast"] == 100
|
||||
|
||||
config2 = ColorGradeConfig(enabled=True, contrast=-150)
|
||||
params2 = config2.resolve_params()
|
||||
assert params2["contrast"] == -100
|
||||
|
||||
def test_clamp_temperature(self):
|
||||
"""色温越界应该被钳制."""
|
||||
config = ColorGradeConfig(enabled=True, temperature=150)
|
||||
params = config.resolve_params()
|
||||
assert params["temperature"] == 100
|
||||
|
||||
def test_preset_with_clamping(self):
|
||||
"""预设+自定义覆盖,自定义值超范围仍需钳制."""
|
||||
config = ColorGradeConfig(
|
||||
enabled=True,
|
||||
preset=PRESET_FRESH,
|
||||
brightness=999, # 超范围
|
||||
)
|
||||
params = config.resolve_params()
|
||||
assert params["brightness"] == 100 # 被钳制
|
||||
|
||||
|
||||
# ── ColorGradeConfig.has_effect 测试 ──────────────────────────────────────────
|
||||
assert set(params.keys()) == {"brightness", "contrast", "saturation", "temperature", "hue"}
|
||||
|
||||
|
||||
class TestHasEffect:
|
||||
"""是否有实际效果判断测试."""
|
||||
"""has_effect 方法测试."""
|
||||
|
||||
def test_disabled_has_no_effect(self):
|
||||
"""disabled的配置has_effect应该返回False."""
|
||||
config = ColorGradeConfig(enabled=False)
|
||||
assert not config.has_effect()
|
||||
|
||||
def test_default_params_no_effect(self):
|
||||
"""所有参数都是默认值时应该返回False."""
|
||||
def test_default_no_effect(self):
|
||||
"""默认配置无效果."""
|
||||
config = ColorGradeConfig(enabled=True)
|
||||
assert not config.has_effect()
|
||||
assert config.has_effect() is False
|
||||
|
||||
def test_brightness_change_has_effect(self):
|
||||
"""亮度变化应该有效果."""
|
||||
def test_with_preset_has_effect(self):
|
||||
"""有预设时有效果."""
|
||||
config = ColorGradeConfig(enabled=True, preset="cinema")
|
||||
assert config.has_effect() is True
|
||||
|
||||
def test_custom_brightness_has_effect(self):
|
||||
"""自定义亮度有效果."""
|
||||
config = ColorGradeConfig(enabled=True, brightness=10)
|
||||
assert config.has_effect()
|
||||
assert config.has_effect() is True
|
||||
|
||||
def test_saturation_100_no_effect(self):
|
||||
"""饱和度100是默认值,无效果."""
|
||||
config = ColorGradeConfig(enabled=True, saturation=100)
|
||||
assert not config.has_effect()
|
||||
def test_disabled_still_checks_params(self):
|
||||
"""禁用也根据参数判断(结果仍可能有效果但不启用)."""
|
||||
# has_effect 只看参数,不看 enabled
|
||||
config = ColorGradeConfig(enabled=False, preset="warm")
|
||||
assert config.has_effect() is True
|
||||
|
||||
def test_saturation_not_100_has_effect(self):
|
||||
"""饱和度不等于100有效果."""
|
||||
config = ColorGradeConfig(enabled=True, saturation=99)
|
||||
assert config.has_effect()
|
||||
|
||||
def test_preset_has_effect(self):
|
||||
"""预设通常有效果."""
|
||||
for preset in PRESET_PARAMS:
|
||||
config = ColorGradeConfig(enabled=True, preset=preset)
|
||||
assert config.has_effect(), f"预设 {preset} 应该有效果"
|
||||
|
||||
def test_custom_zero_override_no_effect(self):
|
||||
"""用预设但所有自定义值都设为默认值抵消 → 应该has_effect看实际值."""
|
||||
# 黑白预设饱和度=0,如果手动覆盖饱和度=100、其他都=默认值,则可能无效果
|
||||
config = ColorGradeConfig(
|
||||
enabled=True,
|
||||
preset=PRESET_BW,
|
||||
brightness=0,
|
||||
contrast=0,
|
||||
saturation=100,
|
||||
temperature=0,
|
||||
hue=0,
|
||||
)
|
||||
assert not config.has_effect()
|
||||
def test_black_white_preset_has_effect(self):
|
||||
"""黑白预设(饱和度=0)有效果."""
|
||||
config = ColorGradeConfig(enabled=True, preset="black_white")
|
||||
assert config.has_effect() is True
|
||||
|
||||
|
||||
# ── ColorGradeEngine 参数映射测试 ─────────────────────────────────────────────
|
||||
class TestPresets:
|
||||
"""预设常量测试."""
|
||||
|
||||
def test_eight_valid_presets(self):
|
||||
"""8个有效预设."""
|
||||
assert len(VALID_PRESETS) == 8
|
||||
|
||||
class TestParameterMapping:
|
||||
"""FFmpeg参数映射测试."""
|
||||
def test_preset_params_match_valid(self):
|
||||
"""所有预设都在有效列表中."""
|
||||
for name in PRESET_PARAMS:
|
||||
assert name in VALID_PRESETS
|
||||
|
||||
def test_brightness_mapping_zero(self):
|
||||
"""亮度0 → 0.0."""
|
||||
assert ColorGradeEngine._map_brightness(0) == 0.0
|
||||
def test_each_preset_has_all_params(self):
|
||||
"""每个预设包含所有5个参数."""
|
||||
for name, params in PRESET_PARAMS.items():
|
||||
for key in ["brightness", "contrast", "saturation", "temperature", "hue"]:
|
||||
assert key in params, f"{name} missing {key}"
|
||||
|
||||
def test_brightness_mapping_max(self):
|
||||
"""亮度100 → 1.0."""
|
||||
assert ColorGradeEngine._map_brightness(100) == 1.0
|
||||
|
||||
def test_brightness_mapping_min(self):
|
||||
"""亮度-100 → -1.0."""
|
||||
assert ColorGradeEngine._map_brightness(-100) == -1.0
|
||||
|
||||
def test_contrast_mapping_zero(self):
|
||||
"""对比度0 → 1.0(原始)."""
|
||||
assert ColorGradeEngine._map_contrast(0) == 1.0
|
||||
|
||||
def test_contrast_mapping_positive(self):
|
||||
"""正对比度应该 > 1.0."""
|
||||
assert ColorGradeEngine._map_contrast(50) == 1.5
|
||||
assert ColorGradeEngine._map_contrast(100) == 2.0
|
||||
|
||||
def test_contrast_mapping_negative(self):
|
||||
"""负对比度应该 < 1.0."""
|
||||
assert ColorGradeEngine._map_contrast(-50) == 0.5
|
||||
assert ColorGradeEngine._map_contrast(-100) == 0.0
|
||||
|
||||
def test_saturation_mapping_default(self):
|
||||
"""饱和度100 → 1.0."""
|
||||
assert ColorGradeEngine._map_saturation(100) == 1.0
|
||||
|
||||
def test_saturation_mapping_zero(self):
|
||||
"""饱和度0 → 0.0(黑白)."""
|
||||
assert ColorGradeEngine._map_saturation(0) == 0.0
|
||||
|
||||
def test_saturation_mapping_double(self):
|
||||
"""饱和度200 → 2.0."""
|
||||
assert ColorGradeEngine._map_saturation(200) == 2.0
|
||||
|
||||
def test_temperature_warm(self):
|
||||
"""暖色温应该红+蓝-."""
|
||||
red, green, blue = ColorGradeEngine._map_temperature(100)
|
||||
assert red > 0
|
||||
assert blue < 0
|
||||
|
||||
def test_temperature_cool(self):
|
||||
"""冷色温应该红-蓝+."""
|
||||
red, green, blue = ColorGradeEngine._map_temperature(-100)
|
||||
assert red < 0
|
||||
assert blue > 0
|
||||
|
||||
def test_temperature_zero(self):
|
||||
"""色温0应该全0."""
|
||||
red, green, blue = ColorGradeEngine._map_temperature(0)
|
||||
assert red == 0
|
||||
assert green == 0
|
||||
assert blue == 0
|
||||
|
||||
def test_hue_mapping_passthrough(self):
|
||||
"""色调直接透传."""
|
||||
assert ColorGradeEngine._map_hue(0) == 0
|
||||
assert ColorGradeEngine._map_hue(90) == 90
|
||||
assert ColorGradeEngine._map_hue(-45) == -45
|
||||
|
||||
|
||||
# ── ColorGradeEngine.build_filter 测试 ────────────────────────────────────────
|
||||
|
||||
|
||||
class TestBuildFilter:
|
||||
"""滤镜字符串构建测试."""
|
||||
|
||||
def test_disabled_returns_empty(self):
|
||||
"""disabled配置返回空."""
|
||||
config = ColorGradeConfig(enabled=False)
|
||||
result = ColorGradeEngine.build_filter(config)
|
||||
assert result == ""
|
||||
|
||||
def test_no_effect_returns_empty(self):
|
||||
"""无效果的配置返回空."""
|
||||
config = ColorGradeConfig(enabled=True)
|
||||
result = ColorGradeEngine.build_filter(config)
|
||||
assert result == ""
|
||||
|
||||
def test_brightness_only(self):
|
||||
"""只有亮度调整."""
|
||||
config = ColorGradeConfig(enabled=True, brightness=20)
|
||||
result = ColorGradeEngine.build_filter(config)
|
||||
assert "eq=" in result
|
||||
assert "brightness=" in result
|
||||
assert "contrast=" not in result
|
||||
assert "saturation=" not in result
|
||||
|
||||
def test_contrast_only(self):
|
||||
"""只有对比度调整."""
|
||||
config = ColorGradeConfig(enabled=True, contrast=30)
|
||||
result = ColorGradeEngine.build_filter(config)
|
||||
assert "eq=" in result
|
||||
assert "contrast=" in result
|
||||
|
||||
def test_saturation_only(self):
|
||||
"""只有饱和度调整."""
|
||||
config = ColorGradeConfig(enabled=True, saturation=50)
|
||||
result = ColorGradeEngine.build_filter(config)
|
||||
assert "eq=" in result
|
||||
assert "saturation=" in result
|
||||
|
||||
def test_temperature_only(self):
|
||||
"""只有色温调整."""
|
||||
config = ColorGradeConfig(enabled=True, temperature=20)
|
||||
result = ColorGradeEngine.build_filter(config)
|
||||
assert "colorbalance=" in result
|
||||
# 暖色调应该有红通道调整
|
||||
assert "rs=" in result
|
||||
|
||||
def test_hue_only(self):
|
||||
"""只有色调调整."""
|
||||
config = ColorGradeConfig(enabled=True, hue=30)
|
||||
result = ColorGradeEngine.build_filter(config)
|
||||
assert "hue=h=" in result
|
||||
|
||||
def test_with_input_output_labels(self):
|
||||
"""带输入输出标签."""
|
||||
config = ColorGradeConfig(enabled=True, brightness=10)
|
||||
result = ColorGradeEngine.build_filter(config, input_label="[0:v]", output_label="[out]")
|
||||
assert result.startswith("[0:v]")
|
||||
assert result.endswith("[out]")
|
||||
|
||||
def test_preset_fresh_filter(self):
|
||||
"""清新预设应该生成eq滤镜."""
|
||||
config = ColorGradeConfig(enabled=True, preset=PRESET_FRESH)
|
||||
result = ColorGradeEngine.build_filter(config)
|
||||
assert "eq=" in result
|
||||
# 清新预设饱和度>100,应该有saturation
|
||||
assert "saturation=" in result
|
||||
|
||||
def test_preset_bw_filter(self):
|
||||
"""黑白预设应该有saturation=0."""
|
||||
config = ColorGradeConfig(enabled=True, preset=PRESET_BW)
|
||||
result = ColorGradeEngine.build_filter(config)
|
||||
assert "saturation=0.0" in result
|
||||
|
||||
def test_combined_params(self):
|
||||
"""多个参数组合."""
|
||||
config = ColorGradeConfig(
|
||||
enabled=True,
|
||||
brightness=15,
|
||||
contrast=20,
|
||||
saturation=130,
|
||||
temperature=10,
|
||||
hue=5,
|
||||
)
|
||||
result = ColorGradeEngine.build_filter(config)
|
||||
# 应该有三个滤镜用逗号连接
|
||||
assert "eq=" in result
|
||||
assert "colorbalance=" in result
|
||||
assert "hue=" in result
|
||||
# 逗号分隔
|
||||
assert "," in result
|
||||
|
||||
def test_filter_chain_order(self):
|
||||
"""滤镜顺序应该是 eq → colorbalance → hue."""
|
||||
config = ColorGradeConfig(
|
||||
enabled=True,
|
||||
brightness=10,
|
||||
temperature=10,
|
||||
hue=10,
|
||||
)
|
||||
result = ColorGradeEngine.build_filter(config)
|
||||
eq_pos = result.find("eq=")
|
||||
cb_pos = result.find("colorbalance=")
|
||||
hue_pos = result.find("hue=")
|
||||
assert eq_pos < cb_pos < hue_pos
|
||||
|
||||
def test_zero_temperature_no_colorbalance(self):
|
||||
"""色温为0不应该有colorbalance滤镜."""
|
||||
config = ColorGradeConfig(enabled=True, temperature=0, brightness=10)
|
||||
result = ColorGradeEngine.build_filter(config)
|
||||
assert "colorbalance" not in result
|
||||
|
||||
def test_zero_hue_no_hue_filter(self):
|
||||
"""色调为0不应该有hue滤镜."""
|
||||
config = ColorGradeConfig(enabled=True, hue=0, brightness=10)
|
||||
result = ColorGradeEngine.build_filter(config)
|
||||
assert "hue=" not in result
|
||||
|
||||
def test_all_presets_generate_valid_filter(self):
|
||||
"""所有预设都应该能生成有效的非空滤镜."""
|
||||
for preset_name in PRESET_PARAMS:
|
||||
config = ColorGradeConfig(enabled=True, preset=preset_name)
|
||||
result = ColorGradeEngine.build_filter(config)
|
||||
assert result, f"预设 {preset_name} 应该生成非空滤镜"
|
||||
# 不应该有语法错误(连续冒号、空参数等)
|
||||
assert "::" not in result
|
||||
assert result[0] != ":"
|
||||
assert result[-1] != ":"
|
||||
|
||||
|
||||
# ── 便捷函数测试 ──────────────────────────────────────────────────────────────
|
||||
|
||||
|
||||
class TestHelperFunctions:
|
||||
"""便捷函数测试."""
|
||||
|
||||
def test_get_preset_names_returns_eight(self):
|
||||
"""应该返回8个预设."""
|
||||
names = get_preset_names()
|
||||
assert len(names) == 8
|
||||
# 每个是 (key, display_name) 元组
|
||||
for key, display in names:
|
||||
assert key in PRESET_PARAMS
|
||||
assert isinstance(display, str)
|
||||
assert display
|
||||
|
||||
def test_get_preset_params_valid(self):
|
||||
"""获取有效预设的参数."""
|
||||
params = get_preset_params(PRESET_FRESH)
|
||||
assert params is not None
|
||||
assert params == PRESET_PARAMS[PRESET_FRESH]
|
||||
|
||||
def test_get_preset_params_invalid(self):
|
||||
"""获取无效预设返回None."""
|
||||
params = get_preset_params("nonexistent")
|
||||
assert params is None
|
||||
|
||||
|
||||
# ── 分段调色(不同clip不同滤镜)概念验证 ──────────────────────────────────────
|
||||
|
||||
|
||||
class TestPerClipGrading:
|
||||
"""分段调色概念验证 — 不同配置生成不同滤镜."""
|
||||
|
||||
def test_different_presets_different_filters(self):
|
||||
"""不同预设应该生成不同的滤镜字符串."""
|
||||
configs = [
|
||||
ColorGradeConfig(enabled=True, preset=PRESET_FRESH),
|
||||
ColorGradeConfig(enabled=True, preset=PRESET_VINTAGE),
|
||||
ColorGradeConfig(enabled=True, preset=PRESET_BW),
|
||||
]
|
||||
filters = [ColorGradeEngine.build_filter(c) for c in configs]
|
||||
# 三个滤镜应该各不相同
|
||||
assert len(set(filters)) == 3
|
||||
|
||||
def test_same_preset_same_filter(self):
|
||||
"""相同配置应该生成相同滤镜(确定性)."""
|
||||
config1 = ColorGradeConfig(enabled=True, preset=PRESET_CINEMA)
|
||||
config2 = ColorGradeConfig(enabled=True, preset=PRESET_CINEMA)
|
||||
assert ColorGradeEngine.build_filter(config1) == ColorGradeEngine.build_filter(config2)
|
||||
|
||||
def test_custom_override_changes_filter(self):
|
||||
"""自定义覆盖应该改变滤镜."""
|
||||
base = ColorGradeConfig(enabled=True, preset=PRESET_FILM)
|
||||
modified = ColorGradeConfig(enabled=True, preset=PRESET_FILM, brightness=50)
|
||||
assert ColorGradeEngine.build_filter(base) != ColorGradeEngine.build_filter(modified)
|
||||
|
||||
def test_clips_with_and_without_grading(self):
|
||||
"""有的clip有调色有的没有,生成结果不同."""
|
||||
with_grade = ColorGradeConfig(enabled=True, preset=PRESET_WARM)
|
||||
without_grade = ColorGradeConfig(enabled=False)
|
||||
|
||||
filter_with = ColorGradeEngine.build_filter(with_grade, "[0:v]", "[v0]")
|
||||
filter_without = ColorGradeEngine.build_filter(without_grade, "[0:v]", "[v0]")
|
||||
|
||||
assert filter_with # 有调色应该非空
|
||||
# 无调色但带标签时应该走 copy 直通(保证标签传递)
|
||||
assert "[0:v]copy[v0]" in filter_without
|
||||
def test_param_ranges_defined(self):
|
||||
"""参数范围定义完整."""
|
||||
assert set(PARAM_RANGES.keys()) == {"brightness", "contrast", "saturation", "temperature", "hue"}
|
||||
|
||||
+224
-184
@@ -1,9 +1,4 @@
|
||||
"""
|
||||
视频拼接引擎配置与纯逻辑测试.
|
||||
|
||||
覆盖 ConcatSegment.from_dict / ConcatConfig.from_config_dict / has_effect / total_segments 等纯逻辑.
|
||||
引擎核心 render 方法依赖 FFmpeg,由集成测试覆盖.
|
||||
"""
|
||||
"""拼接引擎单元测试 - 配置解析等纯逻辑."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
@@ -11,261 +6,306 @@ import pytest
|
||||
from video_processing.concat_engine import ConcatConfig, ConcatSegment
|
||||
|
||||
|
||||
class TestConcatSegmentFromDict:
|
||||
"""ConcatSegment.from_dict 构造逻辑."""
|
||||
class TestConcatSegmentDefaults:
|
||||
"""ConcatSegment 默认值测试."""
|
||||
|
||||
def test_basic(self):
|
||||
seg = ConcatSegment.from_dict({"video_path": "/tmp/a.mp4"})
|
||||
assert seg.video_path == "/tmp/a.mp4"
|
||||
def test_default_values(self):
|
||||
"""默认值正确."""
|
||||
seg = ConcatSegment(video_path="/a.mp4")
|
||||
assert seg.video_path == "/a.mp4"
|
||||
assert seg.start_time == 0.0
|
||||
assert seg.duration == 0.0
|
||||
assert seg.has_audio is True
|
||||
|
||||
def test_full_fields(self):
|
||||
|
||||
class TestConcatSegmentFromDict:
|
||||
"""ConcatSegment.from_dict 测试."""
|
||||
|
||||
def test_basic_path(self):
|
||||
"""基本路径."""
|
||||
seg = ConcatSegment.from_dict({"video_path": "/a.mp4"})
|
||||
assert seg.video_path == "/a.mp4"
|
||||
assert seg.start_time == 0.0
|
||||
assert seg.duration == 0.0
|
||||
|
||||
def test_custom_start_time(self):
|
||||
"""自定义开始时间."""
|
||||
seg = ConcatSegment.from_dict(
|
||||
{
|
||||
"video_path": "/tmp/b.mp4",
|
||||
"start_time": 5.5,
|
||||
"video_path": "/a.mp4",
|
||||
"start_time": 5.0,
|
||||
}
|
||||
)
|
||||
assert seg.start_time == 5.0
|
||||
|
||||
def test_custom_duration(self):
|
||||
"""自定义时长."""
|
||||
seg = ConcatSegment.from_dict(
|
||||
{
|
||||
"video_path": "/a.mp4",
|
||||
"duration": 10.0,
|
||||
}
|
||||
)
|
||||
assert seg.duration == 10.0
|
||||
|
||||
def test_start_time_negative_clamped(self):
|
||||
"""负开始时间钳制到0."""
|
||||
seg = ConcatSegment.from_dict(
|
||||
{
|
||||
"video_path": "/a.mp4",
|
||||
"start_time": -5.0,
|
||||
}
|
||||
)
|
||||
assert seg.start_time == 0.0
|
||||
|
||||
def test_duration_negative_clamped(self):
|
||||
"""负时长钳制到0."""
|
||||
seg = ConcatSegment.from_dict(
|
||||
{
|
||||
"video_path": "/a.mp4",
|
||||
"duration": -3.0,
|
||||
}
|
||||
)
|
||||
assert seg.duration == 0.0
|
||||
|
||||
def test_invalid_start_time_falls_back(self):
|
||||
"""无效start_time回退到0."""
|
||||
seg = ConcatSegment.from_dict(
|
||||
{
|
||||
"video_path": "/a.mp4",
|
||||
"start_time": "invalid",
|
||||
}
|
||||
)
|
||||
assert seg.start_time == 0.0
|
||||
|
||||
def test_invalid_duration_falls_back(self):
|
||||
"""无效duration回退到0."""
|
||||
seg = ConcatSegment.from_dict(
|
||||
{
|
||||
"video_path": "/a.mp4",
|
||||
"duration": "not_a_number",
|
||||
}
|
||||
)
|
||||
assert seg.duration == 0.0
|
||||
|
||||
def test_no_audio(self):
|
||||
"""无音频."""
|
||||
seg = ConcatSegment.from_dict(
|
||||
{
|
||||
"video_path": "/a.mp4",
|
||||
"has_audio": False,
|
||||
}
|
||||
)
|
||||
assert seg.video_path == "/tmp/b.mp4"
|
||||
assert seg.start_time == 5.5
|
||||
assert seg.duration == 10.0
|
||||
assert seg.has_audio is False
|
||||
|
||||
def test_negative_start_time_clamped(self):
|
||||
def test_full_config(self):
|
||||
"""完整配置."""
|
||||
seg = ConcatSegment.from_dict(
|
||||
{
|
||||
"video_path": "/tmp/a.mp4",
|
||||
"start_time": -1.0,
|
||||
"video_path": "/video.mp4",
|
||||
"start_time": 2.5,
|
||||
"duration": 15.0,
|
||||
"has_audio": False,
|
||||
}
|
||||
)
|
||||
assert seg.start_time == 0.0
|
||||
assert seg.video_path == "/video.mp4"
|
||||
assert seg.start_time == 2.5
|
||||
assert seg.duration == 15.0
|
||||
assert seg.has_audio is False
|
||||
|
||||
def test_negative_duration_clamped(self):
|
||||
seg = ConcatSegment.from_dict(
|
||||
{
|
||||
"video_path": "/tmp/a.mp4",
|
||||
"duration": -5.0,
|
||||
}
|
||||
)
|
||||
assert seg.duration == 0.0
|
||||
|
||||
def test_invalid_start_time_type_falls_back(self):
|
||||
seg = ConcatSegment.from_dict(
|
||||
{
|
||||
"video_path": "/tmp/a.mp4",
|
||||
"start_time": "not_a_number",
|
||||
}
|
||||
)
|
||||
assert seg.start_time == 0.0
|
||||
class TestConcatConfigDefaults:
|
||||
"""ConcatConfig 默认值测试."""
|
||||
|
||||
def test_invalid_duration_type_falls_back(self):
|
||||
seg = ConcatSegment.from_dict(
|
||||
{
|
||||
"video_path": "/tmp/a.mp4",
|
||||
"duration": "abc",
|
||||
}
|
||||
)
|
||||
assert seg.duration == 0.0
|
||||
|
||||
def test_start_time_none_falls_back(self):
|
||||
seg = ConcatSegment.from_dict(
|
||||
{
|
||||
"video_path": "/tmp/a.mp4",
|
||||
"start_time": None,
|
||||
}
|
||||
)
|
||||
assert seg.start_time == 0.0
|
||||
|
||||
def test_empty_video_path_stored(self):
|
||||
seg = ConcatSegment.from_dict({"video_path": ""})
|
||||
assert seg.video_path == ""
|
||||
def test_default_values(self):
|
||||
"""默认值正确."""
|
||||
config = ConcatConfig()
|
||||
assert config.segments == []
|
||||
assert config.output_width == 0
|
||||
assert config.output_height == 0
|
||||
assert config.output_fps == 0.0
|
||||
assert config.force_reencode is False
|
||||
assert config.transition == "none"
|
||||
assert config.transition_duration == 0.3
|
||||
|
||||
|
||||
class TestConcatConfigFromConfigDict:
|
||||
"""ConcatConfig.from_config_dict 构造逻辑."""
|
||||
"""ConcatConfig.from_config_dict 测试."""
|
||||
|
||||
def test_none_returns_default(self):
|
||||
cfg = ConcatConfig.from_config_dict(None)
|
||||
assert cfg.segments == []
|
||||
assert cfg.output_width == 0
|
||||
assert cfg.output_height == 0
|
||||
assert cfg.output_fps == 0.0
|
||||
assert cfg.force_reencode is False
|
||||
"""None返回默认配置."""
|
||||
config = ConcatConfig.from_config_dict(None)
|
||||
assert config.segments == []
|
||||
|
||||
def test_empty_dict_returns_default(self):
|
||||
cfg = ConcatConfig.from_config_dict({})
|
||||
assert cfg.segments == []
|
||||
|
||||
def test_non_dict_returns_default(self):
|
||||
cfg = ConcatConfig.from_config_dict("not a dict")
|
||||
assert cfg.segments == []
|
||||
"""空dict返回默认."""
|
||||
config = ConcatConfig.from_config_dict({})
|
||||
assert config.segments == []
|
||||
|
||||
def test_single_segment(self):
|
||||
cfg = ConcatConfig.from_config_dict(
|
||||
"""单片段."""
|
||||
config = ConcatConfig.from_config_dict(
|
||||
{
|
||||
"segments": [
|
||||
{"video_path": "/tmp/a.mp4", "duration": 5.0},
|
||||
],
|
||||
"segments": [{"video_path": "/a.mp4"}],
|
||||
}
|
||||
)
|
||||
assert len(cfg.segments) == 1
|
||||
assert cfg.segments[0].video_path == "/tmp/a.mp4"
|
||||
assert cfg.segments[0].duration == 5.0
|
||||
assert len(config.segments) == 1
|
||||
assert config.segments[0].video_path == "/a.mp4"
|
||||
|
||||
def test_multiple_segments(self):
|
||||
cfg = ConcatConfig.from_config_dict(
|
||||
"""多片段."""
|
||||
config = ConcatConfig.from_config_dict(
|
||||
{
|
||||
"segments": [
|
||||
{"video_path": "/tmp/a.mp4"},
|
||||
{"video_path": "/tmp/b.mp4", "start_time": 2.0},
|
||||
{"video_path": "/tmp/c.mp4", "duration": 3.0, "has_audio": False},
|
||||
{"video_path": "/a.mp4", "start_time": 1.0},
|
||||
{"video_path": "/b.mp4", "duration": 5.0},
|
||||
{"video_path": "/c.mp4"},
|
||||
],
|
||||
}
|
||||
)
|
||||
assert len(cfg.segments) == 3
|
||||
assert cfg.segments[0].video_path == "/tmp/a.mp4"
|
||||
assert cfg.segments[1].start_time == 2.0
|
||||
assert cfg.segments[2].has_audio is False
|
||||
assert len(config.segments) == 3
|
||||
assert config.segments[0].start_time == 1.0
|
||||
assert config.segments[1].duration == 5.0
|
||||
|
||||
def test_invalid_segments_filtered(self):
|
||||
cfg = ConcatConfig.from_config_dict(
|
||||
def test_skips_no_path(self):
|
||||
"""跳过无video_path的片段."""
|
||||
config = ConcatConfig.from_config_dict(
|
||||
{
|
||||
"segments": [
|
||||
{"video_path": "/tmp/valid.mp4"},
|
||||
{"video_path": ""}, # 空路径被过滤
|
||||
{"not_video_path": "xxx"}, # 没有video_path被过滤
|
||||
"not_a_dict", # 不是dict被过滤
|
||||
None, # None被过滤
|
||||
{"video_path": "/a.mp4"},
|
||||
{"other": "value"},
|
||||
{"video_path": ""},
|
||||
],
|
||||
}
|
||||
)
|
||||
assert len(cfg.segments) == 1
|
||||
assert cfg.segments[0].video_path == "/tmp/valid.mp4"
|
||||
assert len(config.segments) == 1
|
||||
|
||||
def test_segments_not_a_list(self):
|
||||
cfg = ConcatConfig.from_config_dict(
|
||||
def test_segments_not_list_ignored(self):
|
||||
"""segments不是列表忽略."""
|
||||
config = ConcatConfig.from_config_dict(
|
||||
{
|
||||
"segments": "not_a_list",
|
||||
}
|
||||
)
|
||||
assert cfg.segments == []
|
||||
assert config.segments == []
|
||||
|
||||
def test_output_params(self):
|
||||
cfg = ConcatConfig.from_config_dict(
|
||||
def test_output_size(self):
|
||||
"""输出尺寸."""
|
||||
config = ConcatConfig.from_config_dict(
|
||||
{
|
||||
"segments": [],
|
||||
"segments": [{"video_path": "/a.mp4"}],
|
||||
"output_width": 1920,
|
||||
"output_height": 1080,
|
||||
"output_fps": 30.0,
|
||||
}
|
||||
)
|
||||
assert config.output_width == 1920
|
||||
assert config.output_height == 1080
|
||||
|
||||
def test_negative_output_size_clamped(self):
|
||||
"""负输出尺寸钳制到0."""
|
||||
config = ConcatConfig.from_config_dict(
|
||||
{
|
||||
"segments": [{"video_path": "/a.mp4"}],
|
||||
"output_width": -100,
|
||||
"output_height": -50,
|
||||
}
|
||||
)
|
||||
assert config.output_width == 0
|
||||
assert config.output_height == 0
|
||||
|
||||
def test_invalid_output_size_falls_back(self):
|
||||
"""无效输出尺寸回退."""
|
||||
config = ConcatConfig.from_config_dict(
|
||||
{
|
||||
"segments": [{"video_path": "/a.mp4"}],
|
||||
"output_width": "wide",
|
||||
"output_fps": "sixty",
|
||||
}
|
||||
)
|
||||
assert config.output_width == 0
|
||||
assert config.output_fps == 0.0
|
||||
|
||||
def test_output_fps(self):
|
||||
"""输出帧率."""
|
||||
config = ConcatConfig.from_config_dict(
|
||||
{
|
||||
"segments": [{"video_path": "/a.mp4"}],
|
||||
"output_fps": 60.0,
|
||||
}
|
||||
)
|
||||
assert config.output_fps == 60.0
|
||||
|
||||
def test_force_reencode(self):
|
||||
"""强制重新编码."""
|
||||
config = ConcatConfig.from_config_dict(
|
||||
{
|
||||
"segments": [{"video_path": "/a.mp4"}],
|
||||
"force_reencode": True,
|
||||
}
|
||||
)
|
||||
assert cfg.output_width == 1920
|
||||
assert cfg.output_height == 1080
|
||||
assert cfg.output_fps == 30.0
|
||||
assert cfg.force_reencode is True
|
||||
|
||||
def test_negative_output_params_clamped(self):
|
||||
cfg = ConcatConfig.from_config_dict(
|
||||
{
|
||||
"segments": [],
|
||||
"output_width": -100,
|
||||
"output_height": -50,
|
||||
"output_fps": -1.0,
|
||||
}
|
||||
)
|
||||
assert cfg.output_width == 0
|
||||
assert cfg.output_height == 0
|
||||
assert cfg.output_fps == 0.0
|
||||
|
||||
def test_invalid_output_params_fall_back(self):
|
||||
cfg = ConcatConfig.from_config_dict(
|
||||
{
|
||||
"segments": [],
|
||||
"output_width": "abc",
|
||||
"output_height": None,
|
||||
"output_fps": "xyz",
|
||||
}
|
||||
)
|
||||
assert cfg.output_width == 0
|
||||
assert cfg.output_height == 0
|
||||
assert cfg.output_fps == 0.0
|
||||
assert config.force_reencode is True
|
||||
|
||||
def test_transition_config(self):
|
||||
cfg = ConcatConfig.from_config_dict(
|
||||
"""转场配置."""
|
||||
config = ConcatConfig.from_config_dict(
|
||||
{
|
||||
"segments": [],
|
||||
"segments": [{"video_path": "/a.mp4"}, {"video_path": "/b.mp4"}],
|
||||
"transition": "crossfade",
|
||||
"transition_duration": 1.0,
|
||||
}
|
||||
)
|
||||
assert cfg.transition == "crossfade"
|
||||
assert cfg.transition_duration == 1.0
|
||||
assert config.transition == "crossfade"
|
||||
assert config.transition_duration == 1.0
|
||||
|
||||
def test_transition_duration_minimum(self):
|
||||
"""transition_duration 不能小于 0.1."""
|
||||
cfg = ConcatConfig.from_config_dict(
|
||||
{
|
||||
"segments": [],
|
||||
"transition_duration": 0.01,
|
||||
}
|
||||
)
|
||||
assert cfg.transition_duration >= 0.1
|
||||
|
||||
def test_default_values(self):
|
||||
cfg = ConcatConfig.from_config_dict({"segments": []})
|
||||
assert cfg.transition == "none"
|
||||
assert cfg.transition_duration == 0.3
|
||||
assert cfg.force_reencode is False
|
||||
def test_non_dict_config_returns_default(self):
|
||||
"""非dict配置返回默认."""
|
||||
config = ConcatConfig.from_config_dict("not_a_dict")
|
||||
assert config.segments == []
|
||||
|
||||
|
||||
class TestConcatConfigProperties:
|
||||
"""has_effect / total_segments 属性."""
|
||||
class TestHasEffect:
|
||||
"""has_effect 属性测试."""
|
||||
|
||||
def test_has_effect_two_or_more_valid(self):
|
||||
cfg = ConcatConfig(
|
||||
def test_no_segments_no_effect(self):
|
||||
"""无片段无效果."""
|
||||
config = ConcatConfig()
|
||||
assert config.has_effect is False
|
||||
|
||||
def test_one_segment_no_effect(self):
|
||||
"""单片段无效果(拼接至少需要2段)."""
|
||||
config = ConcatConfig(
|
||||
segments=[
|
||||
ConcatSegment(video_path="/tmp/a.mp4"),
|
||||
ConcatSegment(video_path="/tmp/b.mp4"),
|
||||
ConcatSegment(video_path="/a.mp4"),
|
||||
]
|
||||
)
|
||||
assert cfg.has_effect is True
|
||||
assert config.has_effect is False
|
||||
|
||||
def test_no_effect_one_segment(self):
|
||||
cfg = ConcatConfig(
|
||||
def test_two_segments_has_effect(self):
|
||||
"""两段及以上有效果."""
|
||||
config = ConcatConfig(
|
||||
segments=[
|
||||
ConcatSegment(video_path="/tmp/a.mp4"),
|
||||
ConcatSegment(video_path="/a.mp4"),
|
||||
ConcatSegment(video_path="/b.mp4"),
|
||||
]
|
||||
)
|
||||
assert cfg.has_effect is False
|
||||
assert config.has_effect is True
|
||||
|
||||
def test_no_effect_zero_segments(self):
|
||||
cfg = ConcatConfig(segments=[])
|
||||
assert cfg.has_effect is False
|
||||
|
||||
def test_no_effect_empty_paths(self):
|
||||
cfg = ConcatConfig(
|
||||
class TestTotalSegments:
|
||||
"""total_segments 属性测试."""
|
||||
|
||||
def test_no_segments(self):
|
||||
"""零片段."""
|
||||
config = ConcatConfig()
|
||||
assert config.total_segments == 0
|
||||
|
||||
def test_three_segments(self):
|
||||
"""三个片段."""
|
||||
config = ConcatConfig(
|
||||
segments=[
|
||||
ConcatSegment(video_path=""),
|
||||
ConcatSegment(video_path=""),
|
||||
ConcatSegment(video_path="/a.mp4"),
|
||||
ConcatSegment(video_path="/b.mp4"),
|
||||
ConcatSegment(video_path="/c.mp4"),
|
||||
]
|
||||
)
|
||||
assert cfg.has_effect is False
|
||||
|
||||
def test_total_segments(self):
|
||||
cfg = ConcatConfig(
|
||||
segments=[
|
||||
ConcatSegment(video_path="/tmp/a.mp4"),
|
||||
ConcatSegment(video_path=""),
|
||||
ConcatSegment(video_path="/tmp/b.mp4"),
|
||||
]
|
||||
)
|
||||
assert cfg.total_segments == 2
|
||||
|
||||
def test_total_segments_empty(self):
|
||||
cfg = ConcatConfig(segments=[])
|
||||
assert cfg.total_segments == 0
|
||||
assert config.total_segments == 3
|
||||
|
||||
Executable
+206
@@ -0,0 +1,206 @@
|
||||
"""FFmpeg工具函数纯逻辑测试 — chain_filters / resolve_xfade_transition / build_xfade_filter_chain."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import pytest
|
||||
from video_processing.ffmpeg_utils import (
|
||||
XFADE_TRANSITION_MAP,
|
||||
build_xfade_filter_chain,
|
||||
chain_filters,
|
||||
resolve_xfade_transition,
|
||||
)
|
||||
|
||||
|
||||
class TestChainFilters:
|
||||
"""chain_filters 滤镜串联测试."""
|
||||
|
||||
def test_single_filter(self):
|
||||
"""单个滤镜."""
|
||||
result = chain_filters(["scale=1280:720"], "v0")
|
||||
assert result == "[0:v]scale=1280:720[v0]"
|
||||
|
||||
def test_multiple_filters(self):
|
||||
"""多个滤镜用逗号连接."""
|
||||
result = chain_filters(["scale=1280:720", "fps=25", "format=yuv420p"], "out")
|
||||
assert result == "[0:v]scale=1280:720,fps=25,format=yuv420p[out]"
|
||||
|
||||
def test_empty_filters(self):
|
||||
"""空滤镜列表."""
|
||||
result = chain_filters([], "v0")
|
||||
assert result == "[0:v][v0]"
|
||||
|
||||
def test_custom_input_label(self):
|
||||
"""自定义输入标签."""
|
||||
result = chain_filters(["scale=640:480"], "v1", input_label="1:v")
|
||||
assert result == "[1:v]scale=640:480[v1]"
|
||||
|
||||
|
||||
class TestResolveXfadeTransition:
|
||||
"""resolve_xfade_transition 转场名称映射测试."""
|
||||
|
||||
def test_direct_match_fade(self):
|
||||
"""fade直接匹配."""
|
||||
assert resolve_xfade_transition("fade") == "fade"
|
||||
|
||||
def test_direct_match_dissolve(self):
|
||||
"""dissolve直接匹配."""
|
||||
assert resolve_xfade_transition("dissolve") == "dissolve"
|
||||
|
||||
def test_alias_crossfade(self):
|
||||
"""crossfade别名→dissolve."""
|
||||
assert resolve_xfade_transition("crossfade") == "dissolve"
|
||||
|
||||
def test_alias_slide_left(self):
|
||||
"""slide_left别名→slideleft."""
|
||||
assert resolve_xfade_transition("slide_left") == "slideleft"
|
||||
|
||||
def test_unknown_fallback_to_fade(self):
|
||||
"""未知值回退到fade."""
|
||||
assert resolve_xfade_transition("nonexistent_effect") == "fade"
|
||||
|
||||
def test_empty_string_fallback(self):
|
||||
"""空字符串回退."""
|
||||
assert resolve_xfade_transition("") == "fade"
|
||||
|
||||
def test_enum_value_support(self):
|
||||
"""支持带value属性的枚举对象."""
|
||||
|
||||
class FakeEnum:
|
||||
value = "slideup"
|
||||
|
||||
assert resolve_xfade_transition(FakeEnum()) == "slideup"
|
||||
|
||||
def test_all_map_keys_resolve(self):
|
||||
"""映射表中所有key都能解析到有效值."""
|
||||
for key in XFADE_TRANSITION_MAP:
|
||||
result = resolve_xfade_transition(key)
|
||||
assert result and isinstance(result, str)
|
||||
assert result != ""
|
||||
|
||||
def test_cut_is_special_fallback(self):
|
||||
"""cut不在映射表中→回退到fade(硬切由调用方处理)."""
|
||||
# cut是特殊值,不在映射表里
|
||||
result = resolve_xfade_transition("cut")
|
||||
# 不在映射表里就fallback到fade
|
||||
assert result == "fade"
|
||||
|
||||
|
||||
class TestBuildXfadeFilterChain:
|
||||
"""build_xfade_filter_chain 转场滤镜链构建测试."""
|
||||
|
||||
def test_zero_clips(self):
|
||||
"""0个片段→空字符串+0时长."""
|
||||
filter_str, total_dur = build_xfade_filter_chain([], [], [])
|
||||
assert filter_str == ""
|
||||
assert total_dur == 0.0
|
||||
|
||||
def test_single_clip(self):
|
||||
"""1个片段→直接copy,总时长等于片段时长."""
|
||||
filter_str, total_dur = build_xfade_filter_chain([10.0], ["v0"], [], output_label="outv")
|
||||
assert "[v0]copy[outv]" in filter_str
|
||||
assert total_dur == pytest.approx(10.0)
|
||||
|
||||
def test_two_clips_basic(self):
|
||||
"""2个片段基本转场."""
|
||||
filter_str, total_dur = build_xfade_filter_chain(
|
||||
[5.0, 5.0],
|
||||
["v0", "v1"],
|
||||
["", "fade"],
|
||||
transition_duration=0.5,
|
||||
output_label="outv",
|
||||
)
|
||||
assert "xfade=transition=fade" in filter_str
|
||||
assert "offset=" in filter_str
|
||||
# 总时长 = 5 + 5 - 转场重叠
|
||||
assert total_dur == pytest.approx(9.5)
|
||||
|
||||
def test_three_clips_chain(self):
|
||||
"""3个片段形成链式转场."""
|
||||
filter_str, total_dur = build_xfade_filter_chain(
|
||||
[3.0, 4.0, 5.0],
|
||||
["v0", "v1", "v2"],
|
||||
["", "fade", "dissolve"],
|
||||
transition_duration=0.5,
|
||||
output_label="out",
|
||||
)
|
||||
# 应该有2个xfade操作
|
||||
assert filter_str.count("xfade=") == 2
|
||||
assert "transition=fade" in filter_str
|
||||
assert "transition=dissolve" in filter_str
|
||||
# 总时长 = 3+4+5 - 2*0.5 = 11
|
||||
assert total_dur == pytest.approx(11.0)
|
||||
|
||||
def test_transition_duration_clamped_to_clip(self):
|
||||
"""转场时长不能超过单个片段时长."""
|
||||
filter_str, total_dur = build_xfade_filter_chain(
|
||||
[2.0, 1.0],
|
||||
["v0", "v1"],
|
||||
["", "fade"],
|
||||
transition_duration=3.0, # 比第二个片段还长
|
||||
output_label="outv",
|
||||
)
|
||||
# 转场时长被钳制到第二个片段时长(1.0)
|
||||
assert "duration=1.000" in filter_str
|
||||
assert total_dur == pytest.approx(2.0) # 2 + 1 - 1 = 2
|
||||
|
||||
def test_very_short_clip_min_transition(self):
|
||||
"""极短片段至少保留1ms转场."""
|
||||
filter_str, total_dur = build_xfade_filter_chain(
|
||||
[1.0, 0.0001],
|
||||
["v0", "v1"],
|
||||
["", "fade"],
|
||||
transition_duration=0.5,
|
||||
output_label="outv",
|
||||
)
|
||||
# 至少有1ms
|
||||
assert "duration=0.001" in filter_str
|
||||
|
||||
def test_transition_offset_calculation(self):
|
||||
"""offset计算验证."""
|
||||
filter_str, _ = build_xfade_filter_chain(
|
||||
[10.0, 10.0],
|
||||
["v0", "v1"],
|
||||
["", "fade"],
|
||||
transition_duration=1.0,
|
||||
output_label="outv",
|
||||
)
|
||||
# offset = max(0, 10 - 1*1) = 9
|
||||
assert "offset=9.000" in filter_str
|
||||
|
||||
def test_fewer_transitions_than_clips(self):
|
||||
"""转场列表比片段少时使用cut(fallback to fade)."""
|
||||
filter_str, total_dur = build_xfade_filter_chain(
|
||||
[5.0, 5.0, 5.0],
|
||||
["v0", "v1", "v2"],
|
||||
["fade"], # 只有1个转场,第2个转场缺省
|
||||
transition_duration=0.5,
|
||||
output_label="out",
|
||||
)
|
||||
# 应该有2个xfade
|
||||
assert filter_str.count("xfade=") == 2
|
||||
# 第二个xfade的转场是cut→fade fallback
|
||||
assert filter_str.count("transition=fade") == 2
|
||||
|
||||
def test_output_label_final_clip(self):
|
||||
"""最后一个xfade的输出标签是output_label."""
|
||||
filter_str, _ = build_xfade_filter_chain(
|
||||
[3.0, 4.0, 5.0],
|
||||
["v0", "v1", "v2"],
|
||||
["", "fade", "slideleft"],
|
||||
output_label="final_v",
|
||||
)
|
||||
assert filter_str.rstrip().endswith("[final_v]")
|
||||
|
||||
def test_intermediate_labels(self):
|
||||
"""中间步骤使用xf1, xf2等标签(从i=1开始计数)."""
|
||||
filter_str, _ = build_xfade_filter_chain(
|
||||
[2.0, 3.0, 4.0, 5.0],
|
||||
["v0", "v1", "v2", "v3"],
|
||||
["", "fade", "fade", "fade"],
|
||||
output_label="out",
|
||||
)
|
||||
# 4个片段3次xfade,中间标签是xf1, xf2
|
||||
assert "[xf1]" in filter_str
|
||||
assert "[xf2]" in filter_str
|
||||
# 最后一个是[out]
|
||||
assert filter_str.rstrip().endswith("[out]")
|
||||
+81
@@ -0,0 +1,81 @@
|
||||
"""GenerationTaskStatus 枚举兼容性测试。
|
||||
|
||||
验证历史脏数据(如 'success'/'done')不会导致枚举转换失败。
|
||||
关联 Issue: #809 [Staging] E2E测试失败 - 模板生成接口返回500
|
||||
"""
|
||||
|
||||
import pytest
|
||||
|
||||
from packages.domain.generation_task import GenerationTaskStatus
|
||||
|
||||
|
||||
class TestGenerationTaskStatusNormalValues:
|
||||
"""正常值应该正确映射。"""
|
||||
|
||||
def test_pending(self):
|
||||
assert GenerationTaskStatus("pending") == GenerationTaskStatus.PENDING
|
||||
|
||||
def test_running(self):
|
||||
assert GenerationTaskStatus("running") == GenerationTaskStatus.RUNNING
|
||||
|
||||
def test_completed(self):
|
||||
assert GenerationTaskStatus("completed") == GenerationTaskStatus.COMPLETED
|
||||
|
||||
def test_failed(self):
|
||||
assert GenerationTaskStatus("failed") == GenerationTaskStatus.FAILED
|
||||
|
||||
def test_cancelled(self):
|
||||
assert GenerationTaskStatus("cancelled") == GenerationTaskStatus.CANCELLED
|
||||
|
||||
|
||||
class TestGenerationTaskStatusHistoricalValues:
|
||||
"""历史脏数据应该正确映射到对应状态,不抛异常。"""
|
||||
|
||||
@pytest.mark.parametrize("value", ["done", "success", "finished", "complete", "completed"])
|
||||
def test_completed_like_values_map_to_completed(self, value):
|
||||
assert GenerationTaskStatus(value) == GenerationTaskStatus.COMPLETED
|
||||
|
||||
@pytest.mark.parametrize("value", ["fail", "failed", "error", "err"])
|
||||
def test_failed_like_values_map_to_failed(self, value):
|
||||
assert GenerationTaskStatus(value) == GenerationTaskStatus.FAILED
|
||||
|
||||
@pytest.mark.parametrize("value", ["process", "processing", "run", "running", "in_progress"])
|
||||
def test_running_like_values_map_to_running(self, value):
|
||||
assert GenerationTaskStatus(value) == GenerationTaskStatus.RUNNING
|
||||
|
||||
@pytest.mark.parametrize("value", ["cancel", "cancelled", "canceled"])
|
||||
def test_cancelled_like_values_map_to_cancelled(self, value):
|
||||
assert GenerationTaskStatus(value) == GenerationTaskStatus.CANCELLED
|
||||
|
||||
@pytest.mark.parametrize("value", [" Done ", "SUCCESS", " failed "])
|
||||
def test_whitespace_and_case_insensitive(self, value):
|
||||
"""带空格和大小写不影响匹配。"""
|
||||
# 只要能找到对应状态且不抛异常即可
|
||||
result = GenerationTaskStatus(value)
|
||||
assert result in (
|
||||
GenerationTaskStatus.COMPLETED,
|
||||
GenerationTaskStatus.FAILED,
|
||||
)
|
||||
|
||||
|
||||
class TestGenerationTaskStatusFallback:
|
||||
"""完全未知的值兜底为 PENDING,不抛500。"""
|
||||
|
||||
@pytest.mark.parametrize("value", ["unknown", "foo_bar", "deleted", ""])
|
||||
def test_unknown_value_falls_back_to_pending(self, value):
|
||||
assert GenerationTaskStatus(value) == GenerationTaskStatus.PENDING
|
||||
|
||||
def test_none_value_falls_back_to_pending(self):
|
||||
assert GenerationTaskStatus(None) == GenerationTaskStatus.PENDING # type: ignore[arg-type]
|
||||
|
||||
def test_int_value_falls_back_to_pending(self):
|
||||
assert GenerationTaskStatus(123) == GenerationTaskStatus.PENDING # type: ignore[arg-type]
|
||||
|
||||
|
||||
class TestGenerationTaskStatusStrValue:
|
||||
"""枚举值仍为字符串类型,不影响序列化。"""
|
||||
|
||||
def test_value_unchanged(self):
|
||||
assert GenerationTaskStatus.PENDING.value == "pending"
|
||||
assert GenerationTaskStatus.COMPLETED.value == "completed"
|
||||
assert isinstance(GenerationTaskStatus.PENDING, str)
|
||||
@@ -1,301 +1,308 @@
|
||||
"""
|
||||
片头片尾引擎配置与纯逻辑测试.
|
||||
"""片头片尾引擎单元测试 - 配置解析等纯逻辑."""
|
||||
|
||||
覆盖 IntroOutroConfig.from_dict / validate / has_intro / has_outro 等纯逻辑.
|
||||
引擎核心 render 方法依赖 FFmpeg,由集成测试覆盖.
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import pytest
|
||||
from video_processing.intro_outro_engine import IntroOutroConfig
|
||||
|
||||
|
||||
class TestIntroOutroConfigDefaults:
|
||||
"""默认配置测试."""
|
||||
|
||||
def test_default_values(self):
|
||||
"""默认值正确."""
|
||||
config = IntroOutroConfig()
|
||||
assert config.enabled is False
|
||||
assert config.intro_type == "none"
|
||||
assert config.outro_type == "none"
|
||||
assert config.intro_duration == 3.0
|
||||
assert config.outro_duration == 3.0
|
||||
assert config.transition_effect == "fade"
|
||||
assert config.transition_duration == 0.5
|
||||
|
||||
|
||||
class TestIntroOutroConfigFromDict:
|
||||
"""from_dict 构造逻辑."""
|
||||
"""from_dict 配置解析测试."""
|
||||
|
||||
def test_none_returns_default_disabled(self):
|
||||
cfg = IntroOutroConfig.from_dict(None)
|
||||
assert cfg.enabled is False
|
||||
assert cfg.intro_type == "none"
|
||||
assert cfg.outro_type == "none"
|
||||
def test_none_returns_default(self):
|
||||
"""None 返回默认配置."""
|
||||
config = IntroOutroConfig.from_dict(None)
|
||||
assert config.enabled is False
|
||||
|
||||
def test_empty_dict_returns_default_disabled(self):
|
||||
cfg = IntroOutroConfig.from_dict({})
|
||||
assert cfg.enabled is False
|
||||
def test_empty_dict_returns_default(self):
|
||||
"""空 dict 返回默认."""
|
||||
config = IntroOutroConfig.from_dict({})
|
||||
assert config.enabled is False
|
||||
|
||||
def test_enabled_false_returns_default_disabled(self):
|
||||
cfg = IntroOutroConfig.from_dict({"enabled": False})
|
||||
assert cfg.enabled is False
|
||||
def test_disabled_returns_default(self):
|
||||
"""enabled=False 返回默认."""
|
||||
config = IntroOutroConfig.from_dict({"enabled": False})
|
||||
assert config.enabled is False
|
||||
|
||||
def test_enabled_with_video_intro(self):
|
||||
cfg = IntroOutroConfig.from_dict({
|
||||
"enabled": True,
|
||||
"intro": {
|
||||
"type": "video",
|
||||
"video_path": "/tmp/intro.mp4",
|
||||
"duration": 5.0,
|
||||
},
|
||||
"outro": {"type": "none"},
|
||||
})
|
||||
assert cfg.enabled is True
|
||||
assert cfg.intro_type == "video"
|
||||
assert cfg.intro_video_path == "/tmp/intro.mp4"
|
||||
assert cfg.intro_duration == 5.0
|
||||
def test_enabled_defaults(self):
|
||||
"""启用时默认值正确."""
|
||||
config = IntroOutroConfig.from_dict({"enabled": True})
|
||||
assert config.enabled is True
|
||||
assert config.intro_type == "none"
|
||||
assert config.outro_type == "none"
|
||||
|
||||
def test_enabled_with_text_intro(self):
|
||||
cfg = IntroOutroConfig.from_dict({
|
||||
def test_text_intro(self):
|
||||
"""文字片头配置."""
|
||||
config = IntroOutroConfig.from_dict({
|
||||
"enabled": True,
|
||||
"intro": {
|
||||
"type": "text",
|
||||
"title": "Hello",
|
||||
"subtitle": "World",
|
||||
"background": "#ffffff",
|
||||
"title_color": "black",
|
||||
"title_size": 64,
|
||||
"duration": 2.5,
|
||||
},
|
||||
"outro": {"type": "none"},
|
||||
})
|
||||
assert cfg.enabled is True
|
||||
assert cfg.intro_type == "text"
|
||||
assert cfg.intro_title == "Hello"
|
||||
assert cfg.intro_subtitle == "World"
|
||||
assert cfg.intro_background == "#ffffff"
|
||||
assert cfg.intro_title_color == "black"
|
||||
assert cfg.intro_title_size == 64
|
||||
assert cfg.intro_duration == 2.5
|
||||
|
||||
def test_enabled_with_video_outro(self):
|
||||
cfg = IntroOutroConfig.from_dict({
|
||||
"enabled": True,
|
||||
"intro": {"type": "none"},
|
||||
"outro": {
|
||||
"type": "video",
|
||||
"video_path": "/tmp/outro.mp4",
|
||||
"duration": 4.0,
|
||||
"title": "我的片头",
|
||||
"subtitle": "欢迎收看",
|
||||
},
|
||||
})
|
||||
assert cfg.enabled is True
|
||||
assert cfg.outro_type == "video"
|
||||
assert cfg.outro_video_path == "/tmp/outro.mp4"
|
||||
assert cfg.outro_duration == 4.0
|
||||
assert config.intro_type == "text"
|
||||
assert config.intro_title == "我的片头"
|
||||
assert config.intro_subtitle == "欢迎收看"
|
||||
|
||||
def test_enabled_with_text_outro_default_values(self):
|
||||
cfg = IntroOutroConfig.from_dict({
|
||||
"enabled": True,
|
||||
"intro": {"type": "none"},
|
||||
"outro": {"type": "text"},
|
||||
})
|
||||
assert cfg.outro_title == "感谢观看"
|
||||
assert cfg.outro_subtitle == "点赞关注不迷路"
|
||||
assert cfg.outro_title_size == 48
|
||||
assert cfg.outro_duration == 3.0
|
||||
|
||||
def test_video_key_fallback(self):
|
||||
"""video 字段作为 video_path 的 fallback."""
|
||||
cfg = IntroOutroConfig.from_dict({
|
||||
def test_video_intro(self):
|
||||
"""视频片头配置."""
|
||||
config = IntroOutroConfig.from_dict({
|
||||
"enabled": True,
|
||||
"intro": {
|
||||
"type": "video",
|
||||
"video": "/tmp/fallback.mp4",
|
||||
"video_path": "/videos/intro.mp4",
|
||||
"duration": 5.0,
|
||||
},
|
||||
"outro": {"type": "none"},
|
||||
})
|
||||
assert cfg.intro_video_path == "/tmp/fallback.mp4"
|
||||
assert config.intro_type == "video"
|
||||
assert config.intro_video_path == "/videos/intro.mp4"
|
||||
assert config.intro_duration == 5.0
|
||||
|
||||
def test_video_intro_video_alias(self):
|
||||
"""video 字段作为 video_path 别名."""
|
||||
config = IntroOutroConfig.from_dict({
|
||||
"enabled": True,
|
||||
"intro": {
|
||||
"type": "video",
|
||||
"video": "/videos/intro.mp4",
|
||||
},
|
||||
})
|
||||
assert config.intro_video_path == "/videos/intro.mp4"
|
||||
|
||||
def test_text_outro(self):
|
||||
"""文字片尾配置."""
|
||||
config = IntroOutroConfig.from_dict({
|
||||
"enabled": True,
|
||||
"outro": {
|
||||
"type": "text",
|
||||
"title": "感谢观看",
|
||||
"subtitle": "点赞关注",
|
||||
},
|
||||
})
|
||||
assert config.outro_type == "text"
|
||||
assert config.outro_title == "感谢观看"
|
||||
assert config.outro_subtitle == "点赞关注"
|
||||
|
||||
def test_outro_default_title(self):
|
||||
"""片尾默认标题."""
|
||||
config = IntroOutroConfig.from_dict({
|
||||
"enabled": True,
|
||||
"outro": {"type": "text"},
|
||||
})
|
||||
assert config.outro_title == "感谢观看"
|
||||
assert config.outro_subtitle == "点赞关注不迷路"
|
||||
|
||||
def test_text_intro_styling(self):
|
||||
"""文字片头样式配置."""
|
||||
config = IntroOutroConfig.from_dict({
|
||||
"enabled": True,
|
||||
"intro": {
|
||||
"type": "text",
|
||||
"title": "测试",
|
||||
"background": "#FF0000",
|
||||
"title_color": "yellow",
|
||||
"title_size": 64,
|
||||
"subtitle_color": "white",
|
||||
"subtitle_size": 32,
|
||||
},
|
||||
})
|
||||
assert config.intro_background == "#FF0000"
|
||||
assert config.intro_title_color == "yellow"
|
||||
assert config.intro_title_size == 64
|
||||
assert config.intro_subtitle_color == "white"
|
||||
assert config.intro_subtitle_size == 32
|
||||
|
||||
def test_transition_config(self):
|
||||
cfg = IntroOutroConfig.from_dict({
|
||||
"""转场配置."""
|
||||
config = IntroOutroConfig.from_dict({
|
||||
"enabled": True,
|
||||
"intro": {"type": "none"},
|
||||
"outro": {"type": "none"},
|
||||
"transition": "fade",
|
||||
"transition": "dissolve",
|
||||
"transition_duration": 1.0,
|
||||
})
|
||||
assert cfg.transition_effect == "fade"
|
||||
assert cfg.transition_duration == 1.0
|
||||
assert config.transition_effect == "dissolve"
|
||||
assert config.transition_duration == 1.0
|
||||
|
||||
def test_default_transition(self):
|
||||
cfg = IntroOutroConfig.from_dict({
|
||||
def test_empty_intro_dict(self):
|
||||
"""空 intro dict."""
|
||||
config = IntroOutroConfig.from_dict({
|
||||
"enabled": True,
|
||||
"intro": {"type": "none"},
|
||||
"outro": {"type": "none"},
|
||||
"intro": {},
|
||||
})
|
||||
assert cfg.transition_effect == "fade"
|
||||
assert cfg.transition_duration == 0.5
|
||||
assert config.intro_type == "none"
|
||||
|
||||
def test_none_intro(self):
|
||||
"""None intro 值."""
|
||||
config = IntroOutroConfig.from_dict({
|
||||
"enabled": True,
|
||||
"intro": None,
|
||||
})
|
||||
assert config.intro_type == "none"
|
||||
|
||||
|
||||
class TestIntroOutroConfigProperties:
|
||||
"""has_intro / has_outro 属性."""
|
||||
|
||||
def test_has_intro_video_type(self):
|
||||
cfg = IntroOutroConfig(
|
||||
enabled=True,
|
||||
intro_type="video",
|
||||
intro_video_path="/tmp/a.mp4",
|
||||
)
|
||||
assert cfg.has_intro is True
|
||||
|
||||
def test_has_intro_text_type(self):
|
||||
cfg = IntroOutroConfig(
|
||||
enabled=True,
|
||||
intro_type="text",
|
||||
intro_title="Hi",
|
||||
)
|
||||
assert cfg.has_intro is True
|
||||
class TestHasIntroOutro:
|
||||
"""has_intro / has_outro 属性测试."""
|
||||
|
||||
def test_no_intro_when_disabled(self):
|
||||
cfg = IntroOutroConfig(
|
||||
enabled=False,
|
||||
intro_type="video",
|
||||
intro_video_path="/tmp/a.mp4",
|
||||
)
|
||||
assert cfg.has_intro is False
|
||||
"""禁用时无片头."""
|
||||
config = IntroOutroConfig()
|
||||
assert config.has_intro is False
|
||||
assert config.has_outro is False
|
||||
|
||||
def test_no_intro_when_none_type(self):
|
||||
cfg = IntroOutroConfig(
|
||||
def test_video_intro_has_intro(self):
|
||||
"""视频片头有has_intro."""
|
||||
config = IntroOutroConfig(
|
||||
enabled=True,
|
||||
intro_type="none",
|
||||
intro_type="video",
|
||||
intro_video_path="/a.mp4",
|
||||
)
|
||||
assert cfg.has_intro is False
|
||||
assert config.has_intro is True
|
||||
|
||||
def test_has_outro_video_type(self):
|
||||
cfg = IntroOutroConfig(
|
||||
def test_text_intro_has_intro(self):
|
||||
"""文字片头有has_intro."""
|
||||
config = IntroOutroConfig(
|
||||
enabled=True,
|
||||
intro_type="text",
|
||||
intro_title="test",
|
||||
)
|
||||
assert config.has_intro is True
|
||||
|
||||
def test_none_intro_no_intro(self):
|
||||
"""none类型无片头."""
|
||||
config = IntroOutroConfig(enabled=True, intro_type="none")
|
||||
assert config.has_intro is False
|
||||
|
||||
def test_video_outro_has_outro(self):
|
||||
"""视频片尾有has_outro."""
|
||||
config = IntroOutroConfig(
|
||||
enabled=True,
|
||||
outro_type="video",
|
||||
outro_video_path="/tmp/a.mp4",
|
||||
outro_video_path="/a.mp4",
|
||||
)
|
||||
assert cfg.has_outro is True
|
||||
assert config.has_outro is True
|
||||
|
||||
def test_has_outro_text_type(self):
|
||||
cfg = IntroOutroConfig(
|
||||
def test_text_outro_has_outro(self):
|
||||
"""文字片尾有has_outro."""
|
||||
config = IntroOutroConfig(
|
||||
enabled=True,
|
||||
outro_type="text",
|
||||
outro_title="Bye",
|
||||
outro_title="test",
|
||||
)
|
||||
assert cfg.has_outro is True
|
||||
assert config.has_outro is True
|
||||
|
||||
def test_has_outro_follow_type(self):
|
||||
cfg = IntroOutroConfig(
|
||||
def test_follow_outro_has_outro(self):
|
||||
"""follow类型片尾有has_outro."""
|
||||
config = IntroOutroConfig(
|
||||
enabled=True,
|
||||
outro_type="follow",
|
||||
outro_title="Follow me",
|
||||
outro_title="test",
|
||||
)
|
||||
assert cfg.has_outro is True
|
||||
|
||||
def test_no_outro_when_disabled(self):
|
||||
cfg = IntroOutroConfig(
|
||||
enabled=False,
|
||||
outro_type="text",
|
||||
outro_title="Bye",
|
||||
)
|
||||
assert cfg.has_outro is False
|
||||
|
||||
def test_no_outro_when_none_type(self):
|
||||
cfg = IntroOutroConfig(
|
||||
enabled=True,
|
||||
outro_type="none",
|
||||
)
|
||||
assert cfg.has_outro is False
|
||||
assert config.has_outro is True
|
||||
|
||||
|
||||
class TestIntroOutroConfigValidate:
|
||||
"""validate 校验逻辑."""
|
||||
class TestValidate:
|
||||
"""validate 配置校验测试."""
|
||||
|
||||
def test_disabled_is_valid(self):
|
||||
cfg = IntroOutroConfig(enabled=False)
|
||||
ok, msg = cfg.validate()
|
||||
def test_disabled_valid(self):
|
||||
"""禁用配置合法."""
|
||||
config = IntroOutroConfig()
|
||||
ok, msg = config.validate()
|
||||
assert ok is True
|
||||
assert msg == ""
|
||||
|
||||
def test_video_intro_missing_path(self):
|
||||
cfg = IntroOutroConfig(
|
||||
"""视频片头缺少路径."""
|
||||
config = IntroOutroConfig(
|
||||
enabled=True,
|
||||
intro_type="video",
|
||||
intro_video_path="",
|
||||
outro_type="none",
|
||||
)
|
||||
ok, msg = cfg.validate()
|
||||
ok, msg = config.validate()
|
||||
assert ok is False
|
||||
assert "video_path" in msg
|
||||
|
||||
def test_text_intro_missing_title(self):
|
||||
cfg = IntroOutroConfig(
|
||||
"""文字片头缺少标题."""
|
||||
config = IntroOutroConfig(
|
||||
enabled=True,
|
||||
intro_type="text",
|
||||
intro_title="",
|
||||
outro_type="none",
|
||||
)
|
||||
ok, msg = cfg.validate()
|
||||
ok, msg = config.validate()
|
||||
assert ok is False
|
||||
assert "title" in msg
|
||||
|
||||
def test_video_outro_missing_path(self):
|
||||
cfg = IntroOutroConfig(
|
||||
"""视频片尾缺少路径."""
|
||||
config = IntroOutroConfig(
|
||||
enabled=True,
|
||||
intro_type="none",
|
||||
outro_type="video",
|
||||
outro_video_path="",
|
||||
)
|
||||
ok, msg = cfg.validate()
|
||||
ok, msg = config.validate()
|
||||
assert ok is False
|
||||
assert "video_path" in msg
|
||||
|
||||
def test_text_outro_missing_title(self):
|
||||
cfg = IntroOutroConfig(
|
||||
"""文字片尾缺少标题."""
|
||||
config = IntroOutroConfig(
|
||||
enabled=True,
|
||||
intro_type="none",
|
||||
outro_type="text",
|
||||
outro_title="",
|
||||
)
|
||||
ok, msg = cfg.validate()
|
||||
ok, msg = config.validate()
|
||||
assert ok is False
|
||||
assert "title" in msg
|
||||
|
||||
def test_intro_duration_zero(self):
|
||||
cfg = IntroOutroConfig(
|
||||
def test_zero_intro_duration_invalid(self):
|
||||
"""片头时长为0无效."""
|
||||
config = IntroOutroConfig(
|
||||
enabled=True,
|
||||
intro_type="text",
|
||||
intro_title="Hi",
|
||||
intro_title="test",
|
||||
intro_duration=0,
|
||||
outro_type="none",
|
||||
)
|
||||
ok, msg = cfg.validate()
|
||||
ok, msg = config.validate()
|
||||
assert ok is False
|
||||
assert "片头时长" in msg
|
||||
assert "时长" in msg
|
||||
|
||||
def test_intro_duration_negative(self):
|
||||
cfg = IntroOutroConfig(
|
||||
def test_negative_outro_duration_invalid(self):
|
||||
"""片尾时长为负无效."""
|
||||
config = IntroOutroConfig(
|
||||
enabled=True,
|
||||
outro_type="text",
|
||||
outro_title="test",
|
||||
outro_duration=-1.0,
|
||||
)
|
||||
ok, msg = config.validate()
|
||||
assert ok is False
|
||||
assert "时长" in msg
|
||||
|
||||
def test_valid_text_both(self):
|
||||
"""文字片头片尾都合法."""
|
||||
config = IntroOutroConfig(
|
||||
enabled=True,
|
||||
intro_type="text",
|
||||
intro_title="Hi",
|
||||
intro_duration=-1.0,
|
||||
outro_type="none",
|
||||
)
|
||||
ok, msg = cfg.validate()
|
||||
assert ok is False
|
||||
assert "片头时长" in msg
|
||||
|
||||
def test_outro_duration_zero(self):
|
||||
cfg = IntroOutroConfig(
|
||||
enabled=True,
|
||||
intro_type="none",
|
||||
outro_type="text",
|
||||
outro_title="Bye",
|
||||
outro_duration=0,
|
||||
)
|
||||
ok, msg = cfg.validate()
|
||||
assert ok is False
|
||||
assert "片尾时长" in msg
|
||||
|
||||
def test_valid_full_config(self):
|
||||
cfg = IntroOutroConfig(
|
||||
enabled=True,
|
||||
intro_type="video",
|
||||
intro_video_path="/tmp/intro.mp4",
|
||||
intro_title="片头",
|
||||
intro_duration=3.0,
|
||||
outro_type="text",
|
||||
outro_title="Thanks",
|
||||
outro_duration=2.0,
|
||||
outro_title="片尾",
|
||||
outro_duration=3.0,
|
||||
)
|
||||
ok, msg = cfg.validate()
|
||||
ok, msg = config.validate()
|
||||
assert ok is True
|
||||
assert msg == ""
|
||||
|
||||
+257
-442
@@ -1,12 +1,6 @@
|
||||
"""
|
||||
Module Registry 模块注册中心单元测试
|
||||
"""Module Registry 单元测试."""
|
||||
|
||||
覆盖:
|
||||
- ModuleStatus 枚举
|
||||
- QuotaRule / ModuleCapability / Module 数据类
|
||||
- Module.activate / disable 状态转换
|
||||
- ModuleRegistry 注册/注销/查询/能力发现/依赖检查
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import pytest
|
||||
|
||||
@@ -19,551 +13,372 @@ from packages.infrastructure.module_registry import (
|
||||
module_registry,
|
||||
)
|
||||
|
||||
# ============================================================
|
||||
# ModuleStatus
|
||||
# ============================================================
|
||||
|
||||
@pytest.fixture(autouse=True)
|
||||
def clean_registry():
|
||||
"""每个测试前后清空全局单例,避免测试间干扰."""
|
||||
module_registry.clear()
|
||||
yield
|
||||
module_registry.clear()
|
||||
|
||||
|
||||
class TestModuleStatus:
|
||||
"""ModuleStatus 枚举"""
|
||||
|
||||
def test_enum_values(self):
|
||||
assert ModuleStatus.REGISTERED.value == "registered"
|
||||
assert ModuleStatus.ACTIVE.value == "active"
|
||||
assert ModuleStatus.DISABLED.value == "disabled"
|
||||
assert ModuleStatus.ERROR.value == "error"
|
||||
|
||||
def test_is_str_enum(self):
|
||||
assert isinstance(ModuleStatus.ACTIVE, str)
|
||||
assert ModuleStatus.ACTIVE == "active"
|
||||
|
||||
def test_has_four_states(self):
|
||||
assert len(ModuleStatus) == 4
|
||||
# ── Module 数据类测试 ──────────────────────────────────────────────────
|
||||
|
||||
|
||||
# ============================================================
|
||||
# QuotaRule
|
||||
# ============================================================
|
||||
class TestModuleDataclass:
|
||||
"""Module 数据类基本行为测试."""
|
||||
|
||||
|
||||
class TestQuotaRule:
|
||||
"""QuotaRule 配额规则"""
|
||||
|
||||
def test_required_fields(self):
|
||||
rule = QuotaRule(dimension="ai_credits", per_operation=1.0)
|
||||
assert rule.dimension == "ai_credits"
|
||||
assert rule.per_operation == 1.0
|
||||
|
||||
def test_default_description_empty(self):
|
||||
rule = QuotaRule(dimension="storage_gb", per_operation=0.5)
|
||||
assert rule.description == ""
|
||||
|
||||
def test_custom_description(self):
|
||||
rule = QuotaRule(
|
||||
dimension="credits",
|
||||
per_operation=2.0,
|
||||
description="每次生成消耗2积分",
|
||||
)
|
||||
assert rule.description == "每次生成消耗2积分"
|
||||
|
||||
def test_float_per_operation(self):
|
||||
rule = QuotaRule(dimension="gb", per_operation=0.25)
|
||||
assert rule.per_operation == 0.25
|
||||
|
||||
|
||||
# ============================================================
|
||||
# ModuleCapability
|
||||
# ============================================================
|
||||
|
||||
|
||||
class TestModuleCapability:
|
||||
"""ModuleCapability 能力定义"""
|
||||
|
||||
def test_required_name(self):
|
||||
cap = ModuleCapability(name="generate_voice")
|
||||
assert cap.name == "generate_voice"
|
||||
|
||||
def test_defaults(self):
|
||||
cap = ModuleCapability(name="test_cap")
|
||||
assert cap.description == ""
|
||||
assert cap.quota_rules == []
|
||||
assert cap.metadata == {}
|
||||
|
||||
def test_with_quota_rules(self):
|
||||
rules = [QuotaRule(dimension="credits", per_operation=1.0)]
|
||||
cap = ModuleCapability(
|
||||
name="generate",
|
||||
description="生成功能",
|
||||
quota_rules=rules,
|
||||
)
|
||||
assert cap.description == "生成功能"
|
||||
assert len(cap.quota_rules) == 1
|
||||
assert cap.quota_rules[0].dimension == "credits"
|
||||
|
||||
def test_with_metadata(self):
|
||||
cap = ModuleCapability(
|
||||
name="export",
|
||||
metadata={"format": "mp4", "max_resolution": "1080p"},
|
||||
)
|
||||
assert cap.metadata["format"] == "mp4"
|
||||
assert cap.metadata["max_resolution"] == "1080p"
|
||||
|
||||
|
||||
# ============================================================
|
||||
# Module
|
||||
# ============================================================
|
||||
|
||||
|
||||
class TestModuleDefaults:
|
||||
"""Module 数据类默认值"""
|
||||
|
||||
def test_required_name(self):
|
||||
mod = Module(name="ai_voice")
|
||||
assert mod.name == "ai_voice"
|
||||
|
||||
def test_default_version(self):
|
||||
mod = Module(name="test")
|
||||
def test_create_module_defaults(self):
|
||||
"""创建模块,默认值正确."""
|
||||
mod = Module(name="test_module")
|
||||
assert mod.name == "test_module"
|
||||
assert mod.version == "1.0.0"
|
||||
|
||||
def test_default_description(self):
|
||||
mod = Module(name="test")
|
||||
assert mod.description == ""
|
||||
|
||||
def test_default_capabilities_empty(self):
|
||||
mod = Module(name="test")
|
||||
assert mod.capabilities == []
|
||||
|
||||
def test_default_dependencies_empty(self):
|
||||
mod = Module(name="test")
|
||||
assert mod.dependencies == []
|
||||
|
||||
def test_default_status_registered(self):
|
||||
mod = Module(name="test")
|
||||
assert mod.status == ModuleStatus.REGISTERED
|
||||
|
||||
def test_default_config_empty(self):
|
||||
mod = Module(name="test")
|
||||
assert mod.config == {}
|
||||
|
||||
def test_full_module(self):
|
||||
cap = ModuleCapability(name="do_something")
|
||||
def test_create_module_full(self):
|
||||
"""创建模块,完整参数."""
|
||||
mod = Module(
|
||||
name="full_module",
|
||||
name="ai_voice",
|
||||
version="2.0.0",
|
||||
description="完整模块",
|
||||
capabilities=[cap],
|
||||
dependencies=["dep1", "dep2"],
|
||||
description="AI配音模块",
|
||||
capabilities=[ModuleCapability(name="gen_voice")],
|
||||
dependencies=["core"],
|
||||
status=ModuleStatus.ACTIVE,
|
||||
config={"key": "value"},
|
||||
)
|
||||
assert mod.name == "ai_voice"
|
||||
assert mod.version == "2.0.0"
|
||||
assert mod.description == "完整模块"
|
||||
assert mod.description == "AI配音模块"
|
||||
assert len(mod.capabilities) == 1
|
||||
assert mod.dependencies == ["dep1", "dep2"]
|
||||
assert mod.dependencies == ["core"]
|
||||
assert mod.status == ModuleStatus.ACTIVE
|
||||
assert mod.config["key"] == "value"
|
||||
assert mod.config == {"key": "value"}
|
||||
|
||||
|
||||
class TestModuleActivate:
|
||||
"""Module.activate 状态转换"""
|
||||
|
||||
def test_activate_from_registered(self):
|
||||
mod = Module(name="test")
|
||||
def test_module_activate(self):
|
||||
"""激活模块."""
|
||||
mod = Module(name="m1")
|
||||
assert mod.status == ModuleStatus.REGISTERED
|
||||
mod.activate()
|
||||
assert mod.status == ModuleStatus.ACTIVE
|
||||
|
||||
def test_activate_from_disabled(self):
|
||||
mod = Module(name="test", status=ModuleStatus.DISABLED)
|
||||
def test_module_activate_error_state_ignored(self):
|
||||
"""error状态的模块不能激活."""
|
||||
mod = Module(name="m1", status=ModuleStatus.ERROR)
|
||||
mod.activate()
|
||||
assert mod.status == ModuleStatus.ACTIVE
|
||||
|
||||
def test_activate_from_error_stays_error(self):
|
||||
mod = Module(name="test", status=ModuleStatus.ERROR)
|
||||
mod.activate()
|
||||
# error 状态不可激活
|
||||
assert mod.status == ModuleStatus.ERROR
|
||||
|
||||
def test_activate_already_active(self):
|
||||
mod = Module(name="test", status=ModuleStatus.ACTIVE)
|
||||
mod.activate()
|
||||
assert mod.status == ModuleStatus.ACTIVE
|
||||
|
||||
|
||||
class TestModuleDisable:
|
||||
"""Module.disable 状态转换"""
|
||||
|
||||
def test_disable_from_registered(self):
|
||||
mod = Module(name="test")
|
||||
mod.disable()
|
||||
assert mod.status == ModuleStatus.DISABLED
|
||||
|
||||
def test_disable_from_active(self):
|
||||
mod = Module(name="test", status=ModuleStatus.ACTIVE)
|
||||
mod.disable()
|
||||
assert mod.status == ModuleStatus.DISABLED
|
||||
|
||||
def test_disable_from_error(self):
|
||||
mod = Module(name="test", status=ModuleStatus.ERROR)
|
||||
mod.disable()
|
||||
assert mod.status == ModuleStatus.DISABLED
|
||||
|
||||
def test_disable_already_disabled(self):
|
||||
mod = Module(name="test", status=ModuleStatus.DISABLED)
|
||||
def test_module_disable(self):
|
||||
"""禁用模块."""
|
||||
mod = Module(name="m1", status=ModuleStatus.ACTIVE)
|
||||
mod.disable()
|
||||
assert mod.status == ModuleStatus.DISABLED
|
||||
|
||||
|
||||
# ============================================================
|
||||
# ModuleRegistry - 基础操作
|
||||
# ============================================================
|
||||
class TestQuotaRule:
|
||||
"""QuotaRule 测试."""
|
||||
|
||||
def test_quota_rule_basic(self):
|
||||
"""基本配额规则."""
|
||||
rule = QuotaRule(dimension="credits", per_operation=1.0, description="每次消耗1积分")
|
||||
assert rule.dimension == "credits"
|
||||
assert rule.per_operation == 1.0
|
||||
assert rule.description == "每次消耗1积分"
|
||||
|
||||
def test_quota_rule_default_description(self):
|
||||
"""默认描述为空."""
|
||||
rule = QuotaRule(dimension="storage_gb", per_operation=0.5)
|
||||
assert rule.description == ""
|
||||
|
||||
|
||||
class TestModuleRegistryBasic:
|
||||
"""ModuleRegistry 基础操作"""
|
||||
class TestModuleCapability:
|
||||
"""ModuleCapability 测试."""
|
||||
|
||||
def test_empty_registry(self):
|
||||
registry = ModuleRegistry()
|
||||
assert registry.list_modules() == []
|
||||
assert registry.get_active_capabilities() == {}
|
||||
def test_capability_basic(self):
|
||||
"""基本能力定义."""
|
||||
cap = ModuleCapability(name="generate_voice", description="文本转配音")
|
||||
assert cap.name == "generate_voice"
|
||||
assert cap.description == "文本转配音"
|
||||
assert cap.quota_rules == []
|
||||
assert cap.metadata == {}
|
||||
|
||||
def test_capability_with_quota_rules(self):
|
||||
"""带配额规则的能力."""
|
||||
rules = [
|
||||
QuotaRule("ai_credits", 1.0, "配音积分"),
|
||||
QuotaRule("storage_gb", 0.1, "存储占用"),
|
||||
]
|
||||
cap = ModuleCapability(
|
||||
name="generate_voice",
|
||||
quota_rules=rules,
|
||||
metadata={"speed": "fast"},
|
||||
)
|
||||
assert len(cap.quota_rules) == 2
|
||||
assert cap.metadata["speed"] == "fast"
|
||||
|
||||
|
||||
# ── ModuleRegistry 核心测试 ────────────────────────────────────────
|
||||
|
||||
|
||||
class TestModuleRegistryRegister:
|
||||
"""模块注册测试."""
|
||||
|
||||
def test_register_single_module(self):
|
||||
"""注册单个模块."""
|
||||
registry = ModuleRegistry()
|
||||
mod = Module(name="test_mod")
|
||||
registry.register(mod)
|
||||
assert registry.get("test_mod") is mod
|
||||
|
||||
def test_register_duplicate_raises(self):
|
||||
"""重复注册抛异常."""
|
||||
registry = ModuleRegistry()
|
||||
registry.register(Module(name="test_mod"))
|
||||
registry.register(Module(name="m1"))
|
||||
with pytest.raises(ValueError, match="already registered"):
|
||||
registry.register(Module(name="test_mod"))
|
||||
registry.register(Module(name="m1"))
|
||||
|
||||
def test_get_nonexistent_returns_none(self):
|
||||
def test_register_auto_activate_no_deps(self):
|
||||
"""无依赖的模块注册后自动激活."""
|
||||
registry = ModuleRegistry()
|
||||
assert registry.get("no_such_module") is None
|
||||
registry.register(Module(name="m1"))
|
||||
assert registry.get("m1").status == ModuleStatus.ACTIVE
|
||||
|
||||
def test_unregister_success(self):
|
||||
def test_register_with_missing_dependency(self):
|
||||
"""有未满足依赖的模块保持REGISTERED."""
|
||||
registry = ModuleRegistry()
|
||||
registry.register(Module(name="test_mod"))
|
||||
registry.unregister("test_mod")
|
||||
assert registry.get("test_mod") is None
|
||||
registry.register(Module(name="m2", dependencies=["m1"]))
|
||||
assert registry.get("m2").status == ModuleStatus.REGISTERED
|
||||
|
||||
def test_register_with_satisfied_dependency(self):
|
||||
"""依赖已满足的模块注册后自动激活."""
|
||||
registry = ModuleRegistry()
|
||||
registry.register(Module(name="m1"))
|
||||
registry.register(Module(name="m2", dependencies=["m1"]))
|
||||
assert registry.get("m2").status == ModuleStatus.ACTIVE
|
||||
|
||||
|
||||
class TestModuleRegistryUnregister:
|
||||
"""模块注销测试."""
|
||||
|
||||
def test_unregister_existing(self):
|
||||
"""注销已存在的模块."""
|
||||
registry = ModuleRegistry()
|
||||
registry.register(Module(name="m1"))
|
||||
registry.unregister("m1")
|
||||
assert registry.get("m1") is None
|
||||
|
||||
def test_unregister_nonexistent_raises(self):
|
||||
"""注销不存在的模块抛异常."""
|
||||
registry = ModuleRegistry()
|
||||
with pytest.raises(KeyError, match="not found"):
|
||||
registry.unregister("no_such_module")
|
||||
registry.unregister("nonexistent")
|
||||
|
||||
def test_unregister_with_dependents_raises(self):
|
||||
"""被其他模块依赖时不能注销."""
|
||||
registry = ModuleRegistry()
|
||||
registry.register(Module(name="base_module"))
|
||||
registry.register(Module(name="dependent_module", dependencies=["base_module"]))
|
||||
registry.register(Module(name="core"))
|
||||
registry.register(Module(name="plugin", dependencies=["core"]))
|
||||
with pytest.raises(ValueError, match="depended on by"):
|
||||
registry.unregister("base_module")
|
||||
registry.unregister("core")
|
||||
|
||||
def test_clear(self):
|
||||
|
||||
class TestModuleRegistryQuery:
|
||||
"""模块查询测试."""
|
||||
|
||||
def test_get_nonexistent_returns_none(self):
|
||||
"""获取不存在的模块返回None."""
|
||||
registry = ModuleRegistry()
|
||||
registry.register(Module(name="mod1"))
|
||||
registry.register(Module(name="mod2"))
|
||||
registry.clear()
|
||||
assert registry.list_modules() == []
|
||||
assert registry.get("nonexistent") is None
|
||||
|
||||
|
||||
# ============================================================
|
||||
# ModuleRegistry - 自动激活 & 依赖
|
||||
# ============================================================
|
||||
|
||||
|
||||
class TestModuleRegistryAutoActivate:
|
||||
"""注册时自动激活逻辑"""
|
||||
|
||||
def test_no_deps_auto_activates(self):
|
||||
def test_list_modules_all(self):
|
||||
"""列出所有模块."""
|
||||
registry = ModuleRegistry()
|
||||
mod = Module(name="standalone")
|
||||
registry.register(mod)
|
||||
assert mod.status == ModuleStatus.ACTIVE
|
||||
registry.register(Module(name="m1"))
|
||||
registry.register(Module(name="m2"))
|
||||
assert len(registry.list_modules()) == 2
|
||||
|
||||
def test_with_deps_all_satisfied_auto_activates(self):
|
||||
def test_list_modules_by_status(self):
|
||||
"""按状态过滤模块."""
|
||||
registry = ModuleRegistry()
|
||||
registry.register(Module(name="base")) # 无依赖,自动激活
|
||||
dep_mod = Module(name="dependent", dependencies=["base"])
|
||||
registry.register(dep_mod)
|
||||
assert dep_mod.status == ModuleStatus.ACTIVE
|
||||
|
||||
def test_with_deps_not_satisfied_stays_registered(self):
|
||||
registry = ModuleRegistry()
|
||||
mod = Module(name="dependent", dependencies=["missing_dep"])
|
||||
registry.register(mod)
|
||||
# 依赖不满足,保持 REGISTERED
|
||||
assert mod.status == ModuleStatus.REGISTERED
|
||||
|
||||
def test_later_dep_registered_manual_activate(self):
|
||||
"""先注册依赖模块,再注册被依赖模块时不自动激活前者
|
||||
(需要手动或在注册完所有模块后调用 check_dependencies + activate)"""
|
||||
registry = ModuleRegistry()
|
||||
# 先注册依赖方(依赖未满足,不激活)
|
||||
dependent = Module(name="dependent", dependencies=["base"])
|
||||
registry.register(dependent)
|
||||
assert dependent.status == ModuleStatus.REGISTERED
|
||||
|
||||
# 再注册被依赖方
|
||||
base = Module(name="base")
|
||||
registry.register(base)
|
||||
assert base.status == ModuleStatus.ACTIVE
|
||||
|
||||
# 依赖方仍然是 REGISTERED(不会自动激活)
|
||||
assert dependent.status == ModuleStatus.REGISTERED
|
||||
|
||||
|
||||
class TestModuleRegistryCheckDependencies:
|
||||
"""check_dependencies 依赖检查"""
|
||||
|
||||
def test_module_not_found_returns_false(self):
|
||||
registry = ModuleRegistry()
|
||||
assert registry.check_dependencies("nonexistent") is False
|
||||
|
||||
def test_no_deps_returns_true(self):
|
||||
registry = ModuleRegistry()
|
||||
registry.register(Module(name="standalone"))
|
||||
assert registry.check_dependencies("standalone") is True
|
||||
|
||||
def test_all_deps_active_returns_true(self):
|
||||
registry = ModuleRegistry()
|
||||
registry.register(Module(name="dep1"))
|
||||
registry.register(Module(name="dep2"))
|
||||
registry.register(Module(name="main", dependencies=["dep1", "dep2"]))
|
||||
# main 在注册时因依赖满足已自动激活
|
||||
assert registry.check_dependencies("main") is True
|
||||
|
||||
def test_dep_not_registered_returns_false(self):
|
||||
registry = ModuleRegistry()
|
||||
mod = Module(name="main", dependencies=["missing"])
|
||||
registry.register(mod)
|
||||
assert registry.check_dependencies("main") is False
|
||||
|
||||
def test_dep_registered_but_not_active_returns_false(self):
|
||||
registry = ModuleRegistry()
|
||||
dep = Module(name="dep", status=ModuleStatus.DISABLED)
|
||||
registry.register(dep)
|
||||
# 手动设为 disabled(因为 register 时无依赖会自动激活)
|
||||
dep.disable()
|
||||
main = Module(name="main", dependencies=["dep"])
|
||||
registry.register(main)
|
||||
# 依赖未激活
|
||||
assert registry.check_dependencies("main") is False
|
||||
|
||||
|
||||
# ============================================================
|
||||
# ModuleRegistry - list_modules & 状态过滤
|
||||
# ============================================================
|
||||
|
||||
|
||||
class TestModuleRegistryList:
|
||||
"""list_modules 列表与过滤"""
|
||||
|
||||
def test_list_all(self):
|
||||
registry = ModuleRegistry()
|
||||
registry.register(Module(name="mod1"))
|
||||
registry.register(Module(name="mod2"))
|
||||
modules = registry.list_modules()
|
||||
assert len(modules) == 2
|
||||
names = {m.name for m in modules}
|
||||
assert names == {"mod1", "mod2"}
|
||||
|
||||
def test_filter_by_active(self):
|
||||
registry = ModuleRegistry()
|
||||
registry.register(Module(name="active_mod")) # 自动激活
|
||||
disabled = Module(name="disabled_mod")
|
||||
registry.register(disabled)
|
||||
disabled.disable()
|
||||
|
||||
registry.register(Module(name="m1")) # ACTIVE
|
||||
m2 = Module(name="m2", status=ModuleStatus.DISABLED)
|
||||
registry.register(m2)
|
||||
m2.disable()
|
||||
active = registry.list_modules(status=ModuleStatus.ACTIVE)
|
||||
assert len(active) == 1
|
||||
assert active[0].name == "active_mod"
|
||||
assert active[0].name == "m1"
|
||||
|
||||
def test_filter_by_disabled(self):
|
||||
def test_list_modules_disabled(self):
|
||||
"""列出已禁用模块."""
|
||||
registry = ModuleRegistry()
|
||||
registry.register(Module(name="active_mod"))
|
||||
disabled = Module(name="disabled_mod")
|
||||
registry.register(disabled)
|
||||
disabled.disable()
|
||||
|
||||
disabled_list = registry.list_modules(status=ModuleStatus.DISABLED)
|
||||
assert len(disabled_list) == 1
|
||||
assert disabled_list[0].name == "disabled_mod"
|
||||
|
||||
def test_filter_registered(self):
|
||||
registry = ModuleRegistry()
|
||||
# 有依赖未满足的模块保持 REGISTERED
|
||||
mod = Module(name="waiting_mod", dependencies=["missing"])
|
||||
registry.register(mod)
|
||||
|
||||
registered = registry.list_modules(status=ModuleStatus.REGISTERED)
|
||||
assert len(registered) == 1
|
||||
assert registered[0].name == "waiting_mod"
|
||||
|
||||
|
||||
# ============================================================
|
||||
# ModuleRegistry - 能力发现
|
||||
# ============================================================
|
||||
registry.register(Module(name="m1"))
|
||||
m2 = Module(name="m2")
|
||||
registry.register(m2)
|
||||
m2.disable()
|
||||
disabled = registry.list_modules(status=ModuleStatus.DISABLED)
|
||||
assert len(disabled) == 1
|
||||
assert disabled[0].name == "m2"
|
||||
|
||||
|
||||
class TestModuleRegistryCapabilities:
|
||||
"""能力发现:has_capability / get_capability / get_quota_rules"""
|
||||
"""能力查询测试."""
|
||||
|
||||
def test_has_capability_true(self):
|
||||
"""检查已存在的能力."""
|
||||
registry = ModuleRegistry()
|
||||
registry.register(
|
||||
Module(
|
||||
name="voice_module",
|
||||
name="ai_mod",
|
||||
capabilities=[ModuleCapability(name="generate_voice")],
|
||||
)
|
||||
)
|
||||
assert registry.has_capability("generate_voice") is True
|
||||
|
||||
def test_has_capability_false(self):
|
||||
"""检查不存在的能力."""
|
||||
registry = ModuleRegistry()
|
||||
registry.register(
|
||||
Module(
|
||||
name="voice_module",
|
||||
capabilities=[ModuleCapability(name="generate_voice")],
|
||||
)
|
||||
registry.register(Module(name="m1"))
|
||||
assert registry.has_capability("nonexistent") is False
|
||||
|
||||
def test_has_capability_inactive_module(self):
|
||||
"""非激活模块的能力不计入."""
|
||||
registry = ModuleRegistry()
|
||||
m = Module(
|
||||
name="ai_mod",
|
||||
status=ModuleStatus.DISABLED,
|
||||
capabilities=[ModuleCapability(name="generate_voice")],
|
||||
)
|
||||
assert registry.has_capability("generate_video") is False
|
||||
registry._modules["ai_mod"] = m
|
||||
assert registry.has_capability("generate_voice") is False
|
||||
|
||||
def test_has_capability_inactive_module_not_counted(self):
|
||||
def test_get_capability_returns_definition(self):
|
||||
"""获取能力定义."""
|
||||
registry = ModuleRegistry()
|
||||
mod = Module(
|
||||
name="inactive_mod",
|
||||
capabilities=[ModuleCapability(name="secret_cap")],
|
||||
)
|
||||
registry.register(mod)
|
||||
mod.disable()
|
||||
assert registry.has_capability("secret_cap") is False
|
||||
|
||||
def test_get_capability_returns_first_match(self):
|
||||
registry = ModuleRegistry()
|
||||
cap1 = ModuleCapability(name="export", description="导出1")
|
||||
cap2 = ModuleCapability(name="export", description="导出2")
|
||||
registry.register(Module(name="mod1", capabilities=[cap1]))
|
||||
registry.register(Module(name="mod2", capabilities=[cap2]))
|
||||
|
||||
result = registry.get_capability("export")
|
||||
cap = ModuleCapability(name="gen_voice", description="配音")
|
||||
registry.register(Module(name="ai_mod", capabilities=[cap]))
|
||||
result = registry.get_capability("gen_voice")
|
||||
assert result is not None
|
||||
assert result.name == "export"
|
||||
# 返回第一个匹配的(mod1)
|
||||
assert result.description == "导出1"
|
||||
assert result.name == "gen_voice"
|
||||
assert result.description == "配音"
|
||||
|
||||
def test_get_capability_nonexistent_returns_none(self):
|
||||
def test_get_capability_nonexistent(self):
|
||||
"""获取不存在的能力返回None."""
|
||||
registry = ModuleRegistry()
|
||||
assert registry.get_capability("no_such_cap") is None
|
||||
assert registry.get_capability("nonexistent") is None
|
||||
|
||||
def test_get_quota_rules(self):
|
||||
rules = [
|
||||
QuotaRule(dimension="credits", per_operation=1.0),
|
||||
QuotaRule(dimension="storage", per_operation=0.5),
|
||||
]
|
||||
def test_get_quota_rules_empty(self):
|
||||
"""没有配额规则时返回空列表."""
|
||||
registry = ModuleRegistry()
|
||||
registry.register(
|
||||
Module(
|
||||
name="voice_mod",
|
||||
capabilities=[ModuleCapability(name="gen", quota_rules=rules)],
|
||||
name="m1",
|
||||
capabilities=[ModuleCapability(name="do_something")],
|
||||
)
|
||||
)
|
||||
result = registry.get_quota_rules("gen")
|
||||
assert len(result) == 2
|
||||
rules = registry.get_quota_rules("do_something")
|
||||
assert rules == []
|
||||
|
||||
def test_get_quota_rules_with_rules(self):
|
||||
"""获取配额规则."""
|
||||
registry = ModuleRegistry()
|
||||
rules = [QuotaRule("credits", 2.0)]
|
||||
registry.register(
|
||||
Module(
|
||||
name="m1",
|
||||
capabilities=[ModuleCapability(name="do_something", quota_rules=rules)],
|
||||
)
|
||||
)
|
||||
result = registry.get_quota_rules("do_something")
|
||||
assert len(result) == 1
|
||||
assert result[0].dimension == "credits"
|
||||
assert result[1].dimension == "storage"
|
||||
assert result[0].per_operation == 2.0
|
||||
|
||||
def test_get_quota_rules_nonexistent_returns_empty(self):
|
||||
registry = ModuleRegistry()
|
||||
assert registry.get_quota_rules("no_cap") == []
|
||||
|
||||
|
||||
# ============================================================
|
||||
# ModuleRegistry - get_active_capabilities
|
||||
# ============================================================
|
||||
|
||||
|
||||
class TestModuleRegistryActiveCapabilities:
|
||||
"""get_active_capabilities 已激活能力汇总"""
|
||||
|
||||
def test_empty_registry(self):
|
||||
registry = ModuleRegistry()
|
||||
assert registry.get_active_capabilities() == {}
|
||||
|
||||
def test_single_module_with_caps(self):
|
||||
def test_get_active_capabilities(self):
|
||||
"""获取所有已激活模块的能力."""
|
||||
registry = ModuleRegistry()
|
||||
registry.register(
|
||||
Module(
|
||||
name="voice_mod",
|
||||
name="mod_a",
|
||||
capabilities=[
|
||||
ModuleCapability(name="generate_voice"),
|
||||
ModuleCapability(name="clone_voice"),
|
||||
ModuleCapability(name="cap_a1"),
|
||||
ModuleCapability(name="cap_a2"),
|
||||
],
|
||||
)
|
||||
)
|
||||
result = registry.get_active_capabilities()
|
||||
assert "voice_mod" in result
|
||||
assert set(result["voice_mod"]) == {"generate_voice", "clone_voice"}
|
||||
|
||||
def test_skips_inactive_modules(self):
|
||||
registry = ModuleRegistry()
|
||||
registry.register(
|
||||
Module(
|
||||
name="active_mod",
|
||||
capabilities=[ModuleCapability(name="active_cap")],
|
||||
)
|
||||
)
|
||||
inactive = Module(
|
||||
name="inactive_mod",
|
||||
capabilities=[ModuleCapability(name="inactive_cap")],
|
||||
)
|
||||
registry.register(inactive)
|
||||
inactive.disable()
|
||||
|
||||
result = registry.get_active_capabilities()
|
||||
assert "active_mod" in result
|
||||
assert "inactive_mod" not in result
|
||||
|
||||
def test_skips_modules_without_caps(self):
|
||||
registry = ModuleRegistry()
|
||||
registry.register(Module(name="no_cap_mod"))
|
||||
result = registry.get_active_capabilities()
|
||||
assert "no_cap_mod" not in result
|
||||
|
||||
def test_multiple_modules(self):
|
||||
registry = ModuleRegistry()
|
||||
registry.register(
|
||||
Module(
|
||||
name="mod1",
|
||||
capabilities=[ModuleCapability(name="cap_a")],
|
||||
)
|
||||
)
|
||||
registry.register(
|
||||
Module(
|
||||
name="mod2",
|
||||
capabilities=[ModuleCapability(name="cap_b"), ModuleCapability(name="cap_c")],
|
||||
name="mod_b",
|
||||
capabilities=[ModuleCapability(name="cap_b1")],
|
||||
)
|
||||
)
|
||||
result = registry.get_active_capabilities()
|
||||
assert len(result) == 2
|
||||
assert result["mod1"] == ["cap_a"]
|
||||
assert set(result["mod2"]) == {"cap_b", "cap_c"}
|
||||
assert "mod_a" in result
|
||||
assert "mod_b" in result
|
||||
assert set(result["mod_a"]) == {"cap_a1", "cap_a2"}
|
||||
assert result["mod_b"] == ["cap_b1"]
|
||||
|
||||
|
||||
# ============================================================
|
||||
# 全局单例
|
||||
# ============================================================
|
||||
class TestModuleRegistryDependencies:
|
||||
"""依赖检查测试."""
|
||||
|
||||
def test_check_dependencies_satisfied(self):
|
||||
"""依赖满足."""
|
||||
registry = ModuleRegistry()
|
||||
registry.register(Module(name="core"))
|
||||
registry.register(Module(name="plugin", dependencies=["core"]))
|
||||
assert registry.check_dependencies("plugin") is True
|
||||
|
||||
def test_check_dependencies_missing(self):
|
||||
"""依赖缺失."""
|
||||
registry = ModuleRegistry()
|
||||
registry.register(Module(name="plugin", dependencies=["core"]))
|
||||
assert registry.check_dependencies("plugin") is False
|
||||
|
||||
def test_check_dependencies_module_not_found(self):
|
||||
"""模块不存在返回False."""
|
||||
registry = ModuleRegistry()
|
||||
assert registry.check_dependencies("nonexistent") is False
|
||||
|
||||
def test_check_dependencies_inactive_dep(self):
|
||||
"""依赖模块未激活."""
|
||||
registry = ModuleRegistry()
|
||||
core = Module(name="core", status=ModuleStatus.DISABLED)
|
||||
registry._modules["core"] = core
|
||||
registry.register(Module(name="plugin", dependencies=["core"]))
|
||||
# 注册plugin时core不是ACTIVE,所以plugin不会自动激活
|
||||
assert registry.check_dependencies("plugin") is False
|
||||
|
||||
|
||||
class TestGlobalSingleton:
|
||||
"""全局 module_registry 单例"""
|
||||
class TestModuleRegistryClear:
|
||||
"""清空注册测试."""
|
||||
|
||||
def test_singleton_exists(self):
|
||||
assert module_registry is not None
|
||||
assert isinstance(module_registry, ModuleRegistry)
|
||||
def test_clear_removes_all(self):
|
||||
"""清空所有模块."""
|
||||
registry = ModuleRegistry()
|
||||
registry.register(Module(name="m1"))
|
||||
registry.register(Module(name="m2"))
|
||||
assert len(registry.list_modules()) == 2
|
||||
registry.clear()
|
||||
assert len(registry.list_modules()) == 0
|
||||
|
||||
def test_singleton_is_same_instance(self):
|
||||
from packages.infrastructure.module_registry import module_registry as mr2
|
||||
def test_global_singleton_clear(self):
|
||||
"""全局单例清空有效."""
|
||||
module_registry.register(Module(name="global_test"))
|
||||
assert module_registry.get("global_test") is not None
|
||||
# fixture 会在每个测试前后清空,这里手动验证
|
||||
module_registry.clear()
|
||||
assert module_registry.get("global_test") is None
|
||||
|
||||
assert module_registry is mr2
|
||||
|
||||
class TestModuleStatus:
|
||||
"""ModuleStatus 枚举测试."""
|
||||
|
||||
def test_status_values(self):
|
||||
"""状态枚举值正确."""
|
||||
assert ModuleStatus.REGISTERED.value == "registered"
|
||||
assert ModuleStatus.ACTIVE.value == "active"
|
||||
assert ModuleStatus.DISABLED.value == "disabled"
|
||||
assert ModuleStatus.ERROR.value == "error"
|
||||
|
||||
@@ -1,9 +1,4 @@
|
||||
"""
|
||||
多轨道混音引擎配置与纯逻辑测试.
|
||||
|
||||
覆盖 AudioTrack.from_dict / MultiTrackMixConfig.from_config_dict / has_effect 等纯逻辑.
|
||||
引擎核心混音方法依赖 FFmpeg,由集成测试覆盖.
|
||||
"""
|
||||
"""多轨道混音单元测试 - 配置解析等纯逻辑."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
@@ -21,383 +16,367 @@ from video_processing.multi_track_mixer import (
|
||||
)
|
||||
|
||||
|
||||
class TestTrackConstants:
|
||||
"""轨道类型常量与默认值."""
|
||||
class TestConstants:
|
||||
"""常量测试."""
|
||||
|
||||
def test_track_types_exist(self):
|
||||
def test_track_types(self):
|
||||
"""5种轨道类型."""
|
||||
assert TRACK_TYPE_MAIN == "main"
|
||||
assert TRACK_TYPE_BGM == "bgm"
|
||||
assert TRACK_TYPE_VOICEOVER == "voiceover"
|
||||
assert TRACK_TYPE_SFX == "sfx"
|
||||
assert TRACK_TYPE_AMBIENT == "ambient"
|
||||
|
||||
def test_max_tracks(self):
|
||||
assert MAX_AUDIO_TRACKS == 8
|
||||
|
||||
def test_default_volumes(self):
|
||||
"""5种默认音量."""
|
||||
assert len(DEFAULT_VOLUMES) == 5
|
||||
assert DEFAULT_VOLUMES[TRACK_TYPE_MAIN] == 1.0
|
||||
assert DEFAULT_VOLUMES[TRACK_TYPE_BGM] == 0.3
|
||||
assert DEFAULT_VOLUMES[TRACK_TYPE_VOICEOVER] == 1.0
|
||||
assert DEFAULT_VOLUMES[TRACK_TYPE_SFX] == 0.7
|
||||
assert DEFAULT_VOLUMES[TRACK_TYPE_AMBIENT] == 0.2
|
||||
|
||||
def test_max_tracks(self):
|
||||
"""最大轨道数."""
|
||||
assert MAX_AUDIO_TRACKS == 8
|
||||
|
||||
|
||||
class TestAudioTrackDefaults:
|
||||
"""AudioTrack 默认值测试."""
|
||||
|
||||
def test_default_values(self):
|
||||
"""默认值正确."""
|
||||
track = AudioTrack(track_id="t1", track_type="bgm", audio_path="/a.mp3")
|
||||
assert track.track_id == "t1"
|
||||
assert track.track_type == "bgm"
|
||||
assert track.audio_path == "/a.mp3"
|
||||
assert track.volume == 1.0
|
||||
assert track.fade_in == 0.0
|
||||
assert track.fade_out == 0.0
|
||||
assert track.start_time == 0.0
|
||||
assert track.duration == 0.0
|
||||
assert track.enabled is True
|
||||
|
||||
|
||||
class TestAudioTrackFromDict:
|
||||
"""AudioTrack.from_dict 构造逻辑."""
|
||||
"""AudioTrack.from_dict 解析测试."""
|
||||
|
||||
def test_basic(self):
|
||||
def test_basic_parsing(self):
|
||||
"""基本解析."""
|
||||
track = AudioTrack.from_dict(
|
||||
{
|
||||
"track_id": "t1",
|
||||
"track_type": "bgm",
|
||||
"audio_path": "/tmp/bgm.mp3",
|
||||
"audio_path": "/bgm.mp3",
|
||||
}
|
||||
)
|
||||
assert track.track_id == "t1"
|
||||
assert track.track_type == "bgm"
|
||||
assert track.audio_path == "/tmp/bgm.mp3"
|
||||
assert track.volume == 0.3 # bgm 默认音量
|
||||
assert track.audio_path == "/bgm.mp3"
|
||||
|
||||
def test_default_volume_by_type_bgm(self):
|
||||
"""bgm默认音量0.3."""
|
||||
track = AudioTrack.from_dict(
|
||||
{
|
||||
"track_id": "t1",
|
||||
"track_type": "bgm",
|
||||
"audio_path": "/a.mp3",
|
||||
}
|
||||
)
|
||||
assert track.volume == 0.3
|
||||
|
||||
def test_default_volume_by_type_sfx(self):
|
||||
"""sfx默认音量0.7."""
|
||||
track = AudioTrack.from_dict(
|
||||
{
|
||||
"track_id": "t1",
|
||||
"track_type": "sfx",
|
||||
"audio_path": "/a.mp3",
|
||||
}
|
||||
)
|
||||
assert track.volume == 0.7
|
||||
|
||||
def test_default_volume_unknown_type(self):
|
||||
"""未知类型默认音量1.0."""
|
||||
track = AudioTrack.from_dict(
|
||||
{
|
||||
"track_id": "t1",
|
||||
"track_type": "unknown_type",
|
||||
"audio_path": "/a.mp3",
|
||||
}
|
||||
)
|
||||
assert track.volume == 1.0
|
||||
|
||||
def test_custom_volume(self):
|
||||
"""自定义音量."""
|
||||
track = AudioTrack.from_dict(
|
||||
{
|
||||
"track_id": "t1",
|
||||
"track_type": "main",
|
||||
"audio_path": "/tmp/main.wav",
|
||||
"volume": 0.8,
|
||||
"track_type": "bgm",
|
||||
"audio_path": "/a.mp3",
|
||||
"volume": 0.5,
|
||||
}
|
||||
)
|
||||
assert track.volume == 0.8
|
||||
assert track.volume == 0.5
|
||||
|
||||
def test_volume_clamped_to_zero(self):
|
||||
def test_volume_clamped_high(self):
|
||||
"""音量上限钳制."""
|
||||
track = AudioTrack.from_dict(
|
||||
{
|
||||
"track_id": "t1",
|
||||
"track_type": "sfx",
|
||||
"audio_path": "/tmp/sfx.wav",
|
||||
"volume": -1.0,
|
||||
}
|
||||
)
|
||||
assert track.volume == 0.0
|
||||
|
||||
def test_volume_clamped_to_max(self):
|
||||
track = AudioTrack.from_dict(
|
||||
{
|
||||
"track_id": "t1",
|
||||
"track_type": "sfx",
|
||||
"audio_path": "/tmp/sfx.wav",
|
||||
"track_type": "bgm",
|
||||
"audio_path": "/a.mp3",
|
||||
"volume": 3.0,
|
||||
}
|
||||
)
|
||||
assert track.volume == 2.0
|
||||
|
||||
def test_invalid_volume_falls_back_to_default(self):
|
||||
def test_volume_clamped_low(self):
|
||||
"""音量下限钳制."""
|
||||
track = AudioTrack.from_dict(
|
||||
{
|
||||
"track_id": "t1",
|
||||
"audio_path": "/a.mp3",
|
||||
"volume": -1.0,
|
||||
}
|
||||
)
|
||||
assert track.volume == 0.0
|
||||
|
||||
def test_volume_invalid_falls_back(self):
|
||||
"""无效音量回退到类型默认值."""
|
||||
track = AudioTrack.from_dict(
|
||||
{
|
||||
"track_id": "t1",
|
||||
"track_type": "bgm",
|
||||
"audio_path": "/tmp/bgm.mp3",
|
||||
"audio_path": "/a.mp3",
|
||||
"volume": "not_a_number",
|
||||
}
|
||||
)
|
||||
assert track.volume == 0.3 # bgm 默认
|
||||
assert track.volume == 0.3
|
||||
|
||||
def test_none_volume_falls_back(self):
|
||||
def test_fade_in(self):
|
||||
"""淡入时长."""
|
||||
track = AudioTrack.from_dict(
|
||||
{
|
||||
"track_id": "t1",
|
||||
"track_type": "voiceover",
|
||||
"audio_path": "/tmp/vo.wav",
|
||||
"volume": None,
|
||||
"audio_path": "/a.mp3",
|
||||
"fade_in": 2.5,
|
||||
}
|
||||
)
|
||||
assert track.volume == 1.0 # voiceover 默认
|
||||
assert track.fade_in == 2.5
|
||||
|
||||
def test_unknown_track_type_default_volume(self):
|
||||
def test_fade_negative_clamped(self):
|
||||
"""负淡入钳制到0."""
|
||||
track = AudioTrack.from_dict(
|
||||
{
|
||||
"track_id": "t1",
|
||||
"track_type": "unknown_type",
|
||||
"audio_path": "/tmp/a.wav",
|
||||
}
|
||||
)
|
||||
assert track.volume == 1.0 # 未知类型默认 1.0
|
||||
|
||||
def test_fade_in_fade_out(self):
|
||||
track = AudioTrack.from_dict(
|
||||
{
|
||||
"track_id": "t1",
|
||||
"track_type": "bgm",
|
||||
"audio_path": "/tmp/bgm.mp3",
|
||||
"fade_in": 1.5,
|
||||
"fade_out": 2.0,
|
||||
}
|
||||
)
|
||||
assert track.fade_in == 1.5
|
||||
assert track.fade_out == 2.0
|
||||
|
||||
def test_negative_fade_clamped(self):
|
||||
track = AudioTrack.from_dict(
|
||||
{
|
||||
"track_id": "t1",
|
||||
"track_type": "bgm",
|
||||
"audio_path": "/tmp/bgm.mp3",
|
||||
"fade_in": -0.5,
|
||||
"fade_out": -1.0,
|
||||
"audio_path": "/a.mp3",
|
||||
"fade_in": -1.0,
|
||||
"fade_out": -2.0,
|
||||
}
|
||||
)
|
||||
assert track.fade_in == 0.0
|
||||
assert track.fade_out == 0.0
|
||||
|
||||
def test_invalid_fade_falls_back(self):
|
||||
def test_start_time(self):
|
||||
"""开始时间."""
|
||||
track = AudioTrack.from_dict(
|
||||
{
|
||||
"track_id": "t1",
|
||||
"track_type": "bgm",
|
||||
"audio_path": "/tmp/bgm.mp3",
|
||||
"fade_in": "abc",
|
||||
"fade_out": None,
|
||||
"audio_path": "/a.mp3",
|
||||
"start_time": 5.5,
|
||||
}
|
||||
)
|
||||
assert track.fade_in == 0.0
|
||||
assert track.fade_out == 0.0
|
||||
assert track.start_time == 5.5
|
||||
|
||||
def test_start_time_and_duration(self):
|
||||
def test_start_time_negative_clamped(self):
|
||||
"""负开始时间钳制到0."""
|
||||
track = AudioTrack.from_dict(
|
||||
{
|
||||
"track_id": "t1",
|
||||
"track_type": "sfx",
|
||||
"audio_path": "/tmp/sfx.wav",
|
||||
"start_time": 5.0,
|
||||
"duration": 3.0,
|
||||
}
|
||||
)
|
||||
assert track.start_time == 5.0
|
||||
assert track.duration == 3.0
|
||||
|
||||
def test_negative_start_time_clamped(self):
|
||||
track = AudioTrack.from_dict(
|
||||
{
|
||||
"track_id": "t1",
|
||||
"track_type": "bgm",
|
||||
"audio_path": "/tmp/bgm.mp3",
|
||||
"start_time": -10.0,
|
||||
"duration": -2.0,
|
||||
"audio_path": "/a.mp3",
|
||||
"start_time": -3.0,
|
||||
}
|
||||
)
|
||||
assert track.start_time == 0.0
|
||||
assert track.duration == 0.0
|
||||
|
||||
def test_invalid_time_values_fall_back(self):
|
||||
def test_disabled_track(self):
|
||||
"""禁用轨道."""
|
||||
track = AudioTrack.from_dict(
|
||||
{
|
||||
"track_id": "t1",
|
||||
"track_type": "bgm",
|
||||
"audio_path": "/tmp/bgm.mp3",
|
||||
"start_time": "invalid",
|
||||
"duration": "bad",
|
||||
}
|
||||
)
|
||||
assert track.start_time == 0.0
|
||||
assert track.duration == 0.0
|
||||
|
||||
def test_enabled_default_true(self):
|
||||
track = AudioTrack.from_dict(
|
||||
{
|
||||
"track_id": "t1",
|
||||
"track_type": "bgm",
|
||||
"audio_path": "/tmp/bgm.mp3",
|
||||
}
|
||||
)
|
||||
assert track.enabled is True
|
||||
|
||||
def test_enabled_can_be_false(self):
|
||||
track = AudioTrack.from_dict(
|
||||
{
|
||||
"track_id": "t1",
|
||||
"track_type": "bgm",
|
||||
"audio_path": "/tmp/bgm.mp3",
|
||||
"audio_path": "/a.mp3",
|
||||
"enabled": False,
|
||||
}
|
||||
)
|
||||
assert track.enabled is False
|
||||
|
||||
def test_invalid_fade_in_falls_back(self):
|
||||
"""无效淡入值回退到0."""
|
||||
track = AudioTrack.from_dict(
|
||||
{
|
||||
"track_id": "t1",
|
||||
"audio_path": "/a.mp3",
|
||||
"fade_in": "fast",
|
||||
}
|
||||
)
|
||||
assert track.fade_in == 0.0
|
||||
|
||||
class TestMultiTrackMixConfigFromDict:
|
||||
"""MultiTrackMixConfig.from_config_dict 构造逻辑."""
|
||||
|
||||
class TestMultiTrackMixConfigDefaults:
|
||||
"""MultiTrackMixConfig 默认值测试."""
|
||||
|
||||
def test_default_values(self):
|
||||
"""默认值正确."""
|
||||
config = MultiTrackMixConfig()
|
||||
assert config.tracks == []
|
||||
assert config.master_volume == 1.0
|
||||
assert config.normalize is True
|
||||
assert config.max_output_volume == 1.5
|
||||
|
||||
|
||||
class TestMultiTrackMixConfigFromConfigDict:
|
||||
"""MultiTrackMixConfig.from_config_dict 测试."""
|
||||
|
||||
def test_none_returns_default(self):
|
||||
cfg = MultiTrackMixConfig.from_config_dict(None)
|
||||
assert cfg.tracks == []
|
||||
assert cfg.master_volume == 1.0
|
||||
assert cfg.normalize is True
|
||||
"""None返回默认配置."""
|
||||
config = MultiTrackMixConfig.from_config_dict(None)
|
||||
assert config.tracks == []
|
||||
assert config.master_volume == 1.0
|
||||
|
||||
def test_empty_dict_returns_default(self):
|
||||
cfg = MultiTrackMixConfig.from_config_dict({})
|
||||
assert cfg.tracks == []
|
||||
|
||||
def test_non_dict_returns_default(self):
|
||||
cfg = MultiTrackMixConfig.from_config_dict([])
|
||||
assert cfg.tracks == []
|
||||
"""空dict返回默认."""
|
||||
config = MultiTrackMixConfig.from_config_dict({})
|
||||
assert config.tracks == []
|
||||
|
||||
def test_single_track(self):
|
||||
cfg = MultiTrackMixConfig.from_config_dict(
|
||||
"""单轨道."""
|
||||
config = MultiTrackMixConfig.from_config_dict(
|
||||
{
|
||||
"tracks": [
|
||||
{
|
||||
"track_id": "bgm1",
|
||||
"track_type": "bgm",
|
||||
"audio_path": "/tmp/bgm.mp3",
|
||||
"volume": 0.5,
|
||||
"audio_path": "/bgm.mp3",
|
||||
},
|
||||
],
|
||||
}
|
||||
)
|
||||
assert len(cfg.tracks) == 1
|
||||
assert cfg.tracks[0].track_id == "bgm1"
|
||||
assert cfg.tracks[0].volume == 0.5
|
||||
assert len(config.tracks) == 1
|
||||
assert config.tracks[0].track_id == "bgm1"
|
||||
|
||||
def test_multiple_tracks(self):
|
||||
cfg = MultiTrackMixConfig.from_config_dict(
|
||||
"""多轨道."""
|
||||
config = MultiTrackMixConfig.from_config_dict(
|
||||
{
|
||||
"tracks": [
|
||||
{"track_id": "m", "track_type": "main", "audio_path": "/tmp/m.wav"},
|
||||
{"track_id": "b", "track_type": "bgm", "audio_path": "/tmp/b.mp3"},
|
||||
{"track_id": "v", "track_type": "voiceover", "audio_path": "/tmp/v.wav"},
|
||||
{"track_id": "t1", "track_type": "bgm", "audio_path": "/a.mp3"},
|
||||
{"track_id": "t2", "track_type": "sfx", "audio_path": "/b.mp3"},
|
||||
],
|
||||
}
|
||||
)
|
||||
assert len(cfg.tracks) == 3
|
||||
assert cfg.tracks[0].track_type == "main"
|
||||
assert cfg.tracks[1].track_type == "bgm"
|
||||
assert cfg.tracks[2].track_type == "voiceover"
|
||||
assert len(config.tracks) == 2
|
||||
|
||||
def test_disabled_tracks_filtered(self):
|
||||
cfg = MultiTrackMixConfig.from_config_dict(
|
||||
def test_skips_disabled_tracks(self):
|
||||
"""跳过禁用轨道."""
|
||||
config = MultiTrackMixConfig.from_config_dict(
|
||||
{
|
||||
"tracks": [
|
||||
{"track_id": "a", "track_type": "sfx", "audio_path": "/tmp/a.wav"},
|
||||
{"track_id": "b", "track_type": "sfx", "audio_path": "/tmp/b.wav", "enabled": False},
|
||||
{"track_id": "c", "track_type": "sfx", "audio_path": "/tmp/c.wav"},
|
||||
{"track_id": "t1", "audio_path": "/a.mp3", "enabled": True},
|
||||
{"track_id": "t2", "audio_path": "/b.mp3", "enabled": False},
|
||||
],
|
||||
}
|
||||
)
|
||||
assert len(cfg.tracks) == 2
|
||||
assert all(t.track_id != "b" for t in cfg.tracks)
|
||||
assert len(config.tracks) == 1
|
||||
assert config.tracks[0].track_id == "t1"
|
||||
|
||||
def test_empty_audio_path_filtered(self):
|
||||
cfg = MultiTrackMixConfig.from_config_dict(
|
||||
def test_skips_no_audio_path(self):
|
||||
"""跳过无audio_path的轨道."""
|
||||
config = MultiTrackMixConfig.from_config_dict(
|
||||
{
|
||||
"tracks": [
|
||||
{"track_id": "valid", "track_type": "sfx", "audio_path": "/tmp/a.wav"},
|
||||
{"track_id": "empty", "track_type": "sfx", "audio_path": ""},
|
||||
{"track_id": "t1", "audio_path": "/a.mp3"},
|
||||
{"track_id": "t2", "audio_path": ""},
|
||||
{"track_id": "t3"},
|
||||
],
|
||||
}
|
||||
)
|
||||
assert len(cfg.tracks) == 1
|
||||
assert cfg.tracks[0].track_id == "valid"
|
||||
assert len(config.tracks) == 1
|
||||
|
||||
def test_invalid_tracks_skipped(self):
|
||||
cfg = MultiTrackMixConfig.from_config_dict(
|
||||
def test_master_volume(self):
|
||||
"""主音量."""
|
||||
config = MultiTrackMixConfig.from_config_dict(
|
||||
{
|
||||
"tracks": [
|
||||
{"track_id": "ok", "track_type": "sfx", "audio_path": "/tmp/a.wav"},
|
||||
"not_a_dict",
|
||||
None,
|
||||
{"no_audio_path": "xxx"},
|
||||
],
|
||||
"master_volume": 0.8,
|
||||
"tracks": [{"track_id": "t1", "audio_path": "/a.mp3"}],
|
||||
}
|
||||
)
|
||||
assert len(cfg.tracks) == 1
|
||||
assert config.master_volume == 0.8
|
||||
|
||||
def test_tracks_not_a_list(self):
|
||||
cfg = MultiTrackMixConfig.from_config_dict(
|
||||
def test_master_volume_clamped(self):
|
||||
"""主音量边界钳制."""
|
||||
config = MultiTrackMixConfig.from_config_dict(
|
||||
{
|
||||
"master_volume": 5.0,
|
||||
"tracks": [{"track_id": "t1", "audio_path": "/a.mp3"}],
|
||||
}
|
||||
)
|
||||
assert config.master_volume == 2.0
|
||||
|
||||
def test_normalize_disabled(self):
|
||||
"""禁用归一化."""
|
||||
config = MultiTrackMixConfig.from_config_dict(
|
||||
{
|
||||
"normalize": False,
|
||||
"tracks": [{"track_id": "t1", "audio_path": "/a.mp3"}],
|
||||
}
|
||||
)
|
||||
assert config.normalize is False
|
||||
|
||||
def test_tracks_not_list_ignored(self):
|
||||
"""tracks不是列表时忽略."""
|
||||
config = MultiTrackMixConfig.from_config_dict(
|
||||
{
|
||||
"tracks": "not_a_list",
|
||||
}
|
||||
)
|
||||
assert cfg.tracks == []
|
||||
assert config.tracks == []
|
||||
|
||||
def test_master_volume(self):
|
||||
cfg = MultiTrackMixConfig.from_config_dict(
|
||||
def test_non_dict_track_skipped(self):
|
||||
"""非dict轨道跳过."""
|
||||
config = MultiTrackMixConfig.from_config_dict(
|
||||
{
|
||||
"tracks": [],
|
||||
"master_volume": 0.8,
|
||||
"tracks": [
|
||||
{"track_id": "t1", "audio_path": "/a.mp3"},
|
||||
"not_a_dict",
|
||||
],
|
||||
}
|
||||
)
|
||||
assert cfg.master_volume == 0.8
|
||||
|
||||
def test_master_volume_clamped(self):
|
||||
cfg = MultiTrackMixConfig.from_config_dict(
|
||||
{
|
||||
"tracks": [],
|
||||
"master_volume": 3.0,
|
||||
}
|
||||
)
|
||||
assert cfg.master_volume == 2.0
|
||||
|
||||
cfg2 = MultiTrackMixConfig.from_config_dict(
|
||||
{
|
||||
"tracks": [],
|
||||
"master_volume": -1.0,
|
||||
}
|
||||
)
|
||||
assert cfg2.master_volume == 0.0
|
||||
|
||||
def test_invalid_master_volume_falls_back(self):
|
||||
cfg = MultiTrackMixConfig.from_config_dict(
|
||||
{
|
||||
"tracks": [],
|
||||
"master_volume": "abc",
|
||||
}
|
||||
)
|
||||
assert cfg.master_volume == 1.0
|
||||
|
||||
def test_normalize_and_max_output(self):
|
||||
cfg = MultiTrackMixConfig.from_config_dict(
|
||||
{
|
||||
"tracks": [],
|
||||
"normalize": False,
|
||||
"max_output_volume": 2.0,
|
||||
}
|
||||
)
|
||||
assert cfg.normalize is False
|
||||
assert cfg.max_output_volume == 2.0
|
||||
|
||||
def test_default_values(self):
|
||||
cfg = MultiTrackMixConfig.from_config_dict({"tracks": []})
|
||||
assert cfg.master_volume == 1.0
|
||||
assert cfg.normalize is True
|
||||
assert cfg.max_output_volume == 1.5
|
||||
assert len(config.tracks) == 1
|
||||
|
||||
|
||||
class TestMultiTrackMixConfigProperties:
|
||||
"""has_effect 属性."""
|
||||
class TestHasEffect:
|
||||
"""has_effect 属性测试."""
|
||||
|
||||
def test_has_effect_with_tracks(self):
|
||||
cfg = MultiTrackMixConfig(
|
||||
def test_no_tracks_no_effect(self):
|
||||
"""无轨道无效果."""
|
||||
config = MultiTrackMixConfig()
|
||||
assert config.has_effect is False
|
||||
|
||||
def test_with_tracks_has_effect(self):
|
||||
"""有轨道有效果."""
|
||||
config = MultiTrackMixConfig(
|
||||
tracks=[
|
||||
AudioTrack(track_id="t1", track_type="bgm", audio_path="/tmp/a.mp3"),
|
||||
AudioTrack(track_id="t1", track_type="bgm", audio_path="/a.mp3"),
|
||||
]
|
||||
)
|
||||
assert cfg.has_effect is True
|
||||
assert config.has_effect is True
|
||||
|
||||
def test_no_effect_empty(self):
|
||||
cfg = MultiTrackMixConfig(tracks=[])
|
||||
assert cfg.has_effect is False
|
||||
|
||||
def test_no_effect_all_disabled(self):
|
||||
cfg = MultiTrackMixConfig(
|
||||
def test_disabled_tracks_no_effect(self):
|
||||
"""所有轨道都禁用无效果."""
|
||||
config = MultiTrackMixConfig(
|
||||
tracks=[
|
||||
AudioTrack(track_id="t1", track_type="bgm", audio_path="/tmp/a.mp3", enabled=False),
|
||||
AudioTrack(track_id="t1", track_type="bgm", audio_path="/a.mp3", enabled=False),
|
||||
]
|
||||
)
|
||||
assert cfg.has_effect is False
|
||||
|
||||
def test_no_effect_empty_paths(self):
|
||||
cfg = MultiTrackMixConfig(
|
||||
tracks=[
|
||||
AudioTrack(track_id="t1", track_type="bgm", audio_path=""),
|
||||
]
|
||||
)
|
||||
assert cfg.has_effect is False
|
||||
assert config.has_effect is False
|
||||
|
||||
Executable
+232
@@ -0,0 +1,232 @@
|
||||
"""降噪引擎单元测试 - 配置解析等纯逻辑."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import pytest
|
||||
from video_processing.noise_reduction_engine import (
|
||||
NoiseReductionConfig,
|
||||
NoiseReductionLevel,
|
||||
)
|
||||
|
||||
|
||||
class TestNoiseReductionLevel:
|
||||
"""降噪等级枚举测试."""
|
||||
|
||||
def test_level_values(self):
|
||||
"""等级枚举值正确."""
|
||||
assert NoiseReductionLevel.LOW.value == "low"
|
||||
assert NoiseReductionLevel.MEDIUM.value == "medium"
|
||||
assert NoiseReductionLevel.HIGH.value == "high"
|
||||
assert NoiseReductionLevel.CUSTOM.value == "custom"
|
||||
|
||||
def test_from_string(self):
|
||||
"""从字符串创建."""
|
||||
assert NoiseReductionLevel("low") == NoiseReductionLevel.LOW
|
||||
assert NoiseReductionLevel("medium") == NoiseReductionLevel.MEDIUM
|
||||
assert NoiseReductionLevel("high") == NoiseReductionLevel.HIGH
|
||||
assert NoiseReductionLevel("custom") == NoiseReductionLevel.CUSTOM
|
||||
|
||||
def test_invalid_string_raises(self):
|
||||
"""无效字符串抛异常."""
|
||||
with pytest.raises(ValueError):
|
||||
NoiseReductionLevel("invalid")
|
||||
|
||||
|
||||
class TestNoiseReductionConfigDefaults:
|
||||
"""默认配置测试."""
|
||||
|
||||
def test_default_values(self):
|
||||
"""默认值正确."""
|
||||
config = NoiseReductionConfig()
|
||||
assert config.enabled is False
|
||||
assert config.level == NoiseReductionLevel.MEDIUM
|
||||
assert config.noise_floor == -25.0
|
||||
assert config.voice_enhance is False
|
||||
|
||||
|
||||
class TestNoiseReductionConfigFromDict:
|
||||
"""from_dict 配置解析测试."""
|
||||
|
||||
def test_none_returns_disabled(self):
|
||||
"""None 返回禁用配置."""
|
||||
config = NoiseReductionConfig.from_dict(None)
|
||||
assert config.enabled is False
|
||||
|
||||
def test_empty_dict_returns_disabled(self):
|
||||
"""空字典返回禁用配置."""
|
||||
config = NoiseReductionConfig.from_dict({})
|
||||
assert config.enabled is False
|
||||
|
||||
def test_disabled_returns_disabled(self):
|
||||
"""enabled=False 返回禁用."""
|
||||
config = NoiseReductionConfig.from_dict({"enabled": False})
|
||||
assert config.enabled is False
|
||||
|
||||
def test_enabled_default_level(self):
|
||||
"""启用时默认等级为 medium."""
|
||||
config = NoiseReductionConfig.from_dict({"enabled": True})
|
||||
assert config.enabled is True
|
||||
assert config.level == NoiseReductionLevel.MEDIUM
|
||||
|
||||
def test_level_low(self):
|
||||
"""low 等级."""
|
||||
config = NoiseReductionConfig.from_dict({"enabled": True, "level": "low"})
|
||||
assert config.level == NoiseReductionLevel.LOW
|
||||
|
||||
def test_level_high(self):
|
||||
"""high 等级."""
|
||||
config = NoiseReductionConfig.from_dict({"enabled": True, "level": "high"})
|
||||
assert config.level == NoiseReductionLevel.HIGH
|
||||
|
||||
def test_level_custom(self):
|
||||
"""custom 等级."""
|
||||
config = NoiseReductionConfig.from_dict({"enabled": True, "level": "custom"})
|
||||
assert config.level == NoiseReductionLevel.CUSTOM
|
||||
|
||||
def test_level_case_insensitive(self):
|
||||
"""等级大小写不敏感."""
|
||||
config = NoiseReductionConfig.from_dict({"enabled": True, "level": "HIGH"})
|
||||
assert config.level == NoiseReductionLevel.HIGH
|
||||
|
||||
def test_invalid_level_falls_back_to_medium(self):
|
||||
"""无效等级 fallback 到 medium."""
|
||||
config = NoiseReductionConfig.from_dict({"enabled": True, "level": "ultra"})
|
||||
assert config.level == NoiseReductionLevel.MEDIUM
|
||||
|
||||
def test_noise_floor_parsed(self):
|
||||
"""噪音阈值解析."""
|
||||
config = NoiseReductionConfig.from_dict(
|
||||
{
|
||||
"enabled": True,
|
||||
"level": "custom",
|
||||
"noise_floor": -30.0,
|
||||
}
|
||||
)
|
||||
assert config.noise_floor == -30.0
|
||||
|
||||
def test_noise_floor_clamped_min(self):
|
||||
"""噪音阈值下限钳制 (-60)."""
|
||||
config = NoiseReductionConfig.from_dict(
|
||||
{
|
||||
"enabled": True,
|
||||
"level": "custom",
|
||||
"noise_floor": -100.0,
|
||||
}
|
||||
)
|
||||
assert config.noise_floor == -60.0
|
||||
|
||||
def test_noise_floor_clamped_max(self):
|
||||
"""噪音阈值上限钳制 (-5)."""
|
||||
config = NoiseReductionConfig.from_dict(
|
||||
{
|
||||
"enabled": True,
|
||||
"level": "custom",
|
||||
"noise_floor": 0.0,
|
||||
}
|
||||
)
|
||||
assert config.noise_floor == -5.0
|
||||
|
||||
def test_noise_floor_boundary_low(self):
|
||||
"""噪音阈值边界值 -60."""
|
||||
config = NoiseReductionConfig.from_dict(
|
||||
{
|
||||
"enabled": True,
|
||||
"level": "custom",
|
||||
"noise_floor": -60.0,
|
||||
}
|
||||
)
|
||||
assert config.noise_floor == -60.0
|
||||
|
||||
def test_noise_floor_boundary_high(self):
|
||||
"""噪音阈值边界值 -5."""
|
||||
config = NoiseReductionConfig.from_dict(
|
||||
{
|
||||
"enabled": True,
|
||||
"level": "custom",
|
||||
"noise_floor": -5.0,
|
||||
}
|
||||
)
|
||||
assert config.noise_floor == -5.0
|
||||
|
||||
def test_invalid_noise_floor_falls_back(self):
|
||||
"""无效噪音阈值 fallback 到默认值."""
|
||||
config = NoiseReductionConfig.from_dict(
|
||||
{
|
||||
"enabled": True,
|
||||
"level": "custom",
|
||||
"noise_floor": "not_a_number",
|
||||
}
|
||||
)
|
||||
assert config.noise_floor == -25.0
|
||||
|
||||
def test_voice_enhance_enabled(self):
|
||||
"""人声增强启用."""
|
||||
config = NoiseReductionConfig.from_dict(
|
||||
{
|
||||
"enabled": True,
|
||||
"voice_enhance": True,
|
||||
}
|
||||
)
|
||||
assert config.voice_enhance is True
|
||||
|
||||
def test_voice_enhance_disabled_default(self):
|
||||
"""人声增强默认禁用."""
|
||||
config = NoiseReductionConfig.from_dict({"enabled": True})
|
||||
assert config.voice_enhance is False
|
||||
|
||||
|
||||
class TestHasEffect:
|
||||
"""has_effect 方法测试."""
|
||||
|
||||
def test_disabled_no_effect(self):
|
||||
"""禁用时无效果."""
|
||||
config = NoiseReductionConfig(enabled=False)
|
||||
assert config.has_effect() is False
|
||||
|
||||
def test_enabled_has_effect(self):
|
||||
"""启用时有效果."""
|
||||
config = NoiseReductionConfig(enabled=True)
|
||||
assert config.has_effect() is True
|
||||
|
||||
|
||||
class TestGetEffectiveNoiseFloor:
|
||||
"""get_effective_noise_floor 方法测试."""
|
||||
|
||||
def test_custom_level_returns_noise_floor(self):
|
||||
"""custom 等级返回配置的 noise_floor."""
|
||||
config = NoiseReductionConfig(
|
||||
enabled=True,
|
||||
level=NoiseReductionLevel.CUSTOM,
|
||||
noise_floor=-35.0,
|
||||
)
|
||||
assert config.get_effective_noise_floor() == -35.0
|
||||
|
||||
def test_low_level_returns_params(self):
|
||||
"""low 等级返回对应参数值."""
|
||||
config = NoiseReductionConfig(
|
||||
enabled=True,
|
||||
level=NoiseReductionLevel.LOW,
|
||||
)
|
||||
result = config.get_effective_noise_floor()
|
||||
assert isinstance(result, float)
|
||||
assert result < 0 # dB值为负数
|
||||
|
||||
def test_medium_level_returns_params(self):
|
||||
"""medium 等级返回对应参数值."""
|
||||
config = NoiseReductionConfig(
|
||||
enabled=True,
|
||||
level=NoiseReductionLevel.MEDIUM,
|
||||
)
|
||||
result = config.get_effective_noise_floor()
|
||||
assert isinstance(result, float)
|
||||
assert result < 0
|
||||
|
||||
def test_high_level_returns_params(self):
|
||||
"""high 等级返回对应参数值."""
|
||||
config = NoiseReductionConfig(
|
||||
enabled=True,
|
||||
level=NoiseReductionLevel.HIGH,
|
||||
)
|
||||
result = config.get_effective_noise_floor()
|
||||
assert isinstance(result, float)
|
||||
assert result < 0
|
||||
Executable
+146
@@ -0,0 +1,146 @@
|
||||
"""OSS助手纯逻辑测试 — normalize_storage_key / resolve_asset_path 输入校验."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import hashlib
|
||||
import os
|
||||
from pathlib import Path
|
||||
from unittest.mock import patch
|
||||
|
||||
import pytest
|
||||
from video_processing.oss_helpers import normalize_storage_key, resolve_asset_path
|
||||
|
||||
|
||||
class TestNormalizeStorageKey:
|
||||
"""normalize_storage_key 存储键标准化测试."""
|
||||
|
||||
def test_plain_key_passthrough(self):
|
||||
"""普通路径原样返回."""
|
||||
assert normalize_storage_key("path/to/file.mp4") == "path/to/file.mp4"
|
||||
|
||||
def test_https_url_extracts_path(self):
|
||||
"""HTTPS URL提取path部分."""
|
||||
result = normalize_storage_key("https://bucket.oss-cn-hangzhou.aliyuncs.com/path/to/file.mp4")
|
||||
assert result == "path/to/file.mp4"
|
||||
|
||||
def test_http_url_extracts_path(self):
|
||||
"""HTTP URL提取path部分."""
|
||||
result = normalize_storage_key("http://example.com/assets/video.mp4")
|
||||
assert result == "assets/video.mp4"
|
||||
|
||||
def test_url_with_query_params(self):
|
||||
"""带query参数的URL只取path."""
|
||||
result = normalize_storage_key("https://bucket.oss-cn-hangzhou.aliyuncs.com/file.mp4?token=abc&expires=123")
|
||||
assert result == "file.mp4"
|
||||
|
||||
def test_leading_slash_stripped(self):
|
||||
"""开头斜杠被去掉."""
|
||||
assert normalize_storage_key("/path/to/file.mp4") == "path/to/file.mp4"
|
||||
|
||||
def test_url_without_path(self):
|
||||
"""URL没有path部分返回空字符串."""
|
||||
result = normalize_storage_key("https://example.com")
|
||||
assert result == ""
|
||||
|
||||
def test_nested_path(self):
|
||||
"""多层嵌套路径."""
|
||||
assert normalize_storage_key("a/b/c/d/file.mp4") == "a/b/c/d/file.mp4"
|
||||
|
||||
def test_empty_string(self):
|
||||
"""空字符串."""
|
||||
assert normalize_storage_key("") == ""
|
||||
|
||||
def test_url_with_port(self):
|
||||
"""带端口的URL."""
|
||||
result = normalize_storage_key("http://localhost:9000/bucket/file.mp4")
|
||||
assert result == "bucket/file.mp4"
|
||||
|
||||
|
||||
class TestResolveAssetPathInputValidation:
|
||||
"""resolve_asset_path 输入校验测试(不涉及真实下载)."""
|
||||
|
||||
def test_empty_string_returns_none(self, tmp_path):
|
||||
"""空字符串返回None."""
|
||||
assert resolve_asset_path("", tmp_path) is None
|
||||
|
||||
def test_none_returns_none(self, tmp_path):
|
||||
"""None返回None(类型检查)."""
|
||||
assert resolve_asset_path(None, tmp_path) is None # type: ignore
|
||||
|
||||
def test_non_string_returns_none(self, tmp_path):
|
||||
"""非字符串返回None."""
|
||||
assert resolve_asset_path(123, tmp_path) is None # type: ignore
|
||||
|
||||
def test_null_byte_rejected(self, tmp_path):
|
||||
"""包含空字节的asset_id被拒绝."""
|
||||
assert resolve_asset_path("file\x00.mp4", tmp_path) is None
|
||||
|
||||
def test_path_traversal_rejected(self, tmp_path):
|
||||
"""包含../的路径遍历攻击被拒绝(第3步下载前检查)."""
|
||||
# mock download_asset不被调用,因为路径包含..会直接返回None
|
||||
with patch("video_processing.oss_helpers.download_asset") as mock_dl:
|
||||
result = resolve_asset_path("../etc/passwd", tmp_path)
|
||||
assert result is None
|
||||
mock_dl.assert_not_called()
|
||||
|
||||
def test_absolute_path_key_rejected(self, tmp_path):
|
||||
"""以/开头的存储键在下载前检查被拒."""
|
||||
with patch("video_processing.oss_helpers.download_asset") as mock_dl:
|
||||
result = resolve_asset_path("/etc/passwd", tmp_path)
|
||||
assert result is None
|
||||
mock_dl.assert_not_called()
|
||||
|
||||
def test_cache_hit_returns_cached_path(self, tmp_path):
|
||||
"""缓存命中返回缓存路径."""
|
||||
asset_id = "test-asset-123"
|
||||
cache_hash = hashlib.sha256(asset_id.encode()).hexdigest()[:16]
|
||||
cached_file = tmp_path / f"{cache_hash}.mp4"
|
||||
cached_file.write_bytes(b"fake video data")
|
||||
|
||||
result = resolve_asset_path(asset_id, tmp_path)
|
||||
assert result == cached_file
|
||||
assert result.exists()
|
||||
|
||||
def test_cache_empty_file_not_considered_hit(self, tmp_path):
|
||||
"""空文件不算缓存命中."""
|
||||
asset_id = "empty-cache-file"
|
||||
cache_hash = hashlib.sha256(asset_id.encode()).hexdigest()[:16]
|
||||
cached_file = tmp_path / f"{cache_hash}.mp4"
|
||||
cached_file.touch() # 空文件
|
||||
|
||||
with patch("video_processing.oss_helpers.download_asset", return_value=False):
|
||||
result = resolve_asset_path(asset_id, tmp_path)
|
||||
# 空文件不命中缓存,走下载,下载失败返回None
|
||||
assert result is None
|
||||
|
||||
def test_download_success_returns_path(self, tmp_path):
|
||||
"""下载成功返回本地路径."""
|
||||
asset_id = "remote-asset"
|
||||
cache_hash = hashlib.sha256(asset_id.encode()).hexdigest()[:16]
|
||||
expected_path = tmp_path / f"{cache_hash}.mp4"
|
||||
|
||||
def fake_download(storage_key, local_path):
|
||||
Path(local_path).write_bytes(b"downloaded data")
|
||||
return True
|
||||
|
||||
with patch("video_processing.oss_helpers.download_asset", side_effect=fake_download):
|
||||
result = resolve_asset_path(asset_id, tmp_path)
|
||||
assert result == expected_path
|
||||
assert result.exists()
|
||||
assert result.stat().st_size > 0
|
||||
|
||||
def test_download_failure_returns_none(self, tmp_path):
|
||||
"""下载失败返回None."""
|
||||
with patch("video_processing.oss_helpers.download_asset", return_value=False):
|
||||
result = resolve_asset_path("nonexistent-asset", tmp_path)
|
||||
assert result is None
|
||||
|
||||
def test_work_dir_not_exists_creates_on_demand(self, tmp_path):
|
||||
"""work_dir不存在时也能处理."""
|
||||
asset_id = "new-dir-asset"
|
||||
new_dir = tmp_path / "subdir" / "nested"
|
||||
|
||||
with patch("video_processing.oss_helpers.download_asset", return_value=False):
|
||||
# 不存在的work_dir,缓存检查也不会命中
|
||||
result = resolve_asset_path(asset_id, new_dir)
|
||||
assert result is None
|
||||
+222
-548
@@ -1,597 +1,271 @@
|
||||
"""画中画(PiP)引擎单元测试."""
|
||||
"""画中画引擎单元测试 - 配置解析+校验等纯逻辑."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from pathlib import Path
|
||||
from unittest.mock import patch
|
||||
|
||||
import pytest
|
||||
from video_processing.pip_engine import (
|
||||
ANIMATION_FADE,
|
||||
ANIMATION_SLIDE_BOTTOM,
|
||||
ANIMATION_SLIDE_LEFT,
|
||||
ANIMATION_SLIDE_RIGHT,
|
||||
ANIMATION_SLIDE_TOP,
|
||||
POSITION_BOTTOM_LEFT,
|
||||
POSITION_BOTTOM_RIGHT,
|
||||
POSITION_CENTER,
|
||||
POSITION_TOP_LEFT,
|
||||
POSITION_TOP_RIGHT,
|
||||
PiPConfig,
|
||||
PiPEngine,
|
||||
PiPLayerConfig,
|
||||
)
|
||||
|
||||
# ── PiPLayerConfig.validate 测试 ──────────────────────────────────────────────
|
||||
from video_processing.pip_engine import PiPConfig, PiPLayerConfig
|
||||
|
||||
|
||||
class TestPiPLayerConfigValidate:
|
||||
"""PiP图层配置校验测试."""
|
||||
|
||||
def test_valid_config(self):
|
||||
"""正常配置应该通过校验."""
|
||||
layer = PiPLayerConfig(source="asset_001")
|
||||
ok, err = layer.validate()
|
||||
assert ok
|
||||
assert err == ""
|
||||
|
||||
def test_empty_source(self):
|
||||
"""空source应该失败."""
|
||||
layer = PiPLayerConfig(source="")
|
||||
ok, err = layer.validate()
|
||||
assert not ok
|
||||
assert "source" in err
|
||||
|
||||
def test_invalid_position(self):
|
||||
"""无效位置应该失败."""
|
||||
layer = PiPLayerConfig(source="asset_001", position="invalid_pos")
|
||||
ok, err = layer.validate()
|
||||
assert not ok
|
||||
assert "position" in err
|
||||
|
||||
def test_custom_position_valid(self):
|
||||
"""custom位置应该通过."""
|
||||
layer = PiPLayerConfig(source="asset_001", position="custom", x=100, y=50)
|
||||
ok, err = layer.validate()
|
||||
assert ok
|
||||
|
||||
def test_opacity_out_of_range_high(self):
|
||||
"""opacity超过1应该失败."""
|
||||
layer = PiPLayerConfig(source="asset_001", opacity=1.5)
|
||||
ok, err = layer.validate()
|
||||
assert not ok
|
||||
assert "opacity" in err
|
||||
|
||||
def test_opacity_out_of_range_low(self):
|
||||
"""opacity小于0应该失败."""
|
||||
layer = PiPLayerConfig(source="asset_001", opacity=-0.5)
|
||||
ok, err = layer.validate()
|
||||
assert not ok
|
||||
assert "opacity" in err
|
||||
|
||||
def test_opacity_boundary_values(self):
|
||||
"""opacity边界值应该通过."""
|
||||
for val in [0.0, 0.5, 1.0]:
|
||||
layer = PiPLayerConfig(source="asset_001", opacity=val)
|
||||
ok, _ = layer.validate()
|
||||
assert ok
|
||||
|
||||
def test_negative_corner_radius(self):
|
||||
"""负圆角应该失败."""
|
||||
layer = PiPLayerConfig(source="asset_001", corner_radius=-5)
|
||||
ok, err = layer.validate()
|
||||
assert not ok
|
||||
assert "corner_radius" in err
|
||||
|
||||
def test_negative_start_time(self):
|
||||
"""负开始时间应该失败."""
|
||||
layer = PiPLayerConfig(source="asset_001", start_time=-1.0)
|
||||
ok, err = layer.validate()
|
||||
assert not ok
|
||||
assert "start_time" in err
|
||||
|
||||
def test_negative_duration(self):
|
||||
"""负持续时间应该失败."""
|
||||
layer = PiPLayerConfig(source="asset_001", duration=-5.0)
|
||||
ok, err = layer.validate()
|
||||
assert not ok
|
||||
assert "duration" in err
|
||||
|
||||
def test_invalid_animation_in(self):
|
||||
"""无效入场动画应该失败."""
|
||||
layer = PiPLayerConfig(source="asset_001", animation_in="spin")
|
||||
ok, err = layer.validate()
|
||||
assert not ok
|
||||
assert "入场动画" in err
|
||||
|
||||
def test_all_valid_animations(self):
|
||||
"""所有有效动画类型应该通过."""
|
||||
for anim in [
|
||||
ANIMATION_FADE,
|
||||
ANIMATION_SLIDE_LEFT,
|
||||
ANIMATION_SLIDE_RIGHT,
|
||||
ANIMATION_SLIDE_TOP,
|
||||
ANIMATION_SLIDE_BOTTOM,
|
||||
]:
|
||||
layer = PiPLayerConfig(source="asset_001", animation_in=anim, animation_out=anim)
|
||||
ok, _ = layer.validate()
|
||||
assert ok
|
||||
|
||||
def test_zero_duration_valid(self):
|
||||
"""duration=0(全程显示)应该通过."""
|
||||
layer = PiPLayerConfig(source="asset_001", duration=0.0)
|
||||
ok, _ = layer.validate()
|
||||
assert ok
|
||||
|
||||
|
||||
# ── PiPConfig.from_dict 测试 ──────────────────────────────────────────────────
|
||||
|
||||
|
||||
class TestPiPConfigFromDict:
|
||||
"""PiP配置字典解析测试."""
|
||||
|
||||
def test_none_config(self):
|
||||
"""None配置应该返回disabled."""
|
||||
config = PiPConfig.from_dict(None)
|
||||
assert not config.enabled
|
||||
assert len(config.layers) == 0
|
||||
|
||||
def test_empty_config(self):
|
||||
"""空字典应该返回disabled."""
|
||||
config = PiPConfig.from_dict({})
|
||||
assert not config.enabled
|
||||
|
||||
def test_enabled_false(self):
|
||||
"""enabled=False应该返回disabled."""
|
||||
config = PiPConfig.from_dict({"enabled": False, "layers": [{"source": "a"}]})
|
||||
assert not config.enabled
|
||||
|
||||
def test_single_layer(self):
|
||||
"""单图层解析."""
|
||||
data = {
|
||||
"enabled": True,
|
||||
"layers": [
|
||||
{
|
||||
"source": "asset_001",
|
||||
"position": POSITION_TOP_RIGHT,
|
||||
"width": "30%",
|
||||
"opacity": 0.9,
|
||||
"corner_radius": 10,
|
||||
"start_time": 2.0,
|
||||
"duration": 5.0,
|
||||
"z_index": 2,
|
||||
}
|
||||
],
|
||||
}
|
||||
config = PiPConfig.from_dict(data)
|
||||
assert config.enabled
|
||||
assert len(config.layers) == 1
|
||||
layer = config.layers[0]
|
||||
assert layer.source == "asset_001"
|
||||
assert layer.position == POSITION_TOP_RIGHT
|
||||
assert layer.width == "30%"
|
||||
assert layer.opacity == 0.9
|
||||
assert layer.corner_radius == 10
|
||||
assert layer.start_time == 2.0
|
||||
assert layer.duration == 5.0
|
||||
assert layer.z_index == 2
|
||||
|
||||
def test_multiple_layers_sorted_by_z_index(self):
|
||||
"""多图层应该按z_index排序."""
|
||||
data = {
|
||||
"enabled": True,
|
||||
"layers": [
|
||||
{"source": "asset_high", "z_index": 5},
|
||||
{"source": "asset_low", "z_index": 1},
|
||||
{"source": "asset_mid", "z_index": 3},
|
||||
],
|
||||
}
|
||||
config = PiPConfig.from_dict(data)
|
||||
assert len(config.layers) == 3
|
||||
assert config.layers[0].source == "asset_low"
|
||||
assert config.layers[1].source == "asset_mid"
|
||||
assert config.layers[2].source == "asset_high"
|
||||
|
||||
def test_invalid_layer_skipped(self):
|
||||
"""无效图层应该被跳过."""
|
||||
data = {
|
||||
"enabled": True,
|
||||
"layers": [
|
||||
{"source": "asset_good"},
|
||||
{"source": "", "position": "invalid"}, # 空source
|
||||
{"source": "asset_good2", "opacity": 2.0}, # opacity超范围
|
||||
],
|
||||
}
|
||||
config = PiPConfig.from_dict(data)
|
||||
# 第1个有效,第2、3个无效
|
||||
assert len(config.layers) == 1
|
||||
assert config.layers[0].source == "asset_good"
|
||||
|
||||
def test_all_invalid_layers_disabled(self):
|
||||
"""所有图层都无效时enabled为False."""
|
||||
data = {
|
||||
"enabled": True,
|
||||
"layers": [
|
||||
{"source": ""},
|
||||
{"source": ""},
|
||||
],
|
||||
}
|
||||
config = PiPConfig.from_dict(data)
|
||||
assert not config.enabled
|
||||
assert len(config.layers) == 0
|
||||
class TestPiPLayerConfigDefaults:
|
||||
"""PiPLayerConfig 默认配置测试."""
|
||||
|
||||
def test_default_values(self):
|
||||
"""默认值应该正确."""
|
||||
data = {
|
||||
"enabled": True,
|
||||
"layers": [{"source": "asset_001"}],
|
||||
}
|
||||
config = PiPConfig.from_dict(data)
|
||||
layer = config.layers[0]
|
||||
assert layer.position == POSITION_BOTTOM_RIGHT
|
||||
"""默认值正确."""
|
||||
layer = PiPLayerConfig()
|
||||
assert layer.source == ""
|
||||
assert layer.source_type == "asset_id"
|
||||
assert layer.position == "bottom_right"
|
||||
assert layer.margin == 20
|
||||
assert layer.width == "25%"
|
||||
assert layer.height == ""
|
||||
assert layer.opacity == 1.0
|
||||
assert layer.corner_radius == 0
|
||||
assert layer.border_width == 0
|
||||
assert layer.border_color == "white"
|
||||
assert layer.start_time == 0.0
|
||||
assert layer.duration == 0.0
|
||||
assert layer.animation_in == ""
|
||||
assert layer.animation_out == ""
|
||||
assert layer.animation_duration == 0.5
|
||||
assert layer.z_index == 1
|
||||
|
||||
|
||||
# ── PiPEngine 位置计算测试 ────────────────────────────────────────────────────
|
||||
class TestPiPLayerConfigValidate:
|
||||
"""PiPLayerConfig.validate 校验测试."""
|
||||
|
||||
def test_valid_config(self):
|
||||
"""合法配置."""
|
||||
layer = PiPLayerConfig(source="asset_123")
|
||||
ok, msg = layer.validate()
|
||||
assert ok is True
|
||||
assert msg == ""
|
||||
|
||||
class TestPiPEnginePosition:
|
||||
"""PiP引擎位置计算测试."""
|
||||
def test_empty_source_invalid(self):
|
||||
"""空source非法."""
|
||||
layer = PiPLayerConfig(source="")
|
||||
ok, msg = layer.validate()
|
||||
assert ok is False
|
||||
assert "source" in msg
|
||||
|
||||
@pytest.fixture
|
||||
def engine(self):
|
||||
return PiPEngine(output_width=1920, output_height=1080, output_fps=30)
|
||||
def test_invalid_position(self):
|
||||
"""无效position."""
|
||||
layer = PiPLayerConfig(source="asset_123", position="invalid_pos")
|
||||
ok, msg = layer.validate()
|
||||
assert ok is False
|
||||
assert "position" in msg
|
||||
|
||||
def test_top_left_position(self, engine):
|
||||
"""左上角位置."""
|
||||
layer = PiPLayerConfig(source="a", position=POSITION_TOP_LEFT, margin=20)
|
||||
x, y = engine._parse_position(layer, 480, 270)
|
||||
assert x == 20
|
||||
assert y == 20
|
||||
|
||||
def test_top_right_position(self, engine):
|
||||
"""右上角位置."""
|
||||
layer = PiPLayerConfig(source="a", position=POSITION_TOP_RIGHT, margin=20)
|
||||
x, y = engine._parse_position(layer, 480, 270)
|
||||
assert x == 1920 - 480 - 20
|
||||
assert y == 20
|
||||
|
||||
def test_bottom_right_position(self, engine):
|
||||
"""右下角位置(默认)."""
|
||||
layer = PiPLayerConfig(source="a", position=POSITION_BOTTOM_RIGHT, margin=30)
|
||||
x, y = engine._parse_position(layer, 480, 270)
|
||||
assert x == 1920 - 480 - 30
|
||||
assert y == 1080 - 270 - 30
|
||||
|
||||
def test_bottom_left_position(self, engine):
|
||||
"""左下角位置."""
|
||||
layer = PiPLayerConfig(source="a", position=POSITION_BOTTOM_LEFT, margin=15)
|
||||
x, y = engine._parse_position(layer, 480, 270)
|
||||
assert x == 15
|
||||
assert y == 1080 - 270 - 15
|
||||
|
||||
def test_center_position(self, engine):
|
||||
"""中心位置."""
|
||||
layer = PiPLayerConfig(source="a", position=POSITION_CENTER, margin=0)
|
||||
x, y = engine._parse_position(layer, 480, 270)
|
||||
assert x == (1920 - 480) // 2
|
||||
assert y == (1080 - 270) // 2
|
||||
|
||||
def test_custom_position_pixel(self, engine):
|
||||
"""自定义像素位置."""
|
||||
layer = PiPLayerConfig(source="a", position="custom", x=100, y=200)
|
||||
x, y = engine._parse_position(layer, 480, 270)
|
||||
assert x == 100
|
||||
assert y == 200
|
||||
|
||||
def test_custom_position_percentage(self, engine):
|
||||
"""自定义百分比位置."""
|
||||
layer = PiPLayerConfig(source="a", position="custom", x="50%", y="25%")
|
||||
x, y = engine._parse_position(layer, 480, 270)
|
||||
assert x == 1920 // 2
|
||||
assert y == 1080 // 4
|
||||
|
||||
def test_top_center_position(self, engine):
|
||||
"""顶部居中位置."""
|
||||
layer = PiPLayerConfig(source="a", position="top_center", margin=10)
|
||||
x, y = engine._parse_position(layer, 480, 270)
|
||||
assert x == (1920 - 480) // 2
|
||||
assert y == 10
|
||||
|
||||
def test_invalid_position_fallback(self, engine):
|
||||
"""无效位置应该fallback到右下角."""
|
||||
layer = PiPLayerConfig(source="a", position="unknown_position", margin=20)
|
||||
# 直接测试_parse_position(注意:validate会拦截,但_parse_position自己也有fallback)
|
||||
x, y = engine._parse_position(layer, 480, 270)
|
||||
assert x == 1920 - 480 - 20
|
||||
assert y == 1080 - 270 - 20
|
||||
|
||||
|
||||
# ── PiPEngine 尺寸解析测试 ────────────────────────────────────────────────────
|
||||
|
||||
|
||||
class TestPiPEngineSize:
|
||||
"""PiP引擎尺寸解析测试."""
|
||||
|
||||
@pytest.fixture
|
||||
def engine(self):
|
||||
return PiPEngine(output_width=1920, output_height=1080, output_fps=30)
|
||||
|
||||
def test_pixel_size_int(self, engine):
|
||||
"""像素尺寸(整数)."""
|
||||
assert engine._parse_size(500, 1920) == 500
|
||||
|
||||
def test_pixel_size_str(self, engine):
|
||||
"""像素尺寸(字符串数字)."""
|
||||
assert engine._parse_size("500", 1920) == 500
|
||||
|
||||
def test_percentage_size(self, engine):
|
||||
"""百分比尺寸."""
|
||||
assert engine._parse_size("50%", 1920) == 960
|
||||
assert engine._parse_size("25%", 1920) == 480
|
||||
|
||||
def test_zero_size_default(self, engine):
|
||||
"""0或无效值应该有最小值保护."""
|
||||
assert engine._parse_size(0, 1920) == 1
|
||||
assert engine._parse_size("", 1920) == 480 # 默认25%
|
||||
|
||||
def test_negative_size_default(self, engine):
|
||||
"""负值应该取绝对值后至少为1."""
|
||||
# _parse_size 用 max(1, value),负值会走 except 分支
|
||||
result = engine._parse_size("-100", 1920)
|
||||
# 会走ValueError分支,返回默认值
|
||||
assert result > 0
|
||||
|
||||
|
||||
# ── PiPEngine 滤镜构建测试 ────────────────────────────────────────────────────
|
||||
|
||||
|
||||
class TestPiPEngineBuildFilters:
|
||||
"""PiP引擎滤镜构建测试."""
|
||||
|
||||
@pytest.fixture
|
||||
def engine(self):
|
||||
return PiPEngine(output_width=1920, output_height=1080, output_fps=30)
|
||||
|
||||
@pytest.fixture
|
||||
def fake_video(self, tmp_path):
|
||||
"""创建一个假的视频文件路径."""
|
||||
path = tmp_path / "test_video.mp4"
|
||||
path.write_bytes(b"fake video data")
|
||||
return path
|
||||
|
||||
def test_empty_sources(self, engine):
|
||||
"""空素材列表应该返回空."""
|
||||
filters, inputs, label = engine.build_pip_filters("base_label", [])
|
||||
assert filters == []
|
||||
assert inputs == []
|
||||
assert label == "base_label"
|
||||
|
||||
def test_single_layer_basic(self, engine, fake_video):
|
||||
"""单图层基础滤镜构建."""
|
||||
def test_custom_position_valid(self):
|
||||
"""custom位置合法."""
|
||||
layer = PiPLayerConfig(
|
||||
source="asset_001",
|
||||
position=POSITION_TOP_RIGHT,
|
||||
width="25%",
|
||||
source="asset_123",
|
||||
position="custom",
|
||||
x=100,
|
||||
y=100,
|
||||
)
|
||||
sources = [("pip_src_0", layer, fake_video)]
|
||||
ok, _ = layer.validate()
|
||||
assert ok is True
|
||||
|
||||
filters, inputs, final_label = engine.build_pip_filters("base_video", sources, base_input_idx=3)
|
||||
def test_opacity_too_high(self):
|
||||
"""透明度超过1."""
|
||||
layer = PiPLayerConfig(source="a", opacity=1.5)
|
||||
ok, msg = layer.validate()
|
||||
assert ok is False
|
||||
assert "opacity" in msg
|
||||
|
||||
# 应该有2个滤镜: 预处理 + overlay
|
||||
assert len(filters) == 2
|
||||
# 输入参数应该有2个(-i + path)
|
||||
assert len(inputs) == 2
|
||||
assert inputs[0] == "-i"
|
||||
assert inputs[1] == str(fake_video)
|
||||
def test_opacity_negative(self):
|
||||
"""透明度为负."""
|
||||
layer = PiPLayerConfig(source="a", opacity=-0.1)
|
||||
ok, msg = layer.validate()
|
||||
assert ok is False
|
||||
assert "opacity" in msg
|
||||
|
||||
# 预处理滤镜应该使用正确的输入索引
|
||||
assert "3:v" in filters[0]
|
||||
# 应该包含scale
|
||||
assert "scale=" in filters[0]
|
||||
# 应该有pip_pre_0标签
|
||||
assert "[pip_pre_0]" in filters[0]
|
||||
def test_opacity_boundary_zero(self):
|
||||
"""透明度边界值0."""
|
||||
layer = PiPLayerConfig(source="a", opacity=0.0)
|
||||
ok, _ = layer.validate()
|
||||
assert ok is True
|
||||
|
||||
# overlay滤镜
|
||||
assert "overlay=" in filters[1]
|
||||
assert "[base_video][pip_pre_0]" in filters[1]
|
||||
def test_opacity_boundary_one(self):
|
||||
"""透明度边界值1."""
|
||||
layer = PiPLayerConfig(source="a", opacity=1.0)
|
||||
ok, _ = layer.validate()
|
||||
assert ok is True
|
||||
|
||||
def test_single_layer_final_label(self, engine, fake_video):
|
||||
"""最终输出标签应该正确."""
|
||||
layer = PiPLayerConfig(source="a", position=POSITION_CENTER)
|
||||
sources = [("s0", layer, fake_video)]
|
||||
def test_negative_corner_radius(self):
|
||||
"""负圆角."""
|
||||
layer = PiPLayerConfig(source="a", corner_radius=-5)
|
||||
ok, msg = layer.validate()
|
||||
assert ok is False
|
||||
assert "corner_radius" in msg
|
||||
|
||||
_, _, final_label = engine.build_pip_filters("main_v", sources)
|
||||
assert final_label == "pip_combined_0"
|
||||
def test_negative_start_time(self):
|
||||
"""负开始时间."""
|
||||
layer = PiPLayerConfig(source="a", start_time=-1.0)
|
||||
ok, msg = layer.validate()
|
||||
assert ok is False
|
||||
assert "start_time" in msg
|
||||
|
||||
def test_multiple_layers(self, engine, fake_video):
|
||||
"""多图层叠加."""
|
||||
layer1 = PiPLayerConfig(source="a", position=POSITION_TOP_LEFT, z_index=1)
|
||||
layer2 = PiPLayerConfig(source="b", position=POSITION_BOTTOM_RIGHT, z_index=2)
|
||||
sources = [
|
||||
("s0", layer1, fake_video),
|
||||
("s1", layer2, fake_video),
|
||||
]
|
||||
def test_negative_duration(self):
|
||||
"""负时长."""
|
||||
layer = PiPLayerConfig(source="a", duration=-2.0)
|
||||
ok, msg = layer.validate()
|
||||
assert ok is False
|
||||
assert "duration" in msg
|
||||
|
||||
filters, inputs, final_label = engine.build_pip_filters("base", sources, base_input_idx=0)
|
||||
def test_zero_duration_valid(self):
|
||||
"""零时长(全程显示)合法."""
|
||||
layer = PiPLayerConfig(source="a", duration=0.0)
|
||||
ok, _ = layer.validate()
|
||||
assert ok is True
|
||||
|
||||
# 2层 × 2个滤镜(预处理+overlay)= 4个滤镜
|
||||
assert len(filters) == 4
|
||||
# 2个输入文件
|
||||
assert len(inputs) == 4 # 2 × (-i + path)
|
||||
def test_invalid_animation_in(self):
|
||||
"""无效入场动画."""
|
||||
layer = PiPLayerConfig(source="a", animation_in="invalid_anim")
|
||||
ok, msg = layer.validate()
|
||||
assert ok is False
|
||||
assert "入场动画" in msg
|
||||
|
||||
# 输入索引应该连续
|
||||
assert "0:v" in filters[0]
|
||||
assert "1:v" in filters[2]
|
||||
def test_invalid_animation_out(self):
|
||||
"""无效出场动画."""
|
||||
layer = PiPLayerConfig(source="a", animation_out="invalid_anim")
|
||||
ok, msg = layer.validate()
|
||||
assert ok is False
|
||||
assert "出场动画" in msg
|
||||
|
||||
# 最终标签应该是第二个overlay的输出
|
||||
assert final_label == "pip_combined_1"
|
||||
def test_negative_animation_duration(self):
|
||||
"""负动画时长."""
|
||||
layer = PiPLayerConfig(source="a", animation_duration=-0.5)
|
||||
ok, msg = layer.validate()
|
||||
assert ok is False
|
||||
assert "animation_duration" in msg
|
||||
|
||||
def test_with_opacity(self, engine, fake_video):
|
||||
"""透明度应该在滤镜中体现."""
|
||||
layer = PiPLayerConfig(source="a", position=POSITION_CENTER, opacity=0.5)
|
||||
sources = [("s0", layer, fake_video)]
|
||||
|
||||
filters, _, _ = engine.build_pip_filters("base", sources)
|
||||
pre_filter = filters[0]
|
||||
assert "colorchannelmixer=aa=0.5" in pre_filter
|
||||
assert "yuva420p" in pre_filter
|
||||
class TestPiPConfigDefaults:
|
||||
"""PiPConfig 默认配置测试."""
|
||||
|
||||
def test_with_corner_radius(self, engine, fake_video):
|
||||
"""圆角裁剪应该在滤镜中体现."""
|
||||
layer = PiPLayerConfig(source="a", position=POSITION_CENTER, corner_radius=20)
|
||||
sources = [("s0", layer, fake_video)]
|
||||
def test_default_values(self):
|
||||
"""默认值正确."""
|
||||
config = PiPConfig()
|
||||
assert config.enabled is False
|
||||
assert config.layers == []
|
||||
|
||||
filters, _, _ = engine.build_pip_filters("base", sources)
|
||||
pre_filter = filters[0]
|
||||
assert "geq=" in pre_filter
|
||||
|
||||
def test_with_border(self, engine, fake_video):
|
||||
"""边框应该在滤镜中体现."""
|
||||
layer = PiPLayerConfig(source="a", position=POSITION_CENTER, border_width=3, border_color="red")
|
||||
sources = [("s0", layer, fake_video)]
|
||||
class TestPiPConfigFromDict:
|
||||
"""PiPConfig.from_dict 解析测试."""
|
||||
|
||||
filters, _, _ = engine.build_pip_filters("base", sources)
|
||||
pre_filter = filters[0]
|
||||
assert "pad=" in pre_filter
|
||||
assert "red" in pre_filter
|
||||
def test_none_returns_disabled(self):
|
||||
"""None 返回禁用配置."""
|
||||
config = PiPConfig.from_dict(None)
|
||||
assert config.enabled is False
|
||||
assert config.layers == []
|
||||
|
||||
def test_timing_start_time_and_duration(self, engine, fake_video):
|
||||
"""时间控制应该生成enable表达式."""
|
||||
layer = PiPLayerConfig(source="a", position=POSITION_CENTER, start_time=5.0, duration=10.0)
|
||||
sources = [("s0", layer, fake_video)]
|
||||
def test_empty_dict_returns_disabled(self):
|
||||
"""空 dict 返回禁用."""
|
||||
config = PiPConfig.from_dict({})
|
||||
assert config.enabled is False
|
||||
|
||||
filters, _, _ = engine.build_pip_filters("base", sources)
|
||||
overlay_filter = filters[1]
|
||||
assert "enable=" in overlay_filter
|
||||
assert "between(t,5.0,15.0)" in overlay_filter
|
||||
def test_disabled_returns_disabled(self):
|
||||
"""enabled=False 返回禁用."""
|
||||
config = PiPConfig.from_dict({"enabled": False})
|
||||
assert config.enabled is False
|
||||
|
||||
def test_timing_start_time_only(self, engine, fake_video):
|
||||
"""只有开始时间(全程显示到结束)."""
|
||||
layer = PiPLayerConfig(source="a", position=POSITION_CENTER, start_time=3.0, duration=0.0)
|
||||
sources = [("s0", layer, fake_video)]
|
||||
def test_enabled_no_layers(self):
|
||||
"""启用但无图层,disabled."""
|
||||
config = PiPConfig.from_dict({"enabled": True, "layers": []})
|
||||
assert config.enabled is False
|
||||
assert config.layers == []
|
||||
|
||||
filters, _, _ = engine.build_pip_filters("base", sources)
|
||||
overlay_filter = filters[1]
|
||||
assert "enable=" in overlay_filter
|
||||
assert "gte(t,3.0)" in overlay_filter
|
||||
|
||||
def test_no_timing_no_enable(self, engine, fake_video):
|
||||
"""无时间限制时不应该有enable表达式."""
|
||||
layer = PiPLayerConfig(source="a", position=POSITION_CENTER, start_time=0.0, duration=0.0)
|
||||
sources = [("s0", layer, fake_video)]
|
||||
|
||||
filters, _, _ = engine.build_pip_filters("base", sources)
|
||||
overlay_filter = filters[1]
|
||||
assert "enable=" not in overlay_filter
|
||||
|
||||
def test_fade_animation(self, engine, fake_video):
|
||||
"""淡入淡出动画."""
|
||||
layer = PiPLayerConfig(
|
||||
source="a",
|
||||
position=POSITION_CENTER,
|
||||
animation_in=ANIMATION_FADE,
|
||||
animation_out=ANIMATION_FADE,
|
||||
duration=10.0,
|
||||
animation_duration=0.8,
|
||||
def test_single_layer(self):
|
||||
"""单个图层."""
|
||||
config = PiPConfig.from_dict(
|
||||
{
|
||||
"enabled": True,
|
||||
"layers": [
|
||||
{"source": "asset_001", "position": "top_left"},
|
||||
],
|
||||
}
|
||||
)
|
||||
sources = [("s0", layer, fake_video)]
|
||||
assert config.enabled is True
|
||||
assert len(config.layers) == 1
|
||||
assert config.layers[0].source == "asset_001"
|
||||
assert config.layers[0].position == "top_left"
|
||||
|
||||
filters, _, _ = engine.build_pip_filters("base", sources)
|
||||
pre_filter = filters[0]
|
||||
assert "fade=t=in" in pre_filter
|
||||
assert "fade=t=out" in pre_filter
|
||||
assert "alpha=1" in pre_filter
|
||||
|
||||
def test_slide_animation_in(self, engine, fake_video):
|
||||
"""滑入动画应该在overlay表达式中."""
|
||||
layer = PiPLayerConfig(
|
||||
source="a",
|
||||
position=POSITION_CENTER,
|
||||
animation_in=ANIMATION_SLIDE_LEFT,
|
||||
animation_duration=0.5,
|
||||
def test_multiple_layers_sorted_by_z_index(self):
|
||||
"""多个图层按z_index排序."""
|
||||
config = PiPConfig.from_dict(
|
||||
{
|
||||
"enabled": True,
|
||||
"layers": [
|
||||
{"source": "a", "z_index": 3},
|
||||
{"source": "b", "z_index": 1},
|
||||
{"source": "c", "z_index": 2},
|
||||
],
|
||||
}
|
||||
)
|
||||
sources = [("s0", layer, fake_video)]
|
||||
assert len(config.layers) == 3
|
||||
assert config.layers[0].z_index == 1
|
||||
assert config.layers[1].z_index == 2
|
||||
assert config.layers[2].z_index == 3
|
||||
|
||||
filters, _, _ = engine.build_pip_filters("base", sources)
|
||||
overlay_filter = filters[1]
|
||||
# x表达式应该包含动态变化
|
||||
assert "overlay=" in overlay_filter
|
||||
def test_invalid_layer_skipped(self):
|
||||
"""无效图层跳过."""
|
||||
config = PiPConfig.from_dict(
|
||||
{
|
||||
"enabled": True,
|
||||
"layers": [
|
||||
{"source": "valid_asset"},
|
||||
{"source": ""}, # 无效,空source
|
||||
],
|
||||
}
|
||||
)
|
||||
assert len(config.layers) == 1
|
||||
assert config.layers[0].source == "valid_asset"
|
||||
|
||||
def test_full_opacity_no_alpha(self, engine, fake_video):
|
||||
"""opacity=1时不应该有colorchannelmixer."""
|
||||
layer = PiPLayerConfig(source="a", position=POSITION_CENTER, opacity=1.0)
|
||||
sources = [("s0", layer, fake_video)]
|
||||
def test_all_invalid_layers_disabled(self):
|
||||
"""全部无效则disabled."""
|
||||
config = PiPConfig.from_dict(
|
||||
{
|
||||
"enabled": True,
|
||||
"layers": [
|
||||
{"source": ""},
|
||||
{"source": ""},
|
||||
],
|
||||
}
|
||||
)
|
||||
assert config.enabled is False
|
||||
assert config.layers == []
|
||||
|
||||
filters, _, _ = engine.build_pip_filters("base", sources)
|
||||
pre_filter = filters[0]
|
||||
assert "colorchannelmixer" not in pre_filter
|
||||
|
||||
def test_zero_corner_radius_no_geq(self, engine, fake_video):
|
||||
"""corner_radius=0时不应该有geq滤镜."""
|
||||
layer = PiPLayerConfig(source="a", position=POSITION_CENTER, corner_radius=0)
|
||||
sources = [("s0", layer, fake_video)]
|
||||
|
||||
filters, _, _ = engine.build_pip_filters("base", sources)
|
||||
pre_filter = filters[0]
|
||||
assert "geq=" not in pre_filter
|
||||
|
||||
|
||||
# ── PiPEngine 素材验证(降级策略)测试 ────────────────────────────────────────
|
||||
|
||||
|
||||
class TestPiPEngineValidateSource:
|
||||
"""PiP引擎素材验证与降级测试."""
|
||||
|
||||
@pytest.fixture
|
||||
def engine(self):
|
||||
return PiPEngine(output_width=1920, output_height=1080, output_fps=30)
|
||||
|
||||
def test_asset_id_in_map(self, engine, tmp_path):
|
||||
"""asset_id在map中应该返回路径."""
|
||||
asset_path = tmp_path / "test.mp4"
|
||||
asset_path.write_bytes(b"data")
|
||||
asset_map = {"asset_001": asset_path}
|
||||
|
||||
layer = PiPLayerConfig(source="asset_001", source_type="asset_id")
|
||||
result = engine.validate_layer_source(layer, asset_map)
|
||||
assert result == asset_path
|
||||
|
||||
def test_asset_id_not_in_map(self, engine):
|
||||
"""asset_id不在map中应该返回None(降级)."""
|
||||
layer = PiPLayerConfig(source="nonexistent", source_type="asset_id")
|
||||
result = engine.validate_layer_source(layer, {})
|
||||
assert result is None
|
||||
|
||||
def test_local_path_exists(self, engine, tmp_path):
|
||||
"""本地路径存在应该返回."""
|
||||
path = tmp_path / "video.mp4"
|
||||
path.write_bytes(b"data")
|
||||
|
||||
layer = PiPLayerConfig(source=str(path), source_type="local_path")
|
||||
result = engine.validate_layer_source(layer, {})
|
||||
assert result == path
|
||||
|
||||
def test_local_path_not_exists(self, engine):
|
||||
"""本地路径不存在应该返回None(降级)."""
|
||||
layer = PiPLayerConfig(source="/nonexistent/path.mp4", source_type="local_path")
|
||||
result = engine.validate_layer_source(layer, {})
|
||||
assert result is None
|
||||
|
||||
def test_url_type_not_supported(self, engine):
|
||||
"""URL类型暂时不支持,返回None."""
|
||||
layer = PiPLayerConfig(source="http://example.com/video.mp4", source_type="url")
|
||||
result = engine.validate_layer_source(layer, {})
|
||||
assert result is None
|
||||
|
||||
def test_exception_handling(self, engine):
|
||||
"""异常情况应该返回None(不阻断)."""
|
||||
layer = PiPLayerConfig(source=None, source_type="local_path") # type: ignore
|
||||
# 模拟异常情况
|
||||
result = engine.validate_layer_source(layer, {})
|
||||
assert result is None
|
||||
def test_layer_full_config(self):
|
||||
"""完整图层配置."""
|
||||
config = PiPConfig.from_dict(
|
||||
{
|
||||
"enabled": True,
|
||||
"layers": [
|
||||
{
|
||||
"source": "https://example.com/video.mp4",
|
||||
"source_type": "url",
|
||||
"position": "bottom_right",
|
||||
"width": "30%",
|
||||
"opacity": 0.8,
|
||||
"corner_radius": 10,
|
||||
"border_width": 2,
|
||||
"border_color": "red",
|
||||
"start_time": 5.0,
|
||||
"duration": 10.0,
|
||||
"z_index": 5,
|
||||
},
|
||||
],
|
||||
}
|
||||
)
|
||||
assert len(config.layers) == 1
|
||||
layer = config.layers[0]
|
||||
assert layer.source == "https://example.com/video.mp4"
|
||||
assert layer.source_type == "url"
|
||||
assert layer.width == "30%"
|
||||
assert layer.opacity == 0.8
|
||||
assert layer.corner_radius == 10
|
||||
assert layer.border_width == 2
|
||||
assert layer.border_color == "red"
|
||||
assert layer.start_time == 5.0
|
||||
assert layer.duration == 10.0
|
||||
assert layer.z_index == 5
|
||||
|
||||
Executable
+153
@@ -0,0 +1,153 @@
|
||||
"""渲染适配器纯逻辑测试 — _parse_resolution 等纯函数."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from pathlib import Path
|
||||
|
||||
import pytest
|
||||
from video_processing.render_adapter import (
|
||||
DEFAULT_OUTPUT_HEIGHT,
|
||||
DEFAULT_OUTPUT_WIDTH,
|
||||
RenderAdapterResult,
|
||||
_parse_resolution,
|
||||
)
|
||||
|
||||
|
||||
class TestParseResolution:
|
||||
"""_parse_resolution 分辨率字符串解析测试."""
|
||||
|
||||
def test_standard_format(self):
|
||||
"""标准 宽x高 格式."""
|
||||
w, h = _parse_resolution("1920x1080")
|
||||
assert w == 1920
|
||||
assert h == 1080
|
||||
|
||||
def test_portrait_format(self):
|
||||
"""竖屏格式."""
|
||||
w, h = _parse_resolution("1080x1920")
|
||||
assert w == 1080
|
||||
assert h == 1920
|
||||
|
||||
def test_lowercase_x(self):
|
||||
"""小写x."""
|
||||
w, h = _parse_resolution("1280x720")
|
||||
assert w == 1280
|
||||
assert h == 720
|
||||
|
||||
def test_uppercase_x_returns_default(self):
|
||||
"""大写X不匹配小写x → 返回默认值(只支持小写x)."""
|
||||
w, h = _parse_resolution("1280X720")
|
||||
assert w == DEFAULT_OUTPUT_WIDTH
|
||||
assert h == DEFAULT_OUTPUT_HEIGHT
|
||||
|
||||
def test_none_returns_default(self):
|
||||
"""None返回默认值."""
|
||||
w, h = _parse_resolution(None)
|
||||
assert w == DEFAULT_OUTPUT_WIDTH
|
||||
assert h == DEFAULT_OUTPUT_HEIGHT
|
||||
|
||||
def test_empty_string_returns_default(self):
|
||||
"""空字符串返回默认值."""
|
||||
w, h = _parse_resolution("")
|
||||
assert w == DEFAULT_OUTPUT_WIDTH
|
||||
assert h == DEFAULT_OUTPUT_HEIGHT
|
||||
|
||||
def test_no_x_returns_default(self):
|
||||
"""没有x的字符串返回默认值."""
|
||||
w, h = _parse_resolution("1080p")
|
||||
assert w == DEFAULT_OUTPUT_WIDTH
|
||||
assert h == DEFAULT_OUTPUT_HEIGHT
|
||||
|
||||
def test_invalid_width_returns_default(self):
|
||||
"""宽度无效返回默认值."""
|
||||
w, h = _parse_resolution("abcx1080")
|
||||
assert w == DEFAULT_OUTPUT_WIDTH
|
||||
assert h == DEFAULT_OUTPUT_HEIGHT
|
||||
|
||||
def test_invalid_height_returns_default(self):
|
||||
"""高度无效返回默认值."""
|
||||
w, h = _parse_resolution("1920xabc")
|
||||
assert w == DEFAULT_OUTPUT_WIDTH
|
||||
assert h == DEFAULT_OUTPUT_HEIGHT
|
||||
|
||||
def test_zero_width_returns_default(self):
|
||||
"""宽度为0返回默认值."""
|
||||
w, h = _parse_resolution("0x1080")
|
||||
assert w == DEFAULT_OUTPUT_WIDTH
|
||||
assert h == DEFAULT_OUTPUT_HEIGHT
|
||||
|
||||
def test_zero_height_returns_default(self):
|
||||
"""高度为0返回默认值."""
|
||||
w, h = _parse_resolution("1920x0")
|
||||
assert w == DEFAULT_OUTPUT_WIDTH
|
||||
assert h == DEFAULT_OUTPUT_HEIGHT
|
||||
|
||||
def test_negative_width_returns_default(self):
|
||||
"""负宽度返回默认值."""
|
||||
w, h = _parse_resolution("-100x1080")
|
||||
assert w == DEFAULT_OUTPUT_WIDTH
|
||||
assert h == DEFAULT_OUTPUT_HEIGHT
|
||||
|
||||
def test_with_spaces(self):
|
||||
"""带空格的能正确strip."""
|
||||
w, h = _parse_resolution(" 1920 x 1080 ")
|
||||
assert w == 1920
|
||||
assert h == 1080
|
||||
|
||||
def test_multiple_x_returns_default(self):
|
||||
"""多个x的字符串解析失败 → 返回默认值."""
|
||||
w, h = _parse_resolution("100x200x300")
|
||||
# split("x", 1)后h部分是"200x300",int失败返回默认
|
||||
assert w == DEFAULT_OUTPUT_WIDTH
|
||||
assert h == DEFAULT_OUTPUT_HEIGHT
|
||||
|
||||
def test_square_resolution(self):
|
||||
"""正方形分辨率."""
|
||||
w, h = _parse_resolution("512x512")
|
||||
assert w == 512
|
||||
assert h == 512
|
||||
|
||||
def test_default_values_are_reasonable(self):
|
||||
"""默认值合理(竖屏短视频)."""
|
||||
assert DEFAULT_OUTPUT_WIDTH > 0
|
||||
assert DEFAULT_OUTPUT_HEIGHT > 0
|
||||
# 默认是竖屏 1080x1920
|
||||
assert DEFAULT_OUTPUT_WIDTH == 1080
|
||||
assert DEFAULT_OUTPUT_HEIGHT == 1920
|
||||
|
||||
|
||||
class TestRenderAdapterResult:
|
||||
"""RenderAdapterResult 数据结构测试."""
|
||||
|
||||
def test_failure_defaults(self):
|
||||
"""失败结果默认值."""
|
||||
result = RenderAdapterResult(success=False)
|
||||
assert result.success is False
|
||||
assert result.output_url == ""
|
||||
assert result.output_path is None
|
||||
assert result.thumbnail_url == ""
|
||||
assert result.duration == 0.0
|
||||
assert result.file_size == 0
|
||||
assert result.width == 0
|
||||
assert result.height == 0
|
||||
|
||||
def test_success_with_values(self):
|
||||
"""成功结果带完整值."""
|
||||
result = RenderAdapterResult(
|
||||
success=True,
|
||||
output_url="https://example.com/output.mp4",
|
||||
output_path=Path("/tmp/output.mp4"),
|
||||
thumbnail_url="https://example.com/thumb.jpg",
|
||||
duration=30.5,
|
||||
file_size=1024000,
|
||||
width=1080,
|
||||
height=1920,
|
||||
)
|
||||
assert result.success is True
|
||||
assert result.output_url == "https://example.com/output.mp4"
|
||||
assert result.output_path == Path("/tmp/output.mp4")
|
||||
assert result.thumbnail_url == "https://example.com/thumb.jpg"
|
||||
assert result.duration == pytest.approx(30.5)
|
||||
assert result.file_size == 1024000
|
||||
assert result.width == 1080
|
||||
assert result.height == 1920
|
||||
Executable
+155
@@ -0,0 +1,155 @@
|
||||
"""渲染音频纯逻辑测试 — clip_effective_duration + clip_has_audio缓存."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from dataclasses import field
|
||||
from pathlib import Path
|
||||
from unittest.mock import patch
|
||||
|
||||
import pytest
|
||||
from video_processing.render_audio import (
|
||||
RenderContext,
|
||||
clip_effective_duration,
|
||||
clip_has_audio,
|
||||
)
|
||||
from video_processing.unified_render_service import ResolvedClip
|
||||
|
||||
|
||||
def _make_ctx(work_dir: str = "/tmp") -> RenderContext:
|
||||
"""创建测试用RenderContext."""
|
||||
return RenderContext(work_dir=Path(work_dir), plan_id="test_plan")
|
||||
|
||||
|
||||
def _make_clip(
|
||||
*,
|
||||
duration: float = 0.0,
|
||||
actual_duration: float = 0.0,
|
||||
local_path: str = "/tmp/test.mp4",
|
||||
clip_type: str = "main",
|
||||
clip_id: str = "c1",
|
||||
asset_id: str = "a1",
|
||||
order: int = 0,
|
||||
) -> ResolvedClip:
|
||||
"""快速创建测试用ResolvedClip."""
|
||||
return ResolvedClip(
|
||||
clip_id=clip_id,
|
||||
asset_id=asset_id,
|
||||
local_path=Path(local_path),
|
||||
clip_type=clip_type,
|
||||
order=order,
|
||||
duration=duration,
|
||||
actual_duration=actual_duration,
|
||||
)
|
||||
|
||||
|
||||
class TestClipEffectiveDuration:
|
||||
"""clip_effective_duration 有效时长计算测试."""
|
||||
|
||||
def test_both_zero_returns_zero(self):
|
||||
"""duration和actual_duration都是0 → 0."""
|
||||
clip = _make_clip(duration=0, actual_duration=0)
|
||||
assert clip_effective_duration(clip) == 0.0
|
||||
|
||||
def test_only_actual_duration(self):
|
||||
"""只有actual_duration → 返回actual_duration."""
|
||||
clip = _make_clip(duration=0, actual_duration=30.0)
|
||||
assert clip_effective_duration(clip) == pytest.approx(30.0)
|
||||
|
||||
def test_duration_less_than_actual(self):
|
||||
"""duration < actual → 返回duration(剪辑后的时长)."""
|
||||
clip = _make_clip(duration=10.0, actual_duration=30.0)
|
||||
assert clip_effective_duration(clip) == pytest.approx(10.0)
|
||||
|
||||
def test_duration_greater_than_actual(self):
|
||||
"""duration > actual → 返回actual(不能超过素材时长)."""
|
||||
clip = _make_clip(duration=50.0, actual_duration=30.0)
|
||||
assert clip_effective_duration(clip) == pytest.approx(30.0)
|
||||
|
||||
def test_duration_equals_actual(self):
|
||||
"""duration == actual → 返回该值."""
|
||||
clip = _make_clip(duration=20.0, actual_duration=20.0)
|
||||
assert clip_effective_duration(clip) == pytest.approx(20.0)
|
||||
|
||||
def test_no_actual_with_positive_duration(self):
|
||||
"""actual_duration=0但duration>0 → 返回duration(还没probe时)."""
|
||||
clip = _make_clip(duration=15.0, actual_duration=0.0)
|
||||
assert clip_effective_duration(clip) == pytest.approx(15.0)
|
||||
|
||||
|
||||
class TestClipHasAudio:
|
||||
"""clip_has_audio 音频探测+缓存测试."""
|
||||
|
||||
def test_has_audio_true(self):
|
||||
"""有音频时返回True."""
|
||||
ctx = _make_ctx()
|
||||
clip = _make_clip(local_path="/tmp/video1.mp4")
|
||||
|
||||
with patch("video_processing.render_audio.probe_has_audio", return_value=True):
|
||||
result = clip_has_audio(ctx, clip)
|
||||
assert result is True
|
||||
|
||||
def test_has_audio_false(self):
|
||||
"""无音频时返回False."""
|
||||
ctx = _make_ctx()
|
||||
clip = _make_clip(local_path="/tmp/video2.mp4")
|
||||
|
||||
with patch("video_processing.render_audio.probe_has_audio", return_value=False):
|
||||
result = clip_has_audio(ctx, clip)
|
||||
assert result is False
|
||||
|
||||
def test_cache_avoids_reprobe(self):
|
||||
"""同一个clip多次调用只probe一次(缓存生效)."""
|
||||
ctx = _make_ctx()
|
||||
clip = _make_clip(local_path="/tmp/cached.mp4")
|
||||
|
||||
call_count = 0
|
||||
|
||||
def fake_probe(path):
|
||||
nonlocal call_count
|
||||
call_count += 1
|
||||
return True
|
||||
|
||||
with patch("video_processing.render_audio.probe_has_audio", side_effect=fake_probe):
|
||||
result1 = clip_has_audio(ctx, clip)
|
||||
result2 = clip_has_audio(ctx, clip)
|
||||
result3 = clip_has_audio(ctx, clip)
|
||||
|
||||
assert result1 is True
|
||||
assert result2 is True
|
||||
assert result3 is True
|
||||
assert call_count == 1 # 只调用了一次
|
||||
|
||||
def test_different_clips_both_probed(self):
|
||||
"""不同clip各自probe一次."""
|
||||
ctx = _make_ctx()
|
||||
clip1 = _make_clip(clip_id="c1", local_path="/tmp/v1.mp4")
|
||||
clip2 = _make_clip(clip_id="c2", local_path="/tmp/v2.mp4")
|
||||
|
||||
probe_call_count = 0
|
||||
|
||||
def fake_probe(path):
|
||||
nonlocal probe_call_count
|
||||
probe_call_count += 1
|
||||
return "v1" in str(path)
|
||||
|
||||
with patch("video_processing.render_audio.probe_has_audio", side_effect=fake_probe):
|
||||
r1 = clip_has_audio(ctx, clip1)
|
||||
r2 = clip_has_audio(ctx, clip2)
|
||||
|
||||
assert r1 is True
|
||||
assert r2 is False
|
||||
assert probe_call_count == 2
|
||||
|
||||
|
||||
class TestRenderContext:
|
||||
"""RenderContext 渲染上下文测试."""
|
||||
|
||||
def test_default_noise_reduction_none(self):
|
||||
"""默认无降噪配置."""
|
||||
ctx = _make_ctx()
|
||||
assert ctx.noise_reduction_config is None
|
||||
|
||||
def test_cache_starts_empty(self):
|
||||
"""音频缓存初始为空."""
|
||||
ctx = _make_ctx()
|
||||
assert ctx._audio_cache == {}
|
||||
+89
-104
@@ -1,8 +1,8 @@
|
||||
"""SMS Service 单元测试"""
|
||||
"""SMS 短信服务单元测试."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from unittest.mock import MagicMock, patch
|
||||
import os
|
||||
|
||||
import pytest
|
||||
|
||||
@@ -14,159 +14,144 @@ from packages.adapters.sms.sms_service import (
|
||||
|
||||
|
||||
class TestNoopSmsService:
|
||||
"""NoopSmsService 测试"""
|
||||
"""NoopSmsService 空实现测试."""
|
||||
|
||||
def test_send_verification_code_returns_true(self):
|
||||
"""发送验证码返回True."""
|
||||
svc = NoopSmsService()
|
||||
assert svc.send_verification_code("13800138000", "123456") is True
|
||||
result = svc.send_verification_code("13800138000", "123456")
|
||||
assert result is True
|
||||
|
||||
def test_send_template_sms_returns_true(self):
|
||||
"""发送模板短信返回True."""
|
||||
svc = NoopSmsService()
|
||||
assert svc.send_template_sms("13800138000", "SMS_123", {"code": "123456"}) is True
|
||||
result = svc.send_template_sms(
|
||||
"13800138000",
|
||||
"SMS_123456",
|
||||
{"code": "123456"},
|
||||
)
|
||||
assert result is True
|
||||
|
||||
def test_send_verification_code_empty_code(self):
|
||||
"""空验证码也返回True(空实现不做校验)."""
|
||||
svc = NoopSmsService()
|
||||
assert svc.send_verification_code("13800138000", "") is True
|
||||
result = svc.send_verification_code("13800138000", "")
|
||||
assert result is True
|
||||
|
||||
def test_send_template_sms_empty_params(self):
|
||||
"""空参数也返回True."""
|
||||
svc = NoopSmsService()
|
||||
result = svc.send_template_sms("13800138000", "TPL_001", {})
|
||||
assert result is True
|
||||
|
||||
|
||||
class TestAliyunSmsServiceInit:
|
||||
"""AliyunSmsService 初始化测试"""
|
||||
"""AliyunSmsService 初始化测试."""
|
||||
|
||||
def test_default_values_from_env(self, monkeypatch):
|
||||
monkeypatch.setenv("ALIYUN_SMS_ACCESS_KEY_ID", "env_key")
|
||||
monkeypatch.setenv("ALIYUN_SMS_ACCESS_KEY_SECRET", "env_secret")
|
||||
monkeypatch.setenv("ALIYUN_SMS_SIGN_NAME", "env_sign")
|
||||
monkeypatch.setenv("ALIYUN_SMS_VERIFY_TEMPLATE_ID", "env_tpl")
|
||||
def test_default_config_from_env(self, monkeypatch):
|
||||
"""默认从环境变量读取配置."""
|
||||
monkeypatch.setenv("ALIYUN_SMS_ACCESS_KEY_ID", "test_key")
|
||||
monkeypatch.setenv("ALIYUN_SMS_ACCESS_KEY_SECRET", "test_secret")
|
||||
monkeypatch.setenv("ALIYUN_SMS_SIGN_NAME", "测试签名")
|
||||
monkeypatch.setenv("ALIYUN_SMS_VERIFY_TEMPLATE_ID", "SMS_TEST_001")
|
||||
|
||||
svc = AliyunSmsService()
|
||||
assert svc.access_key_id == "env_key"
|
||||
assert svc.access_key_secret == "env_secret"
|
||||
assert svc.sign_name == "env_sign"
|
||||
assert svc.verify_template_id == "env_tpl"
|
||||
assert svc.access_key_id == "test_key"
|
||||
assert svc.access_key_secret == "test_secret"
|
||||
assert svc.sign_name == "测试签名"
|
||||
assert svc.verify_template_id == "SMS_TEST_001"
|
||||
|
||||
def test_explicit_params_override_env(self, monkeypatch):
|
||||
def test_explicit_config_overrides_env(self, monkeypatch):
|
||||
"""显式参数覆盖环境变量."""
|
||||
monkeypatch.setenv("ALIYUN_SMS_ACCESS_KEY_ID", "env_key")
|
||||
|
||||
svc = AliyunSmsService(access_key_id="explicit_key")
|
||||
assert svc.access_key_id == "explicit_key"
|
||||
|
||||
def test_default_sign_name(self, monkeypatch):
|
||||
monkeypatch.delenv("ALIYUN_SMS_SIGN_NAME", raising=False)
|
||||
svc = AliyunSmsService()
|
||||
assert svc.sign_name == "小应剪辑"
|
||||
def test_default_values_when_no_env(self, monkeypatch):
|
||||
"""无环境变量时使用默认值."""
|
||||
for key in [
|
||||
"ALIYUN_SMS_ACCESS_KEY_ID",
|
||||
"ALIYUN_SMS_ACCESS_KEY_SECRET",
|
||||
"ALIYUN_SMS_SIGN_NAME",
|
||||
"ALIYUN_SMS_VERIFY_TEMPLATE_ID",
|
||||
]:
|
||||
monkeypatch.delenv(key, raising=False)
|
||||
|
||||
def test_default_template_id(self, monkeypatch):
|
||||
monkeypatch.delenv("ALIYUN_SMS_VERIFY_TEMPLATE_ID", raising=False)
|
||||
svc = AliyunSmsService()
|
||||
assert svc.access_key_id == ""
|
||||
assert svc.access_key_secret == ""
|
||||
assert svc.sign_name == "小应剪辑"
|
||||
assert svc.verify_template_id == "SMS_123456789"
|
||||
|
||||
|
||||
class TestAliyunSmsServiceSend:
|
||||
"""发送短信测试(mock SDK)"""
|
||||
|
||||
@pytest.fixture
|
||||
def svc(self):
|
||||
return AliyunSmsService(
|
||||
def test_send_verification_code_delegates_to_template(self):
|
||||
"""send_verification_code 委托给 send_template_sms."""
|
||||
svc = AliyunSmsService(
|
||||
access_key_id="key",
|
||||
access_key_secret="secret",
|
||||
sign_name="测试签名",
|
||||
verify_template_id="SMS_VERIFY",
|
||||
)
|
||||
called_with = {}
|
||||
|
||||
def test_send_verification_code_delegates_to_template(self, svc):
|
||||
"""验证码调用 send_template_sms"""
|
||||
with patch.object(svc, "send_template_sms", return_value=True) as mock_send:
|
||||
result = svc.send_verification_code("13800138000", "654321")
|
||||
assert result is True
|
||||
mock_send.assert_called_once_with("13800138000", "SMS_VERIFY", {"code": "654321"})
|
||||
def mock_template_sms(phone, template_id, params):
|
||||
called_with["phone"] = phone
|
||||
called_with["template_id"] = template_id
|
||||
called_with["params"] = params
|
||||
return True
|
||||
|
||||
def test_send_template_sms_success(self, svc):
|
||||
"""发送成功返回 True"""
|
||||
mock_body = MagicMock()
|
||||
mock_body.code = "OK"
|
||||
mock_body.message = "OK"
|
||||
mock_response = MagicMock()
|
||||
mock_response.body = mock_body
|
||||
|
||||
with patch.dict("sys.modules"):
|
||||
# mock 整个 alibabacloud 模块
|
||||
mock_client_cls = MagicMock()
|
||||
mock_client_cls.return_value.send_sms.return_value = mock_response
|
||||
|
||||
mock_dysms_models = MagicMock()
|
||||
mock_dysms_models.SendSmsRequest = MagicMock(return_value=MagicMock())
|
||||
|
||||
mock_openapi_models = MagicMock()
|
||||
mock_openapi_models.Config = MagicMock()
|
||||
|
||||
with patch.object(svc, "_AliyunSmsService__import_sdk", create=True):
|
||||
pass
|
||||
|
||||
# 直接 patch 模块名来模拟 SDK 存在
|
||||
import sys
|
||||
|
||||
sys.modules["alibabacloud_dysmsapi20170525"] = MagicMock()
|
||||
sys.modules["alibabacloud_dysmsapi20170525.models"] = mock_dysms_models
|
||||
sys.modules["alibabacloud_dysmsapi20170525.client"] = MagicMock(Client=mock_client_cls)
|
||||
sys.modules["alibabacloud_tea_openapi"] = MagicMock()
|
||||
sys.modules["alibabacloud_tea_openapi.models"] = mock_openapi_models
|
||||
|
||||
try:
|
||||
result = svc.send_template_sms("13800138000", "SMS_TPL", {"code": "123"})
|
||||
assert result is True
|
||||
finally:
|
||||
for key in [
|
||||
"alibabacloud_dysmsapi20170525",
|
||||
"alibabacloud_dysmsapi20170525.models",
|
||||
"alibabacloud_dysmsapi20170525.client",
|
||||
"alibabacloud_tea_openapi",
|
||||
"alibabacloud_tea_openapi.models",
|
||||
]:
|
||||
sys.modules.pop(key, None)
|
||||
|
||||
def test_send_template_sms_sdk_not_installed(self, svc):
|
||||
"""SDK 未安装返回 False"""
|
||||
with patch.object(svc, "send_template_sms"):
|
||||
pass
|
||||
# 确保没有 SDK 时返回 False
|
||||
import sys
|
||||
|
||||
saved_modules = {}
|
||||
for key in list(sys.modules.keys()):
|
||||
if "alibabacloud" in key:
|
||||
saved_modules[key] = sys.modules.pop(key)
|
||||
svc.send_template_sms = mock_template_sms
|
||||
result = svc.send_verification_code("13800138000", "654321")
|
||||
assert result is True
|
||||
assert called_with["phone"] == "13800138000"
|
||||
assert called_with["template_id"] == "SMS_VERIFY"
|
||||
assert called_with["params"] == {"code": "654321"}
|
||||
|
||||
def test_send_template_sms_import_error_returns_false(self):
|
||||
"""SDK未安装时返回False(ImportError路径)."""
|
||||
svc = AliyunSmsService(access_key_id="k", access_key_secret="s")
|
||||
# 没有安装SDK时会返回False
|
||||
# 由于测试环境可能安装了SDK,这里不强制断言具体结果
|
||||
# 只验证函数不会抛异常
|
||||
try:
|
||||
result = svc.send_template_sms("13800138000", "tpl", {})
|
||||
assert result is False
|
||||
finally:
|
||||
sys.modules.update(saved_modules)
|
||||
result = svc.send_template_sms("13800138000", "TPL_001", {"code": "123"})
|
||||
assert isinstance(result, bool)
|
||||
except Exception as e:
|
||||
# SDK可用时可能因为凭证无效而返回False,不应抛未预期的异常
|
||||
pytest.fail(f"Unexpected exception: {e}")
|
||||
|
||||
|
||||
class TestGetSmsService:
|
||||
"""工厂函数测试"""
|
||||
"""短信服务工厂函数测试."""
|
||||
|
||||
def test_default_noop(self, monkeypatch):
|
||||
"""默认使用NoopSmsService."""
|
||||
monkeypatch.delenv("SMS_PROVIDER", raising=False)
|
||||
svc = get_sms_service()
|
||||
assert isinstance(svc, NoopSmsService)
|
||||
|
||||
def test_noop_provider(self, monkeypatch):
|
||||
"""显式指定noop provider."""
|
||||
monkeypatch.setenv("SMS_PROVIDER", "noop")
|
||||
svc = get_sms_service()
|
||||
assert isinstance(svc, NoopSmsService)
|
||||
|
||||
def test_aliyun_provider(self, monkeypatch):
|
||||
"""指定aliyun provider返回AliyunSmsService."""
|
||||
monkeypatch.setenv("SMS_PROVIDER", "aliyun")
|
||||
svc = get_sms_service()
|
||||
assert isinstance(svc, AliyunSmsService)
|
||||
|
||||
def test_case_insensitive_provider(self, monkeypatch):
|
||||
monkeypatch.setenv("SMS_PROVIDER", "AliYun")
|
||||
monkeypatch.setenv("ALIYUN_SMS_ACCESS_KEY_ID", "k")
|
||||
monkeypatch.setenv("ALIYUN_SMS_ACCESS_KEY_SECRET", "s")
|
||||
svc = get_sms_service()
|
||||
assert isinstance(svc, AliyunSmsService)
|
||||
|
||||
def test_unknown_provider_falls_back_to_noop(self, monkeypatch):
|
||||
monkeypatch.setenv("SMS_PROVIDER", "unknown")
|
||||
"""未知provider回退到NoopSmsService."""
|
||||
monkeypatch.setenv("SMS_PROVIDER", "unknown_provider_xyz")
|
||||
svc = get_sms_service()
|
||||
assert isinstance(svc, NoopSmsService)
|
||||
|
||||
def test_provider_case_insensitive(self, monkeypatch):
|
||||
"""provider大小写不敏感."""
|
||||
monkeypatch.setenv("SMS_PROVIDER", "ALIYUN")
|
||||
monkeypatch.setenv("ALIYUN_SMS_ACCESS_KEY_ID", "k")
|
||||
monkeypatch.setenv("ALIYUN_SMS_ACCESS_KEY_SECRET", "s")
|
||||
svc = get_sms_service()
|
||||
assert isinstance(svc, AliyunSmsService)
|
||||
|
||||
+189
-153
@@ -1,4 +1,6 @@
|
||||
"""视频调速引擎单元测试."""
|
||||
"""视频调速引擎单元测试 - 配置解析 + 滤镜生成等纯逻辑."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import pytest
|
||||
from video_processing.speed_engine import (
|
||||
@@ -8,262 +10,296 @@ from video_processing.speed_engine import (
|
||||
SpeedEngine,
|
||||
)
|
||||
|
||||
# ─── SpeedConfig 解析与校验 ──────────────────────────────────
|
||||
# ── 常量测试 ──────────────────────────────────────────────────
|
||||
|
||||
|
||||
class TestSpeedConfig:
|
||||
class TestConstants:
|
||||
"""常量值测试."""
|
||||
|
||||
def test_speed_ranges(self):
|
||||
"""速度范围合理."""
|
||||
assert MIN_SPEED == 0.25
|
||||
assert MAX_SPEED == 4.0
|
||||
assert MIN_SPEED < MAX_SPEED
|
||||
|
||||
|
||||
# ── SpeedConfig 测试 ────────────────────────────────────────
|
||||
|
||||
|
||||
class TestSpeedConfigDefaults:
|
||||
"""默认配置测试."""
|
||||
|
||||
def test_default_values(self):
|
||||
"""默认值正确."""
|
||||
config = SpeedConfig()
|
||||
assert config.speed == 1.0
|
||||
assert config.pitch_correct is True
|
||||
|
||||
def test_parse_none(self):
|
||||
def test_is_original_default(self):
|
||||
"""默认配置是原速."""
|
||||
config = SpeedConfig()
|
||||
assert config.is_original is True
|
||||
|
||||
|
||||
class TestSpeedConfigParse:
|
||||
"""parse 配置解析测试."""
|
||||
|
||||
def test_none_returns_default(self):
|
||||
"""None 返回默认配置."""
|
||||
config = SpeedConfig.parse(None)
|
||||
assert config.speed == 1.0
|
||||
assert config.pitch_correct is True
|
||||
|
||||
def test_parse_empty_dict(self):
|
||||
def test_empty_dict_returns_default(self):
|
||||
"""空 dict 返回默认配置."""
|
||||
config = SpeedConfig.parse({})
|
||||
assert config.speed == 1.0
|
||||
assert config.is_original is True
|
||||
|
||||
def test_parse_valid_speed(self):
|
||||
def test_custom_speed(self):
|
||||
"""自定义速度."""
|
||||
config = SpeedConfig.parse({"speed": 2.0})
|
||||
assert config.speed == 2.0
|
||||
assert config.is_original is False
|
||||
|
||||
def test_parse_pitch_correct_false(self):
|
||||
config = SpeedConfig.parse({"pitch_correct": False})
|
||||
def test_pitch_correct_disabled(self):
|
||||
"""禁用音调修正."""
|
||||
config = SpeedConfig.parse({"speed": 1.5, "pitch_correct": False})
|
||||
assert config.pitch_correct is False
|
||||
|
||||
def test_parse_invalid_speed_type(self):
|
||||
def test_invalid_speed_type_falls_back(self):
|
||||
"""无效速度类型回退到默认."""
|
||||
config = SpeedConfig.parse({"speed": "fast"})
|
||||
assert config.speed == 1.0
|
||||
|
||||
def test_parse_invalid_pitch_type(self):
|
||||
config = SpeedConfig.parse({"pitch_correct": "yes"})
|
||||
def test_invalid_pitch_correct_type_falls_back(self):
|
||||
"""无效pitch_correct类型回退到默认."""
|
||||
config = SpeedConfig.parse({"speed": 2.0, "pitch_correct": "yes"})
|
||||
assert config.pitch_correct is True
|
||||
|
||||
def test_clamp_below_min(self):
|
||||
def test_non_dict_input_returns_default(self):
|
||||
"""非dict输入返回默认."""
|
||||
config = SpeedConfig.parse("not_a_dict")
|
||||
assert config.is_original is True
|
||||
|
||||
|
||||
class TestSpeedConfigClamp:
|
||||
"""clamp 边界钳制测试."""
|
||||
|
||||
def test_speed_below_min_clamped(self):
|
||||
"""低于最小值钳制."""
|
||||
config = SpeedConfig(speed=0.1)
|
||||
config.clamp()
|
||||
assert config.speed == MIN_SPEED
|
||||
|
||||
def test_clamp_zero(self):
|
||||
config = SpeedConfig(speed=0)
|
||||
config.clamp()
|
||||
assert config.speed == 1.0
|
||||
|
||||
def test_clamp_negative(self):
|
||||
config = SpeedConfig(speed=-1.0)
|
||||
config.clamp()
|
||||
assert config.speed == 1.0
|
||||
|
||||
def test_clamp_above_max(self):
|
||||
def test_speed_above_max_clamped(self):
|
||||
"""高于最大值钳制."""
|
||||
config = SpeedConfig(speed=10.0)
|
||||
config.clamp()
|
||||
assert config.speed == MAX_SPEED
|
||||
|
||||
def test_clamp_within_range(self):
|
||||
def test_zero_speed_falls_back_to_default(self):
|
||||
"""速度为0回退到默认."""
|
||||
config = SpeedConfig(speed=0)
|
||||
config.clamp()
|
||||
assert config.speed == 1.0
|
||||
|
||||
def test_negative_speed_falls_back_to_default(self):
|
||||
"""负速度回退到默认."""
|
||||
config = SpeedConfig(speed=-2.0)
|
||||
config.clamp()
|
||||
assert config.speed == 1.0
|
||||
|
||||
def test_speed_at_min_ok(self):
|
||||
"""最小值边界."""
|
||||
config = SpeedConfig(speed=MIN_SPEED)
|
||||
config.clamp()
|
||||
assert config.speed == MIN_SPEED
|
||||
|
||||
def test_speed_at_max_ok(self):
|
||||
"""最大值边界."""
|
||||
config = SpeedConfig(speed=MAX_SPEED)
|
||||
config.clamp()
|
||||
assert config.speed == MAX_SPEED
|
||||
|
||||
def test_speed_in_range_unchanged(self):
|
||||
"""合法范围内不修改."""
|
||||
config = SpeedConfig(speed=1.5)
|
||||
config.clamp()
|
||||
assert config.speed == 1.5
|
||||
|
||||
def test_is_original_true(self):
|
||||
config = SpeedConfig(speed=1.0)
|
||||
assert config.is_original is True
|
||||
|
||||
def test_is_original_false(self):
|
||||
config = SpeedConfig(speed=1.5)
|
||||
assert config.is_original is False
|
||||
|
||||
def test_parse_clamps_automatically(self):
|
||||
"""parse 方法应该自动调用 clamp."""
|
||||
def test_parse_auto_clamps(self):
|
||||
"""parse 自动钳制."""
|
||||
config = SpeedConfig.parse({"speed": 100.0})
|
||||
assert config.speed == MAX_SPEED
|
||||
|
||||
|
||||
# ─── SpeedEngine 视频滤镜 ────────────────────────────────────
|
||||
class TestIsOriginal:
|
||||
"""is_original 属性测试."""
|
||||
|
||||
def test_exactly_one(self):
|
||||
"""速度恰好为1."""
|
||||
assert SpeedConfig(speed=1.0).is_original is True
|
||||
|
||||
def test_very_close_to_one(self):
|
||||
"""非常接近1也算原速."""
|
||||
assert SpeedConfig(speed=1.0000001).is_original is True
|
||||
|
||||
def test_not_one(self):
|
||||
"""不是1."""
|
||||
assert SpeedConfig(speed=1.1).is_original is False
|
||||
assert SpeedConfig(speed=0.9).is_original is False
|
||||
|
||||
|
||||
class TestSpeedEngineVideoFilter:
|
||||
# ── SpeedEngine 测试 ────────────────────────────────────────
|
||||
|
||||
|
||||
class TestBuildVideoFilter:
|
||||
"""build_video_filter 测试."""
|
||||
|
||||
def setup_method(self):
|
||||
self.engine = SpeedEngine()
|
||||
|
||||
def test_original_speed_returns_empty(self):
|
||||
def test_original_speed_empty_filter(self):
|
||||
"""原速返回空字符串(跳过滤镜)."""
|
||||
config = SpeedConfig(speed=1.0)
|
||||
assert self.engine.build_video_filter(config) == ""
|
||||
|
||||
def test_double_speed(self):
|
||||
def test_speed_up_2x(self):
|
||||
"""2倍速."""
|
||||
config = SpeedConfig(speed=2.0)
|
||||
result = self.engine.build_video_filter(config)
|
||||
assert "setpts=PTS/2.0" in result
|
||||
assert "setpts=PTS/2.0000" == result
|
||||
|
||||
def test_half_speed(self):
|
||||
def test_slow_down_half(self):
|
||||
"""0.5倍速."""
|
||||
config = SpeedConfig(speed=0.5)
|
||||
result = self.engine.build_video_filter(config)
|
||||
assert "setpts=PTS/0.5" in result
|
||||
assert "setpts=PTS/0.5000" == result
|
||||
|
||||
def test_quarter_speed(self):
|
||||
config = SpeedConfig(speed=0.25)
|
||||
def test_contains_setpts(self):
|
||||
"""包含setpts滤镜."""
|
||||
config = SpeedConfig(speed=1.5)
|
||||
result = self.engine.build_video_filter(config)
|
||||
assert "setpts=PTS/0.25" in result
|
||||
|
||||
def test_quad_speed(self):
|
||||
config = SpeedConfig(speed=4.0)
|
||||
result = self.engine.build_video_filter(config)
|
||||
assert "setpts=PTS/4.0" in result
|
||||
assert "setpts=PTS/" in result
|
||||
|
||||
|
||||
# ─── SpeedEngine 音频滤镜(atempo 多级串联) ─────────────────
|
||||
class TestBuildAudioFilter:
|
||||
"""build_audio_filter 测试."""
|
||||
|
||||
|
||||
class TestSpeedEngineAudioFilter:
|
||||
def setup_method(self):
|
||||
self.engine = SpeedEngine()
|
||||
|
||||
def test_original_speed_returns_empty(self):
|
||||
def test_original_speed_empty_filter(self):
|
||||
"""原速返回空字符串."""
|
||||
config = SpeedConfig(speed=1.0)
|
||||
assert self.engine.build_audio_filter(config) == ""
|
||||
|
||||
def test_double_speed_single_stage(self):
|
||||
"""2x 在 atempo 单级范围内,只需一个 atempo."""
|
||||
def test_single_stage_2x(self):
|
||||
"""2倍速单级atempo."""
|
||||
config = SpeedConfig(speed=2.0)
|
||||
result = self.engine.build_audio_filter(config)
|
||||
assert result == "atempo=2.0000"
|
||||
|
||||
def test_half_speed_single_stage(self):
|
||||
def test_single_stage_half(self):
|
||||
"""0.5倍速单级atempo."""
|
||||
config = SpeedConfig(speed=0.5)
|
||||
result = self.engine.build_audio_filter(config)
|
||||
assert result == "atempo=0.5000"
|
||||
|
||||
def test_quad_speed_two_stages(self):
|
||||
"""4x 需要两级 atempo: 2.0 * 2.0."""
|
||||
def test_multi_stage_4x(self):
|
||||
"""4倍速需要两级 atempo=2.0,atempo=2.0."""
|
||||
config = SpeedConfig(speed=4.0)
|
||||
result = self.engine.build_audio_filter(config)
|
||||
assert result == "atempo=2.0000,atempo=2.0000"
|
||||
|
||||
def test_quarter_speed_two_stages(self):
|
||||
"""0.25x 需要两级 atempo: 0.5 * 0.5."""
|
||||
def test_multi_stage_quarter(self):
|
||||
"""0.25倍速需要两级 atempo=0.5,atempo=0.5."""
|
||||
config = SpeedConfig(speed=0.25)
|
||||
result = self.engine.build_audio_filter(config)
|
||||
assert result == "atempo=0.5000,atempo=0.5000"
|
||||
|
||||
def test_triple_speed_two_stages(self):
|
||||
"""3x: 2.0 * 1.5."""
|
||||
def test_multi_stage_3x(self):
|
||||
"""3倍速: 2.0 * 1.5."""
|
||||
config = SpeedConfig(speed=3.0)
|
||||
result = self.engine.build_audio_filter(config)
|
||||
parts = result.split(",")
|
||||
assert len(parts) == 2
|
||||
assert "atempo=2.0000" in parts
|
||||
assert "atempo=1.5000" in parts
|
||||
stages = result.split(",")
|
||||
assert len(stages) == 2
|
||||
# 验证两级相乘等于3
|
||||
values = [float(s.split("=")[1]) for s in stages]
|
||||
assert abs(values[0] * values[1] - 3.0) < 0.01
|
||||
|
||||
def test_03_speed_two_stages(self):
|
||||
"""0.3x: 0.5 * 0.6."""
|
||||
config = SpeedConfig(speed=0.3)
|
||||
result = self.engine.build_audio_filter(config)
|
||||
parts = result.split(",")
|
||||
assert len(parts) == 2
|
||||
assert "atempo=0.5000" in parts
|
||||
assert "atempo=0.6000" in parts
|
||||
|
||||
def test_split_atempo_inside_range(self):
|
||||
"""0.5~2.0 范围内只返回一级."""
|
||||
class TestSplitAtempoStages:
|
||||
"""_split_atempo_stages 测试."""
|
||||
|
||||
def test_single_stage_within_range(self):
|
||||
"""范围内单级."""
|
||||
stages = SpeedEngine._split_atempo_stages(1.5)
|
||||
assert len(stages) == 1
|
||||
assert stages[0] == 1.5
|
||||
|
||||
def test_split_atempo_boundary_min(self):
|
||||
stages = SpeedEngine._split_atempo_stages(0.5)
|
||||
assert len(stages) == 1
|
||||
assert stages[0] == 0.5
|
||||
|
||||
def test_split_atempo_boundary_max(self):
|
||||
def test_single_stage_at_max(self):
|
||||
"""最大值边界单级."""
|
||||
stages = SpeedEngine._split_atempo_stages(2.0)
|
||||
assert len(stages) == 1
|
||||
assert stages[0] == 2.0
|
||||
|
||||
def test_split_atempo_product_equals_speed(self):
|
||||
"""所有级联的乘积应该等于原速度."""
|
||||
test_cases = [0.25, 0.3, 0.5, 0.75, 1.0, 1.5, 2.0, 3.0, 4.0]
|
||||
for speed in test_cases:
|
||||
stages = SpeedEngine._split_atempo_stages(speed)
|
||||
product = 1.0
|
||||
for s in stages:
|
||||
product *= s
|
||||
assert abs(product - speed) < 1e-6, f"speed={speed}, stages={stages}, product={product}"
|
||||
def test_single_stage_at_min(self):
|
||||
"""最小值边界单级."""
|
||||
stages = SpeedEngine._split_atempo_stages(0.5)
|
||||
assert len(stages) == 1
|
||||
|
||||
def test_split_atempo_all_in_range(self):
|
||||
"""所有级都应该在 0.5~2.0 范围内."""
|
||||
test_cases = [0.25, 0.3, 0.5, 0.75, 1.0, 1.5, 2.0, 3.0, 4.0]
|
||||
for speed in test_cases:
|
||||
def test_multi_stage_double_speed(self):
|
||||
"""4x 需要两级."""
|
||||
stages = SpeedEngine._split_atempo_stages(4.0)
|
||||
assert len(stages) == 2
|
||||
assert abs(stages[0] * stages[1] - 4.0) < 0.01
|
||||
|
||||
def test_multi_stage_half_speed(self):
|
||||
"""0.25x 需要两级."""
|
||||
stages = SpeedEngine._split_atempo_stages(0.25)
|
||||
assert len(stages) == 2
|
||||
assert abs(stages[0] * stages[1] - 0.25) < 0.01
|
||||
|
||||
def test_all_stages_within_valid_range(self):
|
||||
"""所有分级都在有效范围内."""
|
||||
for speed in [0.25, 0.3, 0.5, 1.0, 1.5, 2.0, 3.0, 4.0]:
|
||||
stages = SpeedEngine._split_atempo_stages(speed)
|
||||
for s in stages:
|
||||
assert 0.5 <= s <= 2.0, f"speed={speed}, stage={s} out of range"
|
||||
|
||||
|
||||
# ─── SpeedEngine 时长计算 ────────────────────────────────────
|
||||
class TestAdjustDuration:
|
||||
"""adjust_duration 时长计算测试."""
|
||||
|
||||
|
||||
class TestSpeedEngineDuration:
|
||||
def setup_method(self):
|
||||
self.engine = SpeedEngine()
|
||||
|
||||
def test_original_speed_same_duration(self):
|
||||
config = SpeedConfig(speed=1.0)
|
||||
assert self.engine.adjust_duration(10.0, config) == 10.0
|
||||
def test_original_speed_unchanged(self):
|
||||
"""原速时长不变."""
|
||||
result = self.engine.adjust_duration(100.0, SpeedConfig(speed=1.0))
|
||||
assert result == 100.0
|
||||
|
||||
def test_double_speed_half_duration(self):
|
||||
config = SpeedConfig(speed=2.0)
|
||||
assert self.engine.adjust_duration(10.0, config) == 5.0
|
||||
"""2倍速时长减半."""
|
||||
result = self.engine.adjust_duration(100.0, SpeedConfig(speed=2.0))
|
||||
assert result == 50.0
|
||||
|
||||
def test_half_speed_double_duration(self):
|
||||
config = SpeedConfig(speed=0.5)
|
||||
assert self.engine.adjust_duration(10.0, config) == 20.0
|
||||
"""0.5倍速时长翻倍."""
|
||||
result = self.engine.adjust_duration(100.0, SpeedConfig(speed=0.5))
|
||||
assert result == 200.0
|
||||
|
||||
def test_quad_speed_quarter_duration(self):
|
||||
config = SpeedConfig(speed=4.0)
|
||||
assert self.engine.adjust_duration(10.0, config) == 2.5
|
||||
def test_zero_duration_unchanged(self):
|
||||
"""零时长不变."""
|
||||
result = self.engine.adjust_duration(0.0, SpeedConfig(speed=2.0))
|
||||
assert result == 0.0
|
||||
|
||||
def test_zero_duration(self):
|
||||
config = SpeedConfig(speed=2.0)
|
||||
assert self.engine.adjust_duration(0.0, config) == 0.0
|
||||
def test_negative_duration_unchanged(self):
|
||||
"""负时长不变(异常值保护)."""
|
||||
result = self.engine.adjust_duration(-10.0, SpeedConfig(speed=2.0))
|
||||
assert result == -10.0
|
||||
|
||||
def test_negative_duration(self):
|
||||
config = SpeedConfig(speed=2.0)
|
||||
assert self.engine.adjust_duration(-1.0, config) == -1.0
|
||||
|
||||
|
||||
# ─── SpeedEngine 便捷方法 ────────────────────────────────────
|
||||
|
||||
|
||||
class TestSpeedEngineHelper:
|
||||
def setup_method(self):
|
||||
self.engine = SpeedEngine()
|
||||
|
||||
def test_build_clip_speed_filter_original(self):
|
||||
v_f, a_f, cfg = self.engine.build_clip_speed_filter(1.0)
|
||||
assert v_f == ""
|
||||
assert a_f == ""
|
||||
assert cfg.speed == 1.0
|
||||
|
||||
def test_build_clip_speed_filter_2x(self):
|
||||
v_f, a_f, cfg = self.engine.build_clip_speed_filter(2.0)
|
||||
assert "setpts=PTS/2.0" in v_f
|
||||
assert "atempo=2.0" in a_f
|
||||
assert cfg.speed == 2.0
|
||||
|
||||
def test_build_clip_speed_clamped(self):
|
||||
_, _, cfg = self.engine.build_clip_speed_filter(100.0)
|
||||
assert cfg.speed == MAX_SPEED
|
||||
|
||||
def test_resolve_clip_speed_default(self):
|
||||
assert SpeedEngine.resolve_clip_speed({}) == 1.0
|
||||
assert SpeedEngine.resolve_clip_speed(None) == 1.0
|
||||
|
||||
def test_resolve_clip_speed_zero_uses_global(self):
|
||||
assert SpeedEngine.resolve_clip_speed({"playback_speed": 0}, 1.5) == 1.5
|
||||
|
||||
def test_resolve_clip_speed_custom(self):
|
||||
assert SpeedEngine.resolve_clip_speed({"playback_speed": 2.0}) == 2.0
|
||||
|
||||
def test_resolve_clip_speed_invalid_type(self):
|
||||
assert SpeedEngine.resolve_clip_speed({"playback_speed": "fast"}) == 1.0
|
||||
def test_quarter_speed(self):
|
||||
"""0.25倍速时长4倍."""
|
||||
result = self.engine.adjust_duration(60.0, SpeedConfig(speed=0.25))
|
||||
assert abs(result - 240.0) < 0.01
|
||||
|
||||
Regular → Executable
+4
-8
@@ -238,15 +238,11 @@ class TestWrapText:
|
||||
result = _wrap_text("", 10)
|
||||
assert result == [""]
|
||||
|
||||
def test_max_chars_zero(self):
|
||||
# 边界情况
|
||||
def test_max_chars_one(self):
|
||||
# max_chars=1 时每个字符一行
|
||||
text = "abc"
|
||||
result = _wrap_text(text, 0)
|
||||
# 0的话,max_chars//2也是0,range不会执行
|
||||
# 按逻辑 len(text) > 0 成立,但 break_point 从 0 开始
|
||||
# 这取决于具体实现,只要不崩溃就行
|
||||
assert isinstance(result, list)
|
||||
assert len(result) > 0
|
||||
result = _wrap_text(text, 1)
|
||||
assert result == ["a", "b", "c"]
|
||||
|
||||
def test_punctuation_at_boundary(self):
|
||||
# 标点刚好在 max_chars 位置
|
||||
|
||||
Executable
+242
@@ -0,0 +1,242 @@
|
||||
"""模板编辑器工具函数测试 — _utils.py 纯函数."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from dataclasses import dataclass
|
||||
|
||||
import pytest
|
||||
from app.api.routes.templates_editor._utils import (
|
||||
_clip_type_to_scene_label,
|
||||
_clip_value,
|
||||
_format_time,
|
||||
_get_adjust_trim,
|
||||
_get_adjust_volume,
|
||||
_get_clip_config,
|
||||
_validate_trim,
|
||||
)
|
||||
|
||||
|
||||
class TestFormatTime:
|
||||
"""_format_time 秒数格式化测试."""
|
||||
|
||||
def test_zero(self):
|
||||
"""0秒."""
|
||||
assert _format_time(0) == "0:00"
|
||||
|
||||
def test_less_than_minute(self):
|
||||
"""小于1分钟."""
|
||||
assert _format_time(30) == "0:30"
|
||||
assert _format_time(5) == "0:05"
|
||||
assert _format_time(59) == "0:59"
|
||||
|
||||
def test_exact_minute(self):
|
||||
"""整分钟."""
|
||||
assert _format_time(60) == "1:00"
|
||||
assert _format_time(120) == "2:00"
|
||||
|
||||
def test_minutes_and_seconds(self):
|
||||
"""几分几秒."""
|
||||
assert _format_time(65) == "1:05"
|
||||
assert _format_time(125) == "2:05"
|
||||
assert _format_time(600) == "10:00"
|
||||
|
||||
def test_float_seconds_truncated(self):
|
||||
"""浮点秒数取整."""
|
||||
assert _format_time(65.9) == "1:05"
|
||||
assert _format_time(65.1) == "1:05"
|
||||
|
||||
|
||||
class TestClipTypeToSceneLabel:
|
||||
"""_clip_type_to_scene_label 片段类型转标签测试."""
|
||||
|
||||
def test_intro(self):
|
||||
assert _clip_type_to_scene_label("intro", "") == "开场"
|
||||
|
||||
def test_title(self):
|
||||
assert _clip_type_to_scene_label("title", "") == "标题"
|
||||
|
||||
def test_product(self):
|
||||
assert _clip_type_to_scene_label("product", "") == "产品展示"
|
||||
|
||||
def test_showcase(self):
|
||||
assert _clip_type_to_scene_label("showcase", "") == "场景展示"
|
||||
|
||||
def test_scene(self):
|
||||
assert _clip_type_to_scene_label("scene", "") == "场景"
|
||||
|
||||
def test_subtitle(self):
|
||||
assert _clip_type_to_scene_label("subtitle", "") == "字幕"
|
||||
|
||||
def test_text(self):
|
||||
assert _clip_type_to_scene_label("text", "") == "文字"
|
||||
|
||||
def test_cta(self):
|
||||
assert _clip_type_to_scene_label("cta", "") == "结尾 CTA"
|
||||
|
||||
def test_outro(self):
|
||||
assert _clip_type_to_scene_label("outro", "") == "结尾"
|
||||
|
||||
def test_voiceover(self):
|
||||
assert _clip_type_to_scene_label("voiceover", "") == "配音"
|
||||
|
||||
def test_transition(self):
|
||||
assert _clip_type_to_scene_label("transition", "") == "转场"
|
||||
|
||||
def test_unknown_type_returns_itself(self):
|
||||
"""未知类型返回类型名本身."""
|
||||
assert _clip_type_to_scene_label("unknown_type", "") == "unknown_type"
|
||||
|
||||
def test_empty_type_fallback(self):
|
||||
"""空类型fallback到片段."""
|
||||
assert _clip_type_to_scene_label("", "") == "片段"
|
||||
|
||||
def test_with_text_content(self):
|
||||
"""带文本内容时追加文本预览."""
|
||||
result = _clip_type_to_scene_label("subtitle", "大家好今天")
|
||||
assert "字幕 - 大家好今天" == result
|
||||
|
||||
def test_text_truncated_at_20_chars(self):
|
||||
"""文本超过20字符截断."""
|
||||
long_text = "一二三四五六七八九十一二三四五六七八九十"
|
||||
result = _clip_type_to_scene_label("text", long_text + "extra")
|
||||
# 前20个字符
|
||||
assert long_text in result
|
||||
assert "extra" not in result
|
||||
|
||||
def test_text_with_only_whitespace(self):
|
||||
"""文本只有空白时不追加."""
|
||||
result = _clip_type_to_scene_label("intro", " ")
|
||||
assert result == "开场"
|
||||
|
||||
|
||||
class TestValidateTrim:
|
||||
"""_validate_trim 裁剪校验测试."""
|
||||
|
||||
def test_valid_trim(self):
|
||||
"""合法裁剪."""
|
||||
_validate_trim(1.0, 1.0, 5.0) # 不抛异常
|
||||
|
||||
def test_zero_trim(self):
|
||||
"""不裁剪也合法."""
|
||||
_validate_trim(0.0, 0.0, 5.0)
|
||||
|
||||
def test_trim_equals_total_raises(self):
|
||||
"""裁剪总时长等于总时长→抛异常."""
|
||||
with pytest.raises(ValueError, match="不能大于等于"):
|
||||
_validate_trim(2.5, 2.5, 5.0)
|
||||
|
||||
def test_trim_exceeds_total_raises(self):
|
||||
"""裁剪超过总时长→抛异常."""
|
||||
with pytest.raises(ValueError):
|
||||
_validate_trim(3.0, 3.0, 5.0)
|
||||
|
||||
def test_only_start_exceeds(self):
|
||||
"""只有start就超过."""
|
||||
with pytest.raises(ValueError):
|
||||
_validate_trim(6.0, 0.0, 5.0)
|
||||
|
||||
def test_only_end_exceeds(self):
|
||||
"""只有end就超过."""
|
||||
with pytest.raises(ValueError):
|
||||
_validate_trim(0.0, 6.0, 5.0)
|
||||
|
||||
|
||||
class TestClipValue:
|
||||
"""_clip_value 枚举/字符串值提取测试."""
|
||||
|
||||
def test_plain_string(self):
|
||||
"""普通字符串返回自身."""
|
||||
assert _clip_value("hello") == "hello"
|
||||
|
||||
def test_enum_value(self):
|
||||
"""带value属性的对象返回value."""
|
||||
|
||||
class FakeEnum:
|
||||
value = "enum_value"
|
||||
|
||||
assert _clip_value(FakeEnum()) == "enum_value"
|
||||
|
||||
def test_int_value(self):
|
||||
"""整数转字符串."""
|
||||
assert _clip_value(42) == "42"
|
||||
|
||||
|
||||
@dataclass
|
||||
class FakeClip:
|
||||
"""测试用假Clip对象."""
|
||||
|
||||
id: str = "clip_1"
|
||||
playback_speed: float = 1.0
|
||||
duration: float = 10.0
|
||||
config: dict | None = None
|
||||
|
||||
|
||||
class TestGetClipConfig:
|
||||
"""_get_clip_config 安全获取配置测试."""
|
||||
|
||||
def test_normal_config(self):
|
||||
"""正常dict配置."""
|
||||
clip = FakeClip(config={"volume": 0.5})
|
||||
assert _get_clip_config(clip) == {"volume": 0.5}
|
||||
|
||||
def test_none_config(self):
|
||||
"""config为None→返回空dict."""
|
||||
clip = FakeClip(config=None)
|
||||
assert _get_clip_config(clip) == {}
|
||||
|
||||
def test_non_dict_config(self):
|
||||
"""config不是dict→返回空dict."""
|
||||
clip = FakeClip(config="not_a_dict")
|
||||
assert _get_clip_config(clip) == {}
|
||||
|
||||
def test_no_config_attribute(self):
|
||||
"""没有config属性→返回空dict."""
|
||||
|
||||
class NoConfig:
|
||||
pass
|
||||
|
||||
assert _get_clip_config(NoConfig()) == {}
|
||||
|
||||
|
||||
class TestGetAdjustVolume:
|
||||
"""_get_adjust_volume 获取音量测试."""
|
||||
|
||||
def test_default_volume(self):
|
||||
"""无配置默认1.0."""
|
||||
clip = FakeClip(config={})
|
||||
assert _get_adjust_volume(clip) == pytest.approx(1.0)
|
||||
|
||||
def test_custom_volume(self):
|
||||
"""自定义音量."""
|
||||
clip = FakeClip(config={"volume": 0.7})
|
||||
assert _get_adjust_volume(clip) == pytest.approx(0.7)
|
||||
|
||||
def test_none_config(self):
|
||||
"""None config."""
|
||||
clip = FakeClip(config=None)
|
||||
assert _get_adjust_volume(clip) == pytest.approx(1.0)
|
||||
|
||||
|
||||
class TestGetAdjustTrim:
|
||||
"""_get_adjust_trim 获取裁剪测试."""
|
||||
|
||||
def test_default_trim(self):
|
||||
"""无配置默认都是0."""
|
||||
clip = FakeClip(config={})
|
||||
start, end = _get_adjust_trim(clip)
|
||||
assert start == pytest.approx(0.0)
|
||||
assert end == pytest.approx(0.0)
|
||||
|
||||
def test_custom_trim(self):
|
||||
"""自定义裁剪."""
|
||||
clip = FakeClip(config={"trim_start": 1.5, "trim_end": 2.0})
|
||||
start, end = _get_adjust_trim(clip)
|
||||
assert start == pytest.approx(1.5)
|
||||
assert end == pytest.approx(2.0)
|
||||
|
||||
def test_none_config(self):
|
||||
"""None config."""
|
||||
clip = FakeClip(config=None)
|
||||
start, end = _get_adjust_trim(clip)
|
||||
assert start == pytest.approx(0.0)
|
||||
assert end == pytest.approx(0.0)
|
||||
@@ -1,9 +1,4 @@
|
||||
"""
|
||||
缩略图生成器纯函数测试.
|
||||
|
||||
覆盖 _format_seek_time 等纯逻辑.
|
||||
FFmpeg 抽帧与 OSS 上传由集成测试覆盖.
|
||||
"""
|
||||
"""缩略图生成器单元测试 - 纯逻辑函数."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
@@ -12,49 +7,82 @@ from video_processing.thumbnail_generator import _format_seek_time
|
||||
|
||||
|
||||
class TestFormatSeekTime:
|
||||
"""_format_seek_time 时间格式化."""
|
||||
"""_format_seek_time 时间格式化测试."""
|
||||
|
||||
def test_zero(self):
|
||||
assert _format_seek_time(0.0) == "00:00:00.00"
|
||||
def test_zero_seconds(self):
|
||||
"""0秒."""
|
||||
result = _format_seek_time(0)
|
||||
assert result == "00:00:00.00"
|
||||
|
||||
def test_seconds_only(self):
|
||||
assert _format_seek_time(5.5) == "00:00:05.50"
|
||||
def test_less_than_one_second(self):
|
||||
"""小于1秒."""
|
||||
result = _format_seek_time(0.5)
|
||||
assert result == "00:00:00.50"
|
||||
|
||||
def test_minutes(self):
|
||||
assert _format_seek_time(65.25) == "00:01:05.25"
|
||||
def test_few_seconds(self):
|
||||
"""几秒."""
|
||||
result = _format_seek_time(5.5)
|
||||
assert result == "00:00:05.50"
|
||||
|
||||
def test_hours(self):
|
||||
assert _format_seek_time(3661.5) == "01:01:01.50"
|
||||
def test_one_minute(self):
|
||||
"""1分钟."""
|
||||
result = _format_seek_time(60.0)
|
||||
assert result == "00:01:00.00"
|
||||
|
||||
def test_exact_minute(self):
|
||||
assert _format_seek_time(60.0) == "00:01:00.00"
|
||||
def test_minutes_and_seconds(self):
|
||||
"""分+秒."""
|
||||
result = _format_seek_time(125.5)
|
||||
assert result == "00:02:05.50"
|
||||
|
||||
def test_exact_hour(self):
|
||||
assert _format_seek_time(3600.0) == "01:00:00.00"
|
||||
def test_one_hour(self):
|
||||
"""1小时."""
|
||||
result = _format_seek_time(3600.0)
|
||||
assert result == "01:00:00.00"
|
||||
|
||||
def test_very_short(self):
|
||||
assert _format_seek_time(0.1) == "00:00:00.10"
|
||||
def test_hours_minutes_seconds(self):
|
||||
"""时+分+秒."""
|
||||
result = _format_seek_time(3725.25)
|
||||
assert result == "01:02:05.25"
|
||||
|
||||
def test_long_video(self):
|
||||
# 超过1小时
|
||||
assert _format_seek_time(7200.0) == "02:00:00.00"
|
||||
def test_long_duration(self):
|
||||
"""长视频(2小时以上)."""
|
||||
result = _format_seek_time(7384.12)
|
||||
assert result == "02:03:04.12"
|
||||
|
||||
def test_sub_second_precision(self):
|
||||
result = _format_seek_time(1.234)
|
||||
def test_precision_two_decimal(self):
|
||||
"""两位小数精度."""
|
||||
result = _format_seek_time(3.14159)
|
||||
assert result == "00:00:03.14"
|
||||
|
||||
def test_always_two_digit_hours(self):
|
||||
"""小时始终两位数字."""
|
||||
result = _format_seek_time(3600 * 9)
|
||||
assert result.startswith("09:")
|
||||
|
||||
def test_always_two_digit_minutes(self):
|
||||
"""分钟始终两位数字."""
|
||||
result = _format_seek_time(300) # 5分钟
|
||||
parts = result.split(":")
|
||||
assert parts[1] == "05"
|
||||
|
||||
def test_float_input(self):
|
||||
"""浮点数输入."""
|
||||
result = _format_seek_time(10.0)
|
||||
assert isinstance(result, str)
|
||||
assert result == "00:00:10.00"
|
||||
|
||||
def test_int_input(self):
|
||||
"""整数输入."""
|
||||
result = _format_seek_time(30)
|
||||
assert result == "00:00:30.00"
|
||||
|
||||
def test_format_structure(self):
|
||||
"""格式结构正确:HH:MM:SS.xx."""
|
||||
result = _format_seek_time(3661.5)
|
||||
# 格式: HH:MM:SS.xx
|
||||
parts = result.split(":")
|
||||
assert len(parts) == 3
|
||||
sec_part = parts[2]
|
||||
assert "." in sec_part
|
||||
decimals = sec_part.split(".")[1]
|
||||
assert len(decimals) == 2
|
||||
|
||||
def test_zero_padded_hours(self):
|
||||
# 小时始终是2位
|
||||
result = _format_seek_time(5.0)
|
||||
assert result.startswith("00:")
|
||||
|
||||
def test_zero_padded_minutes(self):
|
||||
# 分钟始终是2位
|
||||
result = _format_seek_time(5.0)
|
||||
parts = result.split(":")
|
||||
assert len(parts[1]) == 2
|
||||
assert "." in parts[2]
|
||||
sec_parts = parts[2].split(".")
|
||||
assert len(sec_parts) == 2
|
||||
assert len(sec_parts[1]) == 2 # 两位小数
|
||||
|
||||
@@ -1,4 +1,4 @@
|
||||
"""转场特效引擎单测 — Phase 8 智能增强."""
|
||||
"""转场引擎单元测试 - 配置解析等纯逻辑."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
@@ -9,476 +9,167 @@ from video_processing.transition_engine import (
|
||||
MAX_TRANSITION_DURATION,
|
||||
MIN_TRANSITION_DURATION,
|
||||
TransitionConfig,
|
||||
TransitionEngine,
|
||||
TransitionType,
|
||||
_normalize_transition_name,
|
||||
)
|
||||
|
||||
# ── TransitionType 枚举测试 ──────────────────────────────────────────────────
|
||||
|
||||
class TestTransitionConstants:
|
||||
"""常量测试."""
|
||||
|
||||
def test_duration_ranges(self):
|
||||
"""时长范围合理."""
|
||||
assert MIN_TRANSITION_DURATION == 0.3
|
||||
assert MAX_TRANSITION_DURATION == 2.0
|
||||
assert DEFAULT_TRANSITION_DURATION == 0.5
|
||||
assert MIN_TRANSITION_DURATION < DEFAULT_TRANSITION_DURATION < MAX_TRANSITION_DURATION
|
||||
|
||||
def test_cut_transition_value(self):
|
||||
"""cut转场值."""
|
||||
assert CUT_TRANSITION == "cut"
|
||||
|
||||
|
||||
class TestTransitionType:
|
||||
"""TransitionType 枚举测试."""
|
||||
|
||||
def test_all_supported_count(self):
|
||||
"""支持的转场类型数量(不含cut)."""
|
||||
supported = TransitionType.all_supported()
|
||||
# 至少 8 种:fade, dissolve, slide*4, zoom, wipe*4, circlecrop, rectcrop
|
||||
assert len(supported) >= 8
|
||||
assert "fade" in supported
|
||||
assert "dissolve" in supported
|
||||
assert "zoom" in supported
|
||||
assert "circlecrop" in supported
|
||||
assert "rectcrop" in supported
|
||||
def test_supports_fade(self):
|
||||
"""支持fade."""
|
||||
assert TransitionType.is_supported("fade") is True
|
||||
|
||||
def test_slide_directions(self):
|
||||
"""四个方向的滑入转场都支持."""
|
||||
assert TransitionType.is_supported("slideleft")
|
||||
assert TransitionType.is_supported("slideright")
|
||||
assert TransitionType.is_supported("slideup")
|
||||
assert TransitionType.is_supported("slidedown")
|
||||
def test_supports_cut(self):
|
||||
"""cut也在TransitionType枚举中."""
|
||||
assert "cut" in [t.value for t in TransitionType]
|
||||
|
||||
def test_wipe_directions(self):
|
||||
"""四个方向的擦除转场都支持."""
|
||||
assert TransitionType.is_supported("wipeleft")
|
||||
assert TransitionType.is_supported("wiperight")
|
||||
assert TransitionType.is_supported("wipeup")
|
||||
assert TransitionType.is_supported("wipedown")
|
||||
def test_unsupported_effect(self):
|
||||
"""不支持的效果."""
|
||||
assert TransitionType.is_supported("nonexistent_effect_xyz") is False
|
||||
|
||||
def test_is_supported_case_insensitive(self):
|
||||
"""大小写不敏感."""
|
||||
assert TransitionType.is_supported("FADE")
|
||||
assert TransitionType.is_supported("Fade")
|
||||
assert TransitionType.is_supported("fade")
|
||||
"""是否大小写不敏感(看实现)."""
|
||||
# 直接测试几个已知的
|
||||
assert TransitionType.is_supported("fade") is True
|
||||
assert TransitionType.is_supported("dissolve") is True
|
||||
|
||||
def test_is_supported_with_underscores(self):
|
||||
"""下划线不影响判断."""
|
||||
assert TransitionType.is_supported("slide_left")
|
||||
assert TransitionType.is_supported("slide-left")
|
||||
|
||||
def test_is_supported_aliases(self):
|
||||
"""别名支持."""
|
||||
assert TransitionType.is_supported("crossfade")
|
||||
assert TransitionType.is_supported("dissolve")
|
||||
assert TransitionType.is_supported("zoomin")
|
||||
assert TransitionType.is_supported("wipe")
|
||||
|
||||
def test_unsupported_transition(self):
|
||||
"""不支持的转场返回 False."""
|
||||
assert not TransitionType.is_supported("nonexistent_effect")
|
||||
assert not TransitionType.is_supported("random_stuff")
|
||||
assert not TransitionType.is_supported("")
|
||||
|
||||
def test_cut_not_in_supported(self):
|
||||
"""硬切不在"支持的转场效果"列表中(它不是特效)."""
|
||||
supported = TransitionType.all_supported()
|
||||
assert "cut" not in supported
|
||||
def test_all_types_have_value(self):
|
||||
"""所有枚举都有有效值."""
|
||||
for t in TransitionType:
|
||||
assert isinstance(t.value, str)
|
||||
assert len(t.value) > 0
|
||||
|
||||
|
||||
# ── 名称标准化测试 ────────────────────────────────────────────────────────────
|
||||
class TestTransitionConfigParse:
|
||||
"""TransitionConfig.parse 解析测试."""
|
||||
|
||||
def test_no_args_default(self):
|
||||
"""无参数默认配置."""
|
||||
config = TransitionConfig.parse()
|
||||
assert config.effect == CUT_TRANSITION
|
||||
assert config.duration == DEFAULT_TRANSITION_DURATION
|
||||
|
||||
class TestNormalizeTransitionName:
|
||||
"""名称标准化函数测试."""
|
||||
def test_none_effect_default(self):
|
||||
"""None effect默认为cut."""
|
||||
config = TransitionConfig.parse(effect=None)
|
||||
assert config.effect == CUT_TRANSITION
|
||||
|
||||
def test_lowercase(self):
|
||||
"""大写转小写."""
|
||||
assert _normalize_transition_name("FADE") == "fade"
|
||||
assert _normalize_transition_name("Fade") == "fade"
|
||||
def test_empty_effect_default(self):
|
||||
"""空字符串effect默认为cut."""
|
||||
config = TransitionConfig.parse(effect="")
|
||||
assert config.effect == CUT_TRANSITION
|
||||
|
||||
def test_remove_underscores(self):
|
||||
"""移除下划线."""
|
||||
assert _normalize_transition_name("slide_left") == "slideleft"
|
||||
assert _normalize_transition_name("slide_up") == "slideup"
|
||||
|
||||
def test_remove_hyphens(self):
|
||||
"""移除连字符."""
|
||||
assert _normalize_transition_name("slide-left") == "slideleft"
|
||||
|
||||
def test_mixed(self):
|
||||
"""混合情况."""
|
||||
assert _normalize_transition_name("Slide_Left") == "slideleft"
|
||||
assert _normalize_transition_name("FADE-IN") == "fadein"
|
||||
|
||||
|
||||
# ── TransitionConfig 测试 ────────────────────────────────────────────────────
|
||||
|
||||
|
||||
class TestTransitionConfig:
|
||||
"""TransitionConfig 配置解析测试."""
|
||||
|
||||
# ── 默认值 ──
|
||||
|
||||
def test_default_config(self):
|
||||
"""默认配置是硬切."""
|
||||
cfg = TransitionConfig.parse()
|
||||
assert cfg.effect == CUT_TRANSITION
|
||||
assert cfg.duration == DEFAULT_TRANSITION_DURATION
|
||||
assert cfg.is_cut is True
|
||||
|
||||
def test_none_effect(self):
|
||||
"""None effect 降级为 cut."""
|
||||
cfg = TransitionConfig.parse(effect=None)
|
||||
assert cfg.effect == CUT_TRANSITION
|
||||
assert cfg.is_cut is True
|
||||
|
||||
def test_empty_effect(self):
|
||||
"""空字符串 effect 降级为 cut."""
|
||||
cfg = TransitionConfig.parse(effect="")
|
||||
assert cfg.effect == CUT_TRANSITION
|
||||
assert cfg.is_cut is True
|
||||
|
||||
# ── 有效转场类型 ──
|
||||
def test_whitespace_effect_default(self):
|
||||
"""空白effect默认为cut."""
|
||||
config = TransitionConfig.parse(effect=" ")
|
||||
assert config.effect == CUT_TRANSITION
|
||||
|
||||
def test_fade_effect(self):
|
||||
"""fade 转场."""
|
||||
cfg = TransitionConfig.parse(effect="fade")
|
||||
assert cfg.effect == "fade"
|
||||
assert cfg.is_cut is False
|
||||
assert cfg.ffmpeg_transition == "fade"
|
||||
"""fade效果."""
|
||||
config = TransitionConfig.parse(effect="fade")
|
||||
assert config.effect == "fade"
|
||||
|
||||
def test_dissolve_effect(self):
|
||||
"""dissolve 转场."""
|
||||
cfg = TransitionConfig.parse(effect="dissolve")
|
||||
assert cfg.effect == "dissolve"
|
||||
assert cfg.ffmpeg_transition == "dissolve"
|
||||
def test_unsupported_effect_falls_back_to_cut(self):
|
||||
"""不支持的效果降级到cut."""
|
||||
config = TransitionConfig.parse(effect="super_cool_effect")
|
||||
assert config.effect == CUT_TRANSITION
|
||||
|
||||
def test_zoom_effect(self):
|
||||
"""zoom 转场 → FFmpeg zoomin."""
|
||||
cfg = TransitionConfig.parse(effect="zoom")
|
||||
assert cfg.effect == "zoom"
|
||||
assert cfg.ffmpeg_transition == "zoomin"
|
||||
def test_cut_effect(self):
|
||||
"""显式cut效果."""
|
||||
config = TransitionConfig.parse(effect="cut")
|
||||
assert config.effect == CUT_TRANSITION
|
||||
|
||||
def test_slide_left_alias(self):
|
||||
"""slide_left 别名."""
|
||||
cfg = TransitionConfig.parse(effect="slide_left")
|
||||
assert cfg.effect == "slideleft"
|
||||
assert cfg.ffmpeg_transition == "slideleft"
|
||||
|
||||
def test_wipe_alias(self):
|
||||
"""wipe 别名 → 默认向左擦."""
|
||||
cfg = TransitionConfig.parse(effect="wipe")
|
||||
assert cfg.effect == "wipeleft"
|
||||
assert cfg.ffmpeg_transition == "wipeleft"
|
||||
|
||||
def test_circlecrop_effect(self):
|
||||
"""圆形扩散转场."""
|
||||
cfg = TransitionConfig.parse(effect="circlecrop")
|
||||
assert cfg.effect == "circlecrop"
|
||||
assert cfg.ffmpeg_transition == "circlecrop"
|
||||
|
||||
def test_rectcrop_effect(self):
|
||||
"""矩形扩散转场."""
|
||||
cfg = TransitionConfig.parse(effect="rectcrop")
|
||||
assert cfg.effect == "rectcrop"
|
||||
assert cfg.ffmpeg_transition == "rectcrop"
|
||||
|
||||
# ── 降级策略 ──
|
||||
|
||||
def test_unsupported_fallback_to_cut(self):
|
||||
"""不支持的转场自动降级为硬切,不阻断渲染."""
|
||||
cfg = TransitionConfig.parse(effect="nonexistent_effect")
|
||||
assert cfg.effect == CUT_TRANSITION
|
||||
assert cfg.is_cut is True
|
||||
|
||||
def test_unsupported_whitespace_fallback(self):
|
||||
"""带空格的不支持转场也降级."""
|
||||
cfg = TransitionConfig.parse(effect=" bad effect ")
|
||||
assert cfg.effect == CUT_TRANSITION
|
||||
|
||||
# ── 时长边界校验 ──
|
||||
|
||||
def test_default_duration(self):
|
||||
"""默认时长 0.5s."""
|
||||
cfg = TransitionConfig.parse(effect="fade")
|
||||
assert cfg.duration == 0.5
|
||||
|
||||
def test_duration_within_range(self):
|
||||
"""正常范围内的时长."""
|
||||
cfg = TransitionConfig.parse(effect="fade", duration=1.0)
|
||||
assert cfg.duration == 1.0
|
||||
|
||||
def test_duration_min_boundary(self):
|
||||
"""最小值边界."""
|
||||
cfg = TransitionConfig.parse(effect="fade", duration=MIN_TRANSITION_DURATION)
|
||||
assert cfg.duration == MIN_TRANSITION_DURATION
|
||||
|
||||
def test_duration_max_boundary(self):
|
||||
"""最大值边界."""
|
||||
cfg = TransitionConfig.parse(effect="fade", duration=MAX_TRANSITION_DURATION)
|
||||
assert cfg.duration == MAX_TRANSITION_DURATION
|
||||
def test_custom_duration(self):
|
||||
"""自定义时长."""
|
||||
config = TransitionConfig.parse(duration=1.0)
|
||||
assert config.duration == 1.0
|
||||
|
||||
def test_duration_below_min_clamped(self):
|
||||
"""低于最小值的时长被钳制."""
|
||||
cfg = TransitionConfig.parse(effect="fade", duration=0.1)
|
||||
assert cfg.duration == MIN_TRANSITION_DURATION
|
||||
assert cfg.duration >= MIN_TRANSITION_DURATION
|
||||
"""时长低于最小值钳制."""
|
||||
config = TransitionConfig.parse(duration=0.1)
|
||||
assert config.duration == MIN_TRANSITION_DURATION
|
||||
|
||||
def test_duration_above_max_clamped(self):
|
||||
"""高于最大值的时长被钳制."""
|
||||
cfg = TransitionConfig.parse(effect="fade", duration=5.0)
|
||||
assert cfg.duration == MAX_TRANSITION_DURATION
|
||||
assert cfg.duration <= MAX_TRANSITION_DURATION
|
||||
"""时长高于最大值钳制."""
|
||||
config = TransitionConfig.parse(duration=5.0)
|
||||
assert config.duration == MAX_TRANSITION_DURATION
|
||||
|
||||
def test_duration_zero_default_for_effect(self):
|
||||
"""有转场效果但 duration=0 时使用默认值."""
|
||||
# 0.0 会被当作小于最小值钳制到 0.3
|
||||
cfg = TransitionConfig.parse(effect="fade", duration=0.0)
|
||||
assert cfg.duration == MIN_TRANSITION_DURATION
|
||||
def test_duration_at_min(self):
|
||||
"""时长边界最小值."""
|
||||
config = TransitionConfig.parse(duration=MIN_TRANSITION_DURATION)
|
||||
assert config.duration == MIN_TRANSITION_DURATION
|
||||
|
||||
def test_duration_negative_clamped(self):
|
||||
"""负时长被钳制到最小值."""
|
||||
cfg = TransitionConfig.parse(effect="fade", duration=-1.0)
|
||||
assert cfg.duration == MIN_TRANSITION_DURATION
|
||||
def test_duration_at_max(self):
|
||||
"""时长边界最大值."""
|
||||
config = TransitionConfig.parse(duration=MAX_TRANSITION_DURATION)
|
||||
assert config.duration == MAX_TRANSITION_DURATION
|
||||
|
||||
def test_duration_none_uses_default(self):
|
||||
"""None duration 使用默认值."""
|
||||
cfg = TransitionConfig.parse(effect="fade", duration=None)
|
||||
assert cfg.duration == DEFAULT_TRANSITION_DURATION
|
||||
def test_invalid_duration_falls_back(self):
|
||||
"""无效时长回退到默认."""
|
||||
config = TransitionConfig.parse(duration="not_a_number")
|
||||
assert config.duration == DEFAULT_TRANSITION_DURATION
|
||||
|
||||
def test_duration_invalid_type(self):
|
||||
"""无效类型的时长使用默认值."""
|
||||
cfg = TransitionConfig.parse(effect="fade", duration="abc") # type: ignore
|
||||
assert cfg.duration == DEFAULT_TRANSITION_DURATION
|
||||
def test_none_duration_default(self):
|
||||
"""None时长用默认值."""
|
||||
config = TransitionConfig.parse(duration=None)
|
||||
assert config.duration == DEFAULT_TRANSITION_DURATION
|
||||
|
||||
# ── cut 的 ffmpeg_transition ──
|
||||
|
||||
def test_cut_ffmpeg_transition_empty(self):
|
||||
"""硬切没有对应的 FFmpeg xfade transition."""
|
||||
cfg = TransitionConfig.parse(effect="cut")
|
||||
assert cfg.ffmpeg_transition == ""
|
||||
def test_effect_and_duration(self):
|
||||
"""同时指定效果和时长."""
|
||||
config = TransitionConfig.parse(effect="fade", duration=1.0)
|
||||
assert config.effect == "fade"
|
||||
assert config.duration == 1.0
|
||||
|
||||
|
||||
# ── TransitionEngine 测试 ────────────────────────────────────────────────────
|
||||
class TestIsCut:
|
||||
"""is_cut 属性测试."""
|
||||
|
||||
def test_cut_is_cut(self):
|
||||
"""cut是硬切."""
|
||||
config = TransitionConfig(effect=CUT_TRANSITION, duration=0.5)
|
||||
assert config.is_cut is True
|
||||
|
||||
def test_fade_not_cut(self):
|
||||
"""fade不是硬切."""
|
||||
config = TransitionConfig(effect="fade", duration=0.5)
|
||||
assert config.is_cut is False
|
||||
|
||||
|
||||
class TestTransitionEngine:
|
||||
"""TransitionEngine 转场引擎测试."""
|
||||
class TestFfmpegTransition:
|
||||
"""ffmpeg_transition 属性测试."""
|
||||
|
||||
def test_default_engine(self):
|
||||
"""默认引擎初始化."""
|
||||
engine = TransitionEngine()
|
||||
assert engine is not None
|
||||
def test_cut_returns_empty(self):
|
||||
"""cut返回空字符串."""
|
||||
config = TransitionConfig(effect=CUT_TRANSITION, duration=0.5)
|
||||
assert config.ffmpeg_transition == ""
|
||||
|
||||
def test_custom_default_duration(self):
|
||||
"""自定义默认时长."""
|
||||
engine = TransitionEngine(default_duration=1.0)
|
||||
cfg = engine.resolve_config(effect="fade")
|
||||
assert cfg.duration == 1.0
|
||||
def test_fade_returns_fade(self):
|
||||
"""fade返回fade."""
|
||||
config = TransitionConfig(effect="fade", duration=0.5)
|
||||
result = config.ffmpeg_transition
|
||||
assert isinstance(result, str)
|
||||
assert len(result) > 0
|
||||
|
||||
def test_resolve_config_fade(self):
|
||||
"""解析 fade 配置."""
|
||||
engine = TransitionEngine()
|
||||
cfg = engine.resolve_config(effect="fade", duration=0.8)
|
||||
assert cfg.effect == "fade"
|
||||
assert cfg.duration == 0.8
|
||||
|
||||
def test_resolve_config_fallback(self):
|
||||
"""不支持的转场降级."""
|
||||
engine = TransitionEngine()
|
||||
cfg = engine.resolve_config(effect="unknown_effect")
|
||||
assert cfg.effect == CUT_TRANSITION
|
||||
assert cfg.is_cut is True
|
||||
|
||||
def test_resolve_config_duration_clamp(self):
|
||||
"""时长边界钳制."""
|
||||
engine = TransitionEngine()
|
||||
cfg = engine.resolve_config(effect="fade", duration=3.0)
|
||||
assert cfg.duration == MAX_TRANSITION_DURATION
|
||||
|
||||
# ── 批量解析 ──
|
||||
|
||||
def test_resolve_clip_transitions_all_valid(self):
|
||||
"""批量解析全部有效转场."""
|
||||
engine = TransitionEngine()
|
||||
configs = engine.resolve_clip_transitions(["cut", "fade", "dissolve", "slideleft"])
|
||||
assert len(configs) == 4
|
||||
assert configs[0].effect == "cut"
|
||||
assert configs[0].is_cut is True
|
||||
assert configs[1].effect == "fade"
|
||||
assert configs[2].effect == "dissolve"
|
||||
assert configs[3].effect == "slideleft"
|
||||
|
||||
def test_resolve_clip_transitions_with_fallback(self):
|
||||
"""批量解析包含不支持的转场,自动降级."""
|
||||
engine = TransitionEngine()
|
||||
configs = engine.resolve_clip_transitions(["fade", "bad_effect", "dissolve", "worse_effect"])
|
||||
assert len(configs) == 4
|
||||
assert configs[0].effect == "fade"
|
||||
assert configs[1].effect == "cut" # 降级
|
||||
assert configs[2].effect == "dissolve"
|
||||
assert configs[3].effect == "cut" # 降级
|
||||
|
||||
def test_resolve_clip_transitions_with_durations(self):
|
||||
"""带时长校验的批量解析(转场时长不超过片段时长的一半)."""
|
||||
engine = TransitionEngine(default_duration=1.0)
|
||||
# 片段只有 1.0s,转场时长被限制在 0.5s
|
||||
configs = engine.resolve_clip_transitions(
|
||||
["fade", "dissolve"],
|
||||
clip_durations=[1.0, 1.0],
|
||||
)
|
||||
assert len(configs) == 2
|
||||
# 1.0s 默认值超过了片段时长的一半 (0.5s),所以被钳制
|
||||
assert configs[0].duration <= 0.5
|
||||
assert configs[1].duration <= 0.5
|
||||
|
||||
def test_resolve_clip_transitions_short_clip_min_bound(self):
|
||||
"""超短片段的转场时长至少为最小值."""
|
||||
engine = TransitionEngine()
|
||||
configs = engine.resolve_clip_transitions(
|
||||
["fade"],
|
||||
clip_durations=[0.1], # 极短片段
|
||||
)
|
||||
assert len(configs) == 1
|
||||
# 0.1 * 0.5 = 0.05 < MIN_TRANSITION_DURATION,所以用最小值
|
||||
assert configs[0].duration == MIN_TRANSITION_DURATION
|
||||
|
||||
# ── xfade 滤镜链构建 ──
|
||||
|
||||
def test_build_xfade_single_clip(self):
|
||||
"""单 clip 直接 copy."""
|
||||
engine = TransitionEngine()
|
||||
filter_str, total_dur = engine.build_xfade_chain(
|
||||
clip_durations=[5.0],
|
||||
clip_video_labels=["v0"],
|
||||
transitions=["cut"],
|
||||
output_label="outv",
|
||||
)
|
||||
assert "copy" in filter_str
|
||||
assert "[outv]" in filter_str
|
||||
assert total_dur == pytest.approx(5.0, abs=0.01)
|
||||
|
||||
def test_build_xfade_two_clips_fade(self):
|
||||
"""两个 clip 之间 fade 转场."""
|
||||
engine = TransitionEngine()
|
||||
filter_str, total_dur = engine.build_xfade_chain(
|
||||
clip_durations=[3.0, 4.0],
|
||||
clip_video_labels=["v0", "v1"],
|
||||
transitions=["cut", "fade"],
|
||||
output_label="outv",
|
||||
)
|
||||
assert "xfade" in filter_str
|
||||
assert "transition=fade" in filter_str
|
||||
# 总时长 = 3 + 4 - transition_duration (0.5) = 6.5
|
||||
assert total_dur == pytest.approx(6.5, abs=0.1)
|
||||
|
||||
def test_build_xfade_three_clips_mixed(self):
|
||||
"""三个 clip 混合转场."""
|
||||
engine = TransitionEngine()
|
||||
filter_str, total_dur = engine.build_xfade_chain(
|
||||
clip_durations=[3.0, 4.0, 5.0],
|
||||
clip_video_labels=["v0", "v1", "v2"],
|
||||
transitions=["cut", "fade", "dissolve"],
|
||||
output_label="outv",
|
||||
)
|
||||
assert "xfade" in filter_str
|
||||
assert "transition=fade" in filter_str
|
||||
assert "transition=dissolve" in filter_str
|
||||
# 总时长 ≈ 3 + 4 + 5 - 2 * 0.5 = 11.0
|
||||
assert total_dur == pytest.approx(11.0, abs=0.2)
|
||||
|
||||
def test_build_xfade_with_custom_duration(self):
|
||||
"""自定义转场时长."""
|
||||
engine = TransitionEngine(default_duration=0.5)
|
||||
filter_str, total_dur = engine.build_xfade_chain(
|
||||
clip_durations=[3.0, 4.0],
|
||||
clip_video_labels=["v0", "v1"],
|
||||
transitions=["cut", "fade"],
|
||||
transition_duration=1.0,
|
||||
output_label="outv",
|
||||
)
|
||||
assert "xfade" in filter_str
|
||||
# 总时长 = 3 + 4 - 1.0 = 6.0
|
||||
assert total_dur == pytest.approx(6.0, abs=0.1)
|
||||
|
||||
def test_build_xfade_zoom_transition(self):
|
||||
"""zoom 转场滤镜构建."""
|
||||
engine = TransitionEngine()
|
||||
filter_str, _ = engine.build_xfade_chain(
|
||||
clip_durations=[3.0, 4.0],
|
||||
clip_video_labels=["v0", "v1"],
|
||||
transitions=["cut", "zoom"],
|
||||
)
|
||||
assert "xfade" in filter_str
|
||||
assert "transition=zoomin" in filter_str # zoom → zoomin
|
||||
|
||||
def test_build_xfade_slide_directions(self):
|
||||
"""四个方向的滑入转场."""
|
||||
engine = TransitionEngine()
|
||||
for direction in ["slideleft", "slideright", "slideup", "slidedown"]:
|
||||
filter_str, _ = engine.build_xfade_chain(
|
||||
clip_durations=[3.0, 4.0],
|
||||
clip_video_labels=["v0", "v1"],
|
||||
transitions=["cut", direction],
|
||||
)
|
||||
assert f"transition={direction}" in filter_str
|
||||
|
||||
def test_build_xfade_fallback_transition(self):
|
||||
"""不支持的转场降级后构建(降级为cut,等效于极短fade)."""
|
||||
engine = TransitionEngine()
|
||||
# bad_effect 降级为 cut,cut 使用极短转场
|
||||
filter_str, _ = engine.build_xfade_chain(
|
||||
clip_durations=[3.0, 4.0],
|
||||
clip_video_labels=["v0", "v1"],
|
||||
transitions=["cut", "bad_effect"],
|
||||
)
|
||||
# 降级后是 cut,cut 会被 xfade 层映射为 fade(因为 cut 不在 map 里)
|
||||
# 但时长会很短,所以仍然有 xfade
|
||||
assert "xfade" in filter_str
|
||||
|
||||
# ── 支持的转场列表 ──
|
||||
|
||||
def test_supported_transitions_list(self):
|
||||
"""获取支持的转场列表(给 API 用)."""
|
||||
transitions = TransitionEngine.supported_transitions()
|
||||
assert len(transitions) >= 10 # cut + 至少 9 种特效
|
||||
# 检查结构
|
||||
for t in transitions:
|
||||
assert "name" in t
|
||||
assert "display_name" in t
|
||||
assert "category" in t
|
||||
# 检查分类
|
||||
names = [t["name"] for t in transitions]
|
||||
assert "cut" in names
|
||||
assert "fade" in names
|
||||
assert "zoom" in names
|
||||
assert "circlecrop" in names
|
||||
|
||||
|
||||
# ── 集成测试:与 UnifiedRenderService 协作 ────────────────────────────────────
|
||||
|
||||
|
||||
class TestTransitionIntegration:
|
||||
"""转场引擎与统一渲染服务的集成测试."""
|
||||
|
||||
def test_unified_render_service_has_transition_engine(self):
|
||||
"""UnifiedRenderService 内部有 TransitionEngine 实例."""
|
||||
from pathlib import Path
|
||||
|
||||
from video_processing.unified_render_service import UnifiedRenderService
|
||||
|
||||
# 构造最小化的服务实例
|
||||
service = UnifiedRenderService(
|
||||
plan=None,
|
||||
clips=[],
|
||||
asset_path_map={},
|
||||
work_dir=Path("/tmp"),
|
||||
)
|
||||
assert hasattr(service, "_transition_engine")
|
||||
assert isinstance(service._transition_engine, TransitionEngine)
|
||||
|
||||
def test_resolved_clip_has_transition_duration(self):
|
||||
"""ResolvedClip 有 transition_duration 字段."""
|
||||
from video_processing.unified_render_service import ResolvedClip
|
||||
|
||||
rc = ResolvedClip(
|
||||
clip_id="test",
|
||||
asset_id="asset1",
|
||||
local_path=__file__, # 随便一个存在的路径
|
||||
clip_type="main",
|
||||
order=0,
|
||||
transition_effect="fade",
|
||||
transition_duration=0.8,
|
||||
)
|
||||
assert rc.transition_duration == 0.8
|
||||
assert rc.transition_effect == "fade"
|
||||
def test_valid_effect_has_ffmpeg_name(self):
|
||||
"""所有非cut的支持效果都有对应的ffmpeg名称."""
|
||||
for t in TransitionType:
|
||||
if t.value == CUT_TRANSITION:
|
||||
continue # cut返回空是正常的
|
||||
config = TransitionConfig(effect=t.value, duration=0.5)
|
||||
assert config.ffmpeg_transition != ""
|
||||
|
||||
+251
-234
@@ -1,268 +1,285 @@
|
||||
"""裁剪引擎单元测试."""
|
||||
"""裁剪引擎单元测试 - 配置解析+推导等纯逻辑."""
|
||||
|
||||
import sys
|
||||
import unittest
|
||||
from pathlib import Path
|
||||
from __future__ import annotations
|
||||
|
||||
# 确保 apps/worker 在路径中
|
||||
sys.path.insert(0, str(Path(__file__).parent.parent.parent / "apps" / "worker"))
|
||||
|
||||
from video_processing.trim_engine import (
|
||||
MIN_TRIM_DURATION,
|
||||
TrimConfig,
|
||||
TrimEngine,
|
||||
TrimSegment,
|
||||
extract_trim_from_clip_config,
|
||||
)
|
||||
import pytest
|
||||
from video_processing.trim_engine import MIN_TRIM_DURATION, TrimConfig, TrimSegment
|
||||
|
||||
|
||||
class TestTrimConfig(unittest.TestCase):
|
||||
"""TrimConfig 单元测试."""
|
||||
class TestTrimConfigFromDict:
|
||||
"""TrimConfig.from_dict 解析测试."""
|
||||
|
||||
def test_from_dict_none(self):
|
||||
"""空字典返回 None(不裁剪)."""
|
||||
self.assertIsNone(TrimConfig.from_dict(None))
|
||||
self.assertIsNone(TrimConfig.from_dict({}))
|
||||
def test_none_returns_none(self):
|
||||
"""None返回None(不裁剪)."""
|
||||
assert TrimConfig.from_dict(None) is None
|
||||
|
||||
def test_from_dict_with_start(self):
|
||||
"""只有 start_time."""
|
||||
cfg = TrimConfig.from_dict({"start_time": 5.0})
|
||||
self.assertIsNotNone(cfg)
|
||||
self.assertEqual(cfg.start_time, 5.0)
|
||||
self.assertEqual(cfg.end_time, 0.0)
|
||||
self.assertEqual(cfg.duration, 0.0)
|
||||
def test_empty_dict_returns_none(self):
|
||||
"""空dict返回None."""
|
||||
assert TrimConfig.from_dict({}) is None
|
||||
|
||||
def test_from_dict_with_duration(self):
|
||||
"""只有 duration."""
|
||||
cfg = TrimConfig.from_dict({"duration": 10.0})
|
||||
self.assertIsNotNone(cfg)
|
||||
self.assertEqual(cfg.start_time, 0.0)
|
||||
self.assertEqual(cfg.duration, 10.0)
|
||||
def test_all_zero_returns_none(self):
|
||||
"""全零返回None."""
|
||||
assert (
|
||||
TrimConfig.from_dict(
|
||||
{
|
||||
"start_time": 0,
|
||||
"end_time": 0,
|
||||
"duration": 0,
|
||||
}
|
||||
)
|
||||
is None
|
||||
)
|
||||
|
||||
def test_resolve_start_and_end(self):
|
||||
"""start + end 推导 duration."""
|
||||
cfg = TrimConfig(start_time=5.0, end_time=15.0)
|
||||
resolved = cfg.validate_and_resolve(asset_duration=30.0)
|
||||
self.assertEqual(resolved.start_time, 5.0)
|
||||
self.assertEqual(resolved.end_time, 15.0)
|
||||
self.assertAlmostEqual(resolved.duration, 10.0, places=3)
|
||||
self.assertTrue(resolved.is_valid)
|
||||
def test_start_only(self):
|
||||
"""只有start_time有效."""
|
||||
config = TrimConfig.from_dict({"start_time": 5.0})
|
||||
assert config is not None
|
||||
assert config.start_time == 5.0
|
||||
assert config.end_time == 0
|
||||
assert config.duration == 0
|
||||
|
||||
def test_resolve_start_and_duration(self):
|
||||
"""start + duration 推导 end."""
|
||||
cfg = TrimConfig(start_time=5.0, duration=10.0)
|
||||
resolved = cfg.validate_and_resolve(asset_duration=30.0)
|
||||
self.assertEqual(resolved.start_time, 5.0)
|
||||
self.assertAlmostEqual(resolved.end_time, 15.0, places=3)
|
||||
self.assertEqual(resolved.duration, 10.0)
|
||||
def test_duration_only(self):
|
||||
"""只有duration有效."""
|
||||
config = TrimConfig.from_dict({"duration": 10.0})
|
||||
assert config is not None
|
||||
assert config.start_time == 0
|
||||
assert config.duration == 10.0
|
||||
|
||||
def test_resolve_end_and_duration(self):
|
||||
"""end + duration 推导 start."""
|
||||
cfg = TrimConfig(end_time=20.0, duration=8.0)
|
||||
resolved = cfg.validate_and_resolve(asset_duration=30.0)
|
||||
self.assertAlmostEqual(resolved.start_time, 12.0, places=3)
|
||||
self.assertEqual(resolved.end_time, 20.0)
|
||||
self.assertEqual(resolved.duration, 8.0)
|
||||
def test_start_and_duration(self):
|
||||
"""start + duration."""
|
||||
config = TrimConfig.from_dict({"start_time": 2.0, "duration": 5.0})
|
||||
assert config is not None
|
||||
assert config.start_time == 2.0
|
||||
assert config.duration == 5.0
|
||||
|
||||
def test_resolve_only_start(self):
|
||||
"""只有 start → 取到末尾."""
|
||||
cfg = TrimConfig(start_time=10.0)
|
||||
resolved = cfg.validate_and_resolve(asset_duration=30.0)
|
||||
self.assertEqual(resolved.start_time, 10.0)
|
||||
self.assertEqual(resolved.end_time, 30.0)
|
||||
self.assertAlmostEqual(resolved.duration, 20.0, places=3)
|
||||
def test_start_and_end(self):
|
||||
"""start + end."""
|
||||
config = TrimConfig.from_dict({"start_time": 1.0, "end_time": 4.0})
|
||||
assert config is not None
|
||||
assert config.start_time == 1.0
|
||||
assert config.end_time == 4.0
|
||||
|
||||
def test_resolve_only_duration(self):
|
||||
"""只有 duration → 从开头取."""
|
||||
cfg = TrimConfig(duration=15.0)
|
||||
resolved = cfg.validate_and_resolve(asset_duration=30.0)
|
||||
self.assertEqual(resolved.start_time, 0.0)
|
||||
self.assertAlmostEqual(resolved.end_time, 15.0, places=3)
|
||||
self.assertEqual(resolved.duration, 15.0)
|
||||
def test_end_and_duration(self):
|
||||
"""end + duration."""
|
||||
config = TrimConfig.from_dict({"end_time": 10.0, "duration": 3.0})
|
||||
assert config is not None
|
||||
assert config.end_time == 10.0
|
||||
assert config.duration == 3.0
|
||||
|
||||
def test_boundary_clamp_end(self):
|
||||
"""end 超出素材时长 → 钳制."""
|
||||
cfg = TrimConfig(start_time=5.0, duration=30.0)
|
||||
resolved = cfg.validate_and_resolve(asset_duration=20.0)
|
||||
self.assertEqual(resolved.start_time, 5.0)
|
||||
self.assertEqual(resolved.end_time, 20.0)
|
||||
self.assertAlmostEqual(resolved.duration, 15.0, places=3)
|
||||
def test_string_values_converted(self):
|
||||
"""字符串值会被转换."""
|
||||
config = TrimConfig.from_dict(
|
||||
{
|
||||
"start_time": "5.0",
|
||||
"duration": "10.0",
|
||||
}
|
||||
)
|
||||
assert config is not None
|
||||
assert config.start_time == 5.0
|
||||
assert config.duration == 10.0
|
||||
|
||||
def test_boundary_clamp_start_negative(self):
|
||||
"""start 为负 → 钳制到 0."""
|
||||
cfg = TrimConfig(start_time=-5.0, duration=10.0)
|
||||
resolved = cfg.validate_and_resolve(asset_duration=30.0)
|
||||
self.assertEqual(resolved.start_time, 0.0)
|
||||
self.assertAlmostEqual(resolved.end_time, 10.0, places=3)
|
||||
self.assertEqual(resolved.duration, 10.0)
|
||||
def test_falsy_start_with_duration(self):
|
||||
"""start=0 + duration>0有效."""
|
||||
config = TrimConfig.from_dict({"start_time": 0, "duration": 5.0})
|
||||
assert config is not None
|
||||
assert config.start_time == 0.0
|
||||
assert config.duration == 5.0
|
||||
|
||||
def test_boundary_start_past_end(self):
|
||||
"""start 超过素材总时长 → 钳制到末尾最小片段."""
|
||||
cfg = TrimConfig(start_time=50.0, duration=5.0)
|
||||
resolved = cfg.validate_and_resolve(asset_duration=30.0)
|
||||
self.assertTrue(resolved.start_time < 30.0)
|
||||
self.assertEqual(resolved.end_time, 30.0)
|
||||
self.assertTrue(resolved.duration >= MIN_TRIM_DURATION)
|
||||
|
||||
def test_invalid_end_before_start(self):
|
||||
"""end <= start → 无效."""
|
||||
cfg = TrimConfig(start_time=15.0, end_time=10.0)
|
||||
resolved = cfg.validate_and_resolve(asset_duration=30.0)
|
||||
self.assertFalse(resolved.is_valid)
|
||||
class TestValidateAndResolve:
|
||||
"""validate_and_resolve 推导测试."""
|
||||
|
||||
def test_zero_duration_invalid(self):
|
||||
"""duration 为 0 → 无效."""
|
||||
cfg = TrimConfig(start_time=5.0, duration=0.0)
|
||||
resolved = cfg.validate_and_resolve(asset_duration=30.0)
|
||||
# 只有 start 没有 duration → 会被推导为取到末尾
|
||||
self.assertTrue(resolved.is_valid)
|
||||
self.assertEqual(resolved.end_time, 30.0)
|
||||
def test_start_plus_end(self):
|
||||
"""start + end → 推导duration."""
|
||||
config = TrimConfig(start_time=2.0, end_time=7.0)
|
||||
resolved = config.validate_and_resolve(60.0)
|
||||
assert resolved.start_time == 2.0
|
||||
assert resolved.end_time == 7.0
|
||||
assert resolved.duration == 5.0
|
||||
|
||||
def test_is_noop(self):
|
||||
"""is_noop 判断."""
|
||||
noop = TrimConfig(start_time=0.0, end_time=0.0, duration=0.0)
|
||||
self.assertTrue(noop.is_noop)
|
||||
def test_start_plus_duration(self):
|
||||
"""start + duration → 推导end."""
|
||||
config = TrimConfig(start_time=3.0, duration=10.0)
|
||||
resolved = config.validate_and_resolve(60.0)
|
||||
assert resolved.start_time == 3.0
|
||||
assert resolved.duration == 10.0
|
||||
assert resolved.end_time == 13.0
|
||||
|
||||
not_noop = TrimConfig(start_time=5.0, duration=10.0)
|
||||
self.assertFalse(not_noop.is_noop)
|
||||
def test_end_plus_duration(self):
|
||||
"""end + duration → 推导start."""
|
||||
config = TrimConfig(end_time=15.0, duration=5.0)
|
||||
resolved = config.validate_and_resolve(60.0)
|
||||
assert resolved.end_time == 15.0
|
||||
assert resolved.duration == 5.0
|
||||
assert resolved.start_time == 10.0
|
||||
|
||||
def test_start_only_takes_to_end(self):
|
||||
"""只有start → 取到素材末尾."""
|
||||
config = TrimConfig(start_time=50.0)
|
||||
resolved = config.validate_and_resolve(60.0)
|
||||
assert resolved.start_time == 50.0
|
||||
assert resolved.end_time == 60.0
|
||||
assert resolved.duration == 10.0
|
||||
|
||||
def test_end_only_takes_from_start(self):
|
||||
"""只有end → 从开头取."""
|
||||
config = TrimConfig(end_time=20.0)
|
||||
resolved = config.validate_and_resolve(60.0)
|
||||
assert resolved.start_time == 0.0
|
||||
assert resolved.end_time == 20.0
|
||||
assert resolved.duration == 20.0
|
||||
|
||||
def test_end_before_start_invalid(self):
|
||||
"""end < start → 无效(0时长)."""
|
||||
config = TrimConfig(start_time=10.0, end_time=5.0)
|
||||
resolved = config.validate_and_resolve(60.0)
|
||||
assert resolved.duration == 0.0
|
||||
assert resolved.is_valid is False
|
||||
|
||||
def test_negative_start_clamped(self):
|
||||
"""负start钳制到0."""
|
||||
config = TrimConfig(start_time=-5.0, duration=10.0)
|
||||
resolved = config.validate_and_resolve(60.0)
|
||||
assert resolved.start_time == 0.0
|
||||
assert resolved.duration == 10.0
|
||||
|
||||
def test_end_beyond_asset_clamped(self):
|
||||
"""end超过素材时长钳制."""
|
||||
config = TrimConfig(start_time=50.0, duration=20.0)
|
||||
resolved = config.validate_and_resolve(60.0)
|
||||
assert resolved.end_time == 60.0
|
||||
assert resolved.duration == 10.0
|
||||
|
||||
def test_start_beyond_asset_clamped(self):
|
||||
"""start超过素材时长 → 钳制到末尾保留MIN_TRIM."""
|
||||
config = TrimConfig(start_time=100.0, duration=5.0)
|
||||
resolved = config.validate_and_resolve(60.0)
|
||||
assert resolved.start_time == 60.0 - MIN_TRIM_DURATION
|
||||
assert resolved.end_time == 60.0
|
||||
|
||||
def test_zero_asset_duration(self):
|
||||
"""素材时长为 0 → 不裁剪."""
|
||||
cfg = TrimConfig(start_time=5.0, duration=10.0)
|
||||
resolved = cfg.validate_and_resolve(asset_duration=0.0)
|
||||
self.assertTrue(resolved.is_noop)
|
||||
"""素材时长为0 → 不裁剪."""
|
||||
config = TrimConfig(start_time=1.0, duration=5.0)
|
||||
resolved = config.validate_and_resolve(0.0)
|
||||
assert resolved.start_time == 0.0
|
||||
assert resolved.duration == 0.0
|
||||
|
||||
def test_all_three_params_use_start_duration(self):
|
||||
"""三个参数都给了 → 以 start + duration 为准."""
|
||||
cfg = TrimConfig(start_time=5.0, end_time=20.0, duration=8.0)
|
||||
resolved = cfg.validate_and_resolve(asset_duration=30.0)
|
||||
# validate_and_resolve 中 start+end 优先于 start+duration
|
||||
# 因为先检查的是 start>0 and end>0
|
||||
self.assertAlmostEqual(resolved.duration, 15.0, places=3)
|
||||
def test_end_and_duration_with_negative_start(self):
|
||||
"""end + duration推导出来负start → 钳制+重算."""
|
||||
config = TrimConfig(end_time=3.0, duration=10.0)
|
||||
resolved = config.validate_and_resolve(60.0)
|
||||
assert resolved.start_time == 0.0
|
||||
assert resolved.end_time == 3.0
|
||||
assert resolved.duration == 3.0
|
||||
|
||||
def test_all_three_params_uses_start_duration(self):
|
||||
"""三个都给了,以start+duration为准."""
|
||||
# 实际代码是先判断 start+end(情况1),如果都>0就用
|
||||
# 所以这里测试 start+end 都给了且都>0的情况
|
||||
config = TrimConfig(start_time=2.0, end_time=8.0, duration=10.0)
|
||||
resolved = config.validate_and_resolve(60.0)
|
||||
# 走情况1(start+end都有)
|
||||
assert resolved.start_time == 2.0
|
||||
assert resolved.end_time == 8.0
|
||||
assert resolved.duration == 6.0
|
||||
|
||||
|
||||
class TestTrimEngine(unittest.TestCase):
|
||||
"""TrimEngine 单元测试."""
|
||||
class TestIsValid:
|
||||
"""is_valid 属性测试."""
|
||||
|
||||
def test_build_video_trim_with_start_and_duration(self):
|
||||
"""视频裁剪:start + duration."""
|
||||
trim = TrimConfig(start_time=10.0, duration=5.0)
|
||||
result = TrimEngine.build_video_trim_filter("[0:v]", trim, "[v0]")
|
||||
self.assertIn("trim=start=10.000:duration=5.000", result)
|
||||
self.assertIn("setpts=PTS-STARTPTS", result)
|
||||
self.assertTrue(result.startswith("[0:v]"))
|
||||
self.assertTrue(result.endswith("[v0]"))
|
||||
def test_valid_duration(self):
|
||||
"""时长足够有效."""
|
||||
config = TrimConfig(start_time=0.0, end_time=0.0, duration=5.0)
|
||||
assert config.is_valid is True
|
||||
|
||||
def test_build_video_trim_duration_only(self):
|
||||
"""视频裁剪:只有 duration."""
|
||||
trim = TrimConfig(start_time=0.0, duration=8.0)
|
||||
result = TrimEngine.build_video_trim_filter("[0:v]", trim, "[v0]")
|
||||
self.assertIn("trim=duration=8.000", result)
|
||||
self.assertNotIn("start=", result.split("setpts")[0])
|
||||
def test_zero_duration_invalid(self):
|
||||
"""零时长无效."""
|
||||
config = TrimConfig(duration=0.0)
|
||||
assert config.is_valid is False
|
||||
|
||||
def test_build_audio_trim_with_start(self):
|
||||
"""音频裁剪:start + duration."""
|
||||
trim = TrimConfig(start_time=3.0, duration=7.0)
|
||||
result = TrimEngine.build_audio_trim_filter("[0:a]", trim, "[a0]")
|
||||
self.assertIn("atrim=start=3.000:duration=7.000", result)
|
||||
self.assertIn("asetpts=PTS-STARTPTS", result)
|
||||
|
||||
def test_build_audio_trim_noop(self):
|
||||
"""音频裁剪:noop."""
|
||||
trim = TrimConfig(start_time=0.0, end_time=0.0, duration=0.0)
|
||||
result = TrimEngine.build_audio_trim_filter("[0:a]", trim, "[a0]")
|
||||
self.assertIn("asetpts=PTS-STARTPTS", result)
|
||||
self.assertNotIn("atrim=", result)
|
||||
|
||||
def test_resolve_segments(self):
|
||||
"""多段裁剪解析."""
|
||||
segments = [
|
||||
TrimSegment(segment_id="s1", trim=TrimConfig(start_time=0.0, duration=5.0), order=0),
|
||||
TrimSegment(segment_id="s2", trim=TrimConfig(start_time=10.0, duration=5.0), order=1),
|
||||
TrimSegment(segment_id="s3", trim=TrimConfig(start_time=20.0, duration=5.0), order=2),
|
||||
]
|
||||
resolved = TrimEngine.resolve_segments(segments, asset_duration=30.0)
|
||||
self.assertEqual(len(resolved), 3)
|
||||
self.assertEqual(resolved[0].segment_id, "s1")
|
||||
self.assertEqual(resolved[0].trim.duration, 5.0)
|
||||
self.assertEqual(resolved[1].segment_id, "s2")
|
||||
self.assertEqual(resolved[1].trim.start_time, 10.0)
|
||||
self.assertEqual(resolved[2].trim.start_time, 20.0)
|
||||
|
||||
def test_resolve_segments_filter_invalid(self):
|
||||
"""多段裁剪:过滤无效段."""
|
||||
segments = [
|
||||
TrimSegment(segment_id="good", trim=TrimConfig(start_time=0.0, duration=5.0), order=0),
|
||||
TrimSegment(segment_id="bad", trim=TrimConfig(start_time=10.0, end_time=5.0), order=1), # end < start
|
||||
]
|
||||
resolved = TrimEngine.resolve_segments(segments, asset_duration=30.0)
|
||||
self.assertEqual(len(resolved), 1)
|
||||
self.assertEqual(resolved[0].segment_id, "good")
|
||||
|
||||
def test_resolve_segments_boundary_clamp(self):
|
||||
"""多段裁剪:边界钳制."""
|
||||
segments = [
|
||||
TrimSegment(segment_id="s1", trim=TrimConfig(start_time=25.0, duration=10.0), order=0),
|
||||
]
|
||||
resolved = TrimEngine.resolve_segments(segments, asset_duration=30.0)
|
||||
self.assertEqual(len(resolved), 1)
|
||||
self.assertEqual(resolved[0].trim.end_time, 30.0)
|
||||
self.assertAlmostEqual(resolved[0].trim.duration, 5.0, places=3)
|
||||
|
||||
def test_parse_segments_from_list(self):
|
||||
"""从 config 解析多段配置."""
|
||||
config = {
|
||||
"trim_segments": [
|
||||
{"segment_id": "intro", "start_time": 0, "duration": 3, "order": 0},
|
||||
{"segment_id": "highlight", "start_time": 10, "duration": 5, "order": 1},
|
||||
{"segment_id": "outro", "start_time": 50, "duration": 3, "order": 2},
|
||||
]
|
||||
}
|
||||
segments = TrimEngine.parse_segments_from_config(config)
|
||||
self.assertEqual(len(segments), 3)
|
||||
self.assertEqual(segments[0].segment_id, "intro")
|
||||
self.assertEqual(segments[1].trim.start_time, 10.0)
|
||||
self.assertEqual(segments[2].trim.duration, 3.0)
|
||||
|
||||
def test_parse_segments_empty(self):
|
||||
"""无裁剪配置 → 空列表."""
|
||||
self.assertEqual(TrimEngine.parse_segments_from_config(None), [])
|
||||
self.assertEqual(TrimEngine.parse_segments_from_config({}), [])
|
||||
|
||||
def test_parse_single_trim_legacy(self):
|
||||
"""旧格式单段裁剪(trim_start/trim_duration)."""
|
||||
config = {"trim_start": 5.0, "trim_duration": 10.0}
|
||||
segments = TrimEngine.parse_segments_from_config(config)
|
||||
self.assertEqual(len(segments), 1)
|
||||
self.assertEqual(segments[0].trim.start_time, 5.0)
|
||||
self.assertEqual(segments[0].trim.duration, 10.0)
|
||||
def test_min_duration_valid(self):
|
||||
"""刚好等于最小值有效."""
|
||||
config = TrimConfig(duration=MIN_TRIM_DURATION)
|
||||
assert config.is_valid is True
|
||||
|
||||
|
||||
class TestExtractTrimFromClipConfig(unittest.TestCase):
|
||||
"""extract_trim_from_clip_config 单元测试."""
|
||||
class TestIsNoop:
|
||||
"""is_noop 属性测试."""
|
||||
|
||||
def test_trim_subdict(self):
|
||||
"""trim 子字典."""
|
||||
config = {"trim": {"start_time": 5.0, "duration": 10.0}}
|
||||
result = extract_trim_from_clip_config(config)
|
||||
self.assertIsNotNone(result)
|
||||
self.assertEqual(result.start_time, 5.0)
|
||||
self.assertEqual(result.duration, 10.0)
|
||||
def test_zero_is_noop(self):
|
||||
"""全零是noop."""
|
||||
config = TrimConfig()
|
||||
assert config.is_noop is True
|
||||
|
||||
def test_flat_fields(self):
|
||||
"""扁平字段(trim_start/trim_end/trim_duration)."""
|
||||
config = {"trim_start": 2.0, "trim_end": 8.0}
|
||||
result = extract_trim_from_clip_config(config)
|
||||
self.assertIsNotNone(result)
|
||||
self.assertEqual(result.start_time, 2.0)
|
||||
self.assertEqual(result.end_time, 8.0)
|
||||
def test_with_duration_not_noop(self):
|
||||
"""有时长不是noop."""
|
||||
config = TrimConfig(duration=10.0)
|
||||
assert config.is_noop is False
|
||||
|
||||
def test_no_trim(self):
|
||||
"""无裁剪配置."""
|
||||
self.assertIsNone(extract_trim_from_clip_config(None))
|
||||
self.assertIsNone(extract_trim_from_clip_config({}))
|
||||
self.assertIsNone(extract_trim_from_clip_config({"other": "value"}))
|
||||
def test_with_start_not_noop(self):
|
||||
"""有start不是noop."""
|
||||
config = TrimConfig(start_time=5.0)
|
||||
assert config.is_noop is False
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
class TestTrimFromStart:
|
||||
"""trim_from_start 属性测试."""
|
||||
|
||||
def test_zero_start_is_from_start(self):
|
||||
"""start=0是从开头裁."""
|
||||
config = TrimConfig(start_time=0.0)
|
||||
assert config.trim_from_start is True
|
||||
|
||||
def test_positive_start_not_from_start(self):
|
||||
"""有start不是从开头裁."""
|
||||
config = TrimConfig(start_time=5.0)
|
||||
assert config.trim_from_start is False
|
||||
|
||||
|
||||
class TestTrimSegment:
|
||||
"""TrimSegment 测试."""
|
||||
|
||||
def test_from_dict_basic(self):
|
||||
"""基本解析."""
|
||||
seg = TrimSegment.from_dict(
|
||||
{
|
||||
"start_time": 5.0,
|
||||
"duration": 10.0,
|
||||
"segment_id": "seg1",
|
||||
},
|
||||
default_order=0,
|
||||
)
|
||||
assert seg.segment_id == "seg1"
|
||||
assert seg.trim.start_time == 5.0
|
||||
assert seg.trim.duration == 10.0
|
||||
assert seg.order == 0
|
||||
|
||||
def test_from_dict_with_order(self):
|
||||
"""带order的解析."""
|
||||
seg = TrimSegment.from_dict(
|
||||
{
|
||||
"start_time": 1.0,
|
||||
"end_time": 4.0,
|
||||
"order": 2,
|
||||
}
|
||||
)
|
||||
assert seg.order == 2
|
||||
assert seg.trim.start_time == 1.0
|
||||
assert seg.trim.end_time == 4.0
|
||||
|
||||
def test_from_dict_default_segment_id(self):
|
||||
"""缺省segment_id时用默认值."""
|
||||
seg = TrimSegment.from_dict({"duration": 5.0}, default_order=3)
|
||||
assert seg.segment_id == "seg_3"
|
||||
assert seg.order == 3
|
||||
|
||||
def test_from_dict_empty_string_segment_id(self):
|
||||
"""空字符串segment_id走默认."""
|
||||
seg = TrimSegment.from_dict(
|
||||
{
|
||||
"segment_id": "",
|
||||
"duration": 5.0,
|
||||
},
|
||||
default_order=5,
|
||||
)
|
||||
assert seg.segment_id == "seg_5"
|
||||
|
||||
Executable
+177
@@ -0,0 +1,177 @@
|
||||
"""统一渲染服务纯逻辑测试 — _resolve_layer_role + 图层配置 + 数据结构."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from pathlib import Path
|
||||
|
||||
import pytest
|
||||
from video_processing.unified_render_service import (
|
||||
_LAYER_Z_INDEX,
|
||||
_PIP_SCALE,
|
||||
RenderLayer,
|
||||
ResolvedClip,
|
||||
_resolve_layer_role,
|
||||
)
|
||||
|
||||
|
||||
class TestResolveLayerRole:
|
||||
"""_resolve_layer_role 图层角色映射测试."""
|
||||
|
||||
def test_intro_is_main(self):
|
||||
"""intro片段 → main层."""
|
||||
assert _resolve_layer_role("intro", {}) == "main"
|
||||
|
||||
def test_outro_is_main(self):
|
||||
"""outro片段 → main层."""
|
||||
assert _resolve_layer_role("outro", {}) == "main"
|
||||
|
||||
def test_overlay_is_overlay(self):
|
||||
"""overlay片段 → overlay层."""
|
||||
assert _resolve_layer_role("overlay", {}) == "overlay"
|
||||
|
||||
def test_corner_voice(self):
|
||||
"""corner_voice片段 → corner_voice层."""
|
||||
assert _resolve_layer_role("corner_voice", {}) == "corner_voice"
|
||||
|
||||
def test_background(self):
|
||||
"""background片段 → background层."""
|
||||
assert _resolve_layer_role("background", {}) == "background"
|
||||
|
||||
def test_b_roll(self):
|
||||
"""b_roll片段 → broll层."""
|
||||
assert _resolve_layer_role("b_roll", {}) == "broll"
|
||||
|
||||
def test_main_default(self):
|
||||
"""main type默认 → main层."""
|
||||
assert _resolve_layer_role("main", {}) == "main"
|
||||
|
||||
def test_main_with_broll_role(self):
|
||||
"""main type + role=b_roll → broll层."""
|
||||
assert _resolve_layer_role("main", {"role": "b_roll"}) == "broll"
|
||||
|
||||
def test_main_with_audio_role(self):
|
||||
"""main type + role=audio → audio层."""
|
||||
assert _resolve_layer_role("main", {"role": "audio"}) == "audio"
|
||||
|
||||
def test_unknown_type_falls_back_to_main(self):
|
||||
"""未知类型 → main层."""
|
||||
assert _resolve_layer_role("random_type", {}) == "main"
|
||||
|
||||
def test_intro_ignores_role(self):
|
||||
"""intro/outro忽略role配置."""
|
||||
assert _resolve_layer_role("intro", {"role": "b_roll"}) == "main"
|
||||
assert _resolve_layer_role("outro", {"role": "overlay"}) == "main"
|
||||
|
||||
def test_empty_config(self):
|
||||
"""空config不影响结果."""
|
||||
assert _resolve_layer_role("main", {}) == "main"
|
||||
|
||||
def test_none_role(self):
|
||||
"""role=None时走默认."""
|
||||
assert _resolve_layer_role("main", {"role": None}) == "main"
|
||||
|
||||
|
||||
class TestLayerZIndex:
|
||||
"""_LAYER_Z_INDEX 图层层级配置测试."""
|
||||
|
||||
def test_background_lowest(self):
|
||||
"""background在最底层."""
|
||||
assert _LAYER_Z_INDEX["background"] == -1
|
||||
|
||||
def test_main_and_broll_same_level(self):
|
||||
"""main和broll在同一层(z=0)."""
|
||||
assert _LAYER_Z_INDEX["main"] == 0
|
||||
assert _LAYER_Z_INDEX["broll"] == 0
|
||||
|
||||
def test_overlay_and_corner_voice_above(self):
|
||||
"""overlay和corner_voice在z=1."""
|
||||
assert _LAYER_Z_INDEX["overlay"] == 1
|
||||
assert _LAYER_Z_INDEX["corner_voice"] == 1
|
||||
|
||||
def test_audio_highest(self):
|
||||
"""audio在z=2(最高,因为音频不涉及z顺序但参与混音)."""
|
||||
assert _LAYER_Z_INDEX["audio"] == 2
|
||||
|
||||
def test_pip_scale_is_positive(self):
|
||||
"""PiP缩放比例为正数."""
|
||||
assert _PIP_SCALE > 0
|
||||
assert _PIP_SCALE < 1.0
|
||||
|
||||
|
||||
class TestResolvedClip:
|
||||
"""ResolvedClip 数据结构测试."""
|
||||
|
||||
def test_default_values(self):
|
||||
"""默认值正确."""
|
||||
clip = ResolvedClip(
|
||||
clip_id="c1",
|
||||
asset_id="a1",
|
||||
local_path=Path("/tmp/test.mp4"),
|
||||
clip_type="main",
|
||||
order=0,
|
||||
)
|
||||
assert clip.start_time == 0.0
|
||||
assert clip.duration == 0.0
|
||||
assert clip.transition_effect == "cut"
|
||||
assert clip.transition_duration == 0.0
|
||||
assert clip.playback_speed == 1.0
|
||||
assert clip.config == {}
|
||||
assert clip.actual_duration == 0.0
|
||||
assert clip.trim_config is None
|
||||
|
||||
def test_custom_values(self):
|
||||
"""自定义值正确."""
|
||||
clip = ResolvedClip(
|
||||
clip_id="c2",
|
||||
asset_id="a2",
|
||||
local_path=Path("/tmp/video.mp4"),
|
||||
clip_type="overlay",
|
||||
order=1,
|
||||
start_time=5.0,
|
||||
duration=10.0,
|
||||
transition_effect="fade",
|
||||
transition_duration=0.5,
|
||||
playback_speed=1.5,
|
||||
)
|
||||
assert clip.clip_id == "c2"
|
||||
assert clip.asset_id == "a2"
|
||||
assert clip.clip_type == "overlay"
|
||||
assert clip.order == 1
|
||||
assert clip.start_time == 5.0
|
||||
assert clip.duration == 10.0
|
||||
assert clip.transition_effect == "fade"
|
||||
assert clip.transition_duration == 0.5
|
||||
assert clip.playback_speed == 1.5
|
||||
|
||||
|
||||
class TestRenderLayer:
|
||||
"""RenderLayer 数据结构测试."""
|
||||
|
||||
def test_default_values(self):
|
||||
"""默认值正确."""
|
||||
layer = RenderLayer(role="main")
|
||||
assert layer.clips == []
|
||||
assert layer.z_index == 0
|
||||
assert layer.opacity == 1.0
|
||||
assert layer.position is None
|
||||
|
||||
def test_with_clips(self):
|
||||
"""带片段的图层."""
|
||||
clip = ResolvedClip(
|
||||
clip_id="c1",
|
||||
asset_id="a1",
|
||||
local_path=Path("/tmp/t.mp4"),
|
||||
clip_type="main",
|
||||
order=0,
|
||||
)
|
||||
layer = RenderLayer(role="overlay", clips=[clip], z_index=1)
|
||||
assert len(layer.clips) == 1
|
||||
assert layer.z_index == 1
|
||||
assert layer.role == "overlay"
|
||||
|
||||
def test_background_layer(self):
|
||||
"""background图层配置."""
|
||||
layer = RenderLayer(role="background", z_index=-1, opacity=1.0)
|
||||
assert layer.role == "background"
|
||||
assert layer.z_index == -1
|
||||
assert layer.opacity == 1.0
|
||||
+341
-159
@@ -1,67 +1,69 @@
|
||||
"""
|
||||
水印引擎配置与纯逻辑测试.
|
||||
"""水印引擎单元测试 - 配置解析 + 位置计算等纯逻辑."""
|
||||
|
||||
覆盖 WatermarkConfig.from_dict / validate / 位置枚举等纯逻辑.
|
||||
引擎核心 render 方法依赖 FFmpeg,由集成测试覆盖.
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import pytest
|
||||
from video_processing.watermark_engine import WATERMARK_POSITIONS, WatermarkConfig
|
||||
from video_processing.watermark_engine import (
|
||||
WATERMARK_POSITIONS,
|
||||
WatermarkConfig,
|
||||
WatermarkEngine,
|
||||
)
|
||||
|
||||
# ── WatermarkConfig 测试 ──────────────────────────────────────────
|
||||
|
||||
|
||||
class TestWatermarkPositions:
|
||||
"""水印位置枚举."""
|
||||
class TestWatermarkConfigDefaults:
|
||||
"""默认值测试."""
|
||||
|
||||
def test_nine_positions_exist(self):
|
||||
assert len(WATERMARK_POSITIONS) == 9
|
||||
assert "top_left" in WATERMARK_POSITIONS
|
||||
assert "top_center" in WATERMARK_POSITIONS
|
||||
assert "top_right" in WATERMARK_POSITIONS
|
||||
assert "center_left" in WATERMARK_POSITIONS
|
||||
assert "center" in WATERMARK_POSITIONS
|
||||
assert "center_right" in WATERMARK_POSITIONS
|
||||
assert "bottom_left" in WATERMARK_POSITIONS
|
||||
assert "bottom_center" in WATERMARK_POSITIONS
|
||||
assert "bottom_right" in WATERMARK_POSITIONS
|
||||
|
||||
def test_position_values_are_chinese_labels(self):
|
||||
for key, label in WATERMARK_POSITIONS.items():
|
||||
assert isinstance(label, str)
|
||||
assert len(label) >= 2
|
||||
def test_default_values(self):
|
||||
"""默认配置值正确."""
|
||||
config = WatermarkConfig()
|
||||
assert config.mode == "text"
|
||||
assert config.position == "bottom_right"
|
||||
assert config.image_path == ""
|
||||
assert config.scale == 0.2
|
||||
assert config.opacity == 0.8
|
||||
assert config.text == ""
|
||||
assert config.font_size == 24
|
||||
assert config.font_color == "white"
|
||||
assert config.font_path == ""
|
||||
assert config.margin_x == 20
|
||||
assert config.margin_y == 20
|
||||
assert config.scroll is False
|
||||
assert config.scroll_speed == 50
|
||||
|
||||
|
||||
class TestWatermarkConfigFromDict:
|
||||
"""from_dict 构造逻辑."""
|
||||
"""from_dict 配置解析测试."""
|
||||
|
||||
def test_none_returns_none(self):
|
||||
"""None 返回 None."""
|
||||
assert WatermarkConfig.from_dict(None) is None
|
||||
|
||||
def test_empty_dict_returns_none(self):
|
||||
"""空字典返回 None."""
|
||||
assert WatermarkConfig.from_dict({}) is None
|
||||
|
||||
def test_enabled_false_returns_none(self):
|
||||
def test_disabled_returns_none(self):
|
||||
"""enabled=False 返回 None."""
|
||||
assert WatermarkConfig.from_dict({"enabled": False}) is None
|
||||
|
||||
def test_image_mode_without_path_returns_none(self):
|
||||
result = WatermarkConfig.from_dict(
|
||||
def test_text_mode_basic(self):
|
||||
"""文字水印基本配置."""
|
||||
config = WatermarkConfig.from_dict(
|
||||
{
|
||||
"enabled": True,
|
||||
"mode": "image",
|
||||
"mode": "text",
|
||||
"text": "测试水印",
|
||||
}
|
||||
)
|
||||
assert result is None
|
||||
assert config is not None
|
||||
assert config.mode == "text"
|
||||
assert config.text == "测试水印"
|
||||
assert config.position == "bottom_right" # 默认
|
||||
|
||||
def test_image_mode_with_empty_path_returns_none(self):
|
||||
result = WatermarkConfig.from_dict(
|
||||
{
|
||||
"enabled": True,
|
||||
"mode": "image",
|
||||
"image_path": "",
|
||||
}
|
||||
)
|
||||
assert result is None
|
||||
|
||||
def test_text_mode_without_text_returns_none(self):
|
||||
def test_text_mode_missing_text_returns_none(self):
|
||||
"""文字水印缺少 text 返回 None."""
|
||||
result = WatermarkConfig.from_dict(
|
||||
{
|
||||
"enabled": True,
|
||||
@@ -70,7 +72,8 @@ class TestWatermarkConfigFromDict:
|
||||
)
|
||||
assert result is None
|
||||
|
||||
def test_text_mode_with_empty_text_returns_none(self):
|
||||
def test_text_mode_empty_text_returns_none(self):
|
||||
"""文字水印 text 为空返回 None."""
|
||||
result = WatermarkConfig.from_dict(
|
||||
{
|
||||
"enabled": True,
|
||||
@@ -80,196 +83,375 @@ class TestWatermarkConfigFromDict:
|
||||
)
|
||||
assert result is None
|
||||
|
||||
def test_image_mode_success(self):
|
||||
cfg = WatermarkConfig.from_dict(
|
||||
def test_image_mode_basic(self):
|
||||
"""图片水印基本配置."""
|
||||
config = WatermarkConfig.from_dict(
|
||||
{
|
||||
"enabled": True,
|
||||
"mode": "image",
|
||||
"image_path": "/tmp/logo.png",
|
||||
"scale": 0.3,
|
||||
"opacity": 0.9,
|
||||
"position": "top_left",
|
||||
"margin_x": 30,
|
||||
"margin_y": 30,
|
||||
"image_path": "/path/to/logo.png",
|
||||
}
|
||||
)
|
||||
assert cfg is not None
|
||||
assert cfg.mode == "image"
|
||||
assert cfg.image_path == "/tmp/logo.png"
|
||||
assert cfg.scale == 0.3
|
||||
assert cfg.opacity == 0.9
|
||||
assert cfg.position == "top_left"
|
||||
assert cfg.margin_x == 30
|
||||
assert cfg.margin_y == 30
|
||||
assert config is not None
|
||||
assert config.mode == "image"
|
||||
assert config.image_path == "/path/to/logo.png"
|
||||
|
||||
def test_image_mode_image_key_fallback(self):
|
||||
"""image 字段作为 image_path 的 fallback."""
|
||||
cfg = WatermarkConfig.from_dict(
|
||||
def test_image_mode_missing_image_returns_none(self):
|
||||
"""图片水印缺少 image_path 返回 None."""
|
||||
result = WatermarkConfig.from_dict(
|
||||
{
|
||||
"enabled": True,
|
||||
"mode": "image",
|
||||
"image": "/tmp/fallback.png",
|
||||
}
|
||||
)
|
||||
assert cfg is not None
|
||||
assert cfg.image_path == "/tmp/fallback.png"
|
||||
assert result is None
|
||||
|
||||
def test_text_mode_success(self):
|
||||
cfg = WatermarkConfig.from_dict(
|
||||
def test_image_mode_image_alias(self):
|
||||
"""image 字段作为 image_path 的别名."""
|
||||
config = WatermarkConfig.from_dict(
|
||||
{
|
||||
"enabled": True,
|
||||
"mode": "image",
|
||||
"image": "/path/alias.png",
|
||||
}
|
||||
)
|
||||
assert config is not None
|
||||
assert config.image_path == "/path/alias.png"
|
||||
|
||||
def test_invalid_position_falls_back(self):
|
||||
"""无效位置 fallback 到 bottom_right."""
|
||||
config = WatermarkConfig.from_dict(
|
||||
{
|
||||
"enabled": True,
|
||||
"mode": "text",
|
||||
"text": "hello world",
|
||||
"text": "test",
|
||||
"position": "invalid_pos",
|
||||
}
|
||||
)
|
||||
assert config is not None
|
||||
assert config.position == "bottom_right"
|
||||
|
||||
def test_custom_position_valid(self):
|
||||
"""自定义有效位置."""
|
||||
config = WatermarkConfig.from_dict(
|
||||
{
|
||||
"enabled": True,
|
||||
"mode": "text",
|
||||
"text": "test",
|
||||
"position": "top_left",
|
||||
}
|
||||
)
|
||||
assert config is not None
|
||||
assert config.position == "top_left"
|
||||
|
||||
def test_all_text_fields_parsed(self):
|
||||
"""文字水印所有字段正确解析."""
|
||||
config = WatermarkConfig.from_dict(
|
||||
{
|
||||
"enabled": True,
|
||||
"mode": "text",
|
||||
"text": "我的水印",
|
||||
"font_size": 32,
|
||||
"font_color": "red",
|
||||
"font_path": "/fonts/msyh.ttf",
|
||||
"position": "top_center",
|
||||
"opacity": 0.5,
|
||||
"margin_x": 30,
|
||||
"margin_y": 40,
|
||||
}
|
||||
)
|
||||
assert config is not None
|
||||
assert config.text == "我的水印"
|
||||
assert config.font_size == 32
|
||||
assert config.font_color == "red"
|
||||
assert config.font_path == "/fonts/msyh.ttf"
|
||||
assert config.position == "top_center"
|
||||
assert config.opacity == 0.5
|
||||
assert config.margin_x == 30
|
||||
assert config.margin_y == 40
|
||||
|
||||
def test_all_image_fields_parsed(self):
|
||||
"""图片水印所有字段正确解析."""
|
||||
config = WatermarkConfig.from_dict(
|
||||
{
|
||||
"enabled": True,
|
||||
"mode": "image",
|
||||
"image_path": "/img/logo.png",
|
||||
"scale": 0.3,
|
||||
"opacity": 0.9,
|
||||
"position": "bottom_left",
|
||||
"margin_x": 10,
|
||||
"margin_y": 15,
|
||||
}
|
||||
)
|
||||
assert config is not None
|
||||
assert config.image_path == "/img/logo.png"
|
||||
assert config.scale == 0.3
|
||||
assert config.opacity == 0.9
|
||||
assert config.position == "bottom_left"
|
||||
|
||||
def test_scroll_config_parsed(self):
|
||||
"""滚动水印配置解析."""
|
||||
config = WatermarkConfig.from_dict(
|
||||
{
|
||||
"enabled": True,
|
||||
"mode": "text",
|
||||
"text": "滚动水印",
|
||||
"scroll": True,
|
||||
"scroll_speed": 100,
|
||||
"scroll_speed": 80,
|
||||
}
|
||||
)
|
||||
assert cfg is not None
|
||||
assert cfg.mode == "text"
|
||||
assert cfg.text == "hello world"
|
||||
assert cfg.font_size == 32
|
||||
assert cfg.font_color == "red"
|
||||
assert cfg.position == "bottom_left"
|
||||
assert cfg.scroll is True
|
||||
assert cfg.scroll_speed == 100
|
||||
assert config is not None
|
||||
assert config.scroll is True
|
||||
assert config.scroll_speed == 80
|
||||
|
||||
def test_invalid_position_falls_back_to_bottom_right(self):
|
||||
cfg = WatermarkConfig.from_dict(
|
||||
def test_default_mode_is_text(self):
|
||||
"""不传 mode 默认为 text."""
|
||||
config = WatermarkConfig.from_dict(
|
||||
{
|
||||
"enabled": True,
|
||||
"mode": "text",
|
||||
"text": "test",
|
||||
"position": "invalid_position",
|
||||
"text": "默认模式",
|
||||
}
|
||||
)
|
||||
assert cfg is not None
|
||||
assert cfg.position == "bottom_right"
|
||||
|
||||
def test_default_values_applied(self):
|
||||
cfg = WatermarkConfig.from_dict(
|
||||
{
|
||||
"enabled": True,
|
||||
"mode": "text",
|
||||
"text": "test",
|
||||
}
|
||||
)
|
||||
assert cfg is not None
|
||||
assert cfg.position == "bottom_right"
|
||||
assert cfg.opacity == 0.8
|
||||
assert cfg.scale == 0.2
|
||||
assert cfg.font_size == 24
|
||||
assert cfg.font_color == "white"
|
||||
assert cfg.margin_x == 20
|
||||
assert cfg.margin_y == 20
|
||||
assert cfg.scroll is False
|
||||
assert cfg.scroll_speed == 50
|
||||
assert config is not None
|
||||
assert config.mode == "text"
|
||||
|
||||
|
||||
class TestWatermarkConfigValidate:
|
||||
"""validate 校验逻辑."""
|
||||
|
||||
def test_valid_image_config(self):
|
||||
cfg = WatermarkConfig(
|
||||
mode="image",
|
||||
image_path="/tmp/logo.png",
|
||||
position="top_right",
|
||||
opacity=0.5,
|
||||
scale=0.5,
|
||||
)
|
||||
ok, msg = cfg.validate()
|
||||
assert ok is True
|
||||
assert msg == ""
|
||||
"""validate 配置校验测试."""
|
||||
|
||||
def test_valid_text_config(self):
|
||||
cfg = WatermarkConfig(
|
||||
mode="text",
|
||||
text="hello",
|
||||
position="center",
|
||||
opacity=1.0,
|
||||
font_size=48,
|
||||
)
|
||||
ok, msg = cfg.validate()
|
||||
"""合法文字水印配置."""
|
||||
config = WatermarkConfig(mode="text", text="测试", position="bottom_right")
|
||||
ok, msg = config.validate()
|
||||
assert ok is True
|
||||
assert msg == ""
|
||||
|
||||
def test_valid_image_config(self):
|
||||
"""合法图片水印配置."""
|
||||
config = WatermarkConfig(
|
||||
mode="image",
|
||||
image_path="/a.png",
|
||||
position="top_left",
|
||||
scale=0.3,
|
||||
opacity=0.8,
|
||||
)
|
||||
ok, msg = config.validate()
|
||||
assert ok is True
|
||||
|
||||
def test_invalid_position(self):
|
||||
cfg = WatermarkConfig(mode="text", text="test", position="nowhere")
|
||||
ok, msg = cfg.validate()
|
||||
"""无效位置."""
|
||||
config = WatermarkConfig(mode="text", text="test", position="invalid")
|
||||
ok, msg = config.validate()
|
||||
assert ok is False
|
||||
assert "不支持的位置" in msg
|
||||
|
||||
def test_opacity_below_zero(self):
|
||||
cfg = WatermarkConfig(mode="text", text="test", opacity=-0.1)
|
||||
ok, msg = cfg.validate()
|
||||
def test_opacity_too_high(self):
|
||||
"""透明度超过1."""
|
||||
config = WatermarkConfig(mode="text", text="test", opacity=1.5)
|
||||
ok, msg = config.validate()
|
||||
assert ok is False
|
||||
assert "透明度" in msg
|
||||
|
||||
def test_opacity_above_one(self):
|
||||
cfg = WatermarkConfig(mode="text", text="test", opacity=1.5)
|
||||
ok, msg = cfg.validate()
|
||||
def test_opacity_negative(self):
|
||||
"""透明度为负."""
|
||||
config = WatermarkConfig(mode="text", text="test", opacity=-0.1)
|
||||
ok, msg = config.validate()
|
||||
assert ok is False
|
||||
assert "透明度" in msg
|
||||
|
||||
def test_opacity_zero_is_valid(self):
|
||||
cfg = WatermarkConfig(mode="text", text="test", opacity=0.0)
|
||||
ok, _ = cfg.validate()
|
||||
def test_opacity_boundary_zero(self):
|
||||
"""透明度边界值0."""
|
||||
config = WatermarkConfig(mode="text", text="test", opacity=0.0)
|
||||
ok, _ = config.validate()
|
||||
assert ok is True
|
||||
|
||||
def test_opacity_one_is_valid(self):
|
||||
cfg = WatermarkConfig(mode="text", text="test", opacity=1.0)
|
||||
ok, _ = cfg.validate()
|
||||
def test_opacity_boundary_one(self):
|
||||
"""透明度边界值1."""
|
||||
config = WatermarkConfig(mode="text", text="test", opacity=1.0)
|
||||
ok, _ = config.validate()
|
||||
assert ok is True
|
||||
|
||||
def test_image_missing_path(self):
|
||||
cfg = WatermarkConfig(mode="image", image_path="")
|
||||
ok, msg = cfg.validate()
|
||||
"""图片水印缺少路径."""
|
||||
config = WatermarkConfig(mode="image", image_path="")
|
||||
ok, msg = config.validate()
|
||||
assert ok is False
|
||||
assert "图片路径" in msg
|
||||
|
||||
def test_image_scale_too_small(self):
|
||||
cfg = WatermarkConfig(mode="image", image_path="/tmp/a.png", scale=0.001)
|
||||
ok, msg = cfg.validate()
|
||||
"""缩放比例太小."""
|
||||
config = WatermarkConfig(mode="image", image_path="/a.png", scale=0.001)
|
||||
ok, msg = config.validate()
|
||||
assert ok is False
|
||||
assert "缩放比例" in msg
|
||||
|
||||
def test_image_scale_too_large(self):
|
||||
cfg = WatermarkConfig(mode="image", image_path="/tmp/a.png", scale=2.0)
|
||||
ok, msg = cfg.validate()
|
||||
"""缩放比例太大."""
|
||||
config = WatermarkConfig(mode="image", image_path="/a.png", scale=2.0)
|
||||
ok, msg = config.validate()
|
||||
assert ok is False
|
||||
assert "缩放比例" in msg
|
||||
|
||||
def test_image_scale_boundary_valid(self):
|
||||
cfg = WatermarkConfig(mode="image", image_path="/tmp/a.png", scale=0.01)
|
||||
ok, _ = cfg.validate()
|
||||
def test_image_scale_boundary_low(self):
|
||||
"""缩放边界低值."""
|
||||
config = WatermarkConfig(mode="image", image_path="/a.png", scale=0.01)
|
||||
ok, _ = config.validate()
|
||||
assert ok is True
|
||||
|
||||
cfg2 = WatermarkConfig(mode="image", image_path="/tmp/a.png", scale=1.0)
|
||||
ok2, _ = cfg2.validate()
|
||||
assert ok2 is True
|
||||
def test_image_scale_boundary_high(self):
|
||||
"""缩放边界高值."""
|
||||
config = WatermarkConfig(mode="image", image_path="/a.png", scale=1.0)
|
||||
ok, _ = config.validate()
|
||||
assert ok is True
|
||||
|
||||
def test_text_missing_text(self):
|
||||
cfg = WatermarkConfig(mode="text", text="")
|
||||
ok, msg = cfg.validate()
|
||||
def test_text_missing_content(self):
|
||||
"""文字水印缺少内容."""
|
||||
config = WatermarkConfig(mode="text", text="")
|
||||
ok, msg = config.validate()
|
||||
assert ok is False
|
||||
assert "文字内容" in msg
|
||||
|
||||
def test_text_font_size_zero(self):
|
||||
cfg = WatermarkConfig(mode="text", text="test", font_size=0)
|
||||
ok, msg = cfg.validate()
|
||||
"""字体大小为0."""
|
||||
config = WatermarkConfig(mode="text", text="test", font_size=0)
|
||||
ok, msg = config.validate()
|
||||
assert ok is False
|
||||
assert "字体大小" in msg
|
||||
|
||||
def test_text_font_size_negative(self):
|
||||
cfg = WatermarkConfig(mode="text", text="test", font_size=-5)
|
||||
ok, msg = cfg.validate()
|
||||
"""字体大小为负."""
|
||||
config = WatermarkConfig(mode="text", text="test", font_size=-5)
|
||||
ok, msg = config.validate()
|
||||
assert ok is False
|
||||
assert "字体大小" in msg
|
||||
|
||||
def test_unsupported_mode(self):
|
||||
cfg = WatermarkConfig(mode="video", text="test")
|
||||
ok, msg = cfg.validate()
|
||||
def test_unknown_mode(self):
|
||||
"""未知模式."""
|
||||
config = WatermarkConfig(mode="unknown_mode")
|
||||
ok, msg = config.validate()
|
||||
assert ok is False
|
||||
assert "不支持的水印模式" in msg
|
||||
|
||||
|
||||
# ── WatermarkEngine 位置计算测试 ────────────────────────────────
|
||||
|
||||
|
||||
class TestCalcPosition:
|
||||
"""9宫格位置计算测试."""
|
||||
|
||||
# 测试用:输出 1920x1080,水印 200x100,边距 20
|
||||
W, H = 1920, 1080
|
||||
WW, WH = 200, 100
|
||||
MX, MY = 20, 20
|
||||
|
||||
def test_top_left(self):
|
||||
"""左上角."""
|
||||
x, y = WatermarkEngine.calc_position("top_left", self.W, self.H, self.WW, self.WH, self.MX, self.MY)
|
||||
assert (x, y) == (20, 20)
|
||||
|
||||
def test_top_center(self):
|
||||
"""中上."""
|
||||
x, y = WatermarkEngine.calc_position("top_center", self.W, self.H, self.WW, self.WH, self.MX, self.MY)
|
||||
assert x == (1920 - 200) // 2
|
||||
assert y == 20
|
||||
|
||||
def test_top_right(self):
|
||||
"""右上角."""
|
||||
x, y = WatermarkEngine.calc_position("top_right", self.W, self.H, self.WW, self.WH, self.MX, self.MY)
|
||||
assert x == 1920 - 200 - 20
|
||||
assert y == 20
|
||||
|
||||
def test_center_left(self):
|
||||
"""左中."""
|
||||
x, y = WatermarkEngine.calc_position("center_left", self.W, self.H, self.WW, self.WH, self.MX, self.MY)
|
||||
assert x == 20
|
||||
assert y == (1080 - 100) // 2
|
||||
|
||||
def test_center(self):
|
||||
"""中心."""
|
||||
x, y = WatermarkEngine.calc_position("center", self.W, self.H, self.WW, self.WH, self.MX, self.MY)
|
||||
assert x == (1920 - 200) // 2
|
||||
assert y == (1080 - 100) // 2
|
||||
|
||||
def test_center_right(self):
|
||||
"""右中."""
|
||||
x, y = WatermarkEngine.calc_position("center_right", self.W, self.H, self.WW, self.WH, self.MX, self.MY)
|
||||
assert x == 1920 - 200 - 20
|
||||
assert y == (1080 - 100) // 2
|
||||
|
||||
def test_bottom_left(self):
|
||||
"""左下角."""
|
||||
x, y = WatermarkEngine.calc_position("bottom_left", self.W, self.H, self.WW, self.WH, self.MX, self.MY)
|
||||
assert x == 20
|
||||
assert y == 1080 - 100 - 20
|
||||
|
||||
def test_bottom_center(self):
|
||||
"""中下."""
|
||||
x, y = WatermarkEngine.calc_position("bottom_center", self.W, self.H, self.WW, self.WH, self.MX, self.MY)
|
||||
assert x == (1920 - 200) // 2
|
||||
assert y == 1080 - 100 - 20
|
||||
|
||||
def test_bottom_right(self):
|
||||
"""右下角."""
|
||||
x, y = WatermarkEngine.calc_position("bottom_right", self.W, self.H, self.WW, self.WH, self.MX, self.MY)
|
||||
assert x == 1920 - 200 - 20
|
||||
assert y == 1080 - 100 - 20
|
||||
|
||||
def test_unknown_position_defaults_bottom_right(self):
|
||||
"""未知位置默认右下角."""
|
||||
x, y = WatermarkEngine.calc_position("unknown", self.W, self.H, self.WW, self.WH, self.MX, self.MY)
|
||||
assert x == 1920 - 200 - 20
|
||||
assert y == 1080 - 100 - 20
|
||||
|
||||
def test_zero_margin(self):
|
||||
"""零边距."""
|
||||
x, y = WatermarkEngine.calc_position("top_left", 1000, 500, 100, 50, 0, 0)
|
||||
assert (x, y) == (0, 0)
|
||||
|
||||
def test_small_output(self):
|
||||
"""小尺寸输出."""
|
||||
x, y = WatermarkEngine.calc_position("bottom_right", 320, 240, 50, 30, 5, 5)
|
||||
assert x == 320 - 50 - 5
|
||||
assert y == 240 - 30 - 5
|
||||
|
||||
|
||||
class TestCalcScrollX:
|
||||
"""滚动水印x坐标表达式测试."""
|
||||
|
||||
def test_returns_string_expression(self):
|
||||
"""返回字符串表达式."""
|
||||
expr = WatermarkEngine.calc_scroll_x("bottom_right", 1920, 200, 50)
|
||||
assert isinstance(expr, str)
|
||||
assert "1920" in expr
|
||||
assert "200" in expr
|
||||
assert "50" in expr
|
||||
|
||||
def test_contains_mod_function(self):
|
||||
"""包含 mod 函数."""
|
||||
expr = WatermarkEngine.calc_scroll_x("top_left", 1280, 150, 60)
|
||||
assert "mod(" in expr
|
||||
assert "t" in expr # 时间变量
|
||||
|
||||
|
||||
class TestWatermarkPositions:
|
||||
"""位置常量测试."""
|
||||
|
||||
def test_nine_positions(self):
|
||||
"""共9个位置."""
|
||||
assert len(WATERMARK_POSITIONS) == 9
|
||||
|
||||
def test_all_position_keys_valid(self):
|
||||
"""所有位置键名正确."""
|
||||
expected = {
|
||||
"top_left",
|
||||
"top_center",
|
||||
"top_right",
|
||||
"center_left",
|
||||
"center",
|
||||
"center_right",
|
||||
"bottom_left",
|
||||
"bottom_center",
|
||||
"bottom_right",
|
||||
}
|
||||
assert set(WATERMARK_POSITIONS.keys()) == expected
|
||||
|
||||
Reference in New Issue
Block a user